diff --git a/connect/endpoint.go b/connect/endpoint.go index acccaeb6..f7418e87 100644 --- a/connect/endpoint.go +++ b/connect/endpoint.go @@ -25,11 +25,7 @@ type Endpoints struct { func MakeConnectEndpoint(svc ConnectService) endpoint.Endpoint { return func(ctx context.Context, request interface{}) (interface{}, error) { req := request.(mdmConnectRequest) - if req.UserID != nil { - // don't handle user - return mdmConnectResponse{}, nil - } payload, err := svc.Acknowledge(ctx, req.Response) - return mdmConnectResponse{payload, err}, nil + return mdmConnectResponse{payload: payload, Err: err}, nil } } diff --git a/connect/event.go b/connect/event.go index 263f6054..3fda037a 100644 --- a/connect/event.go +++ b/connect/event.go @@ -33,7 +33,7 @@ func MarshalEvent(e *Event) ([]byte, error) { RequestType: e.Response.RequestType, } if e.Response.UserID != nil { - response.Udid = *e.Response.UserID + response.UserId = *e.Response.UserID } return proto.Marshal(&connectproto.Event{ @@ -65,5 +65,8 @@ func UnmarshalEvent(data []byte, e *Event) error { } func strPtr(s string) *string { + if s == "" { + return nil + } return &s } diff --git a/push/db.go b/push/db.go index 6ef51b03..ccf4aff8 100644 --- a/push/db.go +++ b/push/db.go @@ -93,15 +93,16 @@ func (db *DB) pollCheckin(sub pubsub.Subscriber) error { fmt.Println(err) continue } - if ev.Command.UserID != "" { - continue - } info := PushInfo{ UDID: ev.Command.UDID, Token: ev.Command.Token.String(), PushMagic: ev.Command.PushMagic, MDMTopic: ev.Command.Topic, } + if ev.Command.UserID != "" { + // use the GUID if this is a user TokenUpdate. + info.UDID = ev.Command.UserID + } if err := db.Save(&info); err != nil { fmt.Println(err) continue diff --git a/queue/queue.go b/queue/queue.go index 50cb8951..6d5707e4 100644 --- a/queue/queue.go +++ b/queue/queue.go @@ -25,7 +25,12 @@ type Store struct { } func (db *Store) Next(ctx context.Context, resp mdm.Response) (*Command, error) { - dc, err := db.DeviceCommand(resp.UDID) + udid := resp.UDID + if resp.UserID != nil { + // use the user id for user level commands + udid = *resp.UserID + } + dc, err := db.DeviceCommand(udid) if err != nil { if isNotFound(err) { return nil, nil @@ -79,12 +84,12 @@ func (db *Store) Next(ctx context.Context, resp mdm.Response) (*Command, error) } // pop the first command from the queue and add it to the end. - // If the regular queue is empty, send a command that got + // If the regular queue is empty, send a command that got // refused with NotNow before. cmd, dc.Commands = popFirst(dc.Commands) if cmd != nil { dc.Commands = append(dc.Commands, *cmd) - } else if (resp.Status != "NotNow") { + } else if resp.Status != "NotNow" { cmd, dc.NotNow = popFirst(dc.NotNow) if cmd != nil { dc.Commands = append(dc.Commands, *cmd) diff --git a/serve.go b/serve.go index dd679768..c1aaf173 100644 --- a/serve.go +++ b/serve.go @@ -157,7 +157,7 @@ func serve(args []string) error { stdlog.Fatal(err) } - userDB, err := user.NewDB(sm.db, sm.pubclient) + userDB, err := user.NewDB(sm.db, sm.pubclient, log.With(logger, "component", "user db")) if err != nil { stdlog.Fatal(err) } diff --git a/user/db.go b/user/db.go index 5c5b8302..4bd0fa27 100644 --- a/user/db.go +++ b/user/db.go @@ -1,11 +1,16 @@ package user 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/checkin" "github.com/micromdm/micromdm/pubsub" ) @@ -17,9 +22,10 @@ const ( type DB struct { *bolt.DB + logger log.Logger } -func NewDB(db *bolt.DB, pubsubSvc pubsub.PublishSubscriber) (*DB, error) { +func NewDB(db *bolt.DB, pubsubSvc pubsub.PublishSubscriber, logger log.Logger) (*DB, error) { err := db.Update(func(tx *bolt.Tx) error { _, err := tx.CreateBucketIfNotExists([]byte(userIndexBucket)) if err != nil { @@ -33,9 +39,15 @@ func NewDB(db *bolt.DB, pubsubSvc pubsub.PublishSubscriber) (*DB, error) { } datastore := &DB{ - DB: 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 } - return datastore, nil } @@ -105,7 +117,7 @@ func (db *DB) User(uuid string) (*User, error) { return UnmarshalUser(v, &u) }) if err != nil { - return nil, errors.Wrap(err, "get user by user id from bolt") + return nil, errors.Wrap(err, "get user by uuid from bolt") } return &u, nil } @@ -131,6 +143,49 @@ func (db *DB) UserByUserID(userID string) (*User, error) { return &u, nil } +func (db *DB) DeviceUsers(udid string) ([]User, error) { + var users []User + err := db.View(func(tx *bolt.Tx) error { + b := tx.Bucket([]byte(UserBucket)) + c := b.Cursor() + for k, v := c.First(); k != nil; k, v = c.Next() { + var u User + if err := UnmarshalUser(v, &u); err != nil { + return errors.Wrap(err, "unmarshal user for DeviceUsers") + } + if u.UDID == udid { + users = append(users, u) + } + } + return nil + }) + if err != nil { + return nil, errors.Wrap(err, "get device users") + } + return users, nil +} + +func (db *DB) DeleteDeviceUsers(udid string) error { + err := db.Update(func(tx *bolt.Tx) error { + b := tx.Bucket([]byte(UserBucket)) + c := b.Cursor() + for k, v := c.First(); k != nil; k, v = c.Next() { + var u User + if err := UnmarshalUser(v, &u); err != nil { + return errors.Wrap(err, "unmarshal user for DeviceUsers") + } + if u.UDID != udid { + continue + } + if err := b.Delete(k); err != nil { + return errors.Wrapf(err, "delete user %s from device %s", u.UserID, udid) + } + } + return nil + }) + return errors.Wrapf(err, "delete users for UDID %s", udid) +} + type notFound struct { ResourceType string Message string @@ -139,3 +194,71 @@ type notFound struct { func (e *notFound) Error() string { return fmt.Sprintf("not found: %s %s", e.ResourceType, e.Message) } + +func (db *DB) pollCheckin(pubsubSvc pubsub.PublishSubscriber) error { + tokenUpdateEvents, err := pubsubSvc.Subscribe(context.TODO(), "users", checkin.TokenUpdateTopic) + if err != nil { + return errors.Wrapf(err, + "subscribing devices to %s topic", checkin.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) + 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) (checkin.Event, error) { + var ev checkin.Event + if err := checkin.UnmarshalEvent(event.Message, &ev); err != nil { + return checkin.Event{}, err + } + return ev, nil +} + +func isNotFound(err error) bool { + cause := errors.Cause(err) + if _, ok := cause.(*notFound); ok { + return true + } + return false +}