db.go

  1package session
  2
  3import (
  4	"crypto/rand"
  5	"crypto/sha256"
  6	"database/sql"
  7	"encoding/base32"
  8	"errors"
  9	"fmt"
 10	"time"
 11)
 12
 13var (
 14	ErrNoSession error = errors.New("no session")
 15)
 16
 17func Migrate(db *sql.DB) error {
 18	s := `CREATE TABLE IF NOT EXISTS
 19		sessions (
 20			session_hash BLOB NOT NULL,
 21			user_id INTEGER NOT NULL,
 22			expires TIMESTAMP NOT NULL,
 23		FOREIGN KEY (user_id)
 24			REFERENCES users(id)
 25			ON DELETE CASCADE
 26		);
 27	CREATE UNIQUE INDEX IF NOT EXISTS session_idx ON sessions(session_hash);
 28	CREATE TRIGGER IF NOT EXISTS delete_expired_sessions AFTER INSERT ON sessions
 29	BEGIN
 30		DELETE FROM sessions WHERE expires < datetime('now');
 31	END`
 32	if _, err := db.Exec(s); err != nil {
 33		return fmt.Errorf("failed to migrate sessions database: %w", err)
 34	}
 35	return nil
 36}
 37
 38// Persist a session to the database. If session already exists, the expiration time is updated
 39// but the session cannot be assigned to a new userID to prevent session fixation
 40func Persist[T ~int64](db *sql.DB, sessionID string, userID T, expires time.Time) error {
 41	s := `INSERT INTO sessions
 42			(session_hash, user_id, expires)
 43		VALUES (?, ?, ?)
 44		ON CONFLICT (session_hash) DO
 45			UPDATE SET expires=excluded.expires WHERE excluded.user_id = sessions.user_id`
 46	if _, err := db.Exec(s, hash(decode(sessionID)), userID, expires); err != nil {
 47		return fmt.Errorf("failed to persist session to database for user=%d: %w", userID, err)
 48	}
 49	return nil
 50}
 51
 52// DeleteBySessionID deletes a single session
 53func DeleteBySessionID(db *sql.DB, sessionID string) error {
 54	s := `DELETE FROM sessions WHERE session_hash = ?`
 55	if _, err := db.Exec(s, hash(decode(sessionID))); err != nil {
 56		return fmt.Errorf("failed to delete session from database: %w", err)
 57	}
 58	return nil
 59}
 60
 61// DeleteByUserID deletes all sessions associated with a user. This logs out everywhere on multiple
 62// devices.
 63func DeleteByUserID[T ~int64](db *sql.DB, userID T) error {
 64	s := `DELETE FROM sessions WHERE user_id = ?`
 65	if _, err := db.Exec(s, userID); err != nil {
 66		return fmt.Errorf("failed to delete sessions for user=%d: %w", userID, err)
 67	}
 68	return nil
 69}
 70
 71// Read gets the user id associated with a session. Returns ErrNoSession if none found or expired.
 72func Read[T ~int64](db *sql.DB, sessionID string) (userID T, err error) {
 73	s := `SELECT user_id FROM sessions WHERE session_hash = ? AND expires > datetime('now')`
 74	if err := db.QueryRow(s, hash(decode(sessionID))).Scan(&userID); err != nil {
 75		if errors.Is(err, sql.ErrNoRows) {
 76			return 0, ErrNoSession
 77		} else {
 78			return 0, fmt.Errorf("unknown error getting session from database: %w", err)
 79		}
 80	}
 81	return
 82}
 83
 84// NewID returns a cryptographically random session identifier
 85func NewID() string {
 86	// OWASP 2025 recommends at least 64 bits randomness, using 80 for base 32 encoding
 87	var id [10]byte
 88	rand.Read(id[:])
 89	return base32.StdEncoding.EncodeToString(id[:])
 90}
 91
 92func decode(str string) []byte {
 93	b, err := base32.StdEncoding.DecodeString(str)
 94	if err != nil {
 95		return []byte{}
 96	}
 97	return b
 98}
 99
100func hash(in []byte) []byte {
101	h := sha256.Sum256(in)
102	return h[:]
103}