mirror of
https://github.com/micromdm/micromdm/
synced 2026-08-08 18:55:34 +08:00
Added users sessions (#696)
Added login/logout pages and respective database methods for handling users sessions.
This commit is contained in:
9
Makefile
9
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
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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"
|
||||
|
||||
1
go.mod
1
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
|
||||
|
||||
2
go.sum
2
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=
|
||||
|
||||
@@ -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')
|
||||
);
|
||||
|
||||
@@ -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);
|
||||
|
||||
39
internal/data/session/session.go
Normal file
39
internal/data/session/session.go
Normal file
@@ -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
|
||||
}
|
||||
89
internal/data/session/session_postgres.go
Normal file
89
internal/data/session/session_postgres.go
Normal file
@@ -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
|
||||
}
|
||||
119
internal/data/session/session_sqlite.go
Normal file
119
internal/data/session/session_sqlite.go
Normal file
@@ -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()
|
||||
}
|
||||
@@ -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) {
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
104
internal/frontend/account/login.go
Normal file
104
internal/frontend/account/login.go
Normal file
@@ -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
|
||||
}
|
||||
@@ -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))
|
||||
})
|
||||
}
|
||||
|
||||
40
ui/includes/login.tmpl
Normal file
40
ui/includes/login.tmpl
Normal file
@@ -0,0 +1,40 @@
|
||||
{{ define "title" -}}
|
||||
<title>Sign In to {{ .siteName }}</title>
|
||||
{{- end }}
|
||||
{{ define "content" -}}
|
||||
<div class="page">
|
||||
<form class="single-page-form" method="post">
|
||||
<legend>Sign In</legend>
|
||||
{{ .csrfField }}
|
||||
<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>
|
||||
{{with .alert}}
|
||||
<div class="invalid-alert">{{.}}</div>
|
||||
{{end}}
|
||||
<div class="button-container">
|
||||
<button class="btn">Sign In</button>
|
||||
</div>
|
||||
<div class="form-footer">Don't have an account?<a href="/register">Sign Up</a></div>
|
||||
</form>
|
||||
</div>
|
||||
{{- end }}
|
||||
Reference in New Issue
Block a user