From fe50027f15c1e7d030a4505e270cb60745d57003 Mon Sep 17 00:00:00 2001 From: Victor Vrantchan Date: Sun, 13 Nov 2016 22:26:42 -0500 Subject: [PATCH] refactor command service (#50) command: move service implementation to redis sub-package. --- app/config.go | 1 + app/services.go | 32 +++-- command/command.go | 31 +++++ command/datastore.go | 209 --------------------------------- command/datastore_test.go | 44 ------- command/service.go | 61 ---------- command/service/redis/redis.go | 158 +++++++++++++++++++++++++ 7 files changed, 212 insertions(+), 324 deletions(-) create mode 100644 command/command.go delete mode 100644 command/datastore.go delete mode 100644 command/datastore_test.go delete mode 100644 command/service.go create mode 100644 command/service/redis/redis.go diff --git a/app/config.go b/app/config.go index f7cc0f53..6772e28d 100644 --- a/app/config.go +++ b/app/config.go @@ -286,6 +286,7 @@ func (c *PostgresConfig) fromDockerEnv() { type RedisConfig struct { Enabled bool Connection string + Password string // TODO no way to pass this option as a param. } func (c *Config) loadRedis(conn string) { diff --git a/app/services.go b/app/services.go index 8b17e20a..479b79f3 100644 --- a/app/services.go +++ b/app/services.go @@ -5,6 +5,7 @@ import ( pushcertificate "github.com/RobotsAndPencils/buford/certificate" "github.com/RobotsAndPencils/buford/push" + "github.com/garyburd/redigo/redis" "github.com/go-kit/kit/log" "github.com/micromdm/dep" @@ -12,8 +13,10 @@ import ( "github.com/micromdm/micromdm/certificate" "github.com/micromdm/micromdm/checkin" "github.com/micromdm/micromdm/command" + cmdredis "github.com/micromdm/micromdm/command/service/redis" "github.com/micromdm/micromdm/connect" "github.com/micromdm/micromdm/device" + "github.com/micromdm/micromdm/driver" "github.com/micromdm/micromdm/enroll" "github.com/micromdm/micromdm/management" "github.com/micromdm/micromdm/workflow" @@ -23,6 +26,8 @@ import ( // which MicroMDM relies on. func setupServices(config *Config, logger log.Logger) (*serviceManager, error) { sm := &serviceManager{Config: config, logger: logger} + sm.createRedisPool() + sm.setupAppDatastore() sm.setupDeviceDatastore() sm.setupWorkflowDatastore() @@ -44,7 +49,6 @@ func setupServices(config *Config, logger log.Logger) (*serviceManager, error) { // serviceManager knows how to setup the independent components which make up // MicroMDM, mainly Datastores and Services. type serviceManager struct { - CommandDatastore command.Datastore CertificateDatastore certificate.Datastore DeviceDatastore device.Datastore WorkflowDatastore workflow.Datastore @@ -59,6 +63,7 @@ type serviceManager struct { EnrollmentService enroll.Service *Config + pool *redis.Pool logger log.Logger err error } @@ -150,7 +155,10 @@ func (s *serviceManager) setupPushService() { if s.err != nil { return } - cert, key, err := pushcertificate.Load(s.APNS.CertificatePath, s.APNS.PrivateKeyPass) + cert, key, err := pushcertificate.Load( + s.APNS.CertificatePath, + s.APNS.PrivateKeyPass, + ) if err != nil { s.err = err return @@ -215,16 +223,20 @@ func (s *serviceManager) setupDeviceDatastore() { s.DeviceDatastore = db } +func (s *serviceManager) createRedisPool() { + if s.err != nil { + return + } + opts := []driver.ConnOption{driver.Logger(s.logger)} + if s.Redis.Password != "" { + opts = append(opts, driver.WithPassword(s.Redis.Password)) + } + s.pool, s.err = driver.NewRedisPool(s.Redis.Connection, opts...) +} + func (s *serviceManager) setupCommandService() { if s.err != nil { return } - db, err := command.NewDB("redis", s.Redis.Connection, s.logger) - if err != nil { - s.err = err - return - } - s.CommandDatastore = db - s.CommandService = command.NewService(db) - return + s.CommandService, s.err = cmdredis.NewCommandService(s.pool, s.logger) } diff --git a/command/command.go b/command/command.go new file mode 100644 index 00000000..94b8c913 --- /dev/null +++ b/command/command.go @@ -0,0 +1,31 @@ +// Package command manages an MDM Command queue for enrolled devices. +package command + +import ( + "errors" + + "github.com/micromdm/mdm" +) + +// Service defines methods for managing MDM commands in a queue. +type Service interface { + + // NewCommand turns an MDM Command Request into a MDM payload. + NewCommand(*mdm.CommandRequest) (*mdm.Payload, error) + + // NextCommand retrieves the next command in a device's queue. + NextCommand(udid string) ([]byte, int, error) + + // DeleteCommand deletes a previously queued command from a device's queue. + DeleteCommand(deviceUDID, commandUUID string) (int, error) + + // Commands returns all the commands in a device's current queue. + Commands(deviceUDID string) ([]mdm.Payload, error) + + // Find returns a previously queued command. + Find(commandUUID string) (*mdm.Payload, error) +} + +// TODO change this error to a type/interface +// ErrNoKey is returned if there is no key in redis +var ErrNoKey = errors.New("there is no such key in redis.") diff --git a/command/datastore.go b/command/datastore.go deleted file mode 100644 index c3ce20dd..00000000 --- a/command/datastore.go +++ /dev/null @@ -1,209 +0,0 @@ -package command - -import ( - "bytes" - "errors" - "fmt" - "os" - "time" - - "github.com/garyburd/redigo/redis" - kitlog "github.com/go-kit/kit/log" - "github.com/groob/plist" - "github.com/micromdm/mdm" -) - -var ( - // ErrNoKey is returned if there is no key in redis - ErrNoKey = errors.New("There is no such key in redis.") -) - -// Datastore provides methods for saving and retrieving MDM commands -type Datastore interface { - // Saves the payload in redis - // SET CommandUUID plistData - SavePayload(payload *mdm.Payload) error - // Adds MDM commands to a queue in redis list - // LPUSH deviceUDID commandUUID - QueueCommand(deviceUDID, commandUUID string) error - NextCommand(deviceUDID string) ([]byte, int, error) - DeleteCommand(deviceUDID, commandUUID string) (int, error) - Commands(deviceUDID string) ([]mdm.Payload, error) - Find(commandUUID string) (*mdm.Payload, error) -} - -//NewDB creates a Datastore -func NewDB(driver, conn string, logger kitlog.Logger) (Datastore, error) { - var ds Datastore - switch driver { - case "redis": - ds = redisDB{pool: redisPool(conn, logger)} - return ds, nil - default: - return nil, errors.New("unknown driver") - } -} - -type redisDB struct { - pool *redis.Pool -} - -func (rds redisDB) SavePayload(payload *mdm.Payload) error { - var buf bytes.Buffer - // get connection from redis pool - conn := rds.pool.Get() - defer conn.Close() - // encode payload into a plist - err := plist.NewEncoder(&buf).Encode(payload) - if err != nil { - return err - } - // create a commandUUID key with the plist as the value - _, err = conn.Do("set", payload.CommandUUID, buf.String()) - if err != nil { - return err - } - return nil -} - -func (rds redisDB) QueueCommand(deviceUDID, commandUUID string) error { - // get connection from redis pool - conn := rds.pool.Get() - defer conn.Close() - _, err := conn.Do("lpush", deviceUDID, commandUUID) - if err != nil { - return err - } - return nil -} -func (rds redisDB) NextCommand(deviceUDID string) ([]byte, int, error) { - // get connection from redis pool - conn := rds.pool.Get() - defer conn.Close() - // pop the first command - commandUUID, err := redis.String(conn.Do("lpop", deviceUDID)) - if err != nil && err != redis.ErrNil { - return nil, 0, err - } - // if the list is empty - if err == redis.ErrNil { - return []byte{}, 0, nil - } - // push the redis command back to the end of the list - _, err = conn.Do("rpush", deviceUDID, commandUUID) - command, err := redis.String(conn.Do("get", commandUUID)) - if err == redis.ErrNil { - return nil, 0, ErrNoKey - } - - // get a command list length - total, err := redis.Int(conn.Do("llen", deviceUDID)) - if err != nil { - return nil, 0, err - } - return []byte(command), total, err -} - -func (rds redisDB) DeleteCommand(deviceUDID, commandUUID string) (int, error) { - // get connection from redis pool - conn := rds.pool.Get() - defer conn.Close() - // remove from list - _, err := conn.Do("lrem", deviceUDID, 0, commandUUID) - if err != nil { - return 0, err - } - // set the key to expire in an hour - _, err = conn.Do("expire", commandUUID, 3600) - if err != nil { - return 0, err - } - // get a command list length - total, err := redis.Int(conn.Do("llen", deviceUDID)) - if err != nil { - return 0, err - } - return total, nil -} - -func (rds redisDB) Commands(deviceUDID string) ([]mdm.Payload, error) { - conn := rds.pool.Get() - defer conn.Close() - - commandUUIDs, err := redis.Values(conn.Do("LRANGE", deviceUDID, "0", "-1")) - if err != nil { - return nil, err - } - - var payloads []mdm.Payload = make([]mdm.Payload, len(commandUUIDs)) - - for i, commandUUID := range commandUUIDs { - payloadData, err := redis.Bytes(conn.Do("GET", commandUUID)) - if err != nil { - return nil, err - } - - if err := plist.NewDecoder(bytes.NewReader(payloadData)).Decode(&payloads[i]); err != nil { - return nil, err - } - } - - return payloads, nil -} - -func (rds redisDB) Find(commandUUID string) (*mdm.Payload, error) { - conn := rds.pool.Get() - defer conn.Close() - - payloadData, err := redis.Bytes(conn.Do("GET", commandUUID)) - if err != nil { - return nil, err - } - - var payload *mdm.Payload - if err := plist.NewDecoder(bytes.NewReader(payloadData)).Decode(&payload); err != nil { - return nil, err - } - - return payload, nil -} - -func redisPool(conn string, logger kitlog.Logger) *redis.Pool { - pool := &redis.Pool{ - MaxIdle: 3, - IdleTimeout: 240 * time.Second, - Dial: func() (redis.Conn, error) { - c, err := redis.Dial("tcp", conn) - if err != nil { - return nil, err - } - return c, err - }, - TestOnBorrow: func(c redis.Conn, t time.Time) error { - _, err := c.Do("PING") - return err - }, - } - checkRedisConn(pool, logger) - return pool -} - -func checkRedisConn(pool *redis.Pool, logger kitlog.Logger) { - conn := pool.Get() - defer conn.Close() - - var dbError error - maxAttempts := 20 - for attempts := 1; attempts <= maxAttempts; attempts++ { - _, dbError = conn.Do("PING") - if dbError == nil { - break - } - logger.Log("msg", fmt.Sprintf("could not connect to redis: %v", dbError)) - time.Sleep(time.Duration(attempts) * time.Second) - } - if dbError != nil { - logger.Log("err", dbError) - os.Exit(1) - } -} diff --git a/command/datastore_test.go b/command/datastore_test.go deleted file mode 100644 index 0a6b5886..00000000 --- a/command/datastore_test.go +++ /dev/null @@ -1,44 +0,0 @@ -package command - -import ( - "fmt" - "github.com/go-kit/kit/log" - "github.com/micromdm/mdm" - "os" - "testing" -) - -type datastoreFixtures struct { - ds Datastore - logger log.Logger -} - -func setup() (datastoreFixtures, error) { - logger := log.NewLogfmtLogger(os.Stdout) - commandsDb, err := NewDB("redis", "localhost", logger) - if err != nil { - return nil, err - } - - return datastoreFixtures{ds: commandsDb, logger: logger} -} - -func teardown() { - -} - -func TestService_Commands(t *testing.T) { - fixtures, err := setup() - defer teardown() - if err != nil { - t.Errorf("error making new datastore: %v", err) - } - - var commands []mdm.Payload - commands, err = fixtures.ds.Commands("ABCDEF") - if err != nil { - t.Errorf("datastore.Commands returned error: %v", err) - } - - fmt.Printf("%v", commands) -} diff --git a/command/service.go b/command/service.go deleted file mode 100644 index ff2690cf..00000000 --- a/command/service.go +++ /dev/null @@ -1,61 +0,0 @@ -package command - -import "github.com/micromdm/mdm" - -// Service defines methods for managing MDM commands -type Service interface { - NewCommand(*mdm.CommandRequest) (*mdm.Payload, error) - NextCommand(udid string) ([]byte, int, error) - DeleteCommand(deviceUDID, commandUUID string) (int, error) - Commands(deviceUDID string) ([]mdm.Payload, error) - Find(commandUUID string) (*mdm.Payload, error) -} - -// NewService returns a new command service -func NewService(ds Datastore) Service { - return &service{ - db: ds, - } -} - -type service struct { - db Datastore -} - -func (svc service) NewCommand(request *mdm.CommandRequest) (*mdm.Payload, error) { - // create a payload - payload, err := mdm.NewPayload(request) - if err != nil { - return nil, err - } - // save in redis - err = svc.db.SavePayload(payload) - if err != nil { - return nil, err - } - // add command to a queue in redis - err = svc.db.QueueCommand(request.UDID, payload.CommandUUID) - if err != nil { - return nil, err - } - // return created payload to user - return payload, nil -} - -// NextCommand returns an MDM Payload from a list of queued payloads -func (svc service) NextCommand(udid string) ([]byte, int, error) { - return svc.db.NextCommand(udid) -} - -// DeleteCommand returns an MDM Payload from a list of queued payloads -func (svc service) DeleteCommand(deviceUDID, commandUUID string) (int, error) { - return svc.db.DeleteCommand(deviceUDID, commandUUID) -} - -func (svc service) Commands(deviceUDID string) ([]mdm.Payload, error) { - return svc.db.Commands(deviceUDID) -} - -func (svc service) Find(commandUUID string) (*mdm.Payload, error) { - return svc.db.Find(commandUUID) -} diff --git a/command/service/redis/redis.go b/command/service/redis/redis.go new file mode 100644 index 00000000..a14c09b0 --- /dev/null +++ b/command/service/redis/redis.go @@ -0,0 +1,158 @@ +package redis + +import ( + "bytes" + + "github.com/garyburd/redigo/redis" + "github.com/go-kit/kit/log" + "github.com/groob/plist" + + "github.com/micromdm/mdm" + "github.com/micromdm/micromdm/command" +) + +// NewCommandService creates a command.Service backed by redis. +func NewCommandService(pool *redis.Pool, logger log.Logger) (Redis, error) { + return Redis{pool: pool}, nil +} + +// Redis implements command.Service +type Redis struct { + pool *redis.Pool +} + +func (rds Redis) NewCommand(request *mdm.CommandRequest) (*mdm.Payload, error) { + // create a payload + payload, err := mdm.NewPayload(request) + if err != nil { + return nil, err + } + // save in redis + err = rds.SavePayload(payload) + if err != nil { + return nil, err + } + // add command to a queue in redis + err = rds.QueueCommand(request.UDID, payload.CommandUUID) + if err != nil { + return nil, err + } + // return created payload to user + return payload, nil +} + +func (rds Redis) SavePayload(payload *mdm.Payload) error { + var buf bytes.Buffer + // get connection from redis pool + conn := rds.pool.Get() + defer conn.Close() + // encode payload into a plist + err := plist.NewEncoder(&buf).Encode(payload) + if err != nil { + return err + } + // create a commandUUID key with the plist as the value + _, err = conn.Do("set", payload.CommandUUID, buf.String()) + if err != nil { + return err + } + return nil +} + +func (rds Redis) QueueCommand(deviceUDID, commandUUID string) error { + // get connection from redis pool + conn := rds.pool.Get() + defer conn.Close() + _, err := conn.Do("lpush", deviceUDID, commandUUID) + if err != nil { + return err + } + return nil +} +func (rds Redis) NextCommand(deviceUDID string) ([]byte, int, error) { + // get connection from redis pool + conn := rds.pool.Get() + defer conn.Close() + // pop the first command + commandUUID, err := redis.String(conn.Do("lpop", deviceUDID)) + if err != nil && err != redis.ErrNil { + return nil, 0, err + } + // if the list is empty + if err == redis.ErrNil { + return []byte{}, 0, nil + } + // push the redis command back to the end of the list + _, err = conn.Do("rpush", deviceUDID, commandUUID) + cmd, err := redis.String(conn.Do("get", commandUUID)) + if err == redis.ErrNil { + return nil, 0, command.ErrNoKey + } + + // get a command list length + total, err := redis.Int(conn.Do("llen", deviceUDID)) + if err != nil { + return nil, 0, err + } + return []byte(cmd), total, err +} + +func (rds Redis) DeleteCommand(deviceUDID, commandUUID string) (int, error) { + // get connection from redis pool + conn := rds.pool.Get() + defer conn.Close() + // remove from list + _, err := conn.Do("lrem", deviceUDID, 0, commandUUID) + if err != nil { + return 0, err + } + // set the key to expire in an hour + _, err = conn.Do("expire", commandUUID, 3600) + if err != nil { + return 0, err + } + // get a command list length + total, err := redis.Int(conn.Do("llen", deviceUDID)) + if err != nil { + return 0, err + } + return total, nil +} + +func (rds Redis) Commands(deviceUDID string) ([]mdm.Payload, error) { + conn := rds.pool.Get() + defer conn.Close() + + commandUUIDs, err := redis.Values(conn.Do("LRANGE", deviceUDID, "0", "-1")) + if err != nil { + return nil, err + } + + // FIXME this code is going to result in an err if a command is deleted + // mid-loop by another process. + var payloads []mdm.Payload = make([]mdm.Payload, len(commandUUIDs)) + for i, commandUUID := range commandUUIDs { + payloadData, err := redis.Bytes(conn.Do("GET", commandUUID)) + if err != nil { + return nil, err + } + if err := plist.Unmarshal(payloadData, &payloads[i]); err != nil { + return nil, err + } + } + return payloads, nil +} + +func (rds Redis) Find(commandUUID string) (*mdm.Payload, error) { + conn := rds.pool.Get() + defer conn.Close() + + payloadData, err := redis.Bytes(conn.Do("GET", commandUUID)) + if err != nil { + return nil, err + } + + var payload *mdm.Payload + err = plist.Unmarshal(payloadData, payload) + return payload, err +}