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}