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.
This commit is contained in:
Victor Vrantchan
2018-05-30 09:14:41 -07:00
committed by GitHub
parent 35051e7d9b
commit 1ccf1c8ca1
4 changed files with 306 additions and 229 deletions

View File

@@ -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)

View File

@@ -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
}

View File

@@ -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)
}

295
platform/device/worker.go Normal file
View File

@@ -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()
}