From 2f232c194a6381af5b5dd78cae0f8fefb199ba70 Mon Sep 17 00:00:00 2001 From: Victor Vrantchan Date: Thu, 17 Mar 2016 20:17:59 -0400 Subject: [PATCH] first commit --- .docker/data/Dockerfile | 5 + .docker/docker-compose-dev.yml | 60 +++++ .docker/haproxy/Dockerfile | 3 + .../root/usr/local/etc/haproxy/haproxy.cfg | 46 ++++ .docker/pkgr/Dockerfile | 5 + .docker/pkgr/root/etc/nginx/nginx.conf | 41 +++ .docker/prometheus/Dockerfile | 3 + .../root/etc/prometheus/prometheus.yml | 13 + .gitignore | 3 + Dockerfile.dev | 20 ++ checkin/encode_decode.go | 60 +++++ checkin/endpoint.go | 27 ++ checkin/instrumenting.go | 48 ++++ checkin/logging.go | 52 ++++ checkin/request_response.go | 13 + checkin/service.go | 145 ++++++++++ command/datastore.go | 173 ++++++++++++ command/service.go | 128 +++++++++ command/transport.go | 251 ++++++++++++++++++ connect/encode_decode.go | 76 ++++++ connect/endpoint.go | 47 ++++ connect/request_response.go | 14 + connect/service.go | 78 ++++++ device/device.go | 196 ++++++++++++++ glide.lock | 79 ++++++ glide.yaml | 28 ++ main.go | 202 ++++++++++++++ 27 files changed, 1816 insertions(+) create mode 100644 .docker/data/Dockerfile create mode 100644 .docker/docker-compose-dev.yml create mode 100644 .docker/haproxy/Dockerfile create mode 100644 .docker/haproxy/root/usr/local/etc/haproxy/haproxy.cfg create mode 100644 .docker/pkgr/Dockerfile create mode 100644 .docker/pkgr/root/etc/nginx/nginx.conf create mode 100644 .docker/prometheus/Dockerfile create mode 100644 .docker/prometheus/root/etc/prometheus/prometheus.yml create mode 100644 .gitignore create mode 100644 Dockerfile.dev create mode 100644 checkin/encode_decode.go create mode 100644 checkin/endpoint.go create mode 100644 checkin/instrumenting.go create mode 100644 checkin/logging.go create mode 100644 checkin/request_response.go create mode 100644 checkin/service.go create mode 100644 command/datastore.go create mode 100644 command/service.go create mode 100644 command/transport.go create mode 100644 connect/encode_decode.go create mode 100644 connect/endpoint.go create mode 100644 connect/request_response.go create mode 100644 connect/service.go create mode 100644 device/device.go create mode 100644 glide.lock create mode 100644 glide.yaml create mode 100644 main.go diff --git a/.docker/data/Dockerfile b/.docker/data/Dockerfile new file mode 100644 index 00000000..c93079d1 --- /dev/null +++ b/.docker/data/Dockerfile @@ -0,0 +1,5 @@ +FROM alpine:3.3 + +COPY root / + +ENTRYPOINT "/bin/true" diff --git a/.docker/docker-compose-dev.yml b/.docker/docker-compose-dev.yml new file mode 100644 index 00000000..3c3cc592 --- /dev/null +++ b/.docker/docker-compose-dev.yml @@ -0,0 +1,60 @@ +haproxy: + build: haproxy/ + restart: always + ports: + - "80:80" + - "443:443" + volumes_from: + - data + links: + - micromdm + - pkgrepo + - prometheus + +micromdm: + build: ../ + dockerfile: Dockerfile.dev + command: "/micromdm -tls=false" + expose: + - "80" + links: + - postgres + - redis + +postgres: + image: postgres + restart: always + environment: + - POSTGRES_USER=micromdm + - POSTGRES_PASSWORD=micromdm + - POSTGRES_DB=micromdm + - SSLMODE=disable + expose: + - "5432" + +prometheus: + build: prometheus/ + ports: + - "9090:9090" + links: + - micromdm + +redis: + image: redis + restart: always + expose: + - "6379" + +pkgrepo: + build: pkgr/ + expose: + - "80" + volumes_from: + - data + + +data: + build: data/ + volumes: + - /certs.d + - /pkgrepo diff --git a/.docker/haproxy/Dockerfile b/.docker/haproxy/Dockerfile new file mode 100644 index 00000000..09f244ea --- /dev/null +++ b/.docker/haproxy/Dockerfile @@ -0,0 +1,3 @@ +FROM haproxy:1.5 + +COPY root / diff --git a/.docker/haproxy/root/usr/local/etc/haproxy/haproxy.cfg b/.docker/haproxy/root/usr/local/etc/haproxy/haproxy.cfg new file mode 100644 index 00000000..d5a5cbfe --- /dev/null +++ b/.docker/haproxy/root/usr/local/etc/haproxy/haproxy.cfg @@ -0,0 +1,46 @@ +global + maxconn 4096 + +defaults + log global + mode http + option httplog + option dontlognull + timeout connect 5000 + timeout client 50000 + timeout server 50000 + +frontend localnodes + bind *:80 + bind *:443 ssl crt /certs.d ciphers ECDHE-RSA-AES256-SHA:RC4-SHA:RC4:HIGH:!MD5:!aNULL:!EDH:!AESGCM + mode http + # micromdm container + acl is_micromdm path_beg /mdm + use_backend micromdm if is_micromdm + # prometheus + acl is_prometheus path_beg /graph + use_backend prometheus if is_prometheus + #pkgrepo + acl is_pkgrepo path_beg /repo + use_backend pkgrepo if is_pkgrepo + +backend micromdm + balance leastconn + option httpclose + option forwardfor + http-request set-header X-Forwarded-Port %[dst_port] + server micromdm micromdm:80 check + +backend prometheus + balance leastconn + option httpclose + option forwardfor + http-request set-header X-Forwarded-Port %[dst_port] + server prometheus prometheus:9090 check + +backend pkgrepo + balance leastconn + option httpclose + option forwardfor + http-request set-header X-Forwarded-Port %[dst_port] + server pkgrepo pkgrepo:80 check diff --git a/.docker/pkgr/Dockerfile b/.docker/pkgr/Dockerfile new file mode 100644 index 00000000..e196b4ec --- /dev/null +++ b/.docker/pkgr/Dockerfile @@ -0,0 +1,5 @@ +FROM nginx:1.9.9 + +COPY root / + +EXPOSE 80 diff --git a/.docker/pkgr/root/etc/nginx/nginx.conf b/.docker/pkgr/root/etc/nginx/nginx.conf new file mode 100644 index 00000000..f524b51d --- /dev/null +++ b/.docker/pkgr/root/etc/nginx/nginx.conf @@ -0,0 +1,41 @@ +worker_processes 1; + +events { + worker_connections 1024; +} + + +http { + include mime.types; + default_type application/octet-stream; + + sendfile on; + keepalive_timeout 65; + gzip on; + + server { + listen 80 default_server; + + + location / { + root html; + index index.html index.htm; + } + + location /repo { + alias /pkgrepo; + autoindex on; + } + + location /nginx_status { + stub_status on; + access_log off; + } + + error_page 500 502 503 504 /50x.html; + location = /50x.html { + root html; + } + + } +} diff --git a/.docker/prometheus/Dockerfile b/.docker/prometheus/Dockerfile new file mode 100644 index 00000000..69da8b00 --- /dev/null +++ b/.docker/prometheus/Dockerfile @@ -0,0 +1,3 @@ +FROM prom/prometheus + +COPY root / diff --git a/.docker/prometheus/root/etc/prometheus/prometheus.yml b/.docker/prometheus/root/etc/prometheus/prometheus.yml new file mode 100644 index 00000000..0e5f2069 --- /dev/null +++ b/.docker/prometheus/root/etc/prometheus/prometheus.yml @@ -0,0 +1,13 @@ +global: + scrape_interval: 15s + evaluation_interval: 15s + external_labels: + monitor: 'micromdm' + +scrape_configs: + - job_name: 'micromdm' + scrape_interval: 5s + scrape_timeout: 10s + target_groups: + - targets: ['micromdm'] + diff --git a/.gitignore b/.gitignore new file mode 100644 index 00000000..6194e02d --- /dev/null +++ b/.gitignore @@ -0,0 +1,3 @@ +vendor/ +.docker/data/root/ +micromdm diff --git a/Dockerfile.dev b/Dockerfile.dev new file mode 100644 index 00000000..0da2f654 --- /dev/null +++ b/Dockerfile.dev @@ -0,0 +1,20 @@ +FROM golang:alpine + +ENV GO15VENDOREXPERIMENT=1 + +RUN apk --no-cache add curl git && \ + curl -L https://github.com/Masterminds/glide/releases/download/0.9.1/glide-0.9.1-linux-amd64.tar.gz -o glide.tar.gz && \ + tar xzf glide.tar.gz -C /tmp && \ + mv /tmp/linux-amd64/glide /usr/bin/ && \ + rm -f glide.tar.gz && \ + rm -rf /tmp/linux-amd64 && \ + apk del curl + +RUN mkdir -p /go/src/github.com/micromdm/micromdm/ +WORKDIR /go/src/github.com/micromdm/micromdm/ +COPY . /go/src/github.com/micromdm/micromdm/ + +RUN glide install +RUN go build && mv micromdm / + +CMD ["/micromdm"] diff --git a/checkin/encode_decode.go b/checkin/encode_decode.go new file mode 100644 index 00000000..6b82397a --- /dev/null +++ b/checkin/encode_decode.go @@ -0,0 +1,60 @@ +package checkin + +import ( + "log" + "net/http" + + "github.com/groob/plist" +) + +func decodeMDMCheckinRequest(r *http.Request) (interface{}, error) { + var request mdmCheckinRequest + if err := plist.NewDecoder(r.Body).Decode(&request); err != nil { + return nil, err + } + return request, nil +} + +// errorer is implemented by all concrete response types. It allows us to +// change the HTTP response code without needing to trigger an endpoint +// (transport-level) error. For more information, read the big comment in +// endpoint.go. +type errorer interface { + error() error +} + +// encodeResponse is the common method to encode all response types to the +// client. I chose to do it this way because I didn't know if something more +// specific was necessary. It's certainly possible to specialize on a +// per-response (per-method) basis. +func encodeResponse(w http.ResponseWriter, response interface{}) error { + if e, ok := response.(errorer); ok && e.error() != nil { + // Not a Go kit transport error, but a business-logic error. + // Provide those as HTTP errors. + encodeError(w, e.error()) + return nil + } + enc := plist.NewEncoder(w) + enc.Indent(" ") + return enc.Encode(response) +} + +func encodeError(w http.ResponseWriter, err error) { + w.WriteHeader(codeFrom(err)) + response := map[string]interface{}{ + "error": err.Error(), + } + enc := plist.NewEncoder(w) + enc.Indent(" ") + err = enc.Encode(response) + if err != nil { + log.Println(err) + } +} + +func codeFrom(err error) int { + switch err { + default: + return http.StatusInternalServerError + } +} diff --git a/checkin/endpoint.go b/checkin/endpoint.go new file mode 100644 index 00000000..4e9d3a1c --- /dev/null +++ b/checkin/endpoint.go @@ -0,0 +1,27 @@ +package checkin + +import ( + "github.com/go-kit/kit/endpoint" + "golang.org/x/net/context" +) + +func makeCheckinEndpoint(svc MDMCheckinService) endpoint.Endpoint { + return func(ctx context.Context, request interface{}) (interface{}, error) { + req := request.(mdmCheckinRequest) + var err error + switch req.MessageType { + case "Authenticate": + err = svc.Authenticate(req.CheckinCommand) + case "TokenUpdate": + err = svc.TokenUpdate(req.CheckinCommand) + case "CheckOut": + err = svc.Checkout(req.CheckinCommand) + default: + return mdmCheckinResponse{ErrInvalidMessageType}, nil + } + if err != nil { + return mdmCheckinResponse{err}, nil + } + return mdmCheckinResponse{}, nil + } +} diff --git a/checkin/instrumenting.go b/checkin/instrumenting.go new file mode 100644 index 00000000..4b2b6724 --- /dev/null +++ b/checkin/instrumenting.go @@ -0,0 +1,48 @@ +package checkin + +import ( + "fmt" + "time" + + "github.com/go-kit/kit/metrics" + "github.com/micromdm/mdm" +) + +type instrumentingMiddleware struct { + requestCount metrics.Counter + requestLatency metrics.TimeHistogram + MDMCheckinService +} + +func (mw instrumentingMiddleware) Authenticate(cmd mdm.CheckinCommand) (err error) { + defer func(begin time.Time) { + methodField := metrics.Field{Key: "MessageType", Value: cmd.MessageType} + errorField := metrics.Field{Key: "error", Value: fmt.Sprintf("%v", err)} + mw.requestCount.With(methodField).With(errorField).Add(1) + mw.requestLatency.With(methodField).With(errorField).Observe(time.Since(begin)) + }(time.Now()) + err = mw.MDMCheckinService.Authenticate(cmd) + return err +} + +func (mw instrumentingMiddleware) TokenUpdate(cmd mdm.CheckinCommand) (err error) { + defer func(begin time.Time) { + methodField := metrics.Field{Key: "MessageType", Value: cmd.MessageType} + errorField := metrics.Field{Key: "error", Value: fmt.Sprintf("%v", err)} + mw.requestCount.With(methodField).With(errorField).Add(1) + mw.requestLatency.With(methodField).With(errorField).Observe(time.Since(begin)) + }(time.Now()) + err = mw.MDMCheckinService.TokenUpdate(cmd) + return err +} + +func (mw instrumentingMiddleware) Checkout(cmd mdm.CheckinCommand) (err error) { + defer func(begin time.Time) { + methodField := metrics.Field{Key: "MessageType", Value: cmd.MessageType} + errorField := metrics.Field{Key: "error", Value: fmt.Sprintf("%v", err)} + mw.requestCount.With(methodField).With(errorField).Add(1) + mw.requestLatency.With(methodField).With(errorField).Observe(time.Since(begin)) + }(time.Now()) + err = mw.MDMCheckinService.Checkout(cmd) + return err +} diff --git a/checkin/logging.go b/checkin/logging.go new file mode 100644 index 00000000..dd37f6e3 --- /dev/null +++ b/checkin/logging.go @@ -0,0 +1,52 @@ +package checkin + +import ( + "time" + + "github.com/go-kit/kit/log" + "github.com/micromdm/mdm" +) + +type loggingMiddleware struct { + logger log.Logger + MDMCheckinService +} + +func (mw loggingMiddleware) Authenticate(cmd mdm.CheckinCommand) (err error) { + defer func(begin time.Time) { + _ = mw.logger.Log( + "MessageType", cmd.MessageType, + "err", err, + "udid", cmd.UDID, + "took", time.Since(begin), + ) + }(time.Now()) + err = mw.MDMCheckinService.Authenticate(cmd) + return err +} + +func (mw loggingMiddleware) TokenUpdate(cmd mdm.CheckinCommand) (err error) { + defer func(begin time.Time) { + _ = mw.logger.Log( + "MessageType", cmd.MessageType, + "err", err, + "udid", cmd.UDID, + "took", time.Since(begin), + ) + }(time.Now()) + err = mw.MDMCheckinService.TokenUpdate(cmd) + return err +} + +func (mw loggingMiddleware) Checkout(cmd mdm.CheckinCommand) (err error) { + defer func(begin time.Time) { + _ = mw.logger.Log( + "MessageType", cmd.MessageType, + "err", err, + "udid", cmd.UDID, + "took", time.Since(begin), + ) + }(time.Now()) + err = mw.MDMCheckinService.Checkout(cmd) + return err +} diff --git a/checkin/request_response.go b/checkin/request_response.go new file mode 100644 index 00000000..4b8d2e71 --- /dev/null +++ b/checkin/request_response.go @@ -0,0 +1,13 @@ +package checkin + +import "github.com/micromdm/mdm" + +type mdmCheckinRequest struct { + mdm.CheckinCommand +} + +type mdmCheckinResponse struct { + Err error `plist:"error,omitempty"` +} + +func (r mdmCheckinResponse) error() error { return r.Err } diff --git a/checkin/service.go b/checkin/service.go new file mode 100644 index 00000000..0d44cb59 --- /dev/null +++ b/checkin/service.go @@ -0,0 +1,145 @@ +package checkin + +import ( + "errors" + "net/http" + "os" + "time" + + "golang.org/x/net/context" + + "github.com/go-kit/kit/log" + "github.com/go-kit/kit/metrics" + kitprometheus "github.com/go-kit/kit/metrics/prometheus" + httptransport "github.com/go-kit/kit/transport/http" + "github.com/micromdm/mdm" + "github.com/micromdm/micromdm/device" + stdprometheus "github.com/prometheus/client_golang/prometheus" +) + +// ErrInvalidMessageType is an invalid checking command +var ErrInvalidMessageType = errors.New("Invalid MessageType") + +// MDMCheckinService models Apple's MDM Checkin commands +type MDMCheckinService interface { + Authenticate(mdm.CheckinCommand) error + TokenUpdate(mdm.CheckinCommand) error + Checkout(mdm.CheckinCommand) error +} + +// NewCheckinService creates a new MDM Checkin Service +func NewCheckinService(options ...func(*config) error) MDMCheckinService { + conf := &config{} + defaultLogger := log.NewLogfmtLogger(os.Stderr) + for _, option := range options { + if err := option(conf); err != nil { + defaultLogger.Log("err", err) + os.Exit(1) + } + } + var svc MDMCheckinService + svc = mdmCheckinService{conf.db} + if conf.logger != nil { + svc = loggingMiddleware{conf.logger, svc} + } + + fieldKeys := []string{"MessageType", "error"} + requestCount := kitprometheus.NewCounter(stdprometheus.CounterOpts{ + Name: "request_count", + Help: "http request count", + }, fieldKeys) + requestLatency := metrics.NewTimeHistogram(time.Microsecond, kitprometheus.NewSummary(stdprometheus.SummaryOpts{ + Name: "request_latency", + Help: "http request duration", + }, fieldKeys)) + svc = instrumentingMiddleware{requestCount, requestLatency, svc} // add metrics + return svc +} + +// Logger adds a logger to the service +func Logger(logger log.Logger) func(*config) error { + return func(c *config) error { + c.logger = logger + return nil + } +} + +// Datastore adds a db connection to the service +func Datastore(db device.Datastore) func(*config) error { + return func(c *config) error { + c.db = db + return nil + } +} + +type config struct { + logger log.Logger + db device.Datastore +} + +type mdmCheckinService struct { + db device.Datastore +} + +func (svc mdmCheckinService) Authenticate(cmd mdm.CheckinCommand) error { + dev := &device.Device{ + UDID: cmd.UDID, + SerialNumber: &cmd.SerialNumber, + OSVersion: &cmd.OSVersion, + BuildVersion: &cmd.BuildVersion, + ProductName: &cmd.ProductName, + IMEI: &cmd.IMEI, + MEID: &cmd.MEID, + MDMTopic: &cmd.Topic, + } + return svc.db.AddDevice(dev) +} + +func (svc mdmCheckinService) TokenUpdate(cmd mdm.CheckinCommand) error { + token := cmd.Token.String() + unlockToken := cmd.UnlockToken.String() + existing, err := svc.db.GetDeviceByUDID(cmd.UDID) + if err != nil { + return err + } + existing.Token = &token + existing.MDMTopic = &cmd.Topic + existing.PushMagic = &cmd.PushMagic + existing.UnlockToken = &unlockToken + existing.AwaitingConfiguration = &cmd.AwaitingConfiguration + existing.Enrolled = boolPtr(true) + err = svc.db.SaveDevice(existing) + if err != nil { + return err + } + return nil +} + +func (svc mdmCheckinService) Checkout(cmd mdm.CheckinCommand) error { + existing, err := svc.db.GetDeviceByUDID(cmd.UDID) + if err != nil { + return err + } + existing.Enrolled = boolPtr(false) + return nil +} + +// return a pointer to a boolean +func boolPtr(b bool) *bool { + return &b +} + +// ServiceHandler creates an http handler +func ServiceHandler(ctx context.Context, svc MDMCheckinService) http.Handler { + // endpoint + checkin := makeCheckinEndpoint(svc) + + // handler + checkinHandler := httptransport.NewServer( + ctx, + checkin, + decodeMDMCheckinRequest, + encodeResponse, + ) + return checkinHandler +} diff --git a/command/datastore.go b/command/datastore.go new file mode 100644 index 00000000..121252b6 --- /dev/null +++ b/command/datastore.go @@ -0,0 +1,173 @@ +package command + +import ( + "bytes" + "errors" + "fmt" + "os" + "time" + + "github.com/garyburd/redigo/redis" + "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 manages MDM Payloads in redis +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) +} + +type redisDB struct { + pool *redis.Pool +} + +// NewDB creates a new databases connection +func NewDB(driver, conn string, options ...func(*config) error) Datastore { + conf := &config{} + defaultLogger := log.NewLogfmtLogger(os.Stderr) + for _, option := range options { + if err := option(conf); err != nil { + defaultLogger.Log("err", err) + os.Exit(1) + } + } + switch driver { + case "redis": + return redisDB{pool: redisPool(conn, conf.logger)} + default: + conf.logger.Log("err", "unknown driver") + os.Exit(1) + return nil + } +} + +func redisPool(conn string, logger log.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 log.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) + } +} + +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 +} diff --git a/command/service.go b/command/service.go new file mode 100644 index 00000000..d9d11455 --- /dev/null +++ b/command/service.go @@ -0,0 +1,128 @@ +package command + +import ( + "net/http" + "os" + + "golang.org/x/net/context" + + httptransport "github.com/go-kit/kit/transport/http" + + "github.com/go-kit/kit/log" + "github.com/gorilla/mux" + "github.com/micromdm/mdm" +) + +// MDMCommandService allows creating and deleting MDM Command Payloads +type MDMCommandService interface { + NewCommand(*mdm.CommandRequest) (*mdm.Payload, error) + NextCommand(udid string) ([]byte, int, error) + DeleteCommand(deviceUDID, commandUUID string) (int, error) +} + +type mdmCommandService struct { + // a redis datastore + db Datastore +} + +type config struct { + logger log.Logger + db Datastore +} + +// NewCommandService creates a new MDM Command Service +func NewCommandService(options ...func(*config) error) MDMCommandService { + conf := &config{} + defaultLogger := log.NewLogfmtLogger(os.Stderr) + for _, option := range options { + if err := option(conf); err != nil { + defaultLogger.Log("err", err) + os.Exit(1) + } + } + var svc MDMCommandService + svc = mdmCommandService{db: conf.db} + return svc +} + +// Logger adds a logger to the service +func Logger(logger log.Logger) func(*config) error { + return func(c *config) error { + c.logger = logger + return nil + } +} + +// DB adds a db connection to the service +func DB(db Datastore) func(*config) error { + return func(c *config) error { + c.db = db + return nil + } +} + +// ServiceHandler returns an http handler for the command service +func ServiceHandler(ctx context.Context, svc MDMCommandService) http.Handler { + commonOptions := []httptransport.ServerOption{ + httptransport.ServerErrorEncoder(encodeError), + } + newCommandEndpoint := makeNewCommandEndpoint(svc) + newCommandHandler := httptransport.NewServer( + ctx, + newCommandEndpoint, + decodeNewCommandRequest, + encodeResponse, + commonOptions..., + ) + nextCommandEndpoint := makeNextCommandEndpoint(svc) + nextCommandHandler := httptransport.NewServer( + ctx, + nextCommandEndpoint, + decodeNextCommandRequest, + encodeResponse, + commonOptions..., + ) + deleteCommandEndpoint := makeDeleteCommandEndpoint(svc) + deleteCommandHandler := httptransport.NewServer( + ctx, + deleteCommandEndpoint, + decodeDeleteCommandRequest, + encodeResponse, + commonOptions..., + ) + r := mux.NewRouter() + r.Methods("POST").Path("/mdm/commands").Handler(newCommandHandler) + r.Methods("GET").Path("/mdm/commands/{udid}/next").Handler(nextCommandHandler) + r.Methods("DELETE").Path("/mdm/commands/{udid}/{uuid}").Handler(deleteCommandHandler) + return r +} + +func (svc mdmCommandService) 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 mdmCommandService) NextCommand(udid string) ([]byte, int, error) { + return svc.db.NextCommand(udid) +} + +// DeleteCommand returns an MDM Payload from a list of queued payloads +func (svc mdmCommandService) DeleteCommand(deviceUDID, commandUUID string) (int, error) { + return svc.db.DeleteCommand(deviceUDID, commandUUID) +} diff --git a/command/transport.go b/command/transport.go new file mode 100644 index 00000000..2ffde194 --- /dev/null +++ b/command/transport.go @@ -0,0 +1,251 @@ +package command + +import ( + "bytes" + "encoding/json" + "errors" + "fmt" + "io/ioutil" + "log" + "net/http" + + "github.com/go-kit/kit/endpoint" + "github.com/gorilla/mux" + "github.com/micromdm/mdm" + "golang.org/x/net/context" +) + +var ( + // ErrEmptyRequest is returned if the request body is empty + ErrEmptyRequest = errors.New("Request must contain UDID of the device") + errBadRouting = errors.New("inconsistent mapping between route and handler (programmer error)") +) + +// NewCommandRequest represents an HTTP Request for a new MDM Command +type NewCommandRequest struct { + *mdm.CommandRequest +} + +func decodeNewCommandRequest(r *http.Request) (interface{}, error) { + var request NewCommandRequest + err := json.NewDecoder(r.Body).Decode(&request.CommandRequest) + return request, err +} + +// NewCommandResponse is a command reponse +type NewCommandResponse struct { + *mdm.Payload + Err error `json:"error,omitempty"` +} + +func (r NewCommandResponse) error() error { return r.Err } + +// errorer is implemented by all concrete response types. It allows us to +// change the HTTP response code without needing to trigger an endpoint +// (transport-level) error. For more information, read the big comment in +// endpoint.go. +type errorer interface { + error() error +} + +// NextCommandRequest is a request to return the next command in a device queue +type NextCommandRequest struct { + UDID string +} + +func decodeNextCommandRequest(r *http.Request) (interface{}, error) { + vars := mux.Vars(r) + udid, ok := vars["udid"] + if !ok { + return nil, errBadRouting + } + var request NextCommandRequest + request.UDID = udid + return request, nil +} + +// EncodeNextCommandRequest encodes a request for the NextCommand endpoint +func EncodeNextCommandRequest(r *http.Request, request interface{}) error { + req := request.(NextCommandRequest) + path := r.URL.Path + r.URL.Path = fmt.Sprintf("%v/%v/next", path, req.UDID) + var buf bytes.Buffer + if err := json.NewEncoder(&buf).Encode(request); err != nil { + return err + } + r.Body = ioutil.NopCloser(&buf) + return nil +} + +// NextCommandResponse is a response for the next command +type NextCommandResponse struct { + Payload []byte `json:"command_payload"` + Total int `json:"total_payloads"` + Err error `json:"error,omitempty"` +} + +// DecodeNextCommandResponse decodes the response from the provided HTTP response, +// simply by JSON decoding from the response body. It's designed to be used in +// transport/http.Client. +// first decode into map[string]interface{} and check for error in the response +func DecodeNextCommandResponse(resp *http.Response) (interface{}, error) { + var r map[string]interface{} + var response NextCommandResponse + err := json.NewDecoder(resp.Body).Decode(&r) + if rs, ok := r["error"]; ok { + response.Err = errors.New(rs.(string)) + } + if rs, ok := r["total_payloads"]; ok { + response.Total = int(rs.(float64)) + } + if rs, ok := r["command_payload"]; ok { + response.Payload = []byte(rs.(string)) + } + return response, err +} + +func (r NextCommandResponse) error() error { return r.Err } + +// DeleteCommandRequest is a request to delete a command +type DeleteCommandRequest struct { + // device UDID + UDID string + // command UUID + UUID string +} + +// EncodeDeleteCommandRequest encodes a request for the NextCommand endpoint +func EncodeDeleteCommandRequest(r *http.Request, request interface{}) error { + req := request.(DeleteCommandRequest) + path := r.URL.Path + r.URL.Path = fmt.Sprintf("%v/%v/%v", path, req.UDID, req.UUID) + var buf bytes.Buffer + if err := json.NewEncoder(&buf).Encode(request); err != nil { + return err + } + r.Body = ioutil.NopCloser(&buf) + return nil +} + +func decodeDeleteCommandRequest(r *http.Request) (interface{}, error) { + vars := mux.Vars(r) + udid, ok := vars["udid"] + if !ok { + return nil, errBadRouting + } + uuid, ok := vars["uuid"] + if !ok { + return nil, errBadRouting + } + var request DeleteCommandRequest + request.UDID = udid + request.UUID = uuid + return request, nil +} + +// DeleteCommandResponse is a response for a delete request +type DeleteCommandResponse struct { + Total int `json:"remaining_payloads"` + Err error `json:"error,omitempty"` +} + +// DecodeDeleteCommandResponse decodes the response from the provided HTTP response, +// simply by JSON decoding from the response body. It's designed to be used in +// transport/http.Client. +// first decode into map[string]interface{} and check for error in the response +func DecodeDeleteCommandResponse(resp *http.Response) (interface{}, error) { + var r map[string]interface{} + var response DeleteCommandResponse + err := json.NewDecoder(resp.Body).Decode(&r) + if rs, ok := r["error"]; ok { + response.Err = errors.New(rs.(string)) + } + if rs, ok := r["remaining_payloads"]; ok { + response.Total = int(rs.(float64)) + } + return response, err +} + +func makeNewCommandEndpoint(svc MDMCommandService) 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(req.CommandRequest) + if err != nil { + return NewCommandResponse{Err: err}, nil + } + return NewCommandResponse{Payload: payload}, nil + } +} + +func makeNextCommandEndpoint(svc MDMCommandService) endpoint.Endpoint { + return func(ctx context.Context, request interface{}) (interface{}, error) { + req := request.(NextCommandRequest) + if req.UDID == "" { + return NextCommandResponse{Err: ErrEmptyRequest}, nil + } + payload, total, err := svc.NextCommand(req.UDID) + if err != nil { + return NextCommandResponse{Err: err}, nil + } + return NextCommandResponse{Payload: payload, Total: total}, nil + } +} + +func makeDeleteCommandEndpoint(svc MDMCommandService) endpoint.Endpoint { + return func(ctx context.Context, request interface{}) (interface{}, error) { + req := request.(DeleteCommandRequest) + if req.UDID == "" { + return DeleteCommandResponse{Err: ErrEmptyRequest}, nil + } + if req.UUID == "" { + return DeleteCommandResponse{Err: ErrEmptyRequest}, nil + } + total, err := svc.DeleteCommand(req.UDID, req.UUID) + if err != nil { + return DeleteCommandResponse{Err: err}, nil + } + return DeleteCommandResponse{Total: total}, nil + } +} + +// encodeResponse is the common method to encode all response types to the +// client. I chose to do it this way because I didn't know if something more +// specific was necessary. It's certainly possible to specialize on a +// per-response (per-method) basis. +func encodeResponse(w http.ResponseWriter, response interface{}) error { + if e, ok := response.(errorer); ok && e.error() != nil { + // Not a Go kit transport error, but a business-logic error. + // Provide those as HTTP errors. + encodeError(w, e.error()) + return nil + } + jsn, err := json.MarshalIndent(response, "", " ") + if err != nil { + return err + } + w.Write(jsn) + return nil +} + +func encodeError(w http.ResponseWriter, err error) { + w.WriteHeader(codeFrom(err)) + response := map[string]interface{}{ + "error": err.Error(), + } + jsn, err := json.MarshalIndent(response, "", " ") + if err != nil { + log.Println(err) + return + } + w.Write(jsn) +} + +func codeFrom(err error) int { + switch err { + default: + return http.StatusInternalServerError + } +} diff --git a/connect/encode_decode.go b/connect/encode_decode.go new file mode 100644 index 00000000..0d459845 --- /dev/null +++ b/connect/encode_decode.go @@ -0,0 +1,76 @@ +package connect + +import ( + "bytes" + "fmt" + "io/ioutil" + "log" + "net/http" + + "github.com/groob/plist" +) + +func decodeMDMConnectRequest(r *http.Request) (interface{}, error) { + body, _ := ioutil.ReadAll(r.Body) + fmt.Println(string(body)) + reader := bytes.NewReader(body) + + var request mdmConnectRequest + if err := plist.NewDecoder(reader).Decode(&request); err != nil { + return nil, err + } + return request, nil +} + +// errorer is implemented by all concrete response types. It allows us to +// change the HTTP response code without needing to trigger an endpoint +// (transport-level) error. For more information, read the big comment in +// endpoint.go. +type errorer interface { + error() error +} + +// encodeResponse is the common method to encode all response types to the +// client. I chose to do it this way because I didn't know if something more +// specific was necessary. It's certainly possible to specialize on a +// per-response (per-method) basis. +func encodeResponse(w http.ResponseWriter, response interface{}) error { + if e, ok := response.(errorer); ok && e.error() != nil { + // Not a Go kit transport error, but a business-logic error. + // Provide those as HTTP errors. + encodeError(w, e.error()) + return nil + } + resp := response.(mdmConnectResponse) + next := resp.payload + // var decoded = make([]byte, base64.StdEncoding.DecodedLen(len(next))) + // _, err := base64.StdEncoding.Decode(decoded, next) + // if err != nil { + // encodeError(w, err) + // return nil + // } + if len(next) != 0 { + w.Write(next) + } + return nil +} + +func encodeError(w http.ResponseWriter, err error) { + w.WriteHeader(codeFrom(err)) + response := map[string]interface{}{ + "error": err.Error(), + } + enc := plist.NewEncoder(w) + enc.Indent(" ") + err = enc.Encode(response) + if err != nil { + log.Println(err) + } +} + +func codeFrom(err error) int { + switch err { + default: + return http.StatusInternalServerError + } +} diff --git a/connect/endpoint.go b/connect/endpoint.go new file mode 100644 index 00000000..aae8820a --- /dev/null +++ b/connect/endpoint.go @@ -0,0 +1,47 @@ +package connect + +import ( + "errors" + + "github.com/go-kit/kit/endpoint" + "golang.org/x/net/context" +) + +// ErrInvalidMessageType is an invalid checking command +var ErrInvalidMessageType = errors.New("Invalid MessageType") + +func makeConnectEndpoint(svc MDMConnectService) endpoint.Endpoint { + return func(ctx context.Context, request interface{}) (interface{}, error) { + req := request.(mdmConnectRequest) + var err error + switch req.Status { + case "Acknowledged": + total, err := svc.Acknowledge(req.UDID, req.CommandUUID) + if err != nil { + return mdmConnectResponse{Err: err}, nil + } + if total != 0 { + next, _, err := svc.NextCommand(req.UDID) + if err != nil { + return mdmConnectResponse{Err: err}, nil + } + return mdmConnectResponse{payload: next}, nil + } + case "Idle": + next, total, err := svc.NextCommand(req.UDID) + if err != nil { + return mdmConnectResponse{Err: err}, nil + } + if total == 0 { + return mdmConnectResponse{}, nil + } + return mdmConnectResponse{payload: next}, nil + default: + return mdmConnectResponse{Err: ErrInvalidMessageType}, nil + } + if err != nil { + return mdmConnectResponse{Err: err}, nil + } + return mdmConnectResponse{}, nil + } +} diff --git a/connect/request_response.go b/connect/request_response.go new file mode 100644 index 00000000..35c6b32d --- /dev/null +++ b/connect/request_response.go @@ -0,0 +1,14 @@ +package connect + +import "github.com/micromdm/mdm" + +type mdmConnectRequest struct { + mdm.Response +} + +type mdmConnectResponse struct { + payload []byte + Err error `plist:"error,omitempty"` +} + +func (r mdmConnectResponse) error() error { return r.Err } diff --git a/connect/service.go b/connect/service.go new file mode 100644 index 00000000..7566c484 --- /dev/null +++ b/connect/service.go @@ -0,0 +1,78 @@ +package connect + +import ( + "net/http" + "os" + + "golang.org/x/net/context" + + "github.com/go-kit/kit/log" + httptransport "github.com/go-kit/kit/transport/http" + "github.com/micromdm/micromdm/command" +) + +// MDMConnectService ... +type MDMConnectService interface { + Acknowledge(deviceUDID, commandUUID string) (int, error) + NextCommand(deviceUDID string) ([]byte, int, error) +} + +type mdmConnectService struct { + redis command.Datastore +} + +func (svc mdmConnectService) Acknowledge(deviceUDID, commandUUID string) (int, error) { + return svc.redis.DeleteCommand(deviceUDID, commandUUID) + +} + +func (svc mdmConnectService) NextCommand(deviceUDID string) ([]byte, int, error) { + return svc.redis.NextCommand(deviceUDID) +} + +type config struct { + logger log.Logger + redis command.Datastore +} + +// NewConnectService creates a new MDM Connect Service +func NewConnectService(options ...func(*config) error) MDMConnectService { + conf := &config{} + defaultLogger := log.NewLogfmtLogger(os.Stderr) + for _, option := range options { + if err := option(conf); err != nil { + defaultLogger.Log("err", err) + os.Exit(1) + } + } + var svc MDMConnectService + svc = mdmConnectService{conf.redis} + if conf.logger != nil { + // svc = loggingMiddleware{conf.logger, svc} + } + + return svc +} + +// Redis adds a db connection to the service +func Redis(db command.Datastore) func(*config) error { + return func(c *config) error { + c.redis = db + return nil + } +} + +// ServiceHandler creates an http handler +func ServiceHandler(ctx context.Context, svc MDMConnectService) http.Handler { + // endpoint + connect := makeConnectEndpoint(svc) + + // handler + connectHandler := httptransport.NewServer( + ctx, + connect, + decodeMDMConnectRequest, + encodeResponse, + ) + return connectHandler +} diff --git a/device/device.go b/device/device.go new file mode 100644 index 00000000..ae300b6e --- /dev/null +++ b/device/device.go @@ -0,0 +1,196 @@ +package device + +import ( + "errors" + "fmt" + "os" + "time" + + "golang.org/x/net/context" + + "github.com/go-kit/kit/log" + "github.com/jmoiron/sqlx" + _ "github.com/lib/pq" // postgres driver +) + +// ErrNoRowsModified is returned if insert didn't produce results +var ErrNoRowsModified = errors.New("DB: No rows affected") + +// Device represents an iOS or OS X Computer +type Device struct { + // Primary key is UUID + UUID string `json:"uuid"` + UDID string `json:"udid"` + SerialNumber *string `json:"serial_number,omitempty" db:"serial_number,omitempty"` + OSVersion *string `json:"os_version,omitempty" db:"os_version,omitempty"` + BuildVersion *string `json:"build_version,omitempty" db:"build_version,omitempty"` + ProductName *string `json:"product_name,omitempty" db:"product_name,omitempty"` + IMEI *string `json:"imei,omitempty" db:"imei,omitempty"` + MEID *string `json:"meid,omitempty" db:"meid,omitempty"` + //Apple MDM Protocol Topic + MDMTopic *string `json:"mdm_topic,omitempty" db:"apple_mdm_topic,omitempty"` + PushMagic *string `json:"push_magic,omitempty" db:"apple_push_magic,omitempty"` + AwaitingConfiguration *bool `json:"awaiting_configuration,omitempty" db:"awaiting_configuration,omitempty"` + 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"` +} + +// Datastore manages interactions of devices in a database +type Datastore interface { + AddDevice(*Device) error + GetDeviceByUDID(udid string) (*Device, error) + SaveDevice(*Device) error + // RemoveDevice() error + // AllDevices() DeviceList,error +} + +type config struct { + context context.Context + logger log.Logger +} + +// NewDB creates a new databases connection +func NewDB(driver, conn string, options ...func(*config) error) Datastore { + conf := &config{} + defaultLogger := log.NewLogfmtLogger(os.Stderr) + for _, option := range options { + if err := option(conf); err != nil { + defaultLogger.Log("err", err) + os.Exit(1) + } + } + if conf.logger == nil { + conf.logger = defaultLogger + } + switch driver { + case "postgres": + db, err := sqlx.Open(driver, conn) + if err != nil { + conf.logger.Log("err", err) + os.Exit(1) + } + var dbError error + maxAttempts := 20 + for attempts := 1; attempts <= maxAttempts; attempts++ { + dbError = db.Ping() + if dbError == nil { + break + } + conf.logger.Log("msg", fmt.Sprintf("could not connect to postgres: %v", dbError)) + time.Sleep(time.Duration(attempts) * time.Second) + } + if dbError != nil { + conf.logger.Log("err", dbError) + os.Exit(1) + } + migrate(db) + // TODO: configurable with default + db.SetMaxOpenConns(5) + store := pgDatastore{db} + return store + default: + conf.logger.Log("err", "unknown driver") + os.Exit(1) + return nil + } +} + +// Logger adds a logger to the database config +func Logger(logger log.Logger) func(*config) error { + return func(c *config) error { + c.logger = logger + return nil + } +} + +// datastore implementation for postgres +type pgDatastore struct { + *sqlx.DB +} + +func (db pgDatastore) AddDevice(dev *Device) error { + upsert := `INSERT INTO devices + (udid, apple_mdm_topic, os_version, build_version, product_name, serial_number, imei, meid) + VALUES ($1,$2,$3,$4,$5,$6,$7,$8) + ON CONFLICT ON CONSTRAINT devices_udid_key + DO UPDATE SET + apple_mdm_topic=$2, + os_version=$3, + build_version=$4, + product_name=$5, + serial_number=$6, + imei=$7, + meid=$8;` + result, err := db.Exec( + upsert, + dev.UDID, + dev.MDMTopic, + dev.OSVersion, + dev.BuildVersion, + dev.ProductName, + dev.SerialNumber, + dev.IMEI, + dev.MEID, + ) + if err != nil { + return err + } + if res, _ := result.RowsAffected(); res == 0 { + return ErrNoRowsModified + } + return nil +} + +func (db pgDatastore) GetDeviceByUDID(udid string) (*Device, error) { + var device Device + query := `SELECT * FROM devices WHERE udid=$1 LIMIT 1` + return &device, sqlx.Get(db, &device, query, udid) +} + +// SaveDevice updates a device with the latest changes +func (db pgDatastore) SaveDevice(dev *Device) error { + update := `UPDATE devices SET + awaiting_configuration=$2, + apple_push_magic=$3, + apple_mdm_token=$4, + mdm_enrolled=$5 + WHERE uuid=$1` + result, err := db.Exec( + update, + dev.UUID, + dev.AwaitingConfiguration, + dev.PushMagic, + dev.Token, + dev.Enrolled, + ) + if err != nil { + return err + } + if res, _ := result.RowsAffected(); res == 0 { + return ErrNoRowsModified + } + return nil +} + +func migrate(db *sqlx.DB) { + schema := ` + CREATE EXTENSION IF NOT EXISTS "uuid-ossp"; + CREATE TABLE IF NOT EXISTS devices ( + uuid uuid PRIMARY KEY + DEFAULT uuid_generate_v4(), + udid text UNIQUE NOT NULL, + serial_number text, + os_version text, + build_version text, + product_name text, + imei text, + meid text, + apple_mdm_token text, + apple_mdm_topic text, + apple_push_magic text, + mdm_enrolled boolean, + awaiting_configuration boolean + );` + db.MustExec(schema) +} diff --git a/glide.lock b/glide.lock new file mode 100644 index 00000000..4e6adceb --- /dev/null +++ b/glide.lock @@ -0,0 +1,79 @@ +hash: 8fe304e9385e4fbe963eba356a9784b3217ae75b65f50178322a7e1e41950e8d +updated: 2016-03-17T13:39:31.722470317-04:00 +imports: +- name: github.com/beorn7/perks + version: 3ac7bf7a47d159a033b107610db8a1b6575507a4 + subpackages: + - quantile +- name: github.com/garyburd/redigo + version: 4ed1111375cbeb698249ffe48dd463e9b0a63a7a + subpackages: + - redis + - internal +- name: github.com/go-kit/kit + version: 9b8ce4ffb319107a43400cf189dca14ef09255e2 + subpackages: + - endpoint + - log + - metrics + - metrics/prometheus + - transport/http +- name: github.com/golang/protobuf + version: 62e4364d64b32762febb61f2c88c0a29bc49a225 + subpackages: + - proto +- name: github.com/gorilla/context + version: 1ea25387ff6f684839d82767c1733ff4d4d15d0a +- name: github.com/gorilla/mux + version: acf3be1b335c8ce30b2c8d51300984666f0ceefa +- name: github.com/groob/plist + version: 960422c558dc729e57204822eefad0636185ba2b +- name: github.com/jmoiron/sqlx + version: 398dd5876282499cdfd4cb8ea0f31a672abe9495 + subpackages: + - reflectx +- name: github.com/lib/pq + version: 165a3529e799da61ab10faed1fabff3662d6193f + subpackages: + - oid +- name: github.com/matttproud/golang_protobuf_extensions + version: d0c3fe89de86839aecf2e0579c40ba3bb336a453 + subpackages: + - pbutil +- name: github.com/micromdm/mdm + version: 0b08ebf4a9677b95c67c313c7f1189e0455f91d9 +- name: github.com/micromdm/micromdm + version: 94e85e8d676b55b3112e3c10438aef82c6278490 + subpackages: + - checkin + - command + - connect + - device +- name: github.com/prometheus/client_golang + version: 90c15b5efa0dc32a7d259234e02ac9a99e6d3b82 + subpackages: + - prometheus +- name: github.com/prometheus/client_model + version: fa8ad6fec33561be4280a8f0514318c79d7f6cb6 + subpackages: + - go +- name: github.com/prometheus/common + version: e8eabff8812b05acf522b45fdcd725a785188e37 + subpackages: + - expfmt + - internal/bitbucket.org/ww/goautoneg + - model +- name: github.com/prometheus/procfs + version: 406e5b7bfd8201a36e2bb5f7bdae0b03380c2ce8 +- name: github.com/satori/go.uuid + version: e673fdd4dea8a7334adbbe7f57b7e4b00bdc5502 +- name: golang.org/x/net + version: 35b06af0720201bc2f326773a80767387544f8c4 + subpackages: + - context + - context/ctxhttp +- name: gopkg.in/logfmt.v0 + version: ffc984a0eaff44e46e1f42d6e0a65b665587a938 +- name: gopkg.in/stack.v1 + version: 0585967eab0016c8e4e2d55ac20585b469574cec +devImports: [] diff --git a/glide.yaml b/glide.yaml new file mode 100644 index 00000000..a7cec3ce --- /dev/null +++ b/glide.yaml @@ -0,0 +1,28 @@ +package: github.com/micromdm/micromdm +import: +- package: github.com/garyburd/redigo + subpackages: + - redis +- package: github.com/go-kit/kit + subpackages: + - endpoint + - log + - metrics + - metrics/prometheus + - transport/http +- package: github.com/gorilla/mux +- package: github.com/groob/plist +- package: github.com/jmoiron/sqlx +- package: github.com/lib/pq +- package: github.com/micromdm/mdm +- package: github.com/micromdm/micromdm + subpackages: + - checkin + - command + - device +- package: github.com/prometheus/client_golang + subpackages: + - prometheus +- package: golang.org/x/net + subpackages: + - context diff --git a/main.go b/main.go new file mode 100644 index 00000000..1288fc4e --- /dev/null +++ b/main.go @@ -0,0 +1,202 @@ +package main + +import ( + "errors" + "flag" + "fmt" + "net/http" + "os" + + "github.com/go-kit/kit/log" + "github.com/gorilla/mux" + "github.com/micromdm/micromdm/checkin" + "github.com/micromdm/micromdm/command" + "github.com/micromdm/micromdm/connect" + "github.com/micromdm/micromdm/device" + stdprometheus "github.com/prometheus/client_golang/prometheus" + "golang.org/x/net/context" +) + +var ( + // Version info + Version = "unreleased" + gitHash = "unknown" +) + +func main() { + ctx := context.Background() + logger := log.NewLogfmtLogger(os.Stderr) + + //flags + var ( + flPort = flag.String("port", envString("MICROMDM_HTTP_LISTEN_PORT", ""), "port to listen on") + flTLS = flag.Bool("tls", envBool("MICROMDM_USE_TLS"), "use https") + flTLSCert = flag.String("tls-cert", envString("MICROMDM_TLS_CERT", ""), "path to TLS certificate") + flTLSKey = flag.String("tls-key", envString("MICROMDM_TLS_KEY", ""), "path to TLS private key") + flPGconn = flag.String("postgres", envString("MICROMDM_POSTGRES_CONN_URL", ""), "postgres connection url") + flRedisconn = flag.String("redis", envString("MICROMDM_REDIS_CONN_URL", ""), "redis connection url") + flVersion = flag.Bool("version", false, "print version information") + ) + + // set tls to true by default. let user set it to false + *flTLS = true + flag.Parse() + + // -version flag + if *flVersion { + fmt.Printf("micromdm - Version %s\n", Version) + fmt.Printf("Git Hash - %s\n", gitHash) + os.Exit(0) + } + + // check port flag + // if none is provided, default to 80 or 443 + if *flPort == "" { + port := defaultPort(*flTLS) + logger.Log("msg", fmt.Sprintf("No port flag specified. Using %v by default", port)) + *flPort = port + } + + // check cert and key if -tls=true + if *flTLS { + if err := checkTLSFlags(*flTLSKey, *flTLSCert); err != nil { + logger.Log("err", err) + os.Exit(1) + } + } + + pgHostAddr := os.Getenv("POSTGRES_PORT_5432_TCP_ADDR") + if *flPGconn == "" && pgHostAddr != "" { + *flPGconn = getPGConnFromENV(logger, pgHostAddr) + } + + // check database connection + if *flPGconn == "" { + logger.Log("err", "database connection url not specified") + os.Exit(1) + } + + deviceDB := device.NewDB( + "postgres", + *flPGconn, + device.Logger(logger), + ) + + // Checkin Service + checkinSvc := checkin.NewCheckinService( + checkin.Datastore(deviceDB), + checkin.Logger(logger), + ) + checkinHandler := checkin.ServiceHandler(ctx, checkinSvc) + + redisHostAddr := os.Getenv("REDIS_PORT_6379_TCP_ADDR") + if *flRedisconn == "" && redisHostAddr != "" { + *flRedisconn = getRedisConnFromENV(redisHostAddr) + } + + // check database connection + if *flRedisconn == "" { + logger.Log("err", "database connection url not specified") + os.Exit(1) + } + + commandDB := command.NewDB( + "redis", + *flRedisconn, + command.Logger(logger), + ) + + commandSvc := command.NewCommandService( + command.DB(commandDB), + command.Logger(logger), + ) + commandHandler := command.ServiceHandler(ctx, commandSvc) + + connectSvc := connect.NewConnectService( + connect.Redis(commandDB), + ) + connectHandler := connect.ServiceHandler(ctx, connectSvc) + + // router + r := mux.NewRouter() + r.Methods("PUT").Path("/mdm/checkin").Handler(checkinHandler) + r.Methods("PUT").Path("/mdm/connect").Handler(connectHandler) + r.Handle("/mdm/commands", commandHandler) + r.Methods("POST").Path("/mdm/commands").Handler(commandHandler) + r.Methods("GET").Path("/mdm/commands/{udid}/next").Handler(commandHandler) + r.Methods("DELETE").Path("/mdm/commands/{udid}/{uuid}").Handler(commandHandler) + + http.Handle("/", r) + http.Handle("/metrics", stdprometheus.Handler()) + + serve(logger, *flTLS, *flPort, *flTLSKey, *flTLSCert) +} + +// choose http or https +func serve(logger log.Logger, tls bool, port, key, cert string) { + portStr := fmt.Sprintf(":%v", port) + if tls { + logger.Log("msg", "HTTPs", "addr", port) + logger.Log("err", http.ListenAndServeTLS(portStr, cert, key, nil)) + } else { + logger.Log("msg", "HTTP", "addr", port) + logger.Log("err", http.ListenAndServe(portStr, nil)) + } +} + +func envString(key, def string) string { + if env := os.Getenv(key); env != "" { + return env + } + return def +} + +func envBool(key string) bool { + if env := os.Getenv(key); env == "true" { + return true + } + return false +} + +func checkTLSFlags(key, cert string) error { + if key == "" || cert == "" { + return errors.New("You must provide a valid path to a TLS cert and key") + } + return nil +} + +func defaultPort(tls bool) string { + if tls { + return "443" + } + return "80" +} + +// use this in docker container +func getPGConnFromENV(logger log.Logger, host string) string { + user := os.Getenv("POSTGRES_ENV_POSTGRES_USER") + if user == "" { + user = "postgres" + } + dbname := os.Getenv("POSTGRES_ENV_POSTGRES_DB") + if dbname == "" { + dbname = user //same defaults as the docker pgcontainer + } + password := os.Getenv("POSTGRES_ENV_POSTGRES_PASSWORD") + if password == "" { + password = "postgres" + } + sslmode := os.Getenv("POSTGRES_ENV_SSLMODE") + if sslmode == "" { + logger.Log("msg", "POSTGRES_ENV_SSLMODE not specified, using 'require' by default") + sslmode = "require" + } + conn := fmt.Sprintf("user=%v password=%v dbname=%v sslmode=%v host=%v", user, password, dbname, sslmode, host) + return conn +} + +func getRedisConnFromENV(host string) string { + port := os.Getenv("REDIS_PORT_6379_TCP_PORT") + conn := fmt.Sprintf("%v:%v", host, port) + return conn +}