diff --git a/cmd/micromdm/database.go b/cmd/micromdm/database.go new file mode 100644 index 00000000..4fbb8c3d --- /dev/null +++ b/cmd/micromdm/database.go @@ -0,0 +1,124 @@ +package main + +import ( + "context" + "fmt" + "io/ioutil" + "os" + + "crawshaw.io/sqlite" + "crawshaw.io/sqlite/sqlitex" + "github.com/jackc/pgx/v4/pgxpool" + + "micromdm.io/v2/internal/data/user" + "micromdm.io/v2/pkg/log" +) + +// Wrap all the database types here and create access methods. +// The caller can assert a minimal interface for each data package. +// +// Example: +// account.Config{UserStore: db.userdb().(account.UserStore)} +// +// Code generation (or eventual generics) can make this cleaner. +type database struct { + sq *sqlitedb + pg *postgresdb +} + +func (db *database) userdb() interface{} { + if db.sq != nil { + return db.sq.userdb + } + return db.pg.userdb +} + +type sqlitedb struct { + userdb *user.SQLite +} + +func setupSQLite(ctx context.Context, f *cliFlags, logger log.Logger) (*sqlitedb, error) { + conn, err := sqlite.OpenConn(f.databaseURL, 0) + if err != nil { + return nil, fmt.Errorf("open sqlite dbfile: %s", err) + } + + if err := sqliteInit(conn); err != nil { + conn.Close() + return nil, fmt.Errorf("init sqlite: %s", err) + } + + if err := conn.Close(); err != nil { + return nil, fmt.Errorf("sqlite init close: %s", err) + } + + pool, err := sqlitex.Open(f.databaseURL, 0, 24) + if err != nil { + return nil, fmt.Errorf("create sqlite pool: %s", err) + } + + db := &sqlitedb{ + userdb: user.NewSQLite(pool), + } + + log.Debug(logger).Log("msg", "connected to db", "backend", "sqlite") + return db, nil +} + +func sqliteInit(conn *sqlite.Conn) error { + if err := sqlitex.ExecTransient(conn, "PRAGMA journal_mode=WAL;", nil); err != nil { + return err + } + + if err := sqlitex.ExecTransient(conn, "PRAGMA cache_size = -50000;", nil); err != nil { + return err + } + + var ( + migrations []byte + script = "internal/data/migrations/sqlite/initial_tables.sql" + ) + + // only load and run migrations script if the path exists. + // temporary workaround for tests which will have to be + // replaced once the schema needs to exist for the tests. + if _, err := os.Stat(script); err == nil { + migrations, err = ioutil.ReadFile(script) + if err != nil { + return err + } + } + + if err := sqlitex.ExecScript(conn, string(migrations)); err != nil { + return err + } + + return nil +} + +type postgresdb struct { + userdb *user.Postgres +} + +func setupPostgres(ctx context.Context, f *cliFlags, logger log.Logger) (*postgresdb, error) { + //dbpool, err := pgxpool.Connect(ctx, "host=localhost port=5432 user=app dbname=app password=secret sslmode=disable") + dbpool, err := pgxpool.Connect(ctx, f.databaseURL) + if err != nil { + return nil, err + } + + // temporary until migrations are set up + migrations, err := ioutil.ReadFile("internal/data/migrations/postgres/initial_tables.sql") + if _, err := dbpool.Exec(ctx, string(migrations)); err != nil { + dbpool.Close() + return nil, err + } + + db := &postgresdb{ + userdb: user.NewPostgres(dbpool), + } + + log.Debug(logger).Log("msg", "connected to db", "backend", "postgres") + + return db, nil +} diff --git a/cmd/micromdm/micromdm.go b/cmd/micromdm/micromdm.go index 957645f1..a63cc124 100644 --- a/cmd/micromdm/micromdm.go +++ b/cmd/micromdm/micromdm.go @@ -39,6 +39,8 @@ type cliFlags struct { csrfKey string csrfCookieName string csrfFieldName string + + databaseURL string } func micromdm(args []string, stdin io.Reader, stdout, stderr io.Writer) int { @@ -57,6 +59,7 @@ func micromdm(args []string, stdin io.Reader, stdout, stderr io.Writer) int { rootfs.StringVar(&cli.csrfKey, "csrf_key", "", "32 byte long key") rootfs.StringVar(&cli.csrfCookieName, "csrf_cookie_name", "micromdm_csrf", "Name of CSRF Cookie") rootfs.StringVar(&cli.csrfFieldName, "csrf_field_name", "micromdm.csrf", "Name of CSRF field name in HTML input") + rootfs.StringVar(&cli.databaseURL, "database_url", "build/_sqlite.db", "Database URL") // default output is os.Stderr. // setting the output and flag.ContinueOnError overrides allows testing usage. @@ -101,7 +104,7 @@ func micromdm(args []string, stdin io.Reader, stdout, stderr io.Writer) int { return err } - srv, err := setup(cli, logger) + srv, err := setup(ctx, cli, logger) if err != nil { return err } diff --git a/cmd/micromdm/micromdm_test.go b/cmd/micromdm/micromdm_test.go index 08903c5b..c4ed84fc 100644 --- a/cmd/micromdm/micromdm_test.go +++ b/cmd/micromdm/micromdm_test.go @@ -152,7 +152,8 @@ func TestMain(t *testing.T) { stdin: new(bytes.Buffer), stdout: new(bytes.Buffer), stderr: new(bytes.Buffer), - args: []string{"micromdm", "-pidfile", filepath.Join(tmpdir, "sigint.pid")}, + args: []string{"micromdm", "-pidfile", filepath.Join(tmpdir, "sigint.pid"), + "-database_url", filepath.Join(tmpdir, "micromdm.db")}, exit: 1, check: checkExitAfterSignal, signal: exitWith(syscall.SIGTERM), @@ -162,7 +163,8 @@ func TestMain(t *testing.T) { stdin: new(bytes.Buffer), stdout: new(bytes.Buffer), stderr: new(bytes.Buffer), - args: []string{"micromdm", "-pidfile", filepath.Join(tmpdir, "sigusr.pid")}, + args: []string{"micromdm", "-pidfile", filepath.Join(tmpdir, "sigusr.pid"), + "-database_url", filepath.Join(tmpdir, "micromdm.db")}, exit: 1, check: checkLogSwap, signal: logswap(), @@ -172,11 +174,12 @@ func TestMain(t *testing.T) { synchronous: true, }, { - name: "restart on hup", - stdin: new(bytes.Buffer), - stdout: new(bytes.Buffer), - stderr: new(bytes.Buffer), - args: []string{"micromdm", "-pidfile", filepath.Join(tmpdir, "sighup.pid")}, + name: "restart on hup", + stdin: new(bytes.Buffer), + stdout: new(bytes.Buffer), + stderr: new(bytes.Buffer), + args: []string{"micromdm", "-pidfile", filepath.Join(tmpdir, "sighup.pid"), + "-database_url", filepath.Join(tmpdir, "micromdm.db")}, exit: 1, check: checkRestarted, signal: sighup(), diff --git a/cmd/micromdm/server.go b/cmd/micromdm/server.go index 0862debc..896ac800 100644 --- a/cmd/micromdm/server.go +++ b/cmd/micromdm/server.go @@ -1,6 +1,13 @@ package main import ( + "context" + "fmt" + "io/ioutil" + "os" + + "github.com/jackc/pgx/v4" + "micromdm.io/v2/internal/frontend/account" "micromdm.io/v2/pkg/frontend" "micromdm.io/v2/pkg/log" ) @@ -19,12 +26,49 @@ func ui(f *cliFlags, logger log.Logger) (*frontend.Server, error) { }) } -func setup(f *cliFlags, logger log.Logger) (*server, error) { +func setup(ctx context.Context, f *cliFlags, logger log.Logger) (*server, error) { uisrv, err := ui(f, logger) if err != nil { return nil, err } + var db database + switch dbDriver(f.databaseURL) { + case "postgres": + db.pg, err = setupPostgres(ctx, f, logger) + case "sqlite": + db.sq, err = setupSQLite(ctx, f, logger) + default: + err = fmt.Errorf("unsupported database_url value or restricted path %q", f.databaseURL) + } + + if err != nil { + return nil, err + } + srv := &server{ui: uisrv} + + account.HTTP(account.Config{ + HTTP: srv.ui, + UserStore: db.userdb().(account.UserStore), + }) + return srv, nil } + +func dbDriver(dbURL string) string { + if _, err := pgx.ParseConfig(dbURL); err == nil { + return "postgres" + } + + if _, err := os.Stat(dbURL); err == nil { + return "sqlite" + } + + if err := ioutil.WriteFile(dbURL, []byte(""), 0644); err == nil { + os.Remove(dbURL) + return "sqlite" + } + + return "" +} diff --git a/docker-compose.yaml b/docker-compose.yaml new file mode 100644 index 00000000..edb0d920 --- /dev/null +++ b/docker-compose.yaml @@ -0,0 +1,10 @@ +version: "2" +services: + postgresql: + image: postgres:latest + environment: + POSTGRES_DB: app + POSTGRES_USER: app + POSTGRES_PASSWORD: secret + ports: + - 5432:5432 diff --git a/go.mod b/go.mod index 40595491..9c9f84e4 100644 --- a/go.mod +++ b/go.mod @@ -8,6 +8,7 @@ require ( github.com/go-kit/kit v0.10.0 github.com/gorilla/csrf v1.7.0 github.com/gorilla/mux v1.7.4 + github.com/jackc/pgconn v1.6.2 github.com/jackc/pgx/v4 v4.7.2 github.com/oklog/run v1.1.0 github.com/oklog/ulid v1.3.1 // indirect @@ -16,3 +17,5 @@ require ( golang.org/x/crypto v0.0.0-20200709230013-948cd5f35899 rsc.io/goversion v1.2.0 ) + +replace crawshaw.io/sqlite => github.com/groob/sqlite v0.3.3-0.20200721040052-b46ed0907467 diff --git a/go.sum b/go.sum index 0cf24253..b1679aa0 100644 --- a/go.sum +++ b/go.sum @@ -2,8 +2,6 @@ cloud.google.com/go v0.26.0/go.mod h1:aQUYkXzVsufM+DwF1aE+0xfcU+56JwCaLick0ClmMT cloud.google.com/go v0.34.0/go.mod h1:aQUYkXzVsufM+DwF1aE+0xfcU+56JwCaLick0ClmMTw= crawshaw.io/iox v0.0.0-20181124134642-c51c3df30797 h1:yDf7ARQc637HoxDho7xjqdvO5ZA2Yb+xzv/fOnnvZzw= crawshaw.io/iox v0.0.0-20181124134642-c51c3df30797/go.mod h1:sXBiorCo8c46JlQV3oXPKINnZ8mcqnye1EkVkqsectk= -crawshaw.io/sqlite v0.3.2 h1:N6IzTjkiw9FItHAa0jp+ZKC6tuLzXqAYIv+ccIWos1I= -crawshaw.io/sqlite v0.3.2/go.mod h1:igAO5JulrQ1DbdZdtVq48mnZUBAPOeFzer7VhDWNtW4= github.com/BurntSushi/toml v0.3.1/go.mod h1:xHWCNGjB5oqiDr8zfno3MHue2Ht5sIBksp03qcyfWMU= github.com/Knetic/govaluate v3.0.1-0.20171022003610-9aa49832a739+incompatible/go.mod h1:r7JcOSlj0wfOMncg0iLm8Leh48TZaKVeNIfJntJ2wa0= github.com/Shopify/sarama v1.19.0/go.mod h1:FVkBWblsNy7DGZRfXLU0O9RCGt5g3g3yEuWXgklEdEo= @@ -108,6 +106,8 @@ github.com/gorilla/mux v1.7.4/go.mod h1:DVbg23sWSpFRCP0SfiEN6jmj59UnW/n46BH5rLB7 github.com/gorilla/securecookie v1.1.1 h1:miw7JPhV+b/lAHSXz4qd/nN9jRiAFV5FwjeKyCS8BvQ= github.com/gorilla/securecookie v1.1.1/go.mod h1:ra0sb63/xPlUeL+yeDciTfxMRAA+MP+HVt/4epWDjd4= github.com/gorilla/websocket v0.0.0-20170926233335-4201258b820c/go.mod h1:E7qHFY5m1UJ88s3WnNqhKjPHQ0heANvMoAMk2YaljkQ= +github.com/groob/sqlite v0.3.3-0.20200721040052-b46ed0907467 h1:lDQqdenDBPEaTUOdlQOvzZdAdevOZ7WclTROKmgdj2U= +github.com/groob/sqlite v0.3.3-0.20200721040052-b46ed0907467/go.mod h1:igAO5JulrQ1DbdZdtVq48mnZUBAPOeFzer7VhDWNtW4= github.com/grpc-ecosystem/go-grpc-middleware v1.0.1-0.20190118093823-f849b5445de4/go.mod h1:FiyG127CGDf3tlThmgyCl78X/SZQqEOJBCDaAfeWzPs= github.com/grpc-ecosystem/go-grpc-prometheus v1.2.0/go.mod h1:8NvIoxWQoOIhqOTXgfV/d3M/q6VIi02HzZEHgUlZvzk= github.com/grpc-ecosystem/grpc-gateway v1.9.5/go.mod h1:vNeuVxBJEsws4ogUvrchl83t/GYV9WGTSLVdBhOQFDY= diff --git a/internal/data/migrations/postgres/initial_tables.sql b/internal/data/migrations/postgres/initial_tables.sql new file mode 100644 index 00000000..2a001f1c --- /dev/null +++ b/internal/data/migrations/postgres/initial_tables.sql @@ -0,0 +1,16 @@ +DROP TABLE IF EXISTS users; + +CREATE TABLE IF NOT EXISTS users ( + id text PRIMARY KEY NOT NULL, + username text NOT NULL DEFAULT '', + email text NOT NULL DEFAULT '', + password bytea NOT NULL, + salt bytea NOT NULL, + confirmation_hash text, + created_at timestamptz DEFAULT (now() at time zone 'utc'), + updated_at timestamptz DEFAULT (now() at time zone 'utc'), + CONSTRAINT chk_username_not_empty CHECK (username != ''), + CONSTRAINT chk_email_not_empty CHECK (email != ''), + UNIQUE (email), + UNIQUE (username) +); diff --git a/internal/data/migrations/sqlite/initial_tables.sql b/internal/data/migrations/sqlite/initial_tables.sql new file mode 100644 index 00000000..7b02a579 --- /dev/null +++ b/internal/data/migrations/sqlite/initial_tables.sql @@ -0,0 +1,27 @@ +PRAGMA auto_vacuum = INCREMENTAL; + +-- DROP TABLE IF EXISTS users; + +CREATE TABLE IF NOT EXISTS users ( + id text PRIMARY KEY, + username text NOT NULL DEFAULT '', + email text NOT NULL DEFAULT '', + password TEXT NOT NULL, + salt text NOT NULL, + confirmation_hash text, + created_at DATETIME DEFAULT CURRENT_TIMESTAMP, + updated_at DATETIME DEFAULT CURRENT_TIMESTAMP, + CONSTRAINT chk_username_not_empty CHECK (username != ''), + CONSTRAINT chk_email_not_empty CHECK (email != ''), + UNIQUE (email), + UNIQUE (username) +); + +CREATE TRIGGER IF NOT EXISTS tg_users_updated_at + AFTER UPDATE ON users + FOR EACH ROW +BEGIN + UPDATE users SET updated_at = CURRENT_TIMESTAMP +WHERE + id = old.id; +END; diff --git a/internal/data/user/user.go b/internal/data/user/user.go index 1ae6c54a..6d6a8630 100644 --- a/internal/data/user/user.go +++ b/internal/data/user/user.go @@ -60,10 +60,6 @@ func (u *User) ValidatePassword(plaintext string) error { } func (u *User) setPassword(plaintext string) error { - if plaintext == "" { - return errors.New("password cannot be empty") - } - salt, err := random(32, base64.StdEncoding) if err != nil { return err @@ -93,6 +89,24 @@ func random(keySize int, enc *base64.Encoding) ([]byte, error) { // DB func create(username, email, password string) (*User, error) { + val := Error{invalid: make(map[string]string)} + + if username == "" { + val.invalid["username"] = constraints["chk_username_not_empty"]["username"] + } + + if email == "" { + val.invalid["email"] = constraints["chk_email_not_empty"]["email"] + } + + if password == "" { + val.invalid["password"] = constraints["chk_password_not_empty"]["password"] + } + + if len(val.invalid) > 0 { + return nil, val + } + u := &User{ ID: id.New(), Username: username, @@ -112,3 +126,38 @@ func create(username, email, password string) (*User, error) { return u, nil } + +type Error struct { + invalid map[string]string +} + +func (err Error) Invalid() map[string]string { return err.invalid } + +func (err Error) Error() string { + switch len(err.invalid) { + case 0: + return "user validation failed" + case 1: + var key, value string + for k, v := range err.invalid { + key = k + value = v + break + } + return fmt.Sprintf("user validation failed: %s - %s", key, value) + default: + var key, value string + for k, v := range err.invalid { + key = k + value = v + break + } + return fmt.Sprintf("user validation failed: %s - %s and %d other errors", key, value, len(err.invalid)-1) + } +} + +var constraints = map[string]map[string]string{ + "chk_email_not_empty": map[string]string{"email": "You must provide an email address."}, + "chk_username_not_empty": map[string]string{"username": "You must provide a username."}, + "chk_password_not_empty": map[string]string{"password": "You must provide a password."}, +} diff --git a/internal/data/user/user_postgres.go b/internal/data/user/user_postgres.go index 599ff24e..25f0b850 100644 --- a/internal/data/user/user_postgres.go +++ b/internal/data/user/user_postgres.go @@ -5,6 +5,7 @@ import ( "errors" "fmt" + "github.com/jackc/pgconn" "github.com/jackc/pgx/v4/pgxpool" ) @@ -38,7 +39,7 @@ func (d *Postgres) CreateUser(ctx context.Context, username, email, password str u.Salt, u.ConfirmationHash, ).Scan(&u.CreatedAt, &u.UpdatedAt); err != nil { - return nil, fmt.Errorf("store created user in postgres: %w", err) + return nil, fmt.Errorf("store created user in postgres: %w", checkPostgres(err)) } return u, nil @@ -58,3 +59,19 @@ func (d *Postgres) ConfirmUser(ctx context.Context, confirmation string) error { return nil } + +func checkPostgres(err error) error { + var dbErr *pgconn.PgError + if !errors.As(err, &dbErr) { + return err + } + + switch dbErr.Code { + case "23514": + if kv, ok := constraints[dbErr.ConstraintName]; ok { + return Error{invalid: kv} + } + } + + return err +} diff --git a/internal/data/user/user_sqlite.go b/internal/data/user/user_sqlite.go index 366b51ad..52479003 100644 --- a/internal/data/user/user_sqlite.go +++ b/internal/data/user/user_sqlite.go @@ -4,8 +4,10 @@ import ( "context" "errors" "fmt" + "strings" "time" + "crawshaw.io/sqlite" "crawshaw.io/sqlite/sqlitex" ) @@ -42,7 +44,7 @@ func (d *SQLite) CreateUser(ctx context.Context, username, email, password strin stmt.SetBytes("$salt", u.Salt) stmt.SetText("$confirmationHash", *u.ConfirmationHash) if _, err := stmt.Step(); err != nil { - return nil, err + return nil, fmt.Errorf("store created user in sqlite: %w", checkSqlite(err)) } stmt = conn.Prep(`SELECT created_at, updated_at FROM users WHERE id = $id`) @@ -87,3 +89,20 @@ func (d *SQLite) ConfirmUser(ctx context.Context, confirmation string) error { return nil } + +func checkSqlite(err error) error { + var sErr sqlite.Error + if !errors.As(err, &sErr) { + return err + } + + switch sErr.Code { + case sqlite.SQLITE_CONSTRAINT_CHECK: + c := strings.Split(sErr.Msg, ": ") + if kv, ok := constraints[c[len(c)-1]]; ok { + return Error{invalid: kv} + } + } + + return fmt.Errorf("user: %w", err) +} diff --git a/internal/frontend/account/account.go b/internal/frontend/account/account.go new file mode 100644 index 00000000..b410e647 --- /dev/null +++ b/internal/frontend/account/account.go @@ -0,0 +1,36 @@ +// Package account contains web pages for user registration and account management features. +package account + +import ( + "context" + "net/http" + + "micromdm.io/v2/internal/data/user" + "micromdm.io/v2/pkg/frontend" +) + +type UserStore interface { + CreateUser(ctx context.Context, username, email, password string) (*user.User, error) + ConfirmUser(ctx context.Context, token string) error +} + +type server struct { + http frontend.Framework + userdb UserStore +} + +type Config struct { + HTTP frontend.Framework + UserStore UserStore +} + +func HTTP(config Config) { + srv := &server{ + http: config.HTTP, + userdb: config.UserStore, + } + + srv.http.HandleFunc("/register", srv.registerForm, http.MethodGet, http.MethodPost) + srv.http.HandleFunc("/register/done", srv.registerComplete) + srv.http.HandleFunc("/registered/confirm/{token}", srv.registerConfirm) +} diff --git a/internal/frontend/account/register.go b/internal/frontend/account/register.go new file mode 100644 index 00000000..1d132568 --- /dev/null +++ b/internal/frontend/account/register.go @@ -0,0 +1,81 @@ +package account + +import ( + "errors" + "net/http" + + "github.com/gorilla/csrf" + "github.com/gorilla/mux" + + "micromdm.io/v2/pkg/frontend" + "micromdm.io/v2/pkg/log" +) + +func (srv server) registerForm(w http.ResponseWriter, r *http.Request) { + var ( + ctx = r.Context() + logger = log.FromContext(ctx) + username = r.FormValue("username") + email = r.FormValue("email") + password = r.FormValue("password") + data = frontend.Data{ + csrf.TemplateTag: csrf.TemplateField(r), + "form": map[string]string{ + "username": username, + "email": email, + "password": password, + }, + } + ) + + if r.Method == http.MethodGet { + srv.http.RenderTemplate(ctx, w, "register.tmpl", data) + return + } + + usr, err := srv.userdb.CreateUser(ctx, username, email, password) + if err != nil { + srv.http.Fail(ctx, w, err, "register.tmpl", "msg", "creating user") + return + } + + log.Debug(logger).Log( + "msg", "account created", + "username", usr.Username, + "id", usr.ID, + "confirmation_hash", *usr.ConfirmationHash, + ) + + http.Redirect(w, r, "/register/done", http.StatusFound) +} + +func (srv server) registerComplete(w http.ResponseWriter, r *http.Request) { + srv.http.RenderTemplate(r.Context(), w, "register-done.tmpl", frontend.Data{}) +} + +func (srv server) registerConfirm(w http.ResponseWriter, r *http.Request) { + var ( + ctx = r.Context() + logger = log.FromContext(ctx) + vars = mux.Vars(r) + ) + + confirmation, ok := vars["token"] + if !ok { + srv.http.RenderTemplate(ctx, w, "404.tmpl", frontend.Data{}. + WithCode(http.StatusNotFound). + WithLog(errors.New("missing confirmationHash in url")), + ) + return + } + + if err := srv.userdb.ConfirmUser(ctx, confirmation); err != nil { + srv.http.Fail(ctx, w, err, "msg", "confirm user", "confirmation_hash", confirmation) + return + } + + // TODO: create session and stuff + + log.Debug(logger).Log("msg", "user confirmed") + srv.http.RenderTemplate(ctx, w, "register-confirmed.tmpl", frontend.Data{}) +} diff --git a/ui/includes/register-confirmed.tmpl b/ui/includes/register-confirmed.tmpl new file mode 100644 index 00000000..7cd5db71 --- /dev/null +++ b/ui/includes/register-confirmed.tmpl @@ -0,0 +1,14 @@ +{{ define "title"}} +
+ We confirmed your account. You can now use the site. + Sign In +
++ We sent you an email to complete the registration process. You should see it in a bit. +
+