diff --git a/app/http.go b/app/http.go index 5005c078..034309f5 100644 --- a/app/http.go +++ b/app/http.go @@ -14,6 +14,7 @@ import ( "github.com/micromdm/micromdm/checkin" "github.com/micromdm/micromdm/command" "github.com/micromdm/micromdm/connect" + "github.com/micromdm/micromdm/depenroll" "github.com/micromdm/micromdm/enroll" "github.com/micromdm/micromdm/management" "github.com/micromdm/micromdm/version" @@ -28,6 +29,7 @@ func makeHTTPHandler(logger log.Logger, sm *serviceManager) http.Handler { checkinHandler := checkin.ServiceHandler(ctx, sm.CheckinService, httpLogger) connectHandler := connect.ServiceHandler(ctx, sm.ConnectService, httpLogger) enrollHandler := enroll.ServiceHandler(ctx, sm.EnrollmentService, httpLogger) + depenrollHandler := depenroll.ServiceHandler(ctx, sm.DEPEnrollmentService, httpLogger) var handler http.Handler mux := http.NewServeMux() @@ -37,6 +39,7 @@ func makeHTTPHandler(logger log.Logger, sm *serviceManager) http.Handler { mux.Handle("/mdm/checkin", checkinHandler) mux.Handle("/mdm/connect", connectHandler) mux.Handle("/mdm/enroll", enrollHandler) + mux.Handle("/mdm/enroll/dep", depenrollHandler) mux.Handle("/_metrics", prometheus.Handler()) mux.Handle("/_version", version.Handler()) if sm.Server.PackageRepoPath != "" { diff --git a/app/services.go b/app/services.go index ae4ae386..402138be 100644 --- a/app/services.go +++ b/app/services.go @@ -20,6 +20,7 @@ import ( "github.com/micromdm/micromdm/command" cmdredis "github.com/micromdm/micromdm/command/service/redis" "github.com/micromdm/micromdm/connect" + "github.com/micromdm/micromdm/depenroll" "github.com/micromdm/micromdm/device" "github.com/micromdm/micromdm/driver" "github.com/micromdm/micromdm/enroll" @@ -46,6 +47,7 @@ func setupServices(config *Config, logger log.Logger) (*serviceManager, error) { sm.setupCommandService() sm.setupManagementService() sm.setupCheckinService() + sm.setupDEPEnrollmentService() sm.setupConnectService() sm.setupEnrollmentService() if sm.err != nil { @@ -65,11 +67,12 @@ type serviceManager struct { PushService *push.Service pushServiceCert - CommandService command.Service - ManagementService management.Service - CheckinService checkin.Service - ConnectService connect.Service - EnrollmentService enroll.Service + CommandService command.Service + ManagementService management.Service + CheckinService checkin.Service + ConnectService connect.Service + EnrollmentService enroll.Service + DEPEnrollmentService depenroll.Service *Config pool *redis.Pool @@ -176,22 +179,31 @@ func (s *serviceManager) setupCheckinService() { if s.err != nil { return } - var enrollmentProfile []byte + + s.CheckinService = checkin.NewService( + s.DeviceDatastore, + s.ManagementService, + ) + +} + +func (s *serviceManager) setupDEPEnrollmentService() { + if s.err != nil { + return + } // TODO make this optional - checkin.WithEnrollmentProfile([]byte) + var enrollmentProfile []byte if s.DEP.Enabled { enrollmentProfile, s.err = ioutil.ReadFile(s.Enrollment.ProfilePath) if s.err != nil { return } } - - s.CheckinService = checkin.NewService( + s.DEPEnrollmentService = depenroll.NewService( s.DeviceDatastore, - s.ManagementService, s.CommandService, enrollmentProfile, ) - } func (s *serviceManager) setupManagementService() { diff --git a/checkin/endpoint.go b/checkin/endpoint.go index 39379790..c472844d 100644 --- a/checkin/endpoint.go +++ b/checkin/endpoint.go @@ -4,8 +4,9 @@ import ( "errors" "github.com/go-kit/kit/endpoint" - "github.com/micromdm/mdm" "golang.org/x/net/context" + + "github.com/micromdm/mdm" ) // ErrInvalidMessageType is an invalid checking command @@ -41,25 +42,3 @@ func makeCheckinEndpoint(svc Service) endpoint.Endpoint { return mdmCheckinResponse{}, nil } } - -type depEnrollmentRequest struct { - mdm.DEPEnrollmentRequest -} - -type depEnrollmentResponse struct { - Profile []byte // MDM Enrollment Profile - Err error `plist:"error,omitempty"` -} - -func (r depEnrollmentResponse) error() error { return r.Err } - -func makeDEPEnrollmentEndpoint(svc Service) endpoint.Endpoint { - return func(ctx context.Context, request interface{}) (interface{}, error) { - req := request.(depEnrollmentRequest) - profile, err := svc.EnrollDEP(req.DEPEnrollmentRequest.UDID, req.DEPEnrollmentRequest.Serial) - if err != nil { - return depEnrollmentResponse{Err: err}, nil - } - return depEnrollmentResponse{Profile: profile}, nil - } -} diff --git a/checkin/service.go b/checkin/service.go index ba1bc8f1..01ffccfa 100644 --- a/checkin/service.go +++ b/checkin/service.go @@ -1,14 +1,11 @@ package checkin import ( - "errors" - "fmt" + "time" "github.com/micromdm/mdm" - "github.com/micromdm/micromdm/command" "github.com/micromdm/micromdm/device" "github.com/micromdm/micromdm/management" - "time" ) // Service defines methods for and MDM Checkin service @@ -16,27 +13,19 @@ type Service interface { Authenticate(mdm.CheckinCommand) error TokenUpdate(mdm.CheckinCommand) error Checkout(mdm.CheckinCommand) error - // EnrollDEP returns an enrollment profile - // during DEP Enrollment - EnrollDEP(udid, serial string) ([]byte, error) } -// NewService creates a checkin service -// profile holds an enrollment profile -func NewService(devices device.Datastore, ms management.Service, cs command.Service, profile []byte) Service { +// NewService creates a checkin service. +func NewService(devices device.Datastore, ms management.Service) Service { return &service{ - devices: devices, - mgmt: ms, - commands: cs, - profile: profile, + devices: devices, + mgmt: ms, } } type service struct { - devices device.Datastore - mgmt management.Service - commands command.Service - profile []byte + devices device.Datastore + mgmt management.Service } func (svc service) Authenticate(cmd mdm.CheckinCommand) error { @@ -108,40 +97,3 @@ func (svc service) Checkout(cmd mdm.CheckinCommand) error { } return nil } - -func (svc service) EnrollDEP(udid, serial string) ([]byte, error) { - err := svc.initialSetup(udid, serial) - if err != nil { - // TODO: stop ignoring the error there - fmt.Println(err) - } - return svc.profile, nil -} - -func (svc service) initialSetup(deviceUDID, serial string) error { - devs, err := svc.devices.Devices(device.SerialNumber{SerialNumber: serial}) - if err != nil { - return err - } - if len(devs) == 0 { - return errors.New("device not found") - } - dev := devs[0] - if dev.Workflow == "" { - // no workflow, send DeviceConfigured - return svc.sendConfigured(deviceUDID) - } - return nil -} - -func (svc service) sendConfigured(deviceUDID string) error { - cmdRequest := &mdm.CommandRequest{ - UDID: deviceUDID, - RequestType: "DeviceConfigured", - } - _, err := svc.commands.NewCommand(cmdRequest) - if err != nil { - return err - } - return nil -} diff --git a/checkin/transport.go b/checkin/transport.go index 12070746..fcb89725 100644 --- a/checkin/transport.go +++ b/checkin/transport.go @@ -7,7 +7,6 @@ import ( "golang.org/x/net/context" - "github.com/fullsailor/pkcs7" kitlog "github.com/go-kit/kit/log" kithttp "github.com/go-kit/kit/transport/http" "github.com/gorilla/mux" @@ -28,17 +27,9 @@ func ServiceHandler(ctx context.Context, svc Service, logger kitlog.Logger) http encodeResponse, opts..., ) - depEnrollmentHandler := kithttp.NewServer( - ctx, - makeDEPEnrollmentEndpoint(svc), - decodeMDMEnrollmentRequest, - encodeDEPEnrollmentResponse, - opts..., - ) r := mux.NewRouter() r.Handle("/mdm/checkin", checkinHandler).Methods("PUT") - r.Handle("/mdm/checkin", depEnrollmentHandler).Methods("POST") return r } @@ -55,26 +46,6 @@ func decodeMDMCheckinRequest(_ context.Context, r *http.Request) (interface{}, e return request, nil } -// The enrollment request is PkCS7 signed. -// We'll ignore everything but the content for now -func decodeMDMEnrollmentRequest(_ context.Context, r *http.Request) (interface{}, error) { - data, err := ioutil.ReadAll(r.Body) - if err != nil { - return nil, err - } - p7, err := pkcs7.Parse(data) - if err != nil { - return nil, err - } - // TODO: We should verify but not currently possible. Apple - // does no provide a cert for the CA. - var request depEnrollmentRequest - if err := plist.Unmarshal(p7.Content, &request); err != nil { - return nil, err - } - return request, nil -} - type errorer interface { error() error } @@ -87,16 +58,6 @@ type listEncoder interface { encodeList(w http.ResponseWriter) error } -func encodeDEPEnrollmentResponse(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 - } - resp := response.(depEnrollmentResponse) - w.Write(resp.Profile) - return 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) diff --git a/checkin/transport_test.go b/checkin/transport_test.go index 0c6075a1..4b8a07f8 100644 --- a/checkin/transport_test.go +++ b/checkin/transport_test.go @@ -3,17 +3,18 @@ package checkin import ( "bytes" "database/sql" + "io/ioutil" + "net/http" + "net/http/httptest" + "os" + "testing" + "github.com/DavidHuie/gomigrate" "github.com/go-kit/kit/log" "github.com/micromdm/micromdm/command" "github.com/micromdm/micromdm/device" "github.com/micromdm/micromdm/management" "golang.org/x/net/context" - "io/ioutil" - "net/http" - "net/http/httptest" - "os" - "testing" ) var testConn string = "user=postgres password= dbname=travis_ci_test sslmode=disable" @@ -62,8 +63,7 @@ func setup(t *testing.T) *fixtures { } f.mgmt = &mockMgmtService{} - f.profile = []byte{} - f.svc = NewService(f.devices, f.mgmt, f.cmd, f.profile) + f.svc = NewService(f.devices, f.mgmt) f.handler = ServiceHandler(f.ctx, f.svc, f.logger) f.server = httptest.NewServer(f.handler) diff --git a/depenroll/endpoint.go b/depenroll/endpoint.go new file mode 100644 index 00000000..4a158e13 --- /dev/null +++ b/depenroll/endpoint.go @@ -0,0 +1,30 @@ +package depenroll + +import ( + "github.com/go-kit/kit/endpoint" + "golang.org/x/net/context" + + "github.com/micromdm/mdm" +) + +type depEnrollmentRequest struct { + mdm.DEPEnrollmentRequest +} + +type depEnrollmentResponse struct { + Profile []byte // MDM Enrollment Profile + Err error `plist:"error,omitempty"` +} + +func (r depEnrollmentResponse) error() error { return r.Err } + +func makeDEPEnrollmentEndpoint(svc Service) endpoint.Endpoint { + return func(ctx context.Context, request interface{}) (interface{}, error) { + req := request.(depEnrollmentRequest) + profile, err := svc.EnrollDEP(req.DEPEnrollmentRequest.UDID, req.DEPEnrollmentRequest.Serial) + if err != nil { + return depEnrollmentResponse{Err: err}, nil + } + return depEnrollmentResponse{Profile: profile}, nil + } +} diff --git a/depenroll/service.go b/depenroll/service.go new file mode 100644 index 00000000..fdc49ab9 --- /dev/null +++ b/depenroll/service.go @@ -0,0 +1,66 @@ +package depenroll + +import ( + "errors" + "fmt" + + "github.com/micromdm/mdm" + "github.com/micromdm/micromdm/command" + "github.com/micromdm/micromdm/device" +) + +type Service interface { + // EnrollDEP returns an enrollment profile during DEP Enrollment. + EnrollDEP(udid, serial string) ([]byte, error) +} + +func NewService(devices device.Datastore, commands command.Service, profile []byte) Service { + return &service{ + devices: devices, + commands: commands, + profile: profile, + } +} + +type service struct { + devices device.Datastore + commands command.Service + profile []byte +} + +func (svc service) EnrollDEP(udid, serial string) ([]byte, error) { + err := svc.initialSetup(udid, serial) + if err != nil { + // TODO: stop ignoring the error there + fmt.Println(err) + } + return svc.profile, nil +} + +func (svc service) initialSetup(deviceUDID, serial string) error { + devs, err := svc.devices.Devices(device.SerialNumber{SerialNumber: serial}) + if err != nil { + return err + } + if len(devs) == 0 { + return errors.New("device not found") + } + dev := devs[0] + if dev.Workflow == "" { + // no workflow, send DeviceConfigured + return svc.sendConfigured(deviceUDID) + } + return nil +} + +func (svc service) sendConfigured(deviceUDID string) error { + cmdRequest := &mdm.CommandRequest{ + UDID: deviceUDID, + RequestType: "DeviceConfigured", + } + _, err := svc.commands.NewCommand(cmdRequest) + if err != nil { + return err + } + return nil +} diff --git a/depenroll/transport.go b/depenroll/transport.go new file mode 100644 index 00000000..51bc585d --- /dev/null +++ b/depenroll/transport.go @@ -0,0 +1,81 @@ +package depenroll + +import ( + "io/ioutil" + "net/http" + + "github.com/fullsailor/pkcs7" + kitlog "github.com/go-kit/kit/log" + kithttp "github.com/go-kit/kit/transport/http" + "github.com/gorilla/mux" + "github.com/groob/plist" + "golang.org/x/net/context" +) + +// ServiceHandler returns an HTTP Handler for the DEP enrollment service. +func ServiceHandler(ctx context.Context, svc Service, logger kitlog.Logger) http.Handler { + opts := []kithttp.ServerOption{ + kithttp.ServerErrorLogger(logger), + kithttp.ServerErrorEncoder(encodeError), + } + depEnrollmentHandler := kithttp.NewServer( + ctx, + makeDEPEnrollmentEndpoint(svc), + decodeMDMEnrollmentRequest, + encodeDEPEnrollmentResponse, + opts..., + ) + r := mux.NewRouter() + r.Handle("/mdm/enroll/dep", depEnrollmentHandler).Methods("POST") + return r +} + +type errorer interface { + error() error +} + +// The enrollment request is PkCS7 signed. +// We'll ignore everything but the content for now +func decodeMDMEnrollmentRequest(_ context.Context, r *http.Request) (interface{}, error) { + data, err := ioutil.ReadAll(r.Body) + if err != nil { + return nil, err + } + p7, err := pkcs7.Parse(data) + if err != nil { + return nil, err + } + // TODO: We should verify but not currently possible. Apple + // does no provide a cert for the CA. + var request depEnrollmentRequest + if err := plist.Unmarshal(p7.Content, &request); err != nil { + return nil, err + } + return request, nil +} + +func encodeDEPEnrollmentResponse(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 + } + resp := response.(depEnrollmentResponse) + w.Write(resp.Profile) + return nil +} + +// encode errors from business-logic +func encodeError(_ context.Context, err error, w http.ResponseWriter) { + // unwrap if the error is wrapped by kit http in it's own error type + if httperr, ok := err.(kithttp.Error); ok { + err = httperr.Err + } + switch err { + default: + w.WriteHeader(http.StatusInternalServerError) + } + // w.Header().Set("Content-Type", "application/json; charset=utf-8") + plist.NewEncoder(w).Encode(map[string]interface{}{ + "error": err.Error(), + }) +}