diff --git a/Makefile b/Makefile index d80382e7..5fa66c52 100644 --- a/Makefile +++ b/Makefile @@ -1,6 +1,7 @@ .PHONY: build -WORKSPACE := $(dir $(shell go env GOMOD)) +GO := go +WORKSPACE := $(dir $(shell ${GO} env GOMOD)) BUILD_DIR := $(WORKSPACE)/build export GOBIN=$(BUILD_DIR) @@ -22,12 +23,12 @@ BUILD_VERSION = "\ download: @echo "Download dependencies" - @go mod download + @${GO} mod download install-tools: download @echo "Installing tools" - @cat cmd/tools/tools.go | grep _ | awk -F'"' '{print $$2}' | xargs -tI % go install % + @cat cmd/tools/tools.go | grep _ | awk -F'"' '{print $$2}' | xargs -tI % ${GO} install % micromdm: $(eval APP_NAME = micromdm) - go install -race -ldflags ${BUILD_VERSION} ./cmd/micromdm + ${GO} install -race -ldflags ${BUILD_VERSION} ./cmd/micromdm diff --git a/cmd/micromdm/database.go b/cmd/micromdm/database.go index 4fbb8c3d..5e979cdb 100644 --- a/cmd/micromdm/database.go +++ b/cmd/micromdm/database.go @@ -10,6 +10,7 @@ import ( "crawshaw.io/sqlite/sqlitex" "github.com/jackc/pgx/v4/pgxpool" + "micromdm.io/v2/internal/data/session" "micromdm.io/v2/internal/data/user" "micromdm.io/v2/pkg/log" ) @@ -33,8 +34,16 @@ func (db *database) userdb() interface{} { return db.pg.userdb } +func (db *database) sessiondb() interface{} { + if db.sq != nil { + return db.sq.sessiondb + } + return db.pg.sessiondb +} + type sqlitedb struct { - userdb *user.SQLite + userdb *user.SQLite + sessiondb *session.SQLite } func setupSQLite(ctx context.Context, f *cliFlags, logger log.Logger) (*sqlitedb, error) { @@ -58,7 +67,8 @@ func setupSQLite(ctx context.Context, f *cliFlags, logger log.Logger) (*sqlitedb } db := &sqlitedb{ - userdb: user.NewSQLite(pool), + userdb: user.NewSQLite(pool), + sessiondb: session.NewSQLite(pool), } log.Debug(logger).Log("msg", "connected to db", "backend", "sqlite") @@ -97,7 +107,8 @@ func sqliteInit(conn *sqlite.Conn) error { } type postgresdb struct { - userdb *user.Postgres + userdb *user.Postgres + sessiondb *session.Postgres } func setupPostgres(ctx context.Context, f *cliFlags, logger log.Logger) (*postgresdb, error) { @@ -115,7 +126,8 @@ func setupPostgres(ctx context.Context, f *cliFlags, logger log.Logger) (*postgr } db := &postgresdb{ - userdb: user.NewPostgres(dbpool), + userdb: user.NewPostgres(dbpool), + sessiondb: session.NewPostgres(dbpool), } log.Debug(logger).Log("msg", "connected to db", "backend", "postgres") diff --git a/cmd/micromdm/micromdm.go b/cmd/micromdm/micromdm.go index a63cc124..1b597656 100644 --- a/cmd/micromdm/micromdm.go +++ b/cmd/micromdm/micromdm.go @@ -40,6 +40,8 @@ type cliFlags struct { csrfCookieName string csrfFieldName string + authCookieName string + databaseURL string } @@ -60,6 +62,7 @@ func micromdm(args []string, stdin io.Reader, stdout, stderr io.Writer) int { 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") + rootfs.StringVar(&cli.authCookieName, "auth_cookie_name", "micromdm_auth", "Name of authentication cookie") // default output is os.Stderr. // setting the output and flag.ContinueOnError overrides allows testing usage. diff --git a/cmd/micromdm/server.go b/cmd/micromdm/server.go index 896ac800..ddb85f57 100644 --- a/cmd/micromdm/server.go +++ b/cmd/micromdm/server.go @@ -6,6 +6,7 @@ import ( "io/ioutil" "os" + "github.com/gorilla/securecookie" "github.com/jackc/pgx/v4" "micromdm.io/v2/internal/frontend/account" "micromdm.io/v2/pkg/frontend" @@ -16,23 +17,30 @@ type server struct { ui *frontend.Server } -func ui(f *cliFlags, logger log.Logger) (*frontend.Server, error) { +func ui( + f *cliFlags, + logger log.Logger, + sess frontend.SessionStore, + cookie *securecookie.SecureCookie, +) (*frontend.Server, error) { return frontend.New(frontend.Config{ Logger: logger, SiteName: f.siteName, CSRFKey: []byte(f.csrfKey), CSRFCookieName: f.csrfCookieName, CSRFFieldName: f.csrfFieldName, + AuthCookieName: f.authCookieName, + SessionStore: sess, + Cookie: cookie, }) } func setup(ctx context.Context, f *cliFlags, logger log.Logger) (*server, error) { - uisrv, err := ui(f, logger) - if err != nil { - return nil, err - } + var ( + err error + db database + ) - var db database switch dbDriver(f.databaseURL) { case "postgres": db.pg, err = setupPostgres(ctx, f, logger) @@ -46,16 +54,40 @@ func setup(ctx context.Context, f *cliFlags, logger log.Logger) (*server, error) return nil, err } + sc, err := cookie() + if err != nil { + return nil, err + } + + uisrv, err := ui(f, logger, db.sessiondb().(frontend.SessionStore), sc) + if err != nil { + return nil, err + } + srv := &server{ui: uisrv} account.HTTP(account.Config{ - HTTP: srv.ui, - UserStore: db.userdb().(account.UserStore), + HTTP: srv.ui, + UserStore: db.userdb().(account.UserStore), + SessionStore: db.sessiondb().(account.SessionStore), + Cookie: sc, }) return srv, nil } +func cookie() (*securecookie.SecureCookie, error) { + cache := "build/cookie" // TODO: come up with a cache location/use flag + random, err := ioutil.ReadFile(cache) + if err != nil && os.IsNotExist(err) { + random = securecookie.GenerateRandomKey(64) + ioutil.WriteFile(cache, random, 0600) + } else if err != nil { + return nil, err + } + return securecookie.New(random, nil), nil // not encrypted, only signed +} + func dbDriver(dbURL string) string { if _, err := pgx.ParseConfig(dbURL); err == nil { return "postgres" diff --git a/go.mod b/go.mod index 9c9f84e4..91a9953b 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/gorilla/securecookie v1.1.1 github.com/jackc/pgconn v1.6.2 github.com/jackc/pgx/v4 v4.7.2 github.com/oklog/run v1.1.0 diff --git a/go.sum b/go.sum index b1679aa0..aaeeb10a 100644 --- a/go.sum +++ b/go.sum @@ -19,6 +19,7 @@ github.com/armon/go-metrics v0.0.0-20180917152333-f0300d1749da/go.mod h1:Q73ZrmV github.com/armon/go-radix v0.0.0-20180808171621-7fddfc383310/go.mod h1:ufUuZ+zHj4x4TnLV4JWEpy2hxWSpsRywHrMgIH9cCH8= github.com/aryann/difflib v0.0.0-20170710044230-e206f873d14a/go.mod h1:DAHtR1m6lCRdSC2Tm3DSWRPvIPr6xNKyeHdqDQSQT+A= github.com/aws/aws-lambda-go v1.13.3/go.mod h1:4UKl9IzQMoD+QF79YdCuzCwp8VbmG4VAQwij/eHl5CU= +github.com/aws/aws-sdk-go v1.27.0 h1:0xphMHGMLBrPMfxR2AmVjZKcMEESEgWF8Kru94BNByk= github.com/aws/aws-sdk-go v1.27.0/go.mod h1:KmX6BPdI08NWTb3/sm4ZGu5ShLoqVDhKgpiN924inxo= github.com/aws/aws-sdk-go-v2 v0.18.0/go.mod h1:JWVYvqSMppoMJC0x5wdwiImzgXTI9FuZwxzkQq9wy+g= github.com/beorn7/perks v0.0.0-20180321164747-3a771d992973/go.mod h1:Dwedo/Wpr24TaqPxmxbtue+5NUziq4I4S80YR8gNf3Q= @@ -187,6 +188,7 @@ github.com/jackc/puddle v0.0.0-20190608224051-11cab39313c9/go.mod h1:m4B5Dj62Y0f github.com/jackc/puddle v1.1.0/go.mod h1:m4B5Dj62Y0fbyuIc15OsIqK0+JU8nkqQjsgx7dvjSWk= github.com/jackc/puddle v1.1.1 h1:PJAw7H/9hoWC4Kf3J8iNmL1SwA6E8vfsLqBiL+F6CtI= github.com/jackc/puddle v1.1.1/go.mod h1:m4B5Dj62Y0fbyuIc15OsIqK0+JU8nkqQjsgx7dvjSWk= +github.com/jmespath/go-jmespath v0.0.0-20180206201540-c2b33e8439af h1:pmfjZENx5imkbgOkpRUYLnmbU7UEFbjtDA2hxJ1ichM= github.com/jmespath/go-jmespath v0.0.0-20180206201540-c2b33e8439af/go.mod h1:Nht3zPeWKUH0NzdCt2Blrr5ys8VGpn0CEB0cQHVjt7k= github.com/jonboulle/clockwork v0.1.0/go.mod h1:Ii8DK3G1RaLaWxj9trq07+26W01tbo22gdxWY5EU2bo= github.com/json-iterator/go v1.1.6/go.mod h1:+SdeFBvtyEkXs7REEP0seUULqWtbJapLOCVDaaPEHmU= diff --git a/internal/data/migrations/postgres/initial_tables.sql b/internal/data/migrations/postgres/initial_tables.sql index 2a001f1c..8f43c6fb 100644 --- a/internal/data/migrations/postgres/initial_tables.sql +++ b/internal/data/migrations/postgres/initial_tables.sql @@ -1,3 +1,4 @@ +DROP TABLE IF EXISTS sessions; DROP TABLE IF EXISTS users; CREATE TABLE IF NOT EXISTS users ( @@ -14,3 +15,10 @@ CREATE TABLE IF NOT EXISTS users ( UNIQUE (email), UNIQUE (username) ); + +CREATE TABLE IF NOT EXISTS sessions ( + id text PRIMARY KEY NOT NULL, + user_id text REFERENCES users(id) ON DELETE CASCADE, + created_at timestamptz DEFAULT (now() at time zone 'utc'), + accessed_at timestamptz DEFAULT (now() at time zone 'utc') +); diff --git a/internal/data/migrations/sqlite/initial_tables.sql b/internal/data/migrations/sqlite/initial_tables.sql index 7b02a579..44cac4d6 100644 --- a/internal/data/migrations/sqlite/initial_tables.sql +++ b/internal/data/migrations/sqlite/initial_tables.sql @@ -1,6 +1,8 @@ PRAGMA auto_vacuum = INCREMENTAL; +PRAGMA foreign_keys = ON; --- DROP TABLE IF EXISTS users; +DROP TABLE IF EXISTS users; +DROP TABLE IF EXISTS sessions; CREATE TABLE IF NOT EXISTS users ( id text PRIMARY KEY, @@ -25,3 +27,12 @@ BEGIN WHERE id = old.id; END; + +CREATE TABLE IF NOT EXISTS sessions ( + id text PRIMARY KEY NOT NULL, + user_id text REFERENCES users(id) ON DELETE CASCADE, + created_at DATETIME DEFAULT CURRENT_TIMESTAMP, + accessed_at DATETIME DEFAULT CURRENT_TIMESTAMP +); + +CREATE INDEX IF NOT EXISTS session_users_idx ON sessions(user_id); diff --git a/internal/data/session/session.go b/internal/data/session/session.go new file mode 100644 index 00000000..369265a3 --- /dev/null +++ b/internal/data/session/session.go @@ -0,0 +1,39 @@ +// Package session tracks user authenticated sessions. +package session + +import ( + "context" + "errors" + "time" + + "micromdm.io/v2/pkg/id" + "micromdm.io/v2/pkg/viewer" +) + +type Session struct { + ID string + UserID string + CreatedAt time.Time + AccessedAt time.Time +} + +func columns() []string { + return []string{ + "id", + "user_id", + "created_at", + "accessed_at", + } +} + +func create(ctx context.Context) (*Session, error) { + v, ok := viewer.FromContext(ctx) + if !ok { + return nil, errors.New("cannot create session without a viewer in context") + } + + return &Session{ + ID: id.New(), + UserID: v.UserID, + }, nil +} diff --git a/internal/data/session/session_postgres.go b/internal/data/session/session_postgres.go new file mode 100644 index 00000000..488087da --- /dev/null +++ b/internal/data/session/session_postgres.go @@ -0,0 +1,89 @@ +package session + +import ( + "context" + "fmt" + "strings" + + "github.com/jackc/pgx/v4/pgxpool" + "micromdm.io/v2/pkg/viewer" +) + +// Postgres provides methods for creating sessions in PostgreSQL. +type Postgres struct{ db *pgxpool.Pool } + +// NewPostgres creates a Postgres client. +func NewPostgres(db *pgxpool.Pool) *Postgres { + return &Postgres{db: db} +} + +func (d *Postgres) CreateSession(ctx context.Context) (*Session, error) { + s, err := create(ctx) + if err != nil { + return nil, err + } + + q := `INSERT INTO sessions ( id, user_id ) VALUES ( $1, $2 ) RETURNING "created_at", "accessed_at";` + if err := d.db.QueryRow(ctx, q, + s.ID, + s.UserID, + ).Scan(&s.CreatedAt, &s.AccessedAt); err != nil { + return nil, fmt.Errorf("store new session in postgres: %w", err) + } + + return s, nil +} + +func (d *Postgres) DestroySession(ctx context.Context) error { + v, ok := viewer.FromContext(ctx) + if !ok || v.SessionID == "" { + return fmt.Errorf("session: missing valid viewer %v", v) + } + + q := `DELETE FROM sessions WHERE id = $1;` + + if _, err := d.db.Exec(ctx, q, v.SessionID); err != nil { + return err + } + + return nil +} + +func (d *Postgres) FindSession(ctx context.Context, id string) (*Session, error) { + q := fmt.Sprintf(`WITH updated AS ( + UPDATE sessions SET accessed_at = now() at time zone 'utc' + FROM ( + SELECT user_id FROM sessions + WHERE id = $1 + LIMIT 1 + ) sub + JOIN users u on u.id = sub.user_id + WHERE sessions.id = $2 + RETURNING %s + ) + SELECT %s from updated;`, // TODO: add username, full_name to session struct. + strings.Join(prefixTable("sessions", columns()), `, `), + strings.Join(columns(), `, `), + ) + + s := &Session{ID: id} + if err := d.db.QueryRow(ctx, q, id, id).Scan( + &s.ID, + &s.UserID, + &s.CreatedAt, + &s.AccessedAt, + ); err != nil { + return nil, err + } + + return s, nil + +} + +func prefixTable(prefix string, ss []string) []string { + result := make([]string, len(ss), cap(ss)) + for i, s := range ss { + result[i] = fmt.Sprintf("%s.%s", prefix, s) + } + return result +} diff --git a/internal/data/session/session_sqlite.go b/internal/data/session/session_sqlite.go new file mode 100644 index 00000000..a7a8b0a8 --- /dev/null +++ b/internal/data/session/session_sqlite.go @@ -0,0 +1,119 @@ +package session + +import ( + "context" + "errors" + "fmt" + "time" + + "crawshaw.io/sqlite/sqlitex" + "micromdm.io/v2/pkg/viewer" +) + +// SQLite provides methods for creating users in SQLite. +type SQLite struct{ db *sqlitex.Pool } + +// NewSQLite creates a SQLite client. +func NewSQLite(db *sqlitex.Pool) *SQLite { + return &SQLite{db: db} +} + +func (d *SQLite) CreateSession(ctx context.Context) (*Session, error) { + s, err := create(ctx) + if err != nil { + return nil, err + } + + conn := d.db.Get(ctx) + if conn == nil { + return nil, context.Canceled + } + defer d.db.Put(conn) + + stmt := conn.Prep(`INSERT INTO sessions ( id, user_id ) VALUES ( $id, $user_id );`) + stmt.SetText("$id", s.ID) + stmt.SetText("$user_id", s.UserID) + if _, err := stmt.Step(); err != nil { + return nil, fmt.Errorf("store new session in sqlite: %w", err) + } + + stmt = conn.Prep(`SELECT created_at, accessed_at FROM sessions WHERE id = $id`) + stmt.SetText("$id", s.ID) + if _, err := stmt.Step(); err != nil { + return nil, err + } + + s.CreatedAt, err = time.Parse("2006-01-02 15:04:05", stmt.GetText("created_at")) + if err != nil { + return nil, err + } + + s.AccessedAt, err = time.Parse("2006-01-02 15:04:05", stmt.GetText("accessed_at")) + if err != nil { + return nil, err + } + + return s, stmt.Reset() +} + +func (d *SQLite) DestroySession(ctx context.Context) error { + v, ok := viewer.FromContext(ctx) + if !ok || v.SessionID == "" { + return fmt.Errorf("session: missing valid viewer %v", v) + } + + conn := d.db.Get(ctx) + if conn == nil { + return context.Canceled + } + defer d.db.Put(conn) + + stmt := conn.Prep(`DELETE FROM sessions WHERE id = $id;`) + stmt.SetText("$id", v.SessionID) + if _, err := stmt.Step(); err != nil { + return err + } + return nil +} + +func (d *SQLite) FindSession(ctx context.Context, id string) (*Session, error) { + conn := d.db.Get(ctx) + if conn == nil { + return nil, context.Canceled + } + defer d.db.Put(conn) + + stmt := conn.Prep( + `UPDATE sessions + SET accessed_at = CURRENT_TIMESTAMP + WHERE id = $id;`) + stmt.SetText("$id", id) + if _, err := stmt.Step(); err != nil { + return nil, err + } + + if conn.Changes() == 0 { + return nil, errors.New("unknown session.id in sqlite") + } + + stmt = conn.Prep(`SELECT created_at, accessed_at, user_id FROM sessions WHERE id = $id`) + stmt.SetText("$id", id) + if _, err := stmt.Step(); err != nil { + return nil, err + } + + s := Session{ID: id} + var err error + s.CreatedAt, err = time.Parse("2006-01-02 15:04:05", stmt.GetText("created_at")) + if err != nil { + return nil, err + } + + s.AccessedAt, err = time.Parse("2006-01-02 15:04:05", stmt.GetText("accessed_at")) + if err != nil { + return nil, err + } + + s.UserID = stmt.GetText("user_id") + return &s, stmt.Reset() +} diff --git a/internal/data/user/user_postgres.go b/internal/data/user/user_postgres.go index 25f0b850..4f4c2af9 100644 --- a/internal/data/user/user_postgres.go +++ b/internal/data/user/user_postgres.go @@ -4,8 +4,10 @@ import ( "context" "errors" "fmt" + "strings" "github.com/jackc/pgconn" + "github.com/jackc/pgx/v4" "github.com/jackc/pgx/v4/pgxpool" ) @@ -60,6 +62,31 @@ func (d *Postgres) ConfirmUser(ctx context.Context, confirmation string) error { return nil } +func (d *Postgres) FindUserByEmail(ctx context.Context, email string) (*User, error) { + if email == "" { + return nil, Error{invalid: constraints["chk_email_not_empty"]} + } + + u := &User{} + q := fmt.Sprintf(`SELECT %s FROM users WHERE email = $1;`, strings.Join(columns(), `, `)) + if err := d.db.QueryRow(ctx, q, email).Scan( + &u.ID, + &u.Username, + &u.Email, + &u.Password, + &u.Salt, + &u.ConfirmationHash, + &u.CreatedAt, + &u.UpdatedAt, + ); err == pgx.ErrNoRows { + return nil, fmt.Errorf("user (email %q) not found in postgres", email) + } else if err != nil { + return nil, err + } + + return u, nil +} + func checkPostgres(err error) error { var dbErr *pgconn.PgError if !errors.As(err, &dbErr) { diff --git a/internal/data/user/user_sqlite.go b/internal/data/user/user_sqlite.go index 52479003..a871c09c 100644 --- a/internal/data/user/user_sqlite.go +++ b/internal/data/user/user_sqlite.go @@ -47,23 +47,21 @@ func (d *SQLite) CreateUser(ctx context.Context, username, email, password strin 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`) + stmt = conn.Prep(fmt.Sprintf( + `SELECT %s FROM users WHERE id = $id;`, strings.Join(columns(), `, `))) stmt.SetText("$id", u.ID) - if _, err := stmt.Step(); err != nil { + if found, err := stmt.Step(); err != nil { return nil, err + } else if !found { + return nil, fmt.Errorf("user (email %q) not found in sqlite", email) } - u.CreatedAt, err = time.Parse("2006-01-02 15:04:05", stmt.GetText("created_at")) + usr, err := sqliteUser(stmt) if err != nil { return nil, err } - u.UpdatedAt, err = time.Parse("2006-01-02 15:04:05", stmt.GetText("updated_at")) - if err != nil { - return nil, err - } - - return u, stmt.Reset() + return usr, nil } func (d *SQLite) ConfirmUser(ctx context.Context, confirmation string) error { @@ -90,6 +88,62 @@ func (d *SQLite) ConfirmUser(ctx context.Context, confirmation string) error { return nil } +func (d *SQLite) FindUserByEmail(ctx context.Context, email string) (*User, error) { + conn := d.db.Get(ctx) + if conn == nil { + return nil, context.Canceled + } + defer d.db.Put(conn) + + if email == "" { + return nil, Error{invalid: constraints["chk_email_not_empty"]} + } + + stmt := conn.Prep(fmt.Sprintf( + `SELECT %s FROM users WHERE email = $email;`, strings.Join(columns(), `, `))) + stmt.SetText("$email", email) + if found, err := stmt.Step(); err != nil { + return nil, err + } else if !found { + return nil, fmt.Errorf("user (email %q) not found in sqlite", email) + } + + usr, err := sqliteUser(stmt) + if err != nil { + return nil, err + } + return usr, stmt.Reset() +} + +func sqliteUser(stmt *sqlite.Stmt) (*User, error) { + var ( + u = new(User) + err error + ) + + u.ID = stmt.GetText("id") + u.Email = stmt.GetText("email") + u.Password = []byte(stmt.GetText("password")) + u.Salt = []byte(stmt.GetText("salt")) + + // TODO: do we really need the pointer? + if hash := stmt.GetText("confirmation_hash"); hash != "" { + u.ConfirmationHash = &hash + } + + u.CreatedAt, err = time.Parse("2006-01-02 15:04:05", stmt.GetText("created_at")) + if err != nil { + return nil, err + } + + u.UpdatedAt, err = time.Parse("2006-01-02 15:04:05", stmt.GetText("updated_at")) + if err != nil { + return nil, err + } + + return u, stmt.Reset() +} + func checkSqlite(err error) error { var sErr sqlite.Error if !errors.As(err, &sErr) { diff --git a/internal/frontend/account/account.go b/internal/frontend/account/account.go index b410e647..ad9dfbac 100644 --- a/internal/frontend/account/account.go +++ b/internal/frontend/account/account.go @@ -5,6 +5,8 @@ import ( "context" "net/http" + "github.com/gorilla/securecookie" + "micromdm.io/v2/internal/data/session" "micromdm.io/v2/internal/data/user" "micromdm.io/v2/pkg/frontend" ) @@ -12,25 +14,45 @@ import ( type UserStore interface { CreateUser(ctx context.Context, username, email, password string) (*user.User, error) ConfirmUser(ctx context.Context, token string) error + FindUserByEmail(ctx context.Context, email string) (*user.User, error) +} + +type SessionStore interface { + CreateSession(ctx context.Context) (*session.Session, error) + DestroySession(ctx context.Context) error +} + +type CookieAuthFramework interface { + frontend.Framework + AuthCookieName() string } type server struct { - http frontend.Framework - userdb UserStore + http CookieAuthFramework + userdb UserStore + sessiondb SessionStore + cookie *securecookie.SecureCookie } type Config struct { - HTTP frontend.Framework - UserStore UserStore + HTTP CookieAuthFramework + UserStore UserStore + SessionStore SessionStore + Cookie *securecookie.SecureCookie } func HTTP(config Config) { srv := &server{ - http: config.HTTP, - userdb: config.UserStore, + http: config.HTTP, + userdb: config.UserStore, + sessiondb: config.SessionStore, + cookie: config.Cookie, } 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) + + srv.http.HandleFunc("/login", srv.loginForm, http.MethodGet, http.MethodPost) + srv.http.HandleFunc("/logout", srv.logout) } diff --git a/internal/frontend/account/login.go b/internal/frontend/account/login.go new file mode 100644 index 00000000..8bc44a1c --- /dev/null +++ b/internal/frontend/account/login.go @@ -0,0 +1,104 @@ +package account + +import ( + "context" + "fmt" + "net/http" + "time" + + "github.com/gorilla/csrf" + "micromdm.io/v2/pkg/frontend" + "micromdm.io/v2/pkg/log" + "micromdm.io/v2/pkg/viewer" +) + +func (srv server) loginForm(w http.ResponseWriter, r *http.Request) { + var ( + ctx = r.Context() + logger = log.FromContext(ctx) + email = r.FormValue("email") + password = r.FormValue("password") + data = frontend.Data{ + csrf.TemplateTag: csrf.TemplateField(r), + "form": map[string]string{ + "email": email, + "password": password, + }, + } + ) + + if r.Method == http.MethodGet { + srv.http.RenderTemplate(ctx, w, "login.tmpl", data) + return + } + + // TODO handle errors other than 500. + // alert: Username or password incorrect. + + usr, err := srv.userdb.FindUserByEmail(ctx, email) + if err != nil { + srv.http.Fail(ctx, w, err, "msg", "find user for auth") + return + } + + if err := usr.ValidatePassword(password); err != nil { + srv.http.Fail(ctx, w, err, "msg", "auth user") + return + } + + log.Debug(logger).Log("msg", "got user", "user_id", usr.ID) + + if err := srv.createSession(ctx, w, usr.ID); err != nil { + srv.http.Fail(ctx, w, err, "msg", "create session") + return + } + + log.Debug(logger).Log("msg", "logged in", "user_id", usr.ID) +} + +func (srv server) logout(w http.ResponseWriter, r *http.Request) { + var ( + ctx = r.Context() + logger = log.FromContext(ctx) + v, _ = viewer.FromContext(ctx) + ) + + if err := srv.sessiondb.DestroySession(ctx); err != nil { + log.Info(logger).Log("err", err, "msg", "destroy session on logout") + } + + http.SetCookie(w, &http.Cookie{ + Name: srv.http.AuthCookieName(), + Path: "/", + Secure: true, + HttpOnly: true, + MaxAge: -1, + }) + + log.Debug(logger).Log("msg", "user logged out", "user_id", v.UserID, "session_id", v.SessionID) + http.Redirect(w, r, "/login", http.StatusFound) +} + +func (srv server) createSession(ctx context.Context, w http.ResponseWriter, userID string) error { + ctx = viewer.NewContext(ctx, viewer.Viewer{UserID: userID}) + sess, err := srv.sessiondb.CreateSession(ctx) + if err != nil { + return fmt.Errorf("create session for %q: %w", userID, err) + } + + token, err := srv.cookie.Encode(srv.http.AuthCookieName(), map[string]string{"id": sess.ID}) + if err != nil { + return err + } + + http.SetCookie(w, &http.Cookie{ + Name: srv.http.AuthCookieName(), + Path: "/", + Value: token, + Secure: true, + HttpOnly: true, + Expires: time.Now().UTC().Add(30 * time.Minute), + }) + + return nil +} diff --git a/pkg/frontend/frontend.go b/pkg/frontend/frontend.go index 0405a017..102bd67d 100644 --- a/pkg/frontend/frontend.go +++ b/pkg/frontend/frontend.go @@ -12,11 +12,16 @@ import ( "strings" "sync" "text/template" + "time" + "github.com/go-kit/kit/log/level" "github.com/gorilla/csrf" "github.com/gorilla/mux" + "github.com/gorilla/securecookie" + "micromdm.io/v2/internal/data/session" "micromdm.io/v2/pkg/log" + "micromdm.io/v2/pkg/viewer" ) // Data keys. Private, set via helper methods. @@ -60,6 +65,10 @@ type Framework interface { HandleFunc(path string, f func(http.ResponseWriter, *http.Request), methods ...string) } +type SessionStore interface { + FindSession(ctx context.Context, id string) (*session.Session, error) +} + // Server implements Framework. type Server struct { r *mux.Router @@ -72,6 +81,10 @@ type Server struct { csrfKey []byte csrfCookieName string csrfFieldName string + + authCookieName string + sessiondb SessionStore + cookie *securecookie.SecureCookie } // Config parameters to create a new Server. @@ -82,6 +95,10 @@ type Config struct { CSRFKey []byte CSRFCookieName string CSRFFieldName string + + AuthCookieName string + SessionStore SessionStore + Cookie *securecookie.SecureCookie } // New creates a Server. @@ -94,10 +111,14 @@ func New(config Config) (*Server, error) { csrfKey: config.CSRFKey, csrfFieldName: config.CSRFFieldName, csrfCookieName: config.CSRFCookieName, + authCookieName: config.AuthCookieName, + sessiondb: config.SessionStore, + cookie: config.Cookie, } srv.r.Use( log.HTTP(config.Logger), // HTTP logging middleware. + srv.authMW, // check auth, add viewer srv.csrf, // CSRF protection. srv.recoverPanic, // convert any panic into 500 errors. ) @@ -125,6 +146,9 @@ func New(config Config) (*Server, error) { // Handler returns the mux router used by the Server. func (srv *Server) Handler() http.Handler { return srv.r } +// AuthCookieName exposes the cookie name used for authentication. +func (srv *Server) AuthCookieName() string { return srv.authCookieName } + // HandleFunc wraps *mux.Router, allowing other packages to register with the router. func (srv *Server) HandleFunc(path string, f func(http.ResponseWriter, *http.Request), methods ...string) { if len(methods) == 0 { @@ -252,13 +276,13 @@ func (srv *Server) recoverPanic(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { defer func() { if err := recover(); err != nil { - srv.Fail(r.Context(), w, fmt.Errorf("panic: %v", err), "msg", "recover panic") - debug.PrintStack() + srv.Fail(r.Context(), w, fmt.Errorf("panic: %v", err), "msg", "recover panic", "debug_stack", string(debug.Stack())) } }() next.ServeHTTP(w, r) }) } + func csrfDisabled(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { logger := log.FromContext(r.Context()) @@ -306,3 +330,86 @@ func (srv *Server) notFound(w http.ResponseWriter, r *http.Request) { func (srv *Server) indexPage(w http.ResponseWriter, r *http.Request) { srv.RenderTemplate(r.Context(), w, "home.tmpl", Data{}) } + +// SESSION + +var pubicPagePrefix = []string{ + "/login", + "/assets/", + "/forgot", + "/register", +} + +func isPublic(path string) bool { + for _, k := range pubicPagePrefix { + if strings.HasPrefix(path, k) { + return true + } + } + return false +} + +func (srv *Server) sessionFromRequest(r *http.Request) (*session.Session, error) { + ctx := r.Context() + + cookie, err := r.Cookie(srv.authCookieName) + if err != nil { + return nil, err + } + + value := make(map[string]string) + if err := srv.cookie.Decode(srv.authCookieName, cookie.Value, &value); err != nil { + return nil, err + } + + id, ok := value["id"] + if !ok { + return nil, errors.New("auth cookie present but no id value") + } + + sess, err := srv.sessiondb.FindSession(ctx, id) + if err != nil { + return nil, err + } + + if sess.CreatedAt.Before(time.Now().UTC().Add(time.Duration(-30) * time.Minute)) { + return nil, fmt.Errorf("session (id %s) expired", sess.ID) + } + + return sess, nil +} + +func (srv *Server) authMW(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var ( + ctx = r.Context() + logger = log.FromContext(ctx) + ) + sess, err := srv.sessionFromRequest(r) + if err == http.ErrNoCookie && isPublic(r.URL.Path) { + next.ServeHTTP(w, r.WithContext(ctx)) + return + } else if err != nil && isPublic(r.URL.Path) { + srv.Fail(r.Context(), w, err, "msg", "auth mw failed to render") + return + } else if err != nil { + // TODO: handle not found and destroy the session somewhere. + log.Debug(logger).Log("err", err, "msg", "check auth") + http.Redirect(w, r, "/login", http.StatusTemporaryRedirect) + return + } + + ctx = viewer.NewContext(ctx, viewer.Viewer{ + UserID: sess.UserID, + SessionID: sess.ID, + }) + + if strings.HasPrefix(r.URL.Path, "/login") { + level.Debug(logger).Log("msg", "redirect session to dashboard page", "session_id", sess.ID) + http.Redirect(w, r, "/dashboard", http.StatusFound) + return + } + + next.ServeHTTP(w, r.WithContext(ctx)) + }) +} diff --git a/ui/includes/login.tmpl b/ui/includes/login.tmpl new file mode 100644 index 00000000..2471de1e --- /dev/null +++ b/ui/includes/login.tmpl @@ -0,0 +1,40 @@ +{{ define "title" -}} + Sign In to {{ .siteName }} +{{- end }} +{{ define "content" -}} +
+
+ Sign In + {{ .csrfField }} +
+ + + {{- with .errors -}} +
{{ .email }}
+ {{- end }} +
+
+ + + {{ with .errors }} +
{{ .password}}
+ {{ end }} +
+ {{with .alert}} +
{{.}}
+ {{end}} +
+ +
+ +
+
+{{- end }}