From 1ccf1c8ca1cd7e56bea6d3596a45bbd0bbdf1046 Mon Sep 17 00:00:00 2001 From: Victor Vrantchan Date: Wed, 30 May 2018 09:14:41 -0700 Subject: [PATCH] decouple device updates from db CRUD (#427) Move the code for polling topics and updating device records into an independent worker. Cleans up the code and makes it possible to swap the db backend. --- cmd/micromdm/serve.go | 5 +- platform/device/builtin/db.go | 233 +---------------------- platform/device/builtin/db_test.go | 2 +- platform/device/worker.go | 295 +++++++++++++++++++++++++++++ 4 files changed, 306 insertions(+), 229 deletions(-) create mode 100644 platform/device/worker.go diff --git a/cmd/micromdm/serve.go b/cmd/micromdm/serve.go index 4850ed9c..a03df8a2 100644 --- a/cmd/micromdm/serve.go +++ b/cmd/micromdm/serve.go @@ -180,11 +180,14 @@ func serve(args []string) error { removeService = block.LoggingMiddleware(logger)(svc) } - devDB, err := devicebuiltin.NewDB(sm.db, sm.pubclient) + devDB, err := devicebuiltin.NewDB(sm.db) if err != nil { stdlog.Fatal(err) } + 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")) if err != nil { stdlog.Fatal(err) diff --git a/platform/device/builtin/db.go b/platform/device/builtin/db.go index 3e66fb80..c4a384d0 100644 --- a/platform/device/builtin/db.go +++ b/platform/device/builtin/db.go @@ -1,18 +1,12 @@ package builtin import ( - "context" "fmt" - "time" "github.com/boltdb/bolt" "github.com/pkg/errors" - uuid "github.com/satori/go.uuid" - "github.com/micromdm/micromdm/dep/depsync" - "github.com/micromdm/micromdm/mdm" "github.com/micromdm/micromdm/platform/device" - "github.com/micromdm/micromdm/platform/pubsub" ) const ( @@ -27,7 +21,7 @@ type DB struct { *bolt.DB } -func NewDB(db *bolt.DB, pubsubSvc pubsub.PublishSubscriber) (*DB, error) { +func NewDB(db *bolt.DB) (*DB, error) { err := db.Update(func(tx *bolt.Tx) error { _, err := tx.CreateBucketIfNotExists([]byte(deviceIndexBucket)) if err != nil { @@ -39,15 +33,7 @@ func NewDB(db *bolt.DB, pubsubSvc pubsub.PublishSubscriber) (*DB, error) { if err != nil { return nil, errors.Wrapf(err, "creating %s bucket", DeviceBucket) } - datastore := &DB{ - DB: db, - } - if pubsubSvc == nil { // don't start the poller without pubsub. - return datastore, nil - } - if err := datastore.pollCheckin(pubsubSvc); err != nil { - return nil, err - } + datastore := &DB{DB: db} return datastore, nil } @@ -124,6 +110,10 @@ func (e *notFound) Error() string { return fmt.Sprintf("not found: %s %s", e.ResourceType, e.Message) } +func (e *notFound) NotFound() bool { + return true +} + func (db *DB) DeviceByUDID(udid string) (*device.Device, error) { var dev device.Device err := db.View(func(tx *bolt.Tx) error { @@ -165,214 +155,3 @@ func (db *DB) DeviceBySerial(serial string) (*device.Device, error) { } return &dev, nil } - -func isNotFound(err error) bool { - if _, ok := err.(*notFound); ok { - return true - } - return false -} - -func (db *DB) pollCheckin(pubsubSvc pubsub.PublishSubscriber) error { - authenticateEvents, err := pubsubSvc.Subscribe(context.TODO(), "devices", mdm.AuthenticateTopic) - if err != nil { - return errors.Wrapf(err, - "subscribing devices to %s topic", mdm.AuthenticateTopic) - } - tokenUpdateEvents, err := pubsubSvc.Subscribe(context.TODO(), "devices", mdm.TokenUpdateTopic) - if err != nil { - return errors.Wrapf(err, - "subscribing devices to %s topic", mdm.TokenUpdateTopic) - } - checkoutEvents, err := pubsubSvc.Subscribe(context.TODO(), "devices", mdm.CheckoutTopic) - if err != nil { - return errors.Wrapf(err, - "subscribing devices to %s topic", mdm.CheckoutTopic) - } - depSyncEvents, err := pubsubSvc.Subscribe(context.TODO(), "devices", depsync.SyncTopic) - if err != nil { - return errors.Wrapf(err, - "subscribing devices to %s topic", depsync.SyncTopic) - } - connectEvents, err := pubsubSvc.Subscribe(context.TODO(), "devices", mdm.ConnectTopic) - if err != nil { - return errors.Wrapf(err, - "subscribing devices to %s topic", mdm.ConnectTopic) - } - go func() { - for { - select { - case event := <-authenticateEvents: - var ev mdm.CheckinEvent - if err := mdm.UnmarshalCheckinEvent(event.Message, &ev); err != nil { - fmt.Println(err) - continue - } - newDevice := new(device.Device) - bySerial, err := db.DeviceBySerial(ev.Command.SerialNumber) - if err == nil && bySerial != nil { // must be a DEP device - newDevice = bySerial - } - if err != nil && !isNotFound(err) { - fmt.Println(err) // some other issue is going on - continue - } - byUDID, err := db.DeviceByUDID(ev.Command.UDID) - if err != nil && isNotFound(err) { // never checked in - fmt.Printf("checking in new device %s\n", ev.Command.SerialNumber) - } else if err != nil { - fmt.Println(err) - continue - } else if err == nil { - fmt.Printf("re-enrolling device %s\n", ev.Command.SerialNumber) - newDevice = byUDID - newDevice.Enrolled = false - } - - // only create new UUID on initial enrollment. - if newDevice.UUID == "" { - newDevice.UUID = uuid.NewV4().String() - } - newDevice.UDID = ev.Command.UDID - newDevice.OSVersion = ev.Command.OSVersion - newDevice.BuildVersion = ev.Command.BuildVersion - newDevice.ProductName = ev.Command.ProductName - newDevice.SerialNumber = ev.Command.SerialNumber - newDevice.IMEI = ev.Command.IMEI - newDevice.MEID = ev.Command.MEID - newDevice.DeviceName = ev.Command.DeviceName - newDevice.Model = ev.Command.Model - newDevice.ModelName = ev.Command.ModelName - newDevice.LastSeen = time.Now() - // Challenge: ev.Command.Challenge, // FIXME: @groob why is this commented out? - - if err := db.Save(newDevice); err != nil { - fmt.Println(err) - continue - } - case event := <-tokenUpdateEvents: - var ev mdm.CheckinEvent - if err := mdm.UnmarshalCheckinEvent(event.Message, &ev); err != nil { - fmt.Println(err) - continue - } - if ev.Command.UserID != "" { - continue - } - dev, err := db.DeviceByUDID(ev.Command.UDID) - if err != nil { - fmt.Println(err) - continue - } - dev.Token = ev.Command.Token.String() - dev.PushMagic = ev.Command.PushMagic - dev.UnlockToken = ev.Command.UnlockToken.String() - dev.AwaitingConfiguration = ev.Command.AwaitingConfiguration - dev.LastSeen = time.Now() - var newlyEnrolled bool = false - if dev.Enrolled == false { - newlyEnrolled = true - dev.Enrolled = true - } - if err := db.Save(dev); err != nil { - fmt.Println(err) - continue - } - if newlyEnrolled { - fmt.Printf("device %s enrolled\n", ev.Command.UDID) - err := pubsubSvc.Publish(context.TODO(), device.DeviceEnrolledTopic, event.Message) - if err != nil { - fmt.Println(err) - } - } - case event := <-depSyncEvents: - var ev depsync.Event - if err := depsync.UnmarshalEvent(event.Message, &ev); err != nil { - fmt.Println(err) - continue - } - fmt.Printf("got %d devices from DEP\n", len(ev.Devices)) - for _, d := range ev.Devices { - updDevice, updDeviceErr := db.DeviceBySerial(d.SerialNumber) - if updDeviceErr != nil && !isNotFound(updDeviceErr) { - fmt.Printf("error getting device %s: %s\n", d.SerialNumber, err) - continue - } - - if updDeviceErr != nil && isNotFound(updDeviceErr) { - updDevice = new(device.Device) - if d.OpType == "modified" { - fmt.Printf("warning: no existing device for DEP device update: %s\n", d.SerialNumber) - } - } - if updDeviceErr == nil && d.OpType == "added" { - // consider issuing this warning if op_type == "" as well. - // in that case it's likely the device came from a DEP - // fetch (vs. a sync) which could be a re-fetch of devices - fmt.Printf("warning: re-adding existing DEP device: %s\n", d.SerialNumber) - } - if d.OpType == "deleted" { - fmt.Printf("warning: DEP device unassigned: %s\n", d.SerialNumber) - } - - if updDevice.UUID == "" { - // generate UUID for any device that doesn't have one - updDevice.UUID = uuid.NewV4().String() - } - - updDevice.SerialNumber = d.SerialNumber - updDevice.Model = d.Model - updDevice.Description = d.Description - updDevice.Color = d.Color - updDevice.AssetTag = d.AssetTag - updDevice.DEPProfileStatus = device.DEPProfileStatus(d.ProfileStatus) - updDevice.DEPProfileUUID = d.ProfileUUID - updDevice.DEPProfileAssignTime = d.ProfileAssignTime - updDevice.DEPProfileAssignedDate = d.DeviceAssignedDate - updDevice.DEPProfileAssignedBy = d.DeviceAssignedBy - // TODO: support profile_push_time, os, device_family, op_date - - if err := db.Save(updDevice); err != nil { - fmt.Println(err) - continue - } - } - case event := <-connectEvents: - var ev mdm.AcknowledgeEvent - if err := mdm.UnmarshalAcknowledgeEvent(event.Message, &ev); err != nil { - fmt.Println(err) - continue - } - dev, err := db.DeviceByUDID(ev.Response.UDID) - if err != nil { - fmt.Println(err) - continue - } - dev.LastSeen = time.Now() - if err := db.Save(dev); err != nil { - fmt.Println(err) - continue - } - case event := <-checkoutEvents: - var ev mdm.CheckinEvent - if err := mdm.UnmarshalCheckinEvent(event.Message, &ev); err != nil { - fmt.Println(err) - continue - } - dev, err := db.DeviceByUDID(ev.Command.UDID) - if err != nil { - fmt.Println(err) - continue - } - dev.Enrolled = false - dev.LastSeen = time.Now() - if err := db.Save(dev); err != nil { - fmt.Println(err) - continue - } - } - } - }() - - return nil -} diff --git a/platform/device/builtin/db_test.go b/platform/device/builtin/db_test.go index 59ee508c..e5bf3274 100644 --- a/platform/device/builtin/db_test.go +++ b/platform/device/builtin/db_test.go @@ -73,7 +73,7 @@ func setupDB(t *testing.T) *DB { if err != nil { t.Fatalf("couldn't open bolt, err %s\n", err) } - devDB, err := NewDB(db, nil) + devDB, err := NewDB(db) if err != nil { t.Fatalf("couldn't create device DB, err %s\n", err) } diff --git a/platform/device/worker.go b/platform/device/worker.go new file mode 100644 index 00000000..694bc30e --- /dev/null +++ b/platform/device/worker.go @@ -0,0 +1,295 @@ +package device + +import ( + "context" + "time" + + "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/dep/depsync" + "github.com/micromdm/micromdm/mdm" + "github.com/micromdm/micromdm/platform/pubsub" +) + +type DeviceWorkerStore interface { + Save(*Device) error + DeviceByUDID(udid string) (*Device, error) + DeviceBySerial(udid string) (*Device, error) +} + +type Worker struct { + db DeviceWorkerStore + ps pubsub.PublishSubscriber + logger log.Logger +} + +func NewWorker(db DeviceWorkerStore, ps pubsub.PublishSubscriber, logger log.Logger) *Worker { + return &Worker{ + db: db, + ps: ps, + logger: logger, + } +} + +func (w *Worker) Run(ctx context.Context) error { + const subscription = "devices_worker" + authenticateEvents, err := w.ps.Subscribe(ctx, subscription, mdm.AuthenticateTopic) + if err != nil { + return errors.Wrapf(err, "subscribing %s to %s", subscription, mdm.AuthenticateTopic) + } + tokenUpdateEvents, err := w.ps.Subscribe(ctx, subscription, mdm.TokenUpdateTopic) + if err != nil { + return errors.Wrapf(err, "subscribing %s to %s", subscription, mdm.TokenUpdateTopic) + } + checkoutEvents, err := w.ps.Subscribe(ctx, subscription, mdm.CheckoutTopic) + if err != nil { + return errors.Wrapf(err, "subscribing %s to %s", subscription, mdm.CheckoutTopic) + } + depSyncEvents, err := w.ps.Subscribe(ctx, subscription, depsync.SyncTopic) + if err != nil { + return errors.Wrapf(err, "subscribing %s to %s", subscription, mdm.AuthenticateTopic) + } + + connectEvents, err := w.ps.Subscribe(ctx, subscription, mdm.ConnectTopic) + if err != nil { + return errors.Wrapf(err, "subscribing %s to %s", subscription, mdm.ConnectTopic) + } + + for { + var err error + select { + case <-ctx.Done(): + return ctx.Err() + case ev := <-authenticateEvents: + err = w.updateFromAuthenticate(ctx, ev.Message) + case ev := <-tokenUpdateEvents: + err = w.updateFromTokenUpdate(ctx, ev.Message) + case ev := <-checkoutEvents: + err = w.updateFromCheckout(ctx, ev.Message) + case ev := <-depSyncEvents: + err = w.updateFromDEPSync(ctx, ev.Message) + case ev := <-connectEvents: + err = w.updateFromAcknowledge(ctx, ev.Message) + } + if err != nil { + level.Info(w.logger).Log( + "msg", "update device from event", + "err", err, + ) + continue + } + } +} + +func (w *Worker) updateFromDEPSync(ctx context.Context, message []byte) error { + var ev depsync.Event + if err := depsync.UnmarshalEvent(message, &ev); err != nil { + return errors.Wrap(err, "unmarshal depsync event") + } + level.Debug(w.logger).Log( + "msg", "updating devices from DEP", + "device_count", len(ev.Devices), + ) + + for _, dd := range ev.Devices { + dev, err := getOrCreateDeviceBySerial(w.db, dd.SerialNumber) + if err != nil { + return errors.Wrap(err, "get device by serial") + } + + notSeenBefore := dev.UUID == "" + logEvent := dd.OpType == "deleted" || + (notSeenBefore && dd.OpType == "added") || + (!notSeenBefore && dd.OpType == "modified") + if logEvent { + level.Debug(w.logger).Log( + "msg", "updating devices from dep sync", + "op_type", dd.OpType, + "serial", dd.SerialNumber, + "previously_known", !notSeenBefore, + ) + } + + if dev.UUID == "" { + dev.UUID = uuid.NewV4().String() + } + + dev.SerialNumber = dd.SerialNumber + dev.Model = dd.Model + dev.Description = dd.Description + dev.Color = dd.Color + dev.AssetTag = dd.AssetTag + dev.DEPProfileStatus = DEPProfileStatus(dd.ProfileStatus) + dev.DEPProfileUUID = dd.ProfileUUID + dev.DEPProfileAssignTime = dd.ProfileAssignTime + dev.DEPProfileAssignedDate = dd.DeviceAssignedDate + dev.DEPProfileAssignedBy = dd.DeviceAssignedBy + + if err := w.db.Save(dev); err != nil { + return errors.Wrap(err, "save device %s from DEP sync") + } + } + + return nil +} + +func (w *Worker) updateFromAcknowledge(ctx context.Context, message []byte) error { + var ev mdm.AcknowledgeEvent + if err := mdm.UnmarshalAcknowledgeEvent(message, &ev); err != nil { + return errors.Wrap(err, "unmarshal acknowledge event") + } + + dev, err := w.db.DeviceByUDID(ev.Response.UDID) + if err != nil { + return errors.Wrapf(err, "retrieve device with udid %s", ev.Response.UDID) + } + dev.LastSeen = time.Now() + + err = w.db.Save(dev) + return errors.Wrapf(err, "saving updated device for acknowledge event") + +} + +func (w *Worker) updateFromCheckout(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") + } + + dev, err := w.db.DeviceByUDID(ev.Command.UDID) + if err != nil { + return errors.Wrapf(err, "retrieve device with udid %s", ev.Command.UDID) + } + + dev.Enrolled = false + dev.LastSeen = time.Now() + + err = w.db.Save(dev) + return errors.Wrapf(err, "saving updated device for checkout event") + +} + +func (w *Worker) updateFromTokenUpdate(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") + } + + if ev.Command.UserID != "" { + // do not process user checkin events while updating device records. + return nil + } + + dev, err := w.db.DeviceByUDID(ev.Command.UDID) + if err != nil { + return errors.Wrapf(err, "retrieve device with udid %s", ev.Command.UDID) + } + dev.Token = ev.Command.Token.String() + dev.PushMagic = ev.Command.PushMagic + dev.UnlockToken = ev.Command.UnlockToken.String() + dev.AwaitingConfiguration = ev.Command.AwaitingConfiguration + dev.LastSeen = time.Now() + // first TokenUpdate event will have the enrollment status set to false. + newlyEnrolled := dev.Enrolled + dev.Enrolled = true + if err := w.db.Save(dev); err != nil { + return errors.Wrapf(err, "saving updated device for Token event udid=%s", ev.Command.UDID) + } + + if newlyEnrolled { + // notify subscribers of a successful enrollment + // TODO: The enrollment topic needs a custom event. + err = w.ps.Publish(ctx, DeviceEnrolledTopic, message) + return errors.Wrap(err, "publishing new enrollment message") + } + return nil +} + +func (w *Worker) updateFromAuthenticate(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") + } + + device, reenrolling, err := getOrCreateDevice(w.db, ev.Command.SerialNumber, ev.Command.UDID) + if err != nil { + return errors.Wrap(err, "get device for authenticate event") + } + + if reenrolling { + level.Debug(w.logger).Log( + "msg", "re-enrolling device", + "serial", ev.Command.SerialNumber, + ) + } else { + level.Debug(w.logger).Log( + "msg", "enrolling new device", + "serial", ev.Command.SerialNumber, + ) + } + + if device.UUID == "" { + device.UUID = uuid.NewV4().String() + } + device.UDID = ev.Command.UDID + device.OSVersion = ev.Command.OSVersion + device.BuildVersion = ev.Command.BuildVersion + device.ProductName = ev.Command.ProductName + device.SerialNumber = ev.Command.SerialNumber + device.IMEI = ev.Command.IMEI + device.MEID = ev.Command.MEID + device.DeviceName = ev.Command.DeviceName + device.Model = ev.Command.Model + device.ModelName = ev.Command.ModelName + device.LastSeen = time.Now() + err = w.db.Save(device) + return errors.Wrapf(err, "saving updated device for authenticate event") +} + +func getOrCreateDevice(db DeviceWorkerStore, serial, udid string) (dev *Device, reenrolling bool, err error) { + if udid != "" { + // first try to fetch a device by UDID. + // If the device was previously enrolled it will exist. + // In case the device is known, set the enrolled status to false before returning. + byUDID, err := db.DeviceByUDID(udid) + if err == nil { + byUDID.Enrolled = false + return byUDID, true, nil + } + if err != nil && !isNotFound(err) { + return nil, false, errors.Wrapf(err, "retrieve device with udid %s and serial %s", udid, serial) + } + } + + // next try to find the device by serial. If found, it's a DEP device, which contains only the + // serials but not a udid. + dev, err = getOrCreateDeviceBySerial(db, serial) + return dev, false, err +} + +func getOrCreateDeviceBySerial(db DeviceWorkerStore, serial string) (*Device, error) { + bySerial, err := db.DeviceBySerial(serial) + if err == nil && bySerial != nil { + return bySerial, nil + } + if err != nil && !isNotFound(err) { + return nil, errors.Wrapf(err, "retrieve device with serial number %s", serial) + } + + dev := new(Device) + return dev, nil +} + +func isNotFound(err error) bool { + err = errors.Cause(err) + type notFoundErr interface { + error + NotFound() bool + } + + e, ok := err.(notFoundErr) + return ok && e.NotFound() +}