From 175b84e6316c3b7f15c4d65be98c3714fc526ea3 Mon Sep 17 00:00:00 2001 From: Victor Vrantchan Date: Thu, 14 Dec 2017 20:35:50 -0500 Subject: [PATCH] refactor apns API and datastore (#359) Closes #306 --- cmd/micromdm/serve.go | 27 +++------ platform/apns/{ => builtin}/db.go | 15 ++--- platform/apns/endpoint.go | 34 ----------- platform/apns/logging_service.go | 34 ----------- platform/apns/middleware.go | 19 +++++++ platform/apns/push.go | 94 ++++++++++++++++++++++++++++++- platform/apns/server.go | 43 ++++++++++++++ platform/apns/service.go | 53 ++++++----------- platform/apns/transport_http.go | 75 ------------------------ 9 files changed, 187 insertions(+), 207 deletions(-) rename platform/apns/{ => builtin}/db.go (88%) delete mode 100644 platform/apns/endpoint.go delete mode 100644 platform/apns/logging_service.go create mode 100644 platform/apns/middleware.go create mode 100644 platform/apns/server.go delete mode 100644 platform/apns/transport_http.go diff --git a/cmd/micromdm/serve.go b/cmd/micromdm/serve.go index f4e887c3..33ad3fcd 100644 --- a/cmd/micromdm/serve.go +++ b/cmd/micromdm/serve.go @@ -44,6 +44,7 @@ import ( "github.com/micromdm/micromdm/mdm/enroll" "github.com/micromdm/micromdm/pkg/crypto" "github.com/micromdm/micromdm/platform/apns" + apnsbuiltin "github.com/micromdm/micromdm/platform/apns/builtin" "github.com/micromdm/micromdm/platform/appstore" appsbuiltin "github.com/micromdm/micromdm/platform/appstore/builtin" "github.com/micromdm/micromdm/platform/blueprint" @@ -216,18 +217,6 @@ func serve(args []string) error { checkinHandlers = checkin.MakeHTTPHandlers(ctx, e, opts...) } - var pushHandlers apns.HTTPHandlers - { - e := apns.Endpoints{ - PushEndpoint: apns.MakePushEndpoint(sm.pushService), - } - opts := []httptransport.ServerOption{ - httptransport.ServerErrorLogger(httpLogger), - httptransport.ServerErrorEncoder(checkin.EncodeError), - } - pushHandlers = apns.MakeHTTPHandlers(ctx, e, opts...) - } - var commandHandlers command.HTTPHandlers { e := command.Endpoints{ @@ -305,6 +294,8 @@ func serve(args []string) error { } depEndpoints := depapi.MakeServerEndpoints(depsvc) + apnsEndpoints := apns.MakeServerEndpoints(sm.pushService) + connectHandlers := connect.MakeHTTPHandlers(ctx, connectEndpoints, connectOpts...) scepHandler := scep.ServiceHandler(ctx, sm.scepService, httpLogger) @@ -329,6 +320,7 @@ func serve(args []string) error { appsHandler := appstore.MakeHTTPHandler(appEndpoints, logger) deviceHandler := device.MakeHTTPHandler(deviceEndpoints, logger) depHandlers := depapi.MakeHTTPHandler(depEndpoints, logger) + apnsHandlers := apns.MakeHTTPHandler(apnsEndpoints, logger) // API commands. Only handled if the user provides an api key. if *flAPIKey != "" { @@ -346,7 +338,7 @@ func serve(args []string) error { r.Handle("/v1/dep/account", apiAuthMiddleware(*flAPIKey, depHandlers)) r.Handle("/v1/dep/profiles", apiAuthMiddleware(*flAPIKey, depHandlers)) r.Handle("/v1/commands", apiAuthMiddleware(*flAPIKey, commandHandlers.NewCommandHandler)).Methods("POST") - r.Handle("/push/{udid}", apiAuthMiddleware(*flAPIKey, pushHandlers.PushHandler)) + r.Handle("/push/{udid}", apiAuthMiddleware(*flAPIKey, apnsHandlers)) } if *flRepoPath != "" { @@ -669,7 +661,7 @@ func (c *server) setupPushService(logger log.Logger) { } after: - db, err := apns.NewDB(c.db, c.pubclient) + db, err := apnsbuiltin.NewDB(c.db, c.pubclient) if err != nil { c.err = err return @@ -680,10 +672,9 @@ after: c.err = errors.Wrap(err, "starting micromdm push service") return } - c.pushService = apns.NewLoggingService( - service, - log.With(level.Info(logger), "component", "push"), - ) + c.pushService = apns.LoggingMiddleware( + log.With(level.Info(logger), "component", "apns"), + )(service) } func (c *server) setupEnrollmentService() { diff --git a/platform/apns/db.go b/platform/apns/builtin/db.go similarity index 88% rename from platform/apns/db.go rename to platform/apns/builtin/db.go index 0890431d..4dfc705d 100644 --- a/platform/apns/db.go +++ b/platform/apns/builtin/db.go @@ -1,4 +1,4 @@ -package apns +package builtin import ( "context" @@ -8,6 +8,7 @@ import ( "github.com/pkg/errors" "github.com/micromdm/micromdm/mdm/checkin" + "github.com/micromdm/micromdm/platform/apns" "github.com/micromdm/micromdm/platform/pubsub" ) @@ -43,15 +44,15 @@ func (e *notFound) Error() string { return fmt.Sprintf("not found: %s %s", e.ResourceType, e.Message) } -func (db *DB) PushInfo(udid string) (*PushInfo, error) { - var info PushInfo +func (db *DB) PushInfo(udid string) (*apns.PushInfo, error) { + var info apns.PushInfo err := db.View(func(tx *bolt.Tx) error { b := tx.Bucket([]byte(PushBucket)) v := b.Get([]byte(udid)) if v == nil { return ¬Found{"PushInfo", fmt.Sprintf("udid %s", udid)} } - return UnmarshalPushInfo(v, &info) + return apns.UnmarshalPushInfo(v, &info) }) if err != nil { return nil, err @@ -59,7 +60,7 @@ func (db *DB) PushInfo(udid string) (*PushInfo, error) { return &info, nil } -func (db *DB) Save(info *PushInfo) error { +func (db *DB) Save(info *apns.PushInfo) error { tx, err := db.DB.Begin(true) if err != nil { return errors.Wrap(err, "begin transaction") @@ -68,7 +69,7 @@ func (db *DB) Save(info *PushInfo) error { if bkt == nil { return fmt.Errorf("bucket %q not found!", PushBucket) } - pushproto, err := MarshalPushInfo(info) + pushproto, err := apns.MarshalPushInfo(info) if err != nil { return errors.Wrap(err, "marshalling PushInfo") } @@ -94,7 +95,7 @@ func (db *DB) pollCheckin(sub pubsub.Subscriber) error { fmt.Println(err) continue } - info := PushInfo{ + info := apns.PushInfo{ UDID: ev.Command.UDID, Token: ev.Command.Token.String(), PushMagic: ev.Command.PushMagic, diff --git a/platform/apns/endpoint.go b/platform/apns/endpoint.go deleted file mode 100644 index 3f6c9fb4..00000000 --- a/platform/apns/endpoint.go +++ /dev/null @@ -1,34 +0,0 @@ -package apns - -import ( - "context" - - "github.com/go-kit/kit/endpoint" -) - -type Endpoints struct { - PushEndpoint endpoint.Endpoint -} - -type pushRequest struct { - UDID string -} - -type pushResponse struct { - Status string `json:"status,omitempty"` - ID string `json:"push_notification_id,omitempty"` - Err error `json:"error,omitempty"` -} - -func (r pushResponse) error() error { return r.Err } - -func MakePushEndpoint(svc Service) endpoint.Endpoint { - return func(ctx context.Context, request interface{}) (interface{}, error) { - req := request.(pushRequest) - id, err := svc.Push(ctx, req.UDID) - if err != nil { - return pushResponse{Err: err, Status: "failure"}, nil - } - return pushResponse{Status: "success", ID: id}, nil - } -} diff --git a/platform/apns/logging_service.go b/platform/apns/logging_service.go deleted file mode 100644 index 9ce6b460..00000000 --- a/platform/apns/logging_service.go +++ /dev/null @@ -1,34 +0,0 @@ -package apns - -import ( - "context" - "time" - - "github.com/go-kit/kit/log" -) - -type loggingMiddleware struct { - logger log.Logger - next Service -} - -func NewLoggingService(svc Service, logger log.Logger) loggingMiddleware { - return loggingMiddleware{ - next: svc, - logger: logger, - } -} - -func (mw loggingMiddleware) Push(ctx context.Context, udid string) (id string, err error) { - defer func(begin time.Time) { - _ = mw.logger.Log( - "method", "Push", - "udid", udid, - "err", err, - "took", time.Since(begin), - ) - }(time.Now()) - - id, err = mw.next.Push(ctx, udid) - return -} diff --git a/platform/apns/middleware.go b/platform/apns/middleware.go new file mode 100644 index 00000000..368d2077 --- /dev/null +++ b/platform/apns/middleware.go @@ -0,0 +1,19 @@ +package apns + +import "github.com/go-kit/kit/log" + +type Middleware func(Service) Service + +func LoggingMiddleware(logger log.Logger) Middleware { + return func(next Service) Service { + return &loggingMiddleware{ + next: next, + logger: logger, + } + } +} + +type loggingMiddleware struct { + next Service + logger log.Logger +} diff --git a/platform/apns/push.go b/platform/apns/push.go index e27afd97..1a8fed49 100644 --- a/platform/apns/push.go +++ b/platform/apns/push.go @@ -1,7 +1,95 @@ package apns -import "context" +import ( + "context" + "encoding/json" + "net/http" + "strings" + "time" -type Service interface { - Push(ctx context.Context, udid string) (string, error) + "github.com/RobotsAndPencils/buford/payload" + "github.com/RobotsAndPencils/buford/push" + "github.com/go-kit/kit/endpoint" + "github.com/gorilla/mux" + "github.com/pkg/errors" + + "github.com/micromdm/micromdm/pkg/httputil" +) + +func (svc *PushService) Push(ctx context.Context, deviceUDID string) (string, error) { + info, err := svc.store.PushInfo(deviceUDID) + if err != nil { + return "", errors.Wrap(err, "retrieving PushInfo by UDID") + } + + p := payload.MDM{Token: info.PushMagic} + valid := push.IsDeviceTokenValid(info.Token) + if !valid { + return "", errors.New("invalid push token") + } + jsonPayload, err := json.Marshal(p) + if err != nil { + return "", errors.Wrap(err, "marshalling push notification payload") + } + result, err := svc.pushsvc.Push(info.Token, nil, jsonPayload) + if err != nil && strings.HasSuffix(err.Error(), "remote error: tls: internal error") { + // TODO: yuck, error substring searching. see: + // https://github.com/micromdm/micromdm/issues/150 + return result, errors.Wrap(err, "push error: possibly expired or invalid APNs certificate") + } + return result, err +} + +type pushRequest struct { + UDID string +} + +type pushResponse struct { + Status string `json:"status,omitempty"` + ID string `json:"push_notification_id,omitempty"` + Err error `json:"error,omitempty"` +} + +func (r pushResponse) Failed() error { return r.Err } + +func decodePushRequest(ctx context.Context, r *http.Request) (interface{}, error) { + vars := mux.Vars(r) + udid, ok := vars["udid"] + if !ok { + return 0, errors.New("apns: bad route") + } + return pushRequest{ + UDID: udid, + }, nil +} + +func decodePushResponse(_ context.Context, r *http.Response) (interface{}, error) { + var resp pushResponse + err := httputil.DecodeJSONResponse(r, &resp) + return resp, err +} + +func MakePushEndpoint(svc Service) endpoint.Endpoint { + return func(ctx context.Context, request interface{}) (interface{}, error) { + req := request.(pushRequest) + id, err := svc.Push(ctx, req.UDID) + if err != nil { + return pushResponse{Err: err, Status: "failure"}, nil + } + return pushResponse{Status: "success", ID: id}, nil + } +} + +func (mw loggingMiddleware) Push(ctx context.Context, udid string) (id string, err error) { + defer func(begin time.Time) { + _ = mw.logger.Log( + "method", "Push", + "udid", udid, + "err", err, + "took", time.Since(begin), + ) + }(time.Now()) + + id, err = mw.next.Push(ctx, udid) + return } diff --git a/platform/apns/server.go b/platform/apns/server.go new file mode 100644 index 00000000..62b38526 --- /dev/null +++ b/platform/apns/server.go @@ -0,0 +1,43 @@ +package apns + +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 { + PushEndpoint endpoint.Endpoint +} + +func MakeServerEndpoints(s Service) Endpoints { + return Endpoints{ + PushEndpoint: MakePushEndpoint(s), + } +} + +func MakeHTTPHandler(e Endpoints, logger log.Logger) *mux.Router { + r, options := httputil.NewRouter(logger) + + // GET /push/:udid create an APNS Push notification for a managed device or user(deprecated) + // POST /v1/push/:udid create an APNS Push notification for a managed device or user + + r.Methods("GET").Path("/push/{udid}").Handler(httptransport.NewServer( + e.PushEndpoint, + decodePushRequest, + httputil.EncodeJSONResponse, + options..., + )) + + r.Methods("POST").Path("/v1/push/{udid}").Handler(httptransport.NewServer( + e.PushEndpoint, + decodePushRequest, + httputil.EncodeJSONResponse, + options..., + )) + + return r +} diff --git a/platform/apns/service.go b/platform/apns/service.go index 7f35daed..67383048 100644 --- a/platform/apns/service.go +++ b/platform/apns/service.go @@ -3,13 +3,10 @@ package apns import ( "context" "crypto/tls" - "encoding/json" "fmt" "log" - "strings" "sync" - "github.com/RobotsAndPencils/buford/payload" "github.com/RobotsAndPencils/buford/push" "github.com/pkg/errors" @@ -18,8 +15,16 @@ import ( "github.com/micromdm/micromdm/platform/queue" ) -type Push struct { - db *DB +type Service interface { + Push(ctx context.Context, udid string) (string, error) +} + +type Store interface { + PushInfo(udid string) (*PushInfo, error) +} + +type PushService struct { + store Store start chan struct{} provider PushCertificateProvider @@ -31,17 +36,17 @@ type PushCertificateProvider interface { PushCertificate() (*tls.Certificate, error) } -type Option func(*Push) +type Option func(*PushService) func WithPushService(svc *push.Service) Option { - return func(p *Push) { + return func(p *PushService) { p.pushsvc = svc } } -func New(db *DB, provider PushCertificateProvider, sub pubsub.Subscriber, opts ...Option) (*Push, error) { - pushSvc := Push{ - db: db, +func New(db Store, provider PushCertificateProvider, sub pubsub.Subscriber, opts ...Option) (*PushService, error) { + pushSvc := PushService{ + store: db, provider: provider, start: make(chan struct{}), } @@ -60,7 +65,7 @@ func New(db *DB, provider PushCertificateProvider, sub pubsub.Subscriber, opts . return &pushSvc, nil } -func (svc *Push) startQueuedSubscriber(sub pubsub.Subscriber) error { +func (svc *PushService) startQueuedSubscriber(sub pubsub.Subscriber) error { commandQueuedEvents, err := sub.Subscribe(context.TODO(), "push-info", queue.CommandQueuedTopic) if err != nil { return errors.Wrapf(err, @@ -92,7 +97,7 @@ func (svc *Push) startQueuedSubscriber(sub pubsub.Subscriber) error { return nil } -func updateClient(svc *Push, sub pubsub.Subscriber) error { +func updateClient(svc *PushService, sub pubsub.Subscriber) error { configEvents, err := sub.Subscribe(context.TODO(), "push-server-configs", config.ConfigTopic) if err != nil { return errors.Wrap(err, "update push service client") @@ -130,27 +135,3 @@ func NewPushService(provider PushCertificateProvider) (*push.Service, error) { svc := push.NewService(client, push.Production) return svc, nil } - -func (svc *Push) Push(ctx context.Context, deviceUDID string) (string, error) { - info, err := svc.db.PushInfo(deviceUDID) - if err != nil { - return "", errors.Wrap(err, "retrieving PushInfo by UDID") - } - - p := payload.MDM{Token: info.PushMagic} - valid := push.IsDeviceTokenValid(info.Token) - if !valid { - return "", errors.New("invalid push token") - } - jsonPayload, err := json.Marshal(p) - if err != nil { - return "", errors.Wrap(err, "marshalling push notification payload") - } - result, err := svc.pushsvc.Push(info.Token, nil, jsonPayload) - if err != nil && strings.HasSuffix(err.Error(), "remote error: tls: internal error") { - // TODO: yuck, error substring searching. see: - // https://github.com/micromdm/micromdm/issues/150 - return result, errors.Wrap(err, "push error: possibly expired or invalid APNs certificate") - } - return result, err -} diff --git a/platform/apns/transport_http.go b/platform/apns/transport_http.go deleted file mode 100644 index fafc5a7e..00000000 --- a/platform/apns/transport_http.go +++ /dev/null @@ -1,75 +0,0 @@ -package apns - -import ( - "context" - "encoding/json" - "errors" - "net/http" - - httptransport "github.com/go-kit/kit/transport/http" - "github.com/gorilla/mux" -) - -type HTTPHandlers struct { - PushHandler http.Handler -} - -func MakeHTTPHandlers(ctx context.Context, endpoints Endpoints, opts ...httptransport.ServerOption) HTTPHandlers { - h := HTTPHandlers{ - PushHandler: httptransport.NewServer( - endpoints.PushEndpoint, - 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 errBadRoute = errors.New("bad route") - var req pushRequest - vars := mux.Vars(r) - udid, ok := vars["udid"] - if !ok { - return 0, errBadRoute - } - req.UDID = udid - return req, nil -} - -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) -}