package db import ( "crypto/rand" "database/sql" "encoding/hex" "errors" "time" ) // ErrSessionNotFound is returned when a session token doesn't exist or has // expired. var ErrSessionNotFound = errors.New("session not found") // SessionTTL is how long a web UI login session stays valid after creation. const SessionTTL = 7 * 24 * time.Hour // CreateSession generates a new random session token for username and // stores it with an expiry SessionTTL from now. Returns the token to be // set as a cookie value. func (d *DB) CreateSession(username string) (string, error) { token, err := randomToken() if err != nil { return "", err } expires := time.Now().Add(SessionTTL) _, err = d.conn.Exec( `INSERT INTO web_sessions (token, username, expires_at) VALUES (?, ?, ?)`, token, username, expires, ) if err != nil { return "", err } return token, nil } // SessionUser returns the username associated with token, provided it // exists and hasn't expired. Returns ErrSessionNotFound otherwise. func (d *DB) SessionUser(token string) (string, error) { var username string var expiresAt time.Time err := d.conn.QueryRow( `SELECT username, expires_at FROM web_sessions WHERE token = ?`, token, ).Scan(&username, &expiresAt) if errors.Is(err, sql.ErrNoRows) { return "", ErrSessionNotFound } if err != nil { return "", err } if time.Now().After(expiresAt) { _ = d.DeleteSession(token) return "", ErrSessionNotFound } return username, nil } // DeleteSession removes a session (used on logout). It's not an error if // the token doesn't exist. func (d *DB) DeleteSession(token string) error { _, err := d.conn.Exec(`DELETE FROM web_sessions WHERE token = ?`, token) return err } // PruneExpiredSessions deletes all sessions past their expiry. Intended to // be called periodically (e.g. on server startup and via a background // ticker) to keep the table small. func (d *DB) PruneExpiredSessions() error { _, err := d.conn.Exec(`DELETE FROM web_sessions WHERE expires_at < ?`, time.Now()) return err } func randomToken() (string, error) { buf := make([]byte, 32) if _, err := rand.Read(buf); err != nil { return "", err } return hex.EncodeToString(buf), nil }