From c1c1cfb17572c21d760fc73b874c05cd2e10e7bd Mon Sep 17 00:00:00 2001 From: Victor Vrantchan Date: Sat, 12 May 2018 17:04:05 -0400 Subject: [PATCH] refactor command service to use one method per file (#416) Updated the command service to have the same structure as all the other services within platform/... --- cmd/micromdm/serve.go | 18 ++---- platform/command/command.go | 84 ------------------------- platform/command/endpoint.go | 78 ----------------------- platform/command/new_command.go | 71 +++++++++++++++++++++ platform/command/server.go | 34 ++++++++++ platform/command/service.go | 60 +++--------------- platform/command/transport_http.go | 70 --------------------- platform/command/transport_http_test.go | 34 ---------- 8 files changed, 117 insertions(+), 332 deletions(-) delete mode 100644 platform/command/command.go delete mode 100644 platform/command/endpoint.go create mode 100644 platform/command/new_command.go create mode 100644 platform/command/server.go delete mode 100644 platform/command/transport_http.go delete mode 100644 platform/command/transport_http_test.go diff --git a/cmd/micromdm/serve.go b/cmd/micromdm/serve.go index ef2f3a54..59b898ad 100644 --- a/cmd/micromdm/serve.go +++ b/cmd/micromdm/serve.go @@ -217,18 +217,7 @@ func serve(args []string) error { checkinHandlers = checkin.MakeHTTPHandlers(ctx, e, opts...) } - var commandHandlers command.HTTPHandlers - { - e := command.Endpoints{ - NewCommandEndpoint: command.MakeNewCommandEndpoint(sm.commandService), - } - - opts := []httptransport.ServerOption{ - httptransport.ServerErrorLogger(httpLogger), - httptransport.ServerErrorEncoder(connect.EncodeError), - } - commandHandlers = command.MakeHTTPHandlers(ctx, e, opts...) - } + commandEndpoints := command.MakeServerEndpoints(sm.commandService) connectOpts := []httptransport.ServerOption{ httptransport.ServerErrorLogger(httpLogger), @@ -325,6 +314,7 @@ func serve(args []string) error { depHandlers := depapi.MakeHTTPHandler(depEndpoints, logger) apnsHandlers := apns.MakeHTTPHandler(apnsEndpoints, logger) depsyncHandlers := depsync.MakeHTTPHandler(depsyncEndpoints, logger) + commandHandler := command.MakeHTTPHandler(commandEndpoints, logger) // API commands. Only handled if the user provides an api key. if *flAPIKey != "" { @@ -343,7 +333,7 @@ func serve(args []string) error { r.Handle("/v1/dep/profiles", apiAuthMiddleware(*flAPIKey, depHandlers)) r.Handle("/v1/dep/syncnow", apiAuthMiddleware(*flAPIKey, depsyncHandlers)) r.Handle("/v1/dep/autoassigners", apiAuthMiddleware(*flAPIKey, depsyncHandlers)) - r.Handle("/v1/commands", apiAuthMiddleware(*flAPIKey, commandHandlers.NewCommandHandler)).Methods("POST") + r.Handle("/v1/commands", apiAuthMiddleware(*flAPIKey, commandHandler)) r.Handle("/push/{udid}", apiAuthMiddleware(*flAPIKey, apnsHandlers)) } else { mainLogger.Log("msg", "no api key specified") @@ -478,7 +468,7 @@ func (c *server) setupCommandService() { if c.err != nil { return } - c.commandService, c.err = command.New(c.db, c.pubclient) + c.commandService, c.err = command.New(c.pubclient) } func (c *server) setupWebhooks() { diff --git a/platform/command/command.go b/platform/command/command.go deleted file mode 100644 index 8136c108..00000000 --- a/platform/command/command.go +++ /dev/null @@ -1,84 +0,0 @@ -package command - -import ( - "fmt" - - "github.com/boltdb/bolt" - "github.com/micromdm/micromdm/mdm/mdm" - "github.com/pkg/errors" - "golang.org/x/net/context" - - "github.com/micromdm/micromdm/platform/pubsub" -) - -const ( - - // CommandBucket is the *bolt.DB bucket where commands are archived. - CommandBucket = "mdm.Command.ARCHIVE" - - // CommandTopic is a PubSub topic that events are published to. - CommandTopic = "mdm.Command" -) - -type Command struct { - db *bolt.DB - publisher pubsub.Publisher - archiveFn func(int64, []byte) error -} - -func New(db *bolt.DB, pub pubsub.Publisher) (*Command, error) { - err := db.Update(func(tx *bolt.Tx) error { - _, err := tx.CreateBucketIfNotExists([]byte(CommandBucket)) - return err - }) - if err != nil { - return nil, errors.Wrapf(err, "creating %s bucket", CommandBucket) - } - svc := Command{ - db: db, - publisher: pub, - } - svc.archiveFn = svc.archive - return &svc, nil -} - -func (svc *Command) NewCommand(ctx context.Context, request *mdm.CommandRequest) (*mdm.CommandPayload, error) { - if request == nil { - return nil, errors.New("empty CommandRequest") - } - payload, err := mdm.NewCommandPayload(request) - if err != nil { - return nil, errors.Wrap(err, "creating mdm payload") - } - event := NewEvent(payload, request.UDID) - msg, err := MarshalEvent(event) - if err != nil { - return nil, errors.Wrap(err, "marshalling mdm command event") - } - if err := svc.archive(event.Time.UnixNano(), msg); err != nil { - return nil, errors.Wrap(err, "archive mdm command") - } - if err := svc.publisher.Publish(context.TODO(), CommandTopic, msg); err != nil { - return nil, errors.Wrapf(err, "publish mdm command on topic: %s", CommandTopic) - } - return payload, nil -} - -// archive events to BoltDB bucket using timestamp as key to preserve order. -func (svc *Command) archive(nano int64, msg []byte) error { - tx, err := svc.db.Begin(true) - if err != nil { - return errors.Wrap(err, "begin transaction") - } - defer tx.Rollback() - - bkt := tx.Bucket([]byte(CommandBucket)) - if bkt == nil { - return fmt.Errorf("bucket %q not found!", CommandBucket) - } - key := []byte(fmt.Sprintf("%d", nano)) - if err := bkt.Put(key, msg); err != nil { - return errors.Wrap(err, "put command event to boltdb") - } - return tx.Commit() -} diff --git a/platform/command/endpoint.go b/platform/command/endpoint.go deleted file mode 100644 index 011416fd..00000000 --- a/platform/command/endpoint.go +++ /dev/null @@ -1,78 +0,0 @@ -package command - -import ( - "context" - "errors" - "fmt" - "net/http" - "time" - - "github.com/go-kit/kit/endpoint" - "github.com/go-kit/kit/log" - "github.com/go-kit/kit/metrics" - "github.com/micromdm/micromdm/mdm/mdm" -) - -var errEmptyRequest = errors.New("request must contain UDID of the device") - -type Endpoints struct { - NewCommandEndpoint endpoint.Endpoint -} - -// MakeNewCommandEndpoint creates an endpoint which creates new MDM Commands. -func MakeNewCommandEndpoint(svc Service) endpoint.Endpoint { - return func(ctx context.Context, request interface{}) (interface{}, error) { - req := request.(newCommandRequest) - if req.UDID == "" || req.RequestType == "" { - return newCommandResponse{Err: errEmptyRequest}, nil - } - payload, err := svc.NewCommand(ctx, &req.CommandRequest) - if err != nil { - return newCommandResponse{Err: err}, nil - } - return newCommandResponse{Payload: payload}, nil - } -} - -// EndpointInstrumentingMiddleware returns an endpoint middleware that records -// the duration of each invocation to the passed histogram. The middleware adds -// a single field: "success", which is "true" if no error is returned, and -// "false" otherwise. -func EndpointInstrumentingMiddleware(duration metrics.Histogram) endpoint.Middleware { - return func(next endpoint.Endpoint) endpoint.Endpoint { - return func(ctx context.Context, request interface{}) (response interface{}, err error) { - defer func(begin time.Time) { - duration.With("success", fmt.Sprint(err == nil)).Observe(time.Since(begin).Seconds()) - }(time.Now()) - return next(ctx, request) - - } - } -} - -// EndpointLoggingMiddleware returns an endpoint middleware that logs the -// duration of each invocation, and the resulting error, if any. -func EndpointLoggingMiddleware(logger log.Logger) endpoint.Middleware { - return func(next endpoint.Endpoint) endpoint.Endpoint { - return func(ctx context.Context, request interface{}) (response interface{}, err error) { - - defer func(begin time.Time) { - logger.Log("error", err, "took", time.Since(begin)) - }(time.Now()) - return next(ctx, request) - - } - } -} - -type newCommandRequest struct { - mdm.CommandRequest -} - -type newCommandResponse struct { - Payload *mdm.CommandPayload `json:"payload,omitempty"` - Err error `json:"error,omitempty"` -} - -func (r newCommandResponse) error() error { return r.Err } -func (r newCommandResponse) status() int { return http.StatusCreated } diff --git a/platform/command/new_command.go b/platform/command/new_command.go new file mode 100644 index 00000000..6d880e59 --- /dev/null +++ b/platform/command/new_command.go @@ -0,0 +1,71 @@ +package command + +import ( + "net/http" + + "github.com/go-kit/kit/endpoint" + "github.com/pkg/errors" + "golang.org/x/net/context" + + "github.com/micromdm/micromdm/mdm/mdm" + "github.com/micromdm/micromdm/pkg/httputil" +) + +const ( + // CommandTopic is a PubSub topic that events are published to. + CommandTopic = "mdm.Command" +) + +func (svc *CommandService) NewCommand(ctx context.Context, request *mdm.CommandRequest) (*mdm.CommandPayload, error) { + if request == nil { + return nil, errors.New("empty CommandRequest") + } + payload, err := mdm.NewCommandPayload(request) + if err != nil { + return nil, errors.Wrap(err, "creating mdm payload") + } + event := NewEvent(payload, request.UDID) + msg, err := MarshalEvent(event) + if err != nil { + return nil, errors.Wrap(err, "marshalling mdm command event") + } + if err := svc.publisher.Publish(context.TODO(), CommandTopic, msg); err != nil { + return nil, errors.Wrapf(err, "publish mdm command on topic: %s", CommandTopic) + } + return payload, nil +} + +type newCommandRequest struct { + mdm.CommandRequest +} + +type newCommandResponse struct { + Payload *mdm.CommandPayload `json:"payload,omitempty"` + Err error `json:"error,omitempty"` +} + +func (r newCommandResponse) Failed() error { return r.Err } +func (r newCommandResponse) StatusCode() int { return http.StatusCreated } + +func decodeNewCommandRequest(ctx context.Context, r *http.Request) (interface{}, error) { + var req newCommandRequest + err := httputil.DecodeJSONRequest(r, &req) + return req, err +} + +var errEmptyRequest = errors.New("request must contain UDID of the device") + +// MakeNewCommandEndpoint creates an endpoint which creates new MDM Commands. +func MakeNewCommandEndpoint(svc Service) endpoint.Endpoint { + return func(ctx context.Context, request interface{}) (interface{}, error) { + req := request.(newCommandRequest) + if req.UDID == "" || req.RequestType == "" { + return newCommandResponse{Err: errEmptyRequest}, nil + } + payload, err := svc.NewCommand(ctx, &req.CommandRequest) + if err != nil { + return newCommandResponse{Err: err}, nil + } + return newCommandResponse{Payload: payload}, nil + } +} diff --git a/platform/command/server.go b/platform/command/server.go new file mode 100644 index 00000000..1a611d23 --- /dev/null +++ b/platform/command/server.go @@ -0,0 +1,34 @@ +package command + +import ( + "github.com/go-kit/kit/endpoint" + "github.com/go-kit/kit/log" + httptransport "github.com/go-kit/kit/transport/http" + "github.com/gorilla/mux" + "github.com/micromdm/micromdm/pkg/httputil" +) + +type Endpoints struct { + NewCommandEndpoint endpoint.Endpoint +} + +func MakeServerEndpoints(s Service) Endpoints { + return Endpoints{ + NewCommandEndpoint: MakeNewCommandEndpoint(s), + } +} + +func MakeHTTPHandler(e Endpoints, logger log.Logger) *mux.Router { + r, options := httputil.NewRouter(logger) + + // POST /v1/commands Add new MDM Command to device queue. + + r.Methods("POST").Path("/v1/commands").Handler(httptransport.NewServer( + e.NewCommandEndpoint, + decodeNewCommandRequest, + httputil.EncodeJSONResponse, + options..., + )) + + return r +} diff --git a/platform/command/service.go b/platform/command/service.go index eeb9dd6f..cf297432 100644 --- a/platform/command/service.go +++ b/platform/command/service.go @@ -2,11 +2,8 @@ package command import ( - "time" - - "github.com/go-kit/kit/log" - "github.com/go-kit/kit/metrics" "github.com/micromdm/micromdm/mdm/mdm" + "github.com/micromdm/micromdm/platform/pubsub" "golang.org/x/net/context" ) @@ -14,54 +11,13 @@ type Service interface { NewCommand(context.Context, *mdm.CommandRequest) (*mdm.CommandPayload, error) } -// Middleware describes a service (as opposed to endpoint) middleware. -type Middleware func(Service) Service +type CommandService struct { + publisher pubsub.Publisher +} -// ServiceLoggingMiddleware returns a service middleware that logs the -// parameters and result of each method invocation. -func ServiceLoggingMiddleware(logger log.Logger) Middleware { - return func(next Service) Service { - return serviceLoggingMiddleware{ - logger: logger, - next: next, - } +func New(pub pubsub.Publisher) (*CommandService, error) { + svc := CommandService{ + publisher: pub, } -} - -func (mw serviceLoggingMiddleware) NewCommand(ctx context.Context, req *mdm.CommandRequest) (p *mdm.CommandPayload, err error) { - defer func(begin time.Time) { - mw.logger.Log( - "method", "NewCommand", - "error", err, - "took", time.Since(begin), - ) - }(time.Now()) - return mw.next.NewCommand(ctx, req) -} - -type serviceLoggingMiddleware struct { - logger log.Logger - next Service -} - -// ServiceInstrumentingMiddleware returns a service middleware that tracks the -// number of payloads created by the service. -func ServiceInstrumentingMiddleware(p metrics.Counter) Middleware { - return func(next Service) Service { - return serviceInstrumentingMiddleware{ - payloads: p, - next: next, - } - } -} - -type serviceInstrumentingMiddleware struct { - payloads metrics.Counter - next Service -} - -func (mw serviceInstrumentingMiddleware) NewCommand(ctx context.Context, req *mdm.CommandRequest) (*mdm.CommandPayload, error) { - p, err := mw.next.NewCommand(ctx, req) - mw.payloads.Add(1) - return p, err + return &svc, nil } diff --git a/platform/command/transport_http.go b/platform/command/transport_http.go deleted file mode 100644 index aa628ff8..00000000 --- a/platform/command/transport_http.go +++ /dev/null @@ -1,70 +0,0 @@ -package command - -import ( - "context" - "encoding/json" - "io" - "net/http" - - "github.com/pkg/errors" - - httptransport "github.com/go-kit/kit/transport/http" -) - -type HTTPHandlers struct { - NewCommandHandler http.Handler -} - -func MakeHTTPHandlers(ctx context.Context, endpoints Endpoints, opts ...httptransport.ServerOption) HTTPHandlers { - h := HTTPHandlers{ - NewCommandHandler: httptransport.NewServer( - endpoints.NewCommandEndpoint, - decodeRequest, - encodeResponse, - opts..., - ), - } - return h -} - -type errorer interface { - error() error -} - -type statuser interface { - status() int -} - -// EncodeError is used by the HTTP transport to encode service errors in HTTP. -// The EncodeError should be passed to the Go-Kit httptransport as the -// ServerErrorEncoder to encode error responses with JSON. -func EncodeError(ctx context.Context, err error, w http.ResponseWriter) { - enc := json.NewEncoder(w) - enc.SetIndent("", " ") - - enc.Encode(map[string]interface{}{ - "error": err.Error(), - }) -} - -func decodeRequest(ctx context.Context, r *http.Request) (interface{}, error) { - var req newCommandRequest - err := json.NewDecoder(io.LimitReader(r.Body, 1000000)).Decode(&req) - return req, errors.Wrap(err, "decoding command request") -} - -func encodeResponse(ctx context.Context, w http.ResponseWriter, response interface{}) error { - - if e, ok := response.(errorer); ok && e.error() != nil { - EncodeError(ctx, e.error(), w) - return nil - } - - if s, ok := response.(statuser); ok { - w.WriteHeader(s.status()) - } - - enc := json.NewEncoder(w) - enc.SetIndent("", " ") - return enc.Encode(response) -} diff --git a/platform/command/transport_http_test.go b/platform/command/transport_http_test.go deleted file mode 100644 index b22d3789..00000000 --- a/platform/command/transport_http_test.go +++ /dev/null @@ -1,34 +0,0 @@ -package command - -import ( - "context" - "net/http/httptest" - "strings" - "testing" -) - -func TestDecodeRequest(t *testing.T) { - requestData := ` -{ - "request_type": "InstallApplication", - "udid" : "564D38A0-4C3B-AD69-803B-DAC58A298191", - "manifest_url" : "https://mdm.acme.co/repo/munkitools-3.0.0.3298.plist", - "management_flags" : 1 -} -` - req := httptest.NewRequest("POST", "https://mdm.acme.co/v1/commands", strings.NewReader(requestData)) - request, err := decodeRequest(context.Background(), req) - if err != nil { - t.Fatal(err) - } - decoded := request.(newCommandRequest) - - if have, want := decoded.RequestType, "InstallApplication"; have != want { - t.Errorf("have %s, want %s", have, want) - } - - if have, want := decoded.CommandRequest.InstallApplication.ManifestURL, "https://mdm.acme.co/repo/munkitools-3.0.0.3298.plist"; *have != want { - t.Errorf("have %s, want %s", *have, want) - } - -}