Added users sessions (#696)

Added login/logout pages and respective database methods for handling users sessions.
This commit is contained in:
Victor Vrantchan
2020-08-08 17:15:34 -04:00
committed by GitHub
parent 49c3a04637
commit 7d75f581ec
17 changed files with 705 additions and 34 deletions

View File

@@ -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

View File

@@ -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")

View File

@@ -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.

View File

@@ -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
View File

@@ -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
View File

@@ -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=

View File

@@ -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')
);

View File

@@ -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);

View 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
}

View 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
}

View 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()
}

View File

@@ -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) {

View File

@@ -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) {

View File

@@ -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)
}

View 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
}

View File

@@ -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
View 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 }}