From 6367ee23676f016f271a27e799ba9c52105db131 Mon Sep 17 00:00:00 2001 From: Victor Vrantchan Date: Tue, 17 May 2016 23:42:22 -0400 Subject: [PATCH] add dep endpoint --- checkin/endpoint.go | 11 ++++++ checkin/service.go | 64 +++++++++++++++++++++++++++++++---- checkin/transport.go | 20 +++++++++++ connect/service.go | 1 - device/datastore.go | 22 ++++++++---- device/device.go | 2 +- main.go | 8 ++++- management/endpoint_device.go | 23 +++++++++++++ management/service.go | 22 ++++++++++-- management/transport.go | 26 ++++++++++++++ 10 files changed, 181 insertions(+), 18 deletions(-) diff --git a/checkin/endpoint.go b/checkin/endpoint.go index 6ca3b0e1..39379790 100644 --- a/checkin/endpoint.go +++ b/checkin/endpoint.go @@ -52,3 +52,14 @@ type depEnrollmentResponse struct { } 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 d8fbf1d2..2ec00919 100644 --- a/checkin/service.go +++ b/checkin/service.go @@ -1,7 +1,11 @@ package checkin import ( + "errors" + "fmt" + "github.com/micromdm/mdm" + "github.com/micromdm/micromdm/command" "github.com/micromdm/micromdm/device" "github.com/micromdm/micromdm/management" ) @@ -11,20 +15,27 @@ type Service interface { Authenticate(mdm.CheckinCommand) error TokenUpdate(mdm.CheckinCommand) error Checkout(mdm.CheckinCommand) error - //EnrollDEP(udid string) (*device.Profile, error) + // EnrollDEP returns an enrollment profile + // during DEP Enrollment + EnrollDEP(udid, serial string) ([]byte, error) } // NewService creates a checkin service -func NewService(devices device.Datastore, ms management.Service) Service { +// profile holds an enrollment profile +func NewService(devices device.Datastore, ms management.Service, cs command.Service, profile []byte) Service { return &service{ - devices: devices, - mgmt: ms, + devices: devices, + mgmt: ms, + commands: cs, + profile: profile, } } type service struct { - devices device.Datastore - mgmt management.Service + devices device.Datastore + mgmt management.Service + commands command.Service + profile []byte } func (svc service) Authenticate(cmd mdm.CheckinCommand) error { @@ -44,6 +55,10 @@ func (svc service) Authenticate(cmd mdm.CheckinCommand) error { } func (svc service) TokenUpdate(cmd mdm.CheckinCommand) error { + if cmd.UserID != "" { + // don't handle user updates for now + return nil + } token := cmd.Token.String() unlockToken := cmd.UnlockToken.String() existing, err := svc.devices.GetDeviceByUDID(cmd.UDID, []string{"device_uuid"}...) @@ -77,3 +92,40 @@ 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 8d788114..12070746 100644 --- a/checkin/transport.go +++ b/checkin/transport.go @@ -1,6 +1,7 @@ package checkin import ( + "fmt" "io/ioutil" "net/http" @@ -27,9 +28,17 @@ 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 } @@ -38,6 +47,7 @@ func decodeMDMCheckinRequest(_ context.Context, r *http.Request) (interface{}, e if err != nil { return nil, err } + fmt.Println(string(data)) var request mdmCheckinRequest if err := plist.Unmarshal(data, &request); err != nil { return nil, err @@ -77,6 +87,16 @@ 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/connect/service.go b/connect/service.go index b4b5bd03..4c857531 100644 --- a/connect/service.go +++ b/connect/service.go @@ -55,7 +55,6 @@ func (svc service) checkRequeue(deviceUDID string) (int, error) { UDID: deviceUDID, RequestType: "DeviceConfigured", } - _, err := svc.commands.NewCommand(cmdRequest) if err != nil { return 0, err diff --git a/device/datastore.go b/device/datastore.go index 62ea4808..aff3d2ad 100644 --- a/device/datastore.go +++ b/device/datastore.go @@ -92,6 +92,15 @@ func (p UUID) where() string { return fmt.Sprintf("device_uuid = '%s'", p.UUID) } +// SerialNumber is a filter +type SerialNumber struct { + SerialNumber string +} + +func (p SerialNumber) where() string { + return fmt.Sprintf("serial_number = '%s'", p.SerialNumber) +} + type pgStore struct { *sqlx.DB } @@ -106,7 +115,7 @@ func (store pgStore) GetDeviceByUDID(udid string, fields ...string) (*Device, er func (store pgStore) GetDeviceByUUID(uuid string, fields ...string) (*Device, error) { var device Device s := strings.Join(fields, ", ") - query := `SELECT ` + s + ` FROM devices WHERE udid=$1 LIMIT 1` + query := `SELECT ` + s + ` FROM devices WHERE device_uuid=$1 LIMIT 1` return &device, sqlx.Get(store, &device, query, uuid) } @@ -167,11 +176,10 @@ func (store pgStore) Devices(params ...interface{}) ([]Device, error) { func (store pgStore) Save(msg string, dev *Device) error { var stmt string switch msg { - case "assign": - stmt = `INSERT INTO device_workflow - VALUES (:device_uuid, :workflow_uuid) - ON CONFLICT DO NOTHING;` - + case "assignWorkflow": + stmt = `UPDATE devices SET + workflow_uuid=:workflow_uuid + WHERE device_uuid=:device_uuid` case "tokenUpdate": stmt = `UPDATE devices SET awaiting_configuration=:awaiting_configuration, @@ -179,6 +187,7 @@ func (store pgStore) Save(msg string, dev *Device) error { apple_mdm_token=:apple_mdm_token, mdm_enrolled=:mdm_enrolled WHERE device_uuid=:device_uuid` + case "checkout": stmt = `UPDATE devices SET mdm_enrolled=:mdm_enrolled @@ -271,5 +280,6 @@ func migrate(db *sqlx.DB) { awaiting_configuration boolean ); CREATE UNIQUE INDEX IF NOT EXISTS serial_idx ON devices (serial_number);` + db.MustExec(schema) } diff --git a/device/device.go b/device/device.go index bc17d615..a1bba6d5 100644 --- a/device/device.go +++ b/device/device.go @@ -26,7 +26,7 @@ type Device struct { Token string `json:"token,omitempty" db:"apple_mdm_token,omitempty"` UnlockToken string `json:"unlock_token,omitempty" db:"unlock_token,omitempty"` Enrolled bool `json:"enrolled,omitempty" db:"mdm_enrolled,omitempty"` - Workflow string `json:"workflow,omitempty" db:"workflow_uuid,omitempty"` + Workflow string `json:"workflow_uuid,omitempty" db:"workflow_uuid,omitempty"` DEPDevice bool `json:"dep_device,omitempty" db:"dep_device,omitempty"` Description string `json:"description,omitempty" db:"description"` Model string `json:"model,omitempty" db:"model"` diff --git a/main.go b/main.go index d3b9f72f..521e8413 100644 --- a/main.go +++ b/main.go @@ -4,6 +4,7 @@ import ( "errors" "flag" "fmt" + "io/ioutil" "net/http" "os" @@ -78,6 +79,11 @@ func main() { logger.Log("err", "must set path to enrollment profile") os.Exit(1) } + enrollmentProfile, err := ioutil.ReadFile(*flEnrollment) + if err != nil { + logger.Log("err", err) + os.Exit(1) + } // check cert and key if -tls=true if *flTLS { @@ -148,7 +154,7 @@ func main() { dc := depClient(logger, *flDEPCK, *flDEPCS, *flDEPAT, *flDEPAS, *flDEPServerURL, *flDEPsim) mgmtSvc := management.NewService(deviceDB, workflowDB, dc, pushSvc) commandSvc := command.NewService(commandDB) - checkinSvc := checkin.NewService(deviceDB, mgmtSvc) + checkinSvc := checkin.NewService(deviceDB, mgmtSvc, commandSvc, enrollmentProfile) connectSvc := connect.NewService(deviceDB, commandSvc) httpLogger := log.NewContext(logger).With("component", "http") diff --git a/management/endpoint_device.go b/management/endpoint_device.go index 40d5d530..a7f84de8 100644 --- a/management/endpoint_device.go +++ b/management/endpoint_device.go @@ -57,3 +57,26 @@ func makeShowDeviceEndpoint(svc Service) endpoint.Endpoint { return showDeviceResponse{Device: dev}, nil } } + +type updateDeviceRequest struct { + DeviceUUID string `json:"-"` + Workflow *string `json:"workflow_uuid,omitempty" db:"workflow_uuid,omitempty"` +} + +type updateDeviceResponse struct { + Err error `json:"error,omitempty"` +} + +func (r updateDeviceResponse) error() error { return r.Err } +func (r updateDeviceResponse) status() int { return http.StatusNoContent } + +func makeUpdateDeviceEndpoint(svc Service) endpoint.Endpoint { + return func(ctx context.Context, request interface{}) (interface{}, error) { + req := request.(updateDeviceRequest) + var err error + if req.Workflow != nil { + err = svc.AssignWorkflow(req.DeviceUUID, *req.Workflow) + } + return updateDeviceResponse{Err: err}, nil + } +} diff --git a/management/service.go b/management/service.go index d2ec4afd..60f656ba 100644 --- a/management/service.go +++ b/management/service.go @@ -26,10 +26,15 @@ type Service interface { // Devices Devices() ([]device.Device, error) Device(uuid string) (*device.Device, error) - // dep - FetchDEPDevices() error - // push + // AssignWorkflow assigns a workflow to a device + AssignWorkflow(deviceUUID, workflowUUID string) error + + // push sends a new push notification to the device + // returning the notification ID Push(deviceUDID string) (string, error) + + // FetchDEPDevices updates the device datastore with devices from DEP + FetchDEPDevices() error } // NewService creates a management service @@ -141,3 +146,14 @@ func (svc service) Device(uuid string) (*device.Device, error) { dev := devices[0] return &dev, nil } + +func (svc service) AssignWorkflow(deviceUUID, workflowUUID string) error { + dev, err := svc.devices.GetDeviceByUUID(deviceUUID, + []string{"device_uuid"}..., + ) + if err != nil { + return errors.Wrap(err, "management: assign workflow") + } + dev.Workflow = workflowUUID + return svc.devices.Save("assignWorkflow", dev) +} diff --git a/management/transport.go b/management/transport.go index e913cbc8..d17656e1 100644 --- a/management/transport.go +++ b/management/transport.go @@ -87,6 +87,13 @@ func ServiceHandler(ctx context.Context, svc Service, logger kitlog.Logger) http encodeResponse, opts..., ) + updateDeviceHandler := kithttp.NewServer( + ctx, + makeUpdateDeviceEndpoint(svc), + decodeUpdateDeviceRequest, + encodeResponse, + opts..., + ) pushHandler := kithttp.NewServer( ctx, makePushEndpoint(svc), @@ -102,6 +109,7 @@ func ServiceHandler(ctx context.Context, svc Service, logger kitlog.Logger) http //devices r.Handle("/management/v1/devices", listDevicesHandler).Methods("GET") r.Handle("/management/v1/devices/{uuid}", showDeviceHandler).Methods("GET") + r.Handle("/management/v1/devices/{uuid}", updateDeviceHandler).Methods("PATCH") r.Handle("/management/v1/devices/{udid}/push", pushHandler).Methods("POST") // profiles r.Handle("/management/v1/profiles", addProfileHandler).Methods("POST") @@ -209,6 +217,24 @@ func decodePushRequest(_ context.Context, r *http.Request) (interface{}, error) return pushRequest{UDID: udid}, nil } +func decodeUpdateDeviceRequest(_ context.Context, r *http.Request) (interface{}, error) { + vars := mux.Vars(r) + deviceUUID, ok := vars["uuid"] + if !ok { + return nil, errBadRouting + } + // simple validation + if len(deviceUUID) != 36 { + return nil, errBadUUID + } + var request = updateDeviceRequest{DeviceUUID: deviceUUID} + err := json.NewDecoder(r.Body).Decode(&request) + if err == io.EOF { + return nil, errEmptyRequest + } + return request, 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)