mirror of
https://github.com/micromdm/micromdm/
synced 2026-08-04 16:26:23 +08:00
Add account management web pages (#694)
Added internal/frontend/account to implement account management code. Added user registration templates and handlers. Set up a sqlite and postgres database in package main.
This commit is contained in:
124
cmd/micromdm/database.go
Normal file
124
cmd/micromdm/database.go
Normal file
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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 ""
|
||||
}
|
||||
|
||||
10
docker-compose.yaml
Normal file
10
docker-compose.yaml
Normal file
@@ -0,0 +1,10 @@
|
||||
version: "2"
|
||||
services:
|
||||
postgresql:
|
||||
image: postgres:latest
|
||||
environment:
|
||||
POSTGRES_DB: app
|
||||
POSTGRES_USER: app
|
||||
POSTGRES_PASSWORD: secret
|
||||
ports:
|
||||
- 5432:5432
|
||||
3
go.mod
3
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
|
||||
|
||||
4
go.sum
4
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=
|
||||
|
||||
16
internal/data/migrations/postgres/initial_tables.sql
Normal file
16
internal/data/migrations/postgres/initial_tables.sql
Normal file
@@ -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)
|
||||
);
|
||||
27
internal/data/migrations/sqlite/initial_tables.sql
Normal file
27
internal/data/migrations/sqlite/initial_tables.sql
Normal file
@@ -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;
|
||||
@@ -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."},
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
36
internal/frontend/account/account.go
Normal file
36
internal/frontend/account/account.go
Normal file
@@ -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)
|
||||
}
|
||||
81
internal/frontend/account/register.go
Normal file
81
internal/frontend/account/register.go
Normal file
@@ -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{})
|
||||
}
|
||||
14
ui/includes/register-confirmed.tmpl
Normal file
14
ui/includes/register-confirmed.tmpl
Normal file
@@ -0,0 +1,14 @@
|
||||
{{ define "title"}}
|
||||
<title>Confirmed {{ .siteName }} account</title>
|
||||
{{ end }}
|
||||
{{ define "content" }}
|
||||
<div class="page form-completion">
|
||||
<h3>Thank you for signing up</h3>
|
||||
<div class="msg">
|
||||
<p>
|
||||
We confirmed your account. You can now use the site.
|
||||
<a href="/login">Sign In</a>
|
||||
</p>
|
||||
</div>
|
||||
</div>
|
||||
{{end}}
|
||||
13
ui/includes/register-done.tmpl
Normal file
13
ui/includes/register-done.tmpl
Normal file
@@ -0,0 +1,13 @@
|
||||
{{ define "title" -}}
|
||||
<title>Verify your {{.siteName}} account</title>
|
||||
{{- end }}
|
||||
{{ define "content" -}}
|
||||
<div class="page form-completion">
|
||||
<h3>Thank you for signing up</h3>
|
||||
<div class="msg">
|
||||
<p>
|
||||
We sent you an email to complete the registration process. You should see it in a bit.
|
||||
</p>
|
||||
</div>
|
||||
</div>
|
||||
{{- end }}
|
||||
48
ui/includes/register.tmpl
Normal file
48
ui/includes/register.tmpl
Normal file
@@ -0,0 +1,48 @@
|
||||
{{ define "title" -}}
|
||||
<title>Register for {{ .siteName }}</title>
|
||||
{{- end }}
|
||||
{{ define "content" -}}
|
||||
<div class="page">
|
||||
<form class="single-page-form signup" method="post">
|
||||
<legend>Sign up</legend>
|
||||
{{ .csrfField }}
|
||||
<div class="form-group">
|
||||
<label for="username">Username</label>
|
||||
<input
|
||||
type="text"
|
||||
id="username"
|
||||
name="username"
|
||||
autocorrect="off"
|
||||
autocapitalize="none"
|
||||
{{- with .form.username}}value="{{.}}"{{- end}} />
|
||||
{{- with .errors -}}
|
||||
<div class="invalid-input">{{.username}}</div>
|
||||
{{- end }}
|
||||
</div>
|
||||
<div class="form-group">
|
||||
<label for="email">Email</label>
|
||||
<input
|
||||
type="email"
|
||||
id="email"
|
||||
name="email"
|
||||
{{- with .form.email }}value="{{.}}"{{- end }} />
|
||||
{{- with .errors -}}
|
||||
<div class="invalid-input">{{.email}}</div>
|
||||
{{- end }}
|
||||
</div>
|
||||
<div class="form-group">
|
||||
<label for="password">Password</label>
|
||||
<input
|
||||
type="password"
|
||||
id="password"
|
||||
name="password"
|
||||
{{- with .form.password }}value="{{.}}"{{- end }} />
|
||||
{{- with .errors -}}
|
||||
<div class="invalid-input">{{.password}}</div>
|
||||
{{- end }}
|
||||
</div>
|
||||
<div class="button-container"><button class="btn">Sign Up</button></div>
|
||||
<div class="form-footer">Already have an account?<a href="/login">Sign In</a></div>
|
||||
</form>
|
||||
</div>
|
||||
{{- end }}
|
||||
Reference in New Issue
Block a user