diff --git a/cmd/micromdm/serve.go b/cmd/micromdm/serve.go index a03df8a2..8d18913b 100644 --- a/cmd/micromdm/serve.go +++ b/cmd/micromdm/serve.go @@ -188,10 +188,12 @@ func serve(args []string) error { devWorker := device.NewWorker(devDB, sm.pubclient, logger) go devWorker.Run(context.Background()) - userDB, err := userbuiltin.NewDB(sm.db, sm.pubclient, log.With(logger, "component", "user db")) + userDB, err := userbuiltin.NewDB(sm.db) if err != nil { stdlog.Fatal(err) } + userWorker := user.NewWorker(userDB, sm.pubclient, logger) + go userWorker.Run(context.Background()) sm.profileDB, err = profilebuiltin.NewDB(sm.db) if err != nil { diff --git a/platform/user/builtin/db.go b/platform/user/builtin/db.go index 0999b35e..8d7283f0 100644 --- a/platform/user/builtin/db.go +++ b/platform/user/builtin/db.go @@ -1,17 +1,12 @@ package builtin import ( - "context" "fmt" "github.com/boltdb/bolt" "github.com/go-kit/kit/log" - "github.com/go-kit/kit/log/level" "github.com/pkg/errors" - uuid "github.com/satori/go.uuid" - "github.com/micromdm/micromdm/mdm" - "github.com/micromdm/micromdm/platform/pubsub" "github.com/micromdm/micromdm/platform/user" ) @@ -26,7 +21,7 @@ type DB struct { logger log.Logger } -func NewDB(db *bolt.DB, pubsubSvc pubsub.PublishSubscriber, logger log.Logger) (*DB, error) { +func NewDB(db *bolt.DB) (*DB, error) { err := db.Update(func(tx *bolt.Tx) error { _, err := tx.CreateBucketIfNotExists([]byte(userIndexBucket)) if err != nil { @@ -40,14 +35,7 @@ func NewDB(db *bolt.DB, pubsubSvc pubsub.PublishSubscriber, logger log.Logger) ( } datastore := &DB{ - DB: db, - logger: logger, - } - if pubsubSvc == nil { // don't start the poller without pubsub. - return datastore, nil - } - if err := datastore.pollCheckin(pubsubSvc); err != nil { - return nil, err + DB: db, } return datastore, nil } @@ -199,71 +187,3 @@ func (e *notFound) Error() string { func (e *notFound) NotFound() bool { return true } - -func (db *DB) pollCheckin(pubsubSvc pubsub.PublishSubscriber) error { - tokenUpdateEvents, err := pubsubSvc.Subscribe(context.TODO(), "users", mdm.TokenUpdateTopic) - if err != nil { - return errors.Wrapf(err, - "subscribing devices to %s topic", mdm.TokenUpdateTopic) - } - go func() { - for { - select { - case e := <-tokenUpdateEvents: - event, err := unmarshalCheckin(e) - if err != nil { - level.Info(db.logger).Log("err", err, "msg", "unmarshal TokenUpdate event in user db") - break - } - if event.Command.UserID == "" { - break // only interested in user commands - } - newUser := new(user.User) - byGUID, err := db.UserByUserID(event.Command.UserID) - if err != nil && !isNotFound(err) { - level.Info(db.logger).Log("err", err, "msg", "get user from DB") - break - } - if err == nil && byGUID != nil { - newUser = byGUID - } - if newUser.UUID == "" { - if err := db.DeleteDeviceUsers(event.Command.UDID); err != nil { - level.Info(db.logger).Log( - "err", err, - "msg", "delete existing user before creating new one", - ) - } - newUser.UUID = uuid.NewV4().String() - } - newUser.UDID = event.Command.UDID - newUser.UserID = event.Command.UserID - newUser.UserLongname = event.Command.UserLongName - newUser.UserShortname = event.Command.UserShortName - newUser.AuthToken = event.Command.Token.String() - if err := db.Save(newUser); err != nil { - level.Info(db.logger).Log("err", err, "msg", "update user from TokenUpdate") - break - } - } - } - }() - - return nil -} - -func unmarshalCheckin(event pubsub.Event) (mdm.CheckinEvent, error) { - var ev mdm.CheckinEvent - if err := mdm.UnmarshalCheckinEvent(event.Message, &ev); err != nil { - return mdm.CheckinEvent{}, err - } - return ev, nil -} - -func isNotFound(err error) bool { - cause := errors.Cause(err) - if _, ok := cause.(*notFound); ok { - return true - } - return false -} diff --git a/platform/user/worker.go b/platform/user/worker.go new file mode 100644 index 00000000..bf9f547f --- /dev/null +++ b/platform/user/worker.go @@ -0,0 +1,119 @@ +package user + +import ( + "context" + + "github.com/go-kit/kit/log" + "github.com/go-kit/kit/log/level" + "github.com/pkg/errors" + uuid "github.com/satori/go.uuid" + + "github.com/micromdm/micromdm/mdm" + "github.com/micromdm/micromdm/platform/pubsub" +) + +type WorkerStore interface { + Save(*User) error + DeleteDeviceUsers(udid string) error + UserByUserID(userID string) (*User, error) +} + +type Worker struct { + db WorkerStore + sub pubsub.Subscriber + logger log.Logger +} + +func NewWorker(db WorkerStore, subscriber pubsub.Subscriber, logger log.Logger) *Worker { + return &Worker{ + db: db, + sub: subscriber, + logger: logger, + } +} + +func (w *Worker) Run(ctx context.Context) error { + const subscription = "user_worker" + tokenUpdateEvents, err := w.sub.Subscribe(ctx, subscription, mdm.TokenUpdateTopic) + if err != nil { + return errors.Wrapf(err, "subcribe %s to %s", subscription, mdm.TokenUpdateTopic) + } + + for { + var err error + select { + case <-ctx.Done(): + return ctx.Err() + case ev := <-tokenUpdateEvents: + err = w.updateUserFromTokenUpdate(ctx, ev.Message) + } + + if err != nil { + level.Info(w.logger).Log( + "msg", "update user from event", + "err", err, + ) + continue + } + } +} + +func (w *Worker) updateUserFromTokenUpdate(ctx context.Context, message []byte) error { + var ev mdm.CheckinEvent + if err := mdm.UnmarshalCheckinEvent(message, &ev); err != nil { + return errors.Wrap(err, "unmarshal checkin event for user worker") + } + + if ev.Command.UserID == "" { + // only process user events + return nil + } + + usr, err := getOrCreateUser(w.db, ev.Command.UserID) + if err != nil { + return err + } + if usr.UUID == "" { + if err := w.db.DeleteDeviceUsers(ev.Command.UDID); err != nil { + return errors.Wrapf( + err, + "delete users for device %s before re-creating. user_id=%s", + ev.Command.UDID, + ev.Command.UserID, + ) + usr.UUID = uuid.NewV4().String() + } + } + + usr.UDID = ev.Command.UDID + usr.UserID = ev.Command.UserID + usr.UserLongname = ev.Command.UserLongName + usr.UserShortname = ev.Command.UserShortName + usr.AuthToken = ev.Command.Token.String() + err = w.db.Save(usr) + return errors.Wrapf(err, "saving user %s to device %s", ev.Command.UserID, ev.Command.UDID) +} + +func getOrCreateUser(db WorkerStore, userID string) (*User, error) { + byGUID, err := db.UserByUserID(userID) + if err == nil { + return byGUID, nil + } + if err != nil && !isNotFound(err) { + return nil, errors.Wrap(err, "get user by ID") + } + + usr := new(User) + return usr, nil +} + +func isNotFound(err error) bool { + err = errors.Cause(err) + type notFoundErr interface { + error + NotFound() bool + } + + e, ok := err.(notFoundErr) + return ok && e.NotFound() +}