v0.11.0 harden persistence and prepare game replays

This commit is contained in:
2026-09-07 14:39:00 -04:00
parent 21dea47694
commit 5d2be4abb2
42 changed files with 3591 additions and 693 deletions
+426 -86
View File
@@ -1,137 +1,477 @@
package storage
import (
"context"
"database/sql"
"errors"
"fmt"
"log"
"log/slog"
"time"
)
// RecordNewGame asynchronously records a new game
const gameSelectColumns = `
g.game_id, g.initial_fen,
g.white_player_id, g.white_type, g.white_level, g.white_search_time, g.white_claimed_by,
g.black_player_id, g.black_type, g.black_level, g.black_search_time, g.black_claimed_by,
g.result, g.start_time_utc, g.end_time_utc`
// RecordNewGame asynchronously records a new game. Terminal custom-FEN games
// include their result in this insert rather than relying on a second write.
func (s *Store) RecordNewGame(record GameRecord) error {
if !s.healthStatus.Load() {
return nil // Silently drop if degraded
if record.GameID == "" || record.InitialFEN == "" || record.WhitePlayerID == "" || record.BlackPlayerID == "" {
return errors.New("game ID, initial FEN, and player IDs are required")
}
if err := validateResultTime(record.Result, record.EndTimeUTC); err != nil {
return err
}
select {
case s.writeChan <- func(tx *sql.Tx) error {
query := `INSERT INTO games (
game_id, initial_fen,
white_player_id, white_type, white_level, white_search_time,
black_player_id, black_type, black_level, black_search_time,
start_time_utc
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`
return s.enqueue("record_game", record.GameID, func(tx *sql.Tx) error {
const query = `INSERT INTO games (
game_id, initial_fen,
white_player_id, white_type, white_level, white_search_time, white_claimed_by,
black_player_id, black_type, black_level, black_search_time, black_claimed_by,
start_time_utc, result, end_time_utc
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`
_, err := tx.Exec(query,
record.GameID, record.InitialFEN,
record.WhitePlayerID, record.WhiteType, record.WhiteLevel, record.WhiteSearchTime,
nullableString(record.WhiteClaimedBy),
record.BlackPlayerID, record.BlackType, record.BlackLevel, record.BlackSearchTime,
record.StartTimeUTC,
nullableString(record.BlackClaimedBy),
record.StartTimeUTC, nullableString(record.Result), record.EndTimeUTC,
)
return err
}:
return nil
default:
// Channel full, drop write
log.Printf("Storage write queue full, dropping game record")
return nil
}
})
}
// RecordMove asynchronously records a move
func (s *Store) RecordMove(record MoveRecord) error {
if !s.healthStatus.Load() {
return nil // Silently drop if degraded
// RecordMove atomically persists an accepted move and any first-move claim or
// terminal result caused by that move.
func (s *Store) RecordMove(record MovePersistence) error {
if record.Move.GameID == "" || record.Move.MoveNumber < 1 ||
record.Move.MoveUCI == "" || record.Move.FENAfterMove == "" {
return errors.New("move game ID, positive move number, UCI, and resulting FEN are required")
}
if record.Move.PlayerColor != "w" && record.Move.PlayerColor != "b" {
return fmt.Errorf("invalid move color %q", record.Move.PlayerColor)
}
if record.ClaimColor != "" && record.ClaimColor != "w" && record.ClaimColor != "b" {
return fmt.Errorf("invalid claim color %q", record.ClaimColor)
}
if (record.ClaimColor == "") != (record.ClaimedBy == "") {
return errors.New("claim color and claimant must be provided together")
}
if err := validateResultTime(record.Result, record.EndTimeUTC); err != nil {
return err
}
select {
case s.writeChan <- func(tx *sql.Tx) error {
query := `INSERT INTO moves (
return s.enqueue("record_move", record.Move.GameID, func(tx *sql.Tx) error {
const insertMove = `INSERT INTO moves (
game_id, move_number, move_uci, fen_after_move, player_color, move_time_utc
) VALUES (?, ?, ?, ?, ?, ?)`
if _, err := tx.Exec(insertMove,
record.Move.GameID,
record.Move.MoveNumber,
record.Move.MoveUCI,
record.Move.FENAfterMove,
record.Move.PlayerColor,
record.Move.MoveTimeUTC,
); err != nil {
return err
}
_, err := tx.Exec(query,
record.GameID, record.MoveNumber, record.MoveUCI,
record.FENAfterMove, record.PlayerColor, record.MoveTimeUTC,
if record.ClaimedBy != "" {
column := "white_claimed_by"
if record.ClaimColor == "b" {
column = "black_claimed_by"
}
query := `UPDATE games SET ` + column + ` = ?
WHERE game_id = ? AND (` + column + ` IS NULL OR ` + column + ` = '' OR ` + column + ` = ?)`
result, err := tx.Exec(query, record.ClaimedBy, record.Move.GameID, record.ClaimedBy)
if err != nil {
return err
}
if err := requireOneGame(result, record.Move.GameID); err != nil {
return err
}
}
if record.Result != "" {
result, err := tx.Exec(
`UPDATE games SET result = ?, end_time_utc = ? WHERE game_id = ?`,
record.Result, record.EndTimeUTC, record.Move.GameID,
)
if err != nil {
return err
}
return requireOneGame(result, record.Move.GameID)
}
return nil
})
}
// RecordGameResult persists a terminal transition not accompanied by a move,
// such as a no-legal-moves engine response.
func (s *Store) RecordGameResult(gameID, result string, at time.Time) error {
if gameID == "" {
return errors.New("game ID is required")
}
if !isValidResult(result) {
return fmt.Errorf("invalid game result %q", result)
}
if at.IsZero() {
return errors.New("game result time is required")
}
return s.enqueue("record_game_result", gameID, func(tx *sql.Tx) error {
res, err := tx.Exec(
`UPDATE games SET result = ?, end_time_utc = ? WHERE game_id = ?`,
result, at.UTC(), gameID,
)
return err
}:
return nil
default:
// Channel full, drop write
log.Printf("Storage write queue full, dropping move record")
return nil
}
if err != nil {
return err
}
return requireOneGame(res, gameID)
})
}
// DeleteUndoneMoves asynchronously deletes moves after undo
func (s *Store) DeleteUndoneMoves(gameID string, afterMoveNumber int) error {
if !s.healthStatus.Load() {
return nil // Silently drop if degraded
// RecordSlotClaim persists a claim made independently from a move.
func (s *Store) RecordSlotClaim(gameID, color, userID string) error {
if gameID == "" || userID == "" {
return errors.New("game ID and claimant are required")
}
select {
case s.writeChan <- func(tx *sql.Tx) error {
query := `DELETE FROM moves WHERE game_id = ? AND move_number > ?`
_, err := tx.Exec(query, gameID, afterMoveNumber)
return err
}:
return nil
default:
// Channel full, drop write
log.Printf("Storage write queue full, dropping undo operation")
return nil
if color != "w" && color != "b" {
return fmt.Errorf("invalid claim color %q", color)
}
column := "white_claimed_by"
if color == "b" {
column = "black_claimed_by"
}
return s.enqueue("record_slot_claim", gameID, func(tx *sql.Tx) error {
query := `UPDATE games SET ` + column + ` = ?
WHERE game_id = ? AND (` + column + ` IS NULL OR ` + column + ` = '' OR ` + column + ` = ?)`
res, err := tx.Exec(query, userID, gameID, userID)
if err != nil {
return err
}
return requireOneGame(res, gameID)
})
}
// QueryGames retrieves games with optional filtering
// RecordPlayers keeps persisted player configuration aligned with in-memory
// configuration changes.
func (s *Store) RecordPlayers(gameID string, white, black PlayerRecord) error {
if gameID == "" || white.PlayerID == "" || black.PlayerID == "" {
return errors.New("game ID and player IDs are required")
}
return s.enqueue("record_players", gameID, func(tx *sql.Tx) error {
const query = `UPDATE games SET
white_player_id = ?, white_type = ?, white_level = ?, white_search_time = ?, white_claimed_by = ?,
black_player_id = ?, black_type = ?, black_level = ?, black_search_time = ?, black_claimed_by = ?
WHERE game_id = ?`
res, err := tx.Exec(query,
white.PlayerID, white.Type, white.Level, white.SearchTime, nullableString(white.ClaimedBy),
black.PlayerID, black.Type, black.Level, black.SearchTime, nullableString(black.ClaimedBy),
gameID,
)
if err != nil {
return err
}
return requireOneGame(res, gameID)
})
}
// PlayerRecord is the persistence subset of a player configuration.
type PlayerRecord struct {
PlayerID string
Type int
Level int
SearchTime int
ClaimedBy string
}
// RewindGame atomically removes undone moves and clears a previously terminal
// result so replay readers never observe an ongoing line with a stale outcome.
func (s *Store) RewindGame(gameID string, afterMoveNumber int) error {
if gameID == "" || afterMoveNumber < 0 {
return errors.New("game ID and a non-negative move number are required")
}
return s.enqueue("rewind_game", gameID, func(tx *sql.Tx) error {
if _, err := tx.Exec(
`DELETE FROM moves WHERE game_id = ? AND move_number > ?`,
gameID, afterMoveNumber,
); err != nil {
return err
}
res, err := tx.Exec(
`UPDATE games SET result = NULL, end_time_utc = NULL WHERE game_id = ?`,
gameID,
)
if err != nil {
return err
}
return requireOneGame(res, gameID)
})
}
// QueryGames retrieves games with optional filtering. A player filter matches
// both creation-time player IDs and claims made after game creation.
func (s *Store) QueryGames(gameID, playerID string) ([]GameRecord, error) {
query := `SELECT
game_id, initial_fen,
white_player_id, white_type, white_level, white_search_time,
black_player_id, black_type, black_level, black_search_time,
start_time_utc
FROM games WHERE 1=1`
if err := s.flushBeforeRead(); err != nil {
return nil, err
}
started := time.Now()
query := `SELECT ` + gameSelectColumns + ` FROM games g WHERE 1=1`
var args []any
// Handle gameID filtering
if gameID != "" && gameID != "*" {
query += " AND game_id = ?"
query += " AND g.game_id = ?"
args = append(args, gameID)
}
// Handle playerID filtering
if playerID != "" && playerID != "*" {
query += " AND (white_player_id = ? OR black_player_id = ?)"
args = append(args, playerID, playerID)
query += ` AND (g.white_player_id = ? OR g.black_player_id = ?
OR g.white_claimed_by = ? OR g.black_claimed_by = ?)`
args = append(args, playerID, playerID, playerID, playerID)
}
query += " ORDER BY start_time_utc DESC"
query += " ORDER BY g.start_time_utc DESC, g.game_id DESC"
rows, err := s.db.Query(query, args...)
if err != nil {
return nil, fmt.Errorf("query failed: %w", err)
return nil, fmt.Errorf("query games: %w", err)
}
defer rows.Close()
var games []GameRecord
games := make([]GameRecord, 0)
for rows.Next() {
var g GameRecord
err := rows.Scan(
&g.GameID, &g.InitialFEN,
&g.WhitePlayerID, &g.WhiteType, &g.WhiteLevel, &g.WhiteSearchTime,
&g.BlackPlayerID, &g.BlackType, &g.BlackLevel, &g.BlackSearchTime,
&g.StartTimeUTC,
)
if err != nil {
return nil, fmt.Errorf("scan failed: %w", err)
var record GameRecord
if err := scanGame(rows, &record); err != nil {
return nil, fmt.Errorf("scan game: %w", err)
}
games = append(games, g)
games = append(games, record)
}
if err := rows.Err(); err != nil {
return nil, fmt.Errorf("rows iteration failed: %w", err)
return nil, fmt.Errorf("iterate games: %w", err)
}
slog.Debug("storage games queried", "count", len(games), "duration", time.Since(started))
return games, nil
}
}
func (s *Store) GetGameRecord(gameID string) (*GameRecord, error) {
if err := s.flushBeforeRead(); err != nil {
return nil, err
}
return getGameRecord(s.db, gameID)
}
type gameQueryer interface {
Query(query string, args ...any) (*sql.Rows, error)
QueryRow(query string, args ...any) *sql.Row
}
func getGameRecord(queryer gameQueryer, gameID string) (*GameRecord, error) {
var record GameRecord
row := queryer.QueryRow(`SELECT `+gameSelectColumns+` FROM games g WHERE g.game_id = ?`, gameID)
if err := scanGame(row, &record); err != nil {
return nil, err
}
return &record, nil
}
// GetMovesForGame returns the complete, undo-consistent replay line.
func (s *Store) GetMovesForGame(gameID string) ([]MoveRecord, error) {
if err := s.flushBeforeRead(); err != nil {
return nil, err
}
return getMovesForGame(s.db, gameID)
}
func getMovesForGame(queryer gameQueryer, gameID string) ([]MoveRecord, error) {
const query = `SELECT move_id, game_id, move_number, move_uci,
fen_after_move, player_color, move_time_utc
FROM moves WHERE game_id = ? ORDER BY move_number ASC`
rows, err := queryer.Query(query, gameID)
if err != nil {
return nil, fmt.Errorf("query game moves: %w", err)
}
defer rows.Close()
moves := make([]MoveRecord, 0)
for rows.Next() {
var move MoveRecord
if err := rows.Scan(
&move.MoveID, &move.GameID, &move.MoveNumber, &move.MoveUCI,
&move.FENAfterMove, &move.PlayerColor, &move.MoveTimeUTC,
); err != nil {
return nil, fmt.Errorf("scan game move: %w", err)
}
moves = append(moves, move)
}
if err := rows.Err(); err != nil {
return nil, fmt.Errorf("iterate game moves: %w", err)
}
return moves, nil
}
// GetGameHistory uses one write barrier and one read transaction for a
// consistent game-and-moves snapshot.
func (s *Store) GetGameHistory(gameID string) (*GameRecord, []MoveRecord, error) {
if err := s.flushBeforeRead(); err != nil {
return nil, nil, err
}
started := time.Now()
tx, err := s.db.BeginTx(context.Background(), &sql.TxOptions{ReadOnly: true})
if err != nil {
return nil, nil, fmt.Errorf("begin game history read: %w", err)
}
defer tx.Rollback()
record, err := getGameRecord(tx, gameID)
if err != nil {
return nil, nil, err
}
moves, err := getMovesForGame(tx, gameID)
if err != nil {
return nil, nil, err
}
if err := tx.Commit(); err != nil {
return nil, nil, fmt.Errorf("finish game history read: %w", err)
}
slog.Debug("storage game history queried",
"game_id", gameID, "move_count", len(moves), "duration", time.Since(started))
return record, moves, nil
}
func (s *Store) QueryGamesForUser(userID string, limit, offset int) ([]GameSummaryRecord, error) {
if userID == "" || limit < 1 || limit > 101 || offset < 0 {
return nil, errors.New("user ID, limit from 1 to 101, and non-negative offset are required")
}
if err := s.flushBeforeRead(); err != nil {
return nil, err
}
started := time.Now()
query := `SELECT ` + gameSelectColumns + `,
(SELECT COUNT(*) FROM moves m WHERE m.game_id = g.game_id) AS move_count
FROM games g
WHERE g.white_player_id = ? OR g.black_player_id = ?
OR g.white_claimed_by = ? OR g.black_claimed_by = ?
ORDER BY g.start_time_utc DESC, g.game_id DESC
LIMIT ? OFFSET ?`
rows, err := s.db.Query(query, userID, userID, userID, userID, limit, offset)
if err != nil {
return nil, fmt.Errorf("query user games: %w", err)
}
defer rows.Close()
games := make([]GameSummaryRecord, 0)
for rows.Next() {
var summary GameSummaryRecord
if err := scanGameSummary(rows, &summary); err != nil {
return nil, fmt.Errorf("scan user game: %w", err)
}
games = append(games, summary)
}
if err := rows.Err(); err != nil {
return nil, fmt.Errorf("iterate user games: %w", err)
}
slog.Debug("storage user games queried",
"user_id", userID,
"count", len(games),
"limit", limit,
"offset", offset,
"duration", time.Since(started),
)
return games, nil
}
type rowScanner interface {
Scan(dest ...any) error
}
func scanGame(scanner rowScanner, record *GameRecord) error {
var whiteClaimed, blackClaimed, result sql.NullString
var endTime sql.NullTime
if err := scanner.Scan(
&record.GameID, &record.InitialFEN,
&record.WhitePlayerID, &record.WhiteType, &record.WhiteLevel, &record.WhiteSearchTime, &whiteClaimed,
&record.BlackPlayerID, &record.BlackType, &record.BlackLevel, &record.BlackSearchTime, &blackClaimed,
&result, &record.StartTimeUTC, &endTime,
); err != nil {
return err
}
record.WhiteClaimedBy = whiteClaimed.String
record.BlackClaimedBy = blackClaimed.String
record.Result = result.String
if endTime.Valid {
ended := endTime.Time
record.EndTimeUTC = &ended
}
return nil
}
func scanGameSummary(scanner rowScanner, summary *GameSummaryRecord) error {
var whiteClaimed, blackClaimed, result sql.NullString
var endTime sql.NullTime
if err := scanner.Scan(
&summary.GameID, &summary.InitialFEN,
&summary.WhitePlayerID, &summary.WhiteType, &summary.WhiteLevel, &summary.WhiteSearchTime, &whiteClaimed,
&summary.BlackPlayerID, &summary.BlackType, &summary.BlackLevel, &summary.BlackSearchTime, &blackClaimed,
&result, &summary.StartTimeUTC, &endTime, &summary.MoveCount,
); err != nil {
return err
}
summary.WhiteClaimedBy = whiteClaimed.String
summary.BlackClaimedBy = blackClaimed.String
summary.Result = result.String
if endTime.Valid {
ended := endTime.Time
summary.EndTimeUTC = &ended
}
return nil
}
func requireOneGame(result sql.Result, gameID string) error {
rows, err := result.RowsAffected()
if err != nil {
return err
}
if rows != 1 {
return fmt.Errorf("game %s was not updated", gameID)
}
return nil
}
func nullableString(value string) any {
if value == "" {
return nil
}
return value
}
func isValidResult(result string) bool {
switch result {
case "white_wins", "black_wins", "draw", "stalemate":
return true
default:
return false
}
}
func validateResultTime(result string, ended *time.Time) error {
if result == "" {
if ended != nil {
return errors.New("end time requires a game result")
}
return nil
}
if !isValidResult(result) {
return fmt.Errorf("invalid game result %q", result)
}
if ended == nil || ended.IsZero() {
return errors.New("terminal game result requires an end time")
}
return nil
}
// IsGameNotFound keeps callers independent from database/sql details.
func IsGameNotFound(err error) bool {
return errors.Is(err, sql.ErrNoRows)
}
+52 -24
View File
@@ -24,17 +24,26 @@ type SessionRecord struct {
// GameRecord represents a row in the games table
type GameRecord struct {
GameID string `db:"game_id"`
InitialFEN string `db:"initial_fen"`
WhitePlayerID string `db:"white_player_id"`
WhiteType int `db:"white_type"`
WhiteLevel int `db:"white_level"`
WhiteSearchTime int `db:"white_search_time"`
BlackPlayerID string `db:"black_player_id"`
BlackType int `db:"black_type"`
BlackLevel int `db:"black_level"`
BlackSearchTime int `db:"black_search_time"`
StartTimeUTC time.Time `db:"start_time_utc"`
GameID string `db:"game_id"`
InitialFEN string `db:"initial_fen"`
WhitePlayerID string `db:"white_player_id"`
WhiteType int `db:"white_type"`
WhiteLevel int `db:"white_level"`
WhiteSearchTime int `db:"white_search_time"`
WhiteClaimedBy string `db:"white_claimed_by"`
BlackPlayerID string `db:"black_player_id"`
BlackType int `db:"black_type"`
BlackLevel int `db:"black_level"`
BlackSearchTime int `db:"black_search_time"`
BlackClaimedBy string `db:"black_claimed_by"`
Result string `db:"result"`
StartTimeUTC time.Time `db:"start_time_utc"`
EndTimeUTC *time.Time `db:"end_time_utc"`
}
type GameSummaryRecord struct {
GameRecord
MoveCount int `db:"move_count"`
}
// MoveRecord represents a row in the moves table
@@ -48,7 +57,19 @@ type MoveRecord struct {
MoveTimeUTC time.Time `db:"move_time_utc"`
}
// Schema defines the SQLite database structure
// MovePersistence groups changes caused by one accepted move so the move,
// first-move slot claim, and terminal result commit in one transaction.
type MovePersistence struct {
Move MoveRecord
ClaimColor string
ClaimedBy string
Result string
EndTimeUTC *time.Time
}
// Schema defines tables only. Indexes are applied after legacy column
// migrations so upgrading an older games table never references a missing
// column.
const Schema = `
CREATE TABLE IF NOT EXISTS users (
user_id TEXT PRIMARY KEY,
@@ -61,12 +82,6 @@ CREATE TABLE IF NOT EXISTS users (
last_login_at DATETIME
);
CREATE INDEX IF NOT EXISTS idx_users_username ON users(username);
CREATE INDEX IF NOT EXISTS idx_users_email ON users(email);
CREATE INDEX IF NOT EXISTS idx_users_account_type ON users(account_type);
CREATE INDEX IF NOT EXISTS idx_users_expires_at ON users(expires_at);
CREATE UNIQUE INDEX IF NOT EXISTS idx_users_email_unique ON users(email) WHERE email IS NOT NULL AND email != '';
CREATE TABLE IF NOT EXISTS sessions (
session_id TEXT PRIMARY KEY,
user_id TEXT NOT NULL UNIQUE,
@@ -75,9 +90,6 @@ CREATE TABLE IF NOT EXISTS sessions (
FOREIGN KEY (user_id) REFERENCES users(user_id) ON DELETE CASCADE
);
CREATE INDEX IF NOT EXISTS idx_sessions_user_id ON sessions(user_id);
CREATE INDEX IF NOT EXISTS idx_sessions_expires_at ON sessions(expires_at);
CREATE TABLE IF NOT EXISTS games (
game_id TEXT PRIMARY KEY,
initial_fen TEXT NOT NULL,
@@ -89,7 +101,11 @@ CREATE TABLE IF NOT EXISTS games (
black_type INTEGER NOT NULL,
black_level INTEGER NOT NULL DEFAULT 0,
black_search_time INTEGER NOT NULL DEFAULT 1000,
start_time_utc DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
start_time_utc DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
result TEXT CHECK(result IS NULL OR result IN ('white_wins', 'black_wins', 'draw', 'stalemate')),
end_time_utc DATETIME,
white_claimed_by TEXT,
black_claimed_by TEXT
);
CREATE TABLE IF NOT EXISTS moves (
@@ -103,8 +119,20 @@ CREATE TABLE IF NOT EXISTS moves (
FOREIGN KEY (game_id) REFERENCES games(game_id) ON DELETE CASCADE,
UNIQUE(game_id, move_number)
);
`
CREATE INDEX IF NOT EXISTS idx_moves_game_id ON moves(game_id);
const Indexes = `
CREATE UNIQUE INDEX IF NOT EXISTS idx_users_email_unique
ON users(email) WHERE email IS NOT NULL AND email != '';
CREATE INDEX IF NOT EXISTS idx_users_temp_created_at
ON users(created_at) WHERE account_type = 'temp';
CREATE INDEX IF NOT EXISTS idx_users_temp_expires_at
ON users(expires_at) WHERE account_type = 'temp' AND expires_at IS NOT NULL;
CREATE INDEX IF NOT EXISTS idx_sessions_expires_at ON sessions(expires_at);
CREATE INDEX IF NOT EXISTS idx_games_white_player ON games(white_player_id);
CREATE INDEX IF NOT EXISTS idx_games_black_player ON games(black_player_id);
`
CREATE INDEX IF NOT EXISTS idx_games_white_claimed ON games(white_claimed_by)
WHERE white_claimed_by IS NOT NULL;
CREATE INDEX IF NOT EXISTS idx_games_black_claimed ON games(black_claimed_by)
WHERE black_claimed_by IS NOT NULL;
`
+44 -23
View File
@@ -2,30 +2,23 @@ package storage
import (
"fmt"
"log/slog"
"time"
)
// CreateSession creates or replaces the session for a user (single session per user)
func (s *Store) CreateSession(record SessionRecord) error {
tx, err := s.db.Begin()
if err != nil {
return fmt.Errorf("failed to begin transaction: %w", err)
}
defer tx.Rollback()
// Delete any existing session for this user
deleteQuery := `DELETE FROM sessions WHERE user_id = ?`
if _, err := tx.Exec(deleteQuery, record.UserID); err != nil {
return fmt.Errorf("failed to delete existing session: %w", err)
}
// Insert new session
insertQuery := `INSERT INTO sessions (session_id, user_id, created_at, expires_at) VALUES (?, ?, ?, ?)`
if _, err := tx.Exec(insertQuery, record.SessionID, record.UserID, record.CreatedAt, record.ExpiresAt); err != nil {
const query = `INSERT INTO sessions (session_id, user_id, created_at, expires_at)
VALUES (?, ?, ?, ?)
ON CONFLICT(user_id) DO UPDATE SET
session_id = excluded.session_id,
created_at = excluded.created_at,
expires_at = excluded.expires_at`
if _, err := s.db.Exec(query, record.SessionID, record.UserID, record.CreatedAt, record.ExpiresAt); err != nil {
return fmt.Errorf("failed to create session: %w", err)
}
return tx.Commit()
slog.Debug("storage session created", "user_id", record.UserID, "expires_at", record.ExpiresAt)
return nil
}
// GetSession retrieves a session by ID
@@ -60,6 +53,9 @@ func (s *Store) GetSessionByUserID(userID string) (*SessionRecord, error) {
func (s *Store) DeleteSession(sessionID string) error {
query := `DELETE FROM sessions WHERE session_id = ?`
_, err := s.db.Exec(query, sessionID)
if err == nil {
slog.Debug("storage session deleted")
}
return err
}
@@ -67,6 +63,9 @@ func (s *Store) DeleteSession(sessionID string) error {
func (s *Store) DeleteSessionByUserID(userID string) error {
query := `DELETE FROM sessions WHERE user_id = ?`
_, err := s.db.Exec(query, userID)
if err == nil {
slog.Debug("storage user sessions deleted", "user_id", userID)
}
return err
}
@@ -77,16 +76,38 @@ func (s *Store) DeleteExpiredSessions() (int64, error) {
if err != nil {
return 0, err
}
return result.RowsAffected()
deleted, err := result.RowsAffected()
if err == nil && deleted > 0 {
slog.Debug("storage expired sessions deleted", "count", deleted)
}
return deleted, err
}
// IsSessionValid checks if a session exists and is not expired
func (s *Store) IsSessionValid(sessionID string) (bool, error) {
var count int
query := `SELECT COUNT(*) FROM sessions WHERE session_id = ? AND expires_at > ?`
err := s.db.QueryRow(query, sessionID, time.Now().UTC()).Scan(&count)
var valid bool
const query = `SELECT EXISTS(
SELECT 1 FROM sessions WHERE session_id = ? AND expires_at > ?
)`
err := s.db.QueryRow(query, sessionID, time.Now().UTC()).Scan(&valid)
if err != nil {
return false, err
}
return count > 0, nil
}
return valid, nil
}
// IsSessionValidForUser verifies both expiry and the binding between a JWT
// subject and its persisted session. Checking only the session ID would allow
// a malformed server-issued token to authenticate as the wrong subject.
func (s *Store) IsSessionValidForUser(sessionID, userID string) (bool, error) {
var valid bool
const query = `SELECT EXISTS(
SELECT 1 FROM sessions
WHERE session_id = ? AND user_id = ? AND expires_at > ?
)`
err := s.db.QueryRow(query, sessionID, userID, time.Now().UTC()).Scan(&valid)
if err != nil {
return false, err
}
return valid, nil
}
+305 -65
View File
@@ -3,9 +3,11 @@ package storage
import (
"context"
"database/sql"
"errors"
"fmt"
"log"
"log/slog"
"os"
"strings"
"sync"
"sync/atomic"
"time"
@@ -13,48 +15,67 @@ import (
_ "github.com/mattn/go-sqlite3"
)
const (
writeQueueCapacity = 1000
flushTimeout = 5 * time.Second
schemaVersion = 2
)
var memoryStoreCounter atomic.Uint64
var (
ErrStorageDegraded = errors.New("storage is degraded")
ErrStoreClosed = errors.New("storage is closed")
ErrWriteQueueFull = errors.New("storage write queue is full")
)
type writeRequest struct {
operation string
gameID string
run func(*sql.Tx) error
barrier chan error
}
// Store handles SQLite database operations with async writes for games and sync writes for auth
type Store struct {
db *sql.DB
path string
writeChan chan func(*sql.Tx) error
writeChan chan writeRequest
healthStatus atomic.Bool
writeFailed atomic.Bool
closed atomic.Bool
ctx context.Context
cancel context.CancelFunc
wg sync.WaitGroup
enqueueMu sync.RWMutex
closeOnce sync.Once
closeErr error
}
// NewStore creates a new storage instance with async writer
func NewStore(dataSourceName string, devMode bool) (*Store, error) {
db, err := sql.Open("sqlite3", dataSourceName)
dsn := sqliteDSN(dataSourceName)
db, err := sql.Open("sqlite3", dsn)
if err != nil {
return nil, fmt.Errorf("failed to open database: %w", err)
}
// Enable WAL mode in development for better concurrency
if devMode {
if _, err := db.Exec("PRAGMA journal_mode=WAL"); err != nil {
db.Close()
return nil, fmt.Errorf("failed to enable WAL mode: %w", err)
}
}
// Enable foreign keys
if _, err := db.Exec("PRAGMA foreign_keys = ON"); err != nil {
if err := db.Ping(); err != nil {
db.Close()
return nil, fmt.Errorf("failed to enable foreign keys: %w", err)
return nil, fmt.Errorf("failed to connect to database: %w", err)
}
// Configure connection pool
db.SetMaxOpenConns(25)
db.SetMaxIdleConns(5)
// SQLite benefits from a small pool. WAL and busy_timeout are configured in
// the DSN for every connection, unlike connection-local PRAGMA calls.
db.SetMaxOpenConns(8)
db.SetMaxIdleConns(4)
ctx, cancel := context.WithCancel(context.Background())
s := &Store{
db: db,
path: dataSourceName,
writeChan: make(chan func(*sql.Tx) error, 1000), // Buffered for async writes
writeChan: make(chan writeRequest, writeQueueCapacity),
ctx: ctx,
cancel: cancel,
}
@@ -65,10 +86,34 @@ func NewStore(dataSourceName string, devMode bool) (*Store, error) {
// Start async writer
s.wg.Add(1)
go s.writerLoop()
slog.Debug("storage opened",
"path", dataSourceName,
"dev_mode", devMode,
"write_queue_capacity", writeQueueCapacity,
"max_open_connections", 8,
)
return s, nil
}
func sqliteDSN(dataSourceName string) string {
if dataSourceName == ":memory:" {
// Each Store needs a private shared-cache database: shared cache keeps
// that Store's pooled connections on one database, while the unique name
// prevents independent in-memory stores from leaking into each other.
dataSourceName = fmt.Sprintf(
"file:chess-memory-%d?mode=memory&cache=shared",
memoryStoreCounter.Add(1),
)
}
separator := "?"
if strings.Contains(dataSourceName, "?") {
separator = "&"
}
return dataSourceName + separator +
"_foreign_keys=on&_busy_timeout=5000&_journal_mode=WAL&_synchronous=NORMAL"
}
// IsHealthy returns true if the storage is operational
func (s *Store) IsHealthy() bool {
return s.healthStatus.Load()
@@ -77,96 +122,291 @@ func (s *Store) IsHealthy() bool {
// writerLoop processes async write operations
func (s *Store) writerLoop() {
defer s.wg.Done()
slog.Debug("storage writer started")
defer slog.Debug("storage writer stopped")
for {
select {
case <-s.ctx.Done():
// Drain remaining writes with timeout
deadline := time.After(2 * time.Second)
// Every accepted operation is drained before shutdown. The queue is
// bounded, so shutdown remains bounded by actual database work rather
// than an arbitrary timer that can discard replay history.
for {
select {
case fn := <-s.writeChan:
if s.healthStatus.Load() {
s.executeWrite(fn)
}
case <-deadline:
return
case req := <-s.writeChan:
s.handleWrite(req)
default:
return
}
}
case fn := <-s.writeChan:
// Skip if already degraded
if !s.healthStatus.Load() {
continue
}
s.executeWrite(fn)
case req := <-s.writeChan:
s.handleWrite(req)
}
}
}
// executeWrite runs a transactional write operation
func (s *Store) executeWrite(fn func(*sql.Tx) error) {
tx, err := s.db.Begin()
if err != nil {
log.Printf("Storage degraded: failed to begin transaction: %v", err)
s.healthStatus.Store(false)
func (s *Store) handleWrite(req writeRequest) {
if req.run == nil {
if req.barrier != nil {
var err error
if !s.healthStatus.Load() {
err = ErrStorageDegraded
}
req.barrier <- err
close(req.barrier)
}
return
}
if err := fn(tx); err != nil {
tx.Rollback()
log.Printf("Storage degraded: write operation failed: %v", err)
s.healthStatus.Store(false)
if s.writeFailed.Load() {
slog.Error("storage write skipped after earlier transaction failure",
"operation", req.operation, "game_id", req.gameID)
if req.barrier != nil {
req.barrier <- ErrStorageDegraded
close(req.barrier)
}
return
}
err := s.executeWrite(req)
if req.barrier != nil {
req.barrier <- err
close(req.barrier)
}
}
// executeWrite runs a transactional write operation
func (s *Store) executeWrite(req writeRequest) error {
started := time.Now()
tx, err := s.db.Begin()
if err != nil {
s.writeFailed.Store(true)
s.healthStatus.Store(false)
slog.Error("storage degraded: failed to begin transaction",
"operation", req.operation, "game_id", req.gameID, "error", err)
return err
}
if err := req.run(tx); err != nil {
rollbackErr := tx.Rollback()
s.writeFailed.Store(true)
s.healthStatus.Store(false)
slog.Error("storage degraded: write operation failed",
"operation", req.operation,
"game_id", req.gameID,
"error", err,
"rollback_error", rollbackErr,
)
return err
}
if err := tx.Commit(); err != nil {
log.Printf("Storage degraded: failed to commit: %v", err)
s.writeFailed.Store(true)
s.healthStatus.Store(false)
return
slog.Error("storage degraded: failed to commit",
"operation", req.operation, "game_id", req.gameID, "error", err)
return err
}
slog.Debug("storage write committed",
"operation", req.operation,
"game_id", req.gameID,
"duration", time.Since(started),
"queue_depth", len(s.writeChan),
)
return nil
}
func (s *Store) enqueue(operation, gameID string, fn func(*sql.Tx) error) error {
s.enqueueMu.RLock()
defer s.enqueueMu.RUnlock()
if s.closed.Load() {
return ErrStoreClosed
}
if !s.healthStatus.Load() {
return ErrStorageDegraded
}
select {
case s.writeChan <- writeRequest{operation: operation, gameID: gameID, run: fn}:
slog.Debug("storage write queued",
"operation", operation,
"game_id", gameID,
"queue_depth", len(s.writeChan),
)
return nil
default:
s.healthStatus.Store(false)
slog.Error("storage degraded: write queue full",
"operation", operation,
"game_id", gameID,
"queue_capacity", cap(s.writeChan),
)
return ErrWriteQueueFull
}
}
// Flush waits until every write queued before this call has completed. Replay
// reads use this barrier to provide read-after-write consistency while normal
// gameplay retains the low-latency async write path.
func (s *Store) Flush(ctx context.Context) error {
s.enqueueMu.RLock()
if s.closed.Load() {
s.enqueueMu.RUnlock()
return ErrStoreClosed
}
if !s.healthStatus.Load() {
s.enqueueMu.RUnlock()
return ErrStorageDegraded
}
done := make(chan error, 1)
select {
case s.writeChan <- writeRequest{operation: "flush", barrier: done}:
s.enqueueMu.RUnlock()
case <-ctx.Done():
s.enqueueMu.RUnlock()
return ctx.Err()
}
select {
case err := <-done:
return err
case <-ctx.Done():
return ctx.Err()
}
}
func (s *Store) flushBeforeRead() error {
ctx, cancel := context.WithTimeout(context.Background(), flushTimeout)
defer cancel()
return s.Flush(ctx)
}
// Close gracefully closes the database connection
func (s *Store) Close() error {
// Signal writer to stop
s.cancel()
s.closeOnce.Do(func() {
s.enqueueMu.Lock()
s.closed.Store(true)
s.cancel()
s.enqueueMu.Unlock()
// Wait for writer with timeout
done := make(chan struct{})
go func() {
s.wg.Wait()
close(done)
}()
select {
case <-done:
// Writer finished cleanly
case <-time.After(2 * time.Second):
log.Printf("Warning: storage writer shutdown timeout, some writes may be lost")
}
if s.db != nil {
return s.db.Close()
}
return nil
if s.db != nil {
s.closeErr = s.db.Close()
}
})
return s.closeErr
}
// InitDB creates the database schema
func (s *Store) InitDB() error {
started := time.Now()
tx, err := s.db.Begin()
if err != nil {
return fmt.Errorf("failed to begin transaction: %w", err)
}
defer tx.Rollback()
var currentVersion int
if err := tx.QueryRow("PRAGMA user_version").Scan(&currentVersion); err != nil {
return fmt.Errorf("failed to read schema version: %w", err)
}
if currentVersion > schemaVersion {
return fmt.Errorf(
"database schema version %d is newer than supported version %d",
currentVersion,
schemaVersion,
)
}
if _, err := tx.Exec(Schema); err != nil {
return fmt.Errorf("failed to create schema: %w", err)
}
return tx.Commit()
columns := []struct {
name string
definition string
}{
{"result", "TEXT CHECK(result IS NULL OR result IN ('white_wins', 'black_wins', 'draw', 'stalemate'))"},
{"end_time_utc", "DATETIME"},
{"white_claimed_by", "TEXT"},
{"black_claimed_by", "TEXT"},
}
for _, column := range columns {
if err := ensureColumn(tx, "games", column.name, column.definition); err != nil {
return err
}
}
// These indexes duplicate UNIQUE constraints or are superseded by targeted
// partial/composite indexes. Drop them during upgrades as well as omitting
// them from new databases.
for _, name := range []string{
"idx_users_username",
"idx_users_email",
"idx_users_account_type",
"idx_users_expires_at",
"idx_sessions_user_id",
"idx_moves_game_id",
"idx_games_finished_end_time",
} {
if _, err := tx.Exec("DROP INDEX IF EXISTS " + name); err != nil {
return fmt.Errorf("failed to remove redundant index %s: %w", name, err)
}
}
if _, err := tx.Exec(Indexes); err != nil {
return fmt.Errorf("failed to create indexes: %w", err)
}
if _, err := tx.Exec(fmt.Sprintf("PRAGMA user_version = %d", schemaVersion)); err != nil {
return fmt.Errorf("failed to record schema version: %w", err)
}
if err := tx.Commit(); err != nil {
return fmt.Errorf("failed to commit schema: %w", err)
}
slog.Debug("storage schema ready", "version", schemaVersion, "duration", time.Since(started))
return nil
}
func ensureColumn(tx *sql.Tx, table, column, definition string) error {
rows, err := tx.Query("PRAGMA table_info(" + table + ")")
if err != nil {
return fmt.Errorf("failed to inspect %s schema: %w", table, err)
}
found := false
for rows.Next() {
var cid int
var name, columnType string
var notNull, primaryKey int
var defaultValue any
if err := rows.Scan(&cid, &name, &columnType, &notNull, &defaultValue, &primaryKey); err != nil {
rows.Close()
return fmt.Errorf("failed to inspect %s column: %w", table, err)
}
if name == column {
found = true
}
}
if err := rows.Close(); err != nil {
return fmt.Errorf("failed to close %s schema rows: %w", table, err)
}
if err := rows.Err(); err != nil {
return fmt.Errorf("failed to inspect %s schema: %w", table, err)
}
if found {
return nil
}
if _, err := tx.Exec("ALTER TABLE " + table + " ADD COLUMN " + column + " " + definition); err != nil {
return fmt.Errorf("failed to add %s.%s: %w", table, column, err)
}
slog.Debug("storage schema column added", "table", table, "column", column)
return nil
}
// DeleteDB removes the database file
@@ -182,4 +422,4 @@ func (s *Store) DeleteDB() error {
}
return nil
}
}
+525
View File
@@ -0,0 +1,525 @@
package storage
import (
"context"
"database/sql"
"errors"
"path/filepath"
"strings"
"testing"
"time"
_ "github.com/mattn/go-sqlite3"
)
func TestInitDBMigratesLegacySchemaAndRemovesRedundantIndexes(t *testing.T) {
path := filepath.Join(t.TempDir(), "legacy.db")
db, err := sql.Open("sqlite3", path)
if err != nil {
t.Fatal(err)
}
const legacy = `
CREATE TABLE users (
user_id TEXT PRIMARY KEY,
username TEXT UNIQUE NOT NULL COLLATE NOCASE,
email TEXT COLLATE NOCASE,
password_hash TEXT NOT NULL,
account_type TEXT NOT NULL DEFAULT 'temp',
created_at DATETIME NOT NULL,
expires_at DATETIME,
last_login_at DATETIME
);
CREATE TABLE sessions (
session_id TEXT PRIMARY KEY,
user_id TEXT NOT NULL UNIQUE,
created_at DATETIME NOT NULL,
expires_at DATETIME NOT NULL,
FOREIGN KEY (user_id) REFERENCES users(user_id) ON DELETE CASCADE
);
CREATE TABLE games (
game_id TEXT PRIMARY KEY,
initial_fen TEXT NOT NULL,
white_player_id TEXT NOT NULL,
white_type INTEGER NOT NULL,
white_level INTEGER NOT NULL DEFAULT 0,
white_search_time INTEGER NOT NULL DEFAULT 1000,
black_player_id TEXT NOT NULL,
black_type INTEGER NOT NULL,
black_level INTEGER NOT NULL DEFAULT 0,
black_search_time INTEGER NOT NULL DEFAULT 1000,
start_time_utc DATETIME NOT NULL
);
CREATE TABLE moves (
move_id INTEGER PRIMARY KEY AUTOINCREMENT,
game_id TEXT NOT NULL,
move_number INTEGER NOT NULL,
move_uci TEXT NOT NULL,
fen_after_move TEXT NOT NULL,
player_color TEXT NOT NULL,
move_time_utc DATETIME NOT NULL,
FOREIGN KEY (game_id) REFERENCES games(game_id) ON DELETE CASCADE,
UNIQUE(game_id, move_number)
);
CREATE INDEX idx_users_username ON users(username);
CREATE INDEX idx_users_email ON users(email);
CREATE INDEX idx_users_account_type ON users(account_type);
CREATE INDEX idx_users_expires_at ON users(expires_at);
CREATE INDEX idx_sessions_user_id ON sessions(user_id);
CREATE INDEX idx_moves_game_id ON moves(game_id);`
if _, err := db.Exec(legacy); err != nil {
t.Fatalf("create legacy schema: %v", err)
}
if err := db.Close(); err != nil {
t.Fatal(err)
}
store, err := NewStore(path, false)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = store.Close() })
if err := store.InitDB(); err != nil {
t.Fatalf("migrate schema: %v", err)
}
columns := tableColumns(t, store.db, "games")
for _, column := range []string{"result", "end_time_utc", "white_claimed_by", "black_claimed_by"} {
if !columns[column] {
t.Errorf("migration did not add games.%s", column)
}
}
indexes := schemaIndexes(t, store.db)
for _, obsolete := range []string{
"idx_users_username", "idx_users_email", "idx_users_account_type",
"idx_users_expires_at", "idx_sessions_user_id", "idx_moves_game_id",
"idx_games_finished_end_time",
} {
if indexes[obsolete] {
t.Errorf("redundant index %s remains", obsolete)
}
}
for _, required := range []string{
"idx_users_email_unique", "idx_users_temp_created_at", "idx_users_temp_expires_at",
"idx_sessions_expires_at", "idx_games_white_player", "idx_games_black_player",
"idx_games_white_claimed", "idx_games_black_claimed",
} {
if !indexes[required] {
t.Errorf("required index %s is missing", required)
}
}
var version, foreignKeys int
if err := store.db.QueryRow("PRAGMA user_version").Scan(&version); err != nil {
t.Fatal(err)
}
if version != 2 {
t.Errorf("schema version = %d, want 2", version)
}
if err := store.db.QueryRow("PRAGMA foreign_keys").Scan(&foreignKeys); err != nil {
t.Fatal(err)
}
if foreignKeys != 1 {
t.Errorf("foreign_keys = %d, want 1", foreignKeys)
}
}
func TestReplayPersistenceIsAtomicAndReadAfterWriteConsistent(t *testing.T) {
store := newTestStore(t)
started := time.Date(2026, 9, 7, 1, 2, 3, 0, time.UTC)
ended := started.Add(5 * time.Minute)
if err := store.RecordNewGame(GameRecord{
GameID: "game-1", InitialFEN: "initial",
WhitePlayerID: "anonymous-white", WhiteType: 1,
BlackPlayerID: "black", BlackType: 1,
StartTimeUTC: started,
}); err != nil {
t.Fatal(err)
}
if err := store.RecordMove(MovePersistence{
Move: MoveRecord{
GameID: "game-1", MoveNumber: 1, MoveUCI: "e2e4",
FENAfterMove: "after-e2e4", PlayerColor: "w", MoveTimeUTC: ended,
},
ClaimColor: "w", ClaimedBy: "user-1",
Result: "white_wins", EndTimeUTC: &ended,
}); err != nil {
t.Fatal(err)
}
// GetGameHistory must observe both queued writes without sleeps or polling.
gameRecord, moves, err := store.GetGameHistory("game-1")
if err != nil {
t.Fatal(err)
}
if gameRecord.WhiteClaimedBy != "user-1" || gameRecord.Result != "white_wins" || gameRecord.EndTimeUTC == nil {
t.Fatalf("durable game mutation incomplete: %+v", gameRecord)
}
if len(moves) != 1 || moves[0].MoveNumber != 1 || moves[0].FENAfterMove != "after-e2e4" {
t.Fatalf("moves = %+v", moves)
}
owned, err := store.QueryGamesForUser("user-1", 10, 0)
if err != nil {
t.Fatal(err)
}
if len(owned) != 1 || owned[0].MoveCount != 1 {
t.Fatalf("claimed game lookup = %+v", owned)
}
if err := store.RewindGame("game-1", 0); err != nil {
t.Fatal(err)
}
gameRecord, moves, err = store.GetGameHistory("game-1")
if err != nil {
t.Fatal(err)
}
if gameRecord.Result != "" || gameRecord.EndTimeUTC != nil || len(moves) != 0 {
t.Fatalf("rewind left stale replay data: game=%+v moves=%+v", gameRecord, moves)
}
}
func TestConnectionSettingsApplyAcrossPool(t *testing.T) {
store := newTestStore(t)
ctx := context.Background()
connections := make([]*sql.Conn, 0, 8)
defer func() {
for _, connection := range connections {
_ = connection.Close()
}
}()
// Keep each connection checked out so the pool must create eight distinct
// SQLite connections, then verify connection-local PRAGMAs on every one.
for range 8 {
connection, err := store.db.Conn(ctx)
if err != nil {
t.Fatal(err)
}
connections = append(connections, connection)
}
for i, connection := range connections {
var foreignKeys, busyTimeout, synchronous int
var journalMode string
if err := connection.QueryRowContext(ctx, "PRAGMA foreign_keys").Scan(&foreignKeys); err != nil {
t.Fatal(err)
}
if err := connection.QueryRowContext(ctx, "PRAGMA busy_timeout").Scan(&busyTimeout); err != nil {
t.Fatal(err)
}
if err := connection.QueryRowContext(ctx, "PRAGMA synchronous").Scan(&synchronous); err != nil {
t.Fatal(err)
}
if err := connection.QueryRowContext(ctx, "PRAGMA journal_mode").Scan(&journalMode); err != nil {
t.Fatal(err)
}
if foreignKeys != 1 || busyTimeout != 5000 || synchronous != 1 || !strings.EqualFold(journalMode, "wal") {
t.Errorf(
"connection %d settings: foreign_keys=%d busy_timeout=%d synchronous=%d journal_mode=%s",
i, foreignKeys, busyTimeout, synchronous, journalMode,
)
}
}
}
func TestAcceptedWriteCommitsAfterQueueAdmissionFailureMarksHealthDegraded(t *testing.T) {
store := newTestStore(t)
store.healthStatus.Store(false) // Queue saturation rejects new work but accepted work must drain.
done := make(chan error, 1)
store.handleWrite(writeRequest{
operation: "accepted_before_saturation",
run: func(tx *sql.Tx) error {
_, err := tx.Exec(`INSERT INTO users
(user_id, username, password_hash, account_type, created_at)
VALUES ('user-1', 'alice', 'hash', 'permanent', ?)`, time.Now().UTC())
return err
},
barrier: done,
})
if err := <-done; err != nil {
t.Fatalf("accepted write was discarded: %v", err)
}
if _, err := store.GetUserByID("user-1"); err != nil {
t.Fatalf("accepted write was not committed: %v", err)
}
}
func TestWritesAfterTransactionFailureAreSkippedExplicitly(t *testing.T) {
store := newTestStore(t)
failed := make(chan error, 1)
store.handleWrite(writeRequest{
operation: "forced_failure",
run: func(*sql.Tx) error { return errors.New("forced failure") },
barrier: failed,
})
if err := <-failed; err == nil {
t.Fatal("forced transaction failure was not reported")
}
ran := false
skipped := make(chan error, 1)
store.handleWrite(writeRequest{
operation: "after_failure",
run: func(*sql.Tx) error {
ran = true
return nil
},
barrier: skipped,
})
if err := <-skipped; !errors.Is(err, ErrStorageDegraded) {
t.Fatalf("skipped write error = %v, want ErrStorageDegraded", err)
}
if ran {
t.Fatal("write ran after an earlier transaction broke ordering")
}
}
func TestNewerSchemaVersionIsRejected(t *testing.T) {
store := newTestStore(t)
if _, err := store.db.Exec("PRAGMA user_version = 3"); err != nil {
t.Fatal(err)
}
if err := store.InitDB(); err == nil || !strings.Contains(err.Error(), "newer than supported") {
t.Fatalf("InitDB error = %v, want newer-version rejection", err)
}
var version int
if err := store.db.QueryRow("PRAGMA user_version").Scan(&version); err != nil {
t.Fatal(err)
}
if version != 3 {
t.Fatalf("newer schema version was overwritten: %d", version)
}
}
func TestQueryPlansUseOnlyPurposeBuiltOrConstraintIndexes(t *testing.T) {
store := newTestStore(t)
now := time.Now().UTC()
tests := []struct {
name string
query string
args []any
want []string
}{
{
name: "email partial uniqueness",
query: `SELECT user_id FROM users
WHERE email = ? COLLATE NOCASE AND email IS NOT NULL AND email != ''`,
args: []any{"alice@example.com"}, want: []string{"idx_users_email_unique"},
},
{
name: "temporary expiry cleanup",
query: `SELECT user_id FROM users
WHERE account_type = 'temp' AND expires_at IS NOT NULL AND expires_at < ?`,
args: []any{now}, want: []string{"idx_users_temp_expires_at"},
},
{
name: "oldest temporary account",
query: `SELECT user_id FROM users
WHERE account_type = 'temp' ORDER BY created_at ASC LIMIT 1`,
want: []string{"idx_users_temp_created_at"},
},
{
name: "ordered moves use composite unique constraint",
query: `SELECT move_uci FROM moves
WHERE game_id = ? ORDER BY move_number ASC`,
args: []any{"game-1"}, want: []string{"sqlite_autoindex_moves_1"},
},
{
name: "all user association branches",
query: `SELECT game_id,
(SELECT COUNT(*) FROM moves m WHERE m.game_id = games.game_id)
FROM games WHERE white_player_id = ? OR black_player_id = ?
OR white_claimed_by = ? OR black_claimed_by = ?
ORDER BY start_time_utc DESC, game_id DESC LIMIT ? OFFSET ?`,
args: []any{"user-1", "user-1", "user-1", "user-1", 50, 0},
want: []string{
"idx_games_white_player", "idx_games_black_player",
"idx_games_white_claimed", "idx_games_black_claimed",
},
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
plan := explainQueryPlan(t, store.db, test.query, test.args...)
for _, index := range test.want {
if !strings.Contains(plan, index) {
t.Errorf("query plan does not use %s:\n%s", index, plan)
}
}
})
}
}
func TestForeignKeyCascadeAppliesToSessions(t *testing.T) {
store := newTestStore(t)
now := time.Now().UTC()
if err := store.CreateUser(UserRecord{
UserID: "user-1", Username: "user1", PasswordHash: "hash",
AccountType: "permanent", CreatedAt: now,
}); err != nil {
t.Fatal(err)
}
if err := store.CreateSession(SessionRecord{
SessionID: "session-1", UserID: "user-1", CreatedAt: now, ExpiresAt: now.Add(time.Hour),
}); err != nil {
t.Fatal(err)
}
if err := store.DeleteUser("user-1"); err != nil {
t.Fatal(err)
}
if _, err := store.GetSession("session-1"); !errors.Is(err, sql.ErrNoRows) {
t.Fatalf("session survived user cascade: %v", err)
}
}
func TestLimitedUserCreationIsAtomicWithInitialSession(t *testing.T) {
store := newTestStore(t)
now := time.Now().UTC()
limits := UserLimits{MaxUsers: 1, PermanentSlots: 1}
first := UserRecord{
UserID: "user-1", Username: "alice", PasswordHash: "hash",
AccountType: "temp", CreatedAt: now, ExpiresAt: timePointer(now.Add(time.Hour)),
}
firstSession := SessionRecord{
SessionID: "session-1", UserID: first.UserID,
CreatedAt: now, ExpiresAt: now.Add(time.Hour),
}
if err := store.CreateUserWithinLimits(first, &firstSession, limits); err != nil {
t.Fatal(err)
}
duplicate := first
duplicate.UserID = "user-duplicate"
if err := store.CreateUserWithinLimits(duplicate, nil, limits); !errors.Is(err, ErrUserAlreadyExists) {
t.Fatalf("duplicate error = %v, want ErrUserAlreadyExists", err)
}
if _, err := store.GetUserByID(first.UserID); err != nil {
t.Fatalf("duplicate registration evicted existing user: %v", err)
}
second := UserRecord{
UserID: "user-2", Username: "bob", PasswordHash: "hash",
AccountType: "temp", CreatedAt: now.Add(time.Minute),
ExpiresAt: timePointer(now.Add(2 * time.Hour)),
}
secondSession := SessionRecord{
SessionID: "session-2", UserID: second.UserID,
CreatedAt: now.Add(time.Minute), ExpiresAt: now.Add(2 * time.Hour),
}
if err := store.CreateUserWithinLimits(second, &secondSession, limits); err != nil {
t.Fatal(err)
}
if _, err := store.GetUserByID(first.UserID); !errors.Is(err, sql.ErrNoRows) {
t.Fatalf("oldest temporary user was not replaced: %v", err)
}
if _, err := store.GetSession(firstSession.SessionID); !errors.Is(err, sql.ErrNoRows) {
t.Fatalf("evicted user's session survived cascade: %v", err)
}
if _, err := store.GetSession(secondSession.SessionID); err != nil {
t.Fatalf("initial session not committed with user: %v", err)
}
third := UserRecord{
UserID: "user-3", Username: "charlie", PasswordHash: "hash",
AccountType: "temp", CreatedAt: now.Add(2 * time.Minute),
}
conflictingSession := SessionRecord{
SessionID: secondSession.SessionID, UserID: third.UserID,
CreatedAt: now, ExpiresAt: now.Add(time.Hour),
}
wideLimits := UserLimits{MaxUsers: 10, PermanentSlots: 2}
if err := store.CreateUserWithinLimits(third, &conflictingSession, wideLimits); err == nil {
t.Fatal("expected duplicate session failure")
}
if _, err := store.GetUserByID(third.UserID); !errors.Is(err, sql.ErrNoRows) {
t.Fatalf("session failure did not roll back user: %v", err)
}
}
func timePointer(value time.Time) *time.Time {
return &value
}
func newTestStore(t *testing.T) *Store {
t.Helper()
store, err := NewStore(filepath.Join(t.TempDir(), "chess.db"), false)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = store.Close() })
if err := store.InitDB(); err != nil {
t.Fatal(err)
}
return store
}
func tableColumns(t *testing.T, db *sql.DB, table string) map[string]bool {
t.Helper()
rows, err := db.Query("PRAGMA table_info(" + table + ")")
if err != nil {
t.Fatal(err)
}
defer rows.Close()
columns := make(map[string]bool)
for rows.Next() {
var cid, notNull, primaryKey int
var name, columnType string
var defaultValue any
if err := rows.Scan(&cid, &name, &columnType, &notNull, &defaultValue, &primaryKey); err != nil {
t.Fatal(err)
}
columns[name] = true
}
if err := rows.Err(); err != nil {
t.Fatal(err)
}
return columns
}
func schemaIndexes(t *testing.T, db *sql.DB) map[string]bool {
t.Helper()
rows, err := db.Query(`SELECT name FROM sqlite_schema WHERE type = 'index'`)
if err != nil {
t.Fatal(err)
}
defer rows.Close()
indexes := make(map[string]bool)
for rows.Next() {
var name string
if err := rows.Scan(&name); err != nil {
t.Fatal(err)
}
indexes[name] = true
}
if err := rows.Err(); err != nil {
t.Fatal(err)
}
return indexes
}
func explainQueryPlan(t *testing.T, db *sql.DB, query string, args ...any) string {
t.Helper()
rows, err := db.Query("EXPLAIN QUERY PLAN "+query, args...)
if err != nil {
t.Fatal(err)
}
defer rows.Close()
var details []string
for rows.Next() {
var id, parent, notUsed int
var detail string
if err := rows.Scan(&id, &parent, &notUsed, &detail); err != nil {
t.Fatal(err)
}
details = append(details, detail)
}
if err := rows.Err(); err != nil {
t.Fatal(err)
}
return strings.Join(details, "\n")
}
+94 -26
View File
@@ -2,16 +2,22 @@ package storage
import (
"database/sql"
"errors"
"fmt"
"log"
"log/slog"
"time"
)
var (
ErrUserAlreadyExists = errors.New("username or email already exists")
ErrUserCapacity = errors.New("user capacity reached")
ErrPermanentCapacity = errors.New("permanent user capacity reached")
)
// UserLimits defines registration constraints
type UserLimits struct {
MaxUsers int
PermanentSlots int
TempTTL time.Duration
}
// DefaultUserLimits returns default POC limits
@@ -19,7 +25,6 @@ func DefaultUserLimits() UserLimits {
return UserLimits{
MaxUsers: 100,
PermanentSlots: 10,
TempTTL: 24 * time.Hour,
}
}
@@ -59,16 +64,38 @@ func (s *Store) GetOldestTempUser() (*UserRecord, error) {
// DeleteExpiredTempUsers removes temporary users past their expiry
func (s *Store) DeleteExpiredTempUsers() (int64, error) {
query := `DELETE FROM users WHERE account_type = 'temp' AND expires_at < ?`
query := `DELETE FROM users
WHERE account_type = 'temp' AND expires_at IS NOT NULL AND expires_at < ?`
result, err := s.db.Exec(query, time.Now().UTC())
if err != nil {
return 0, err
}
return result.RowsAffected()
deleted, err := result.RowsAffected()
if err == nil && deleted > 0 {
slog.Debug("storage expired temporary users deleted", "count", deleted)
}
return deleted, err
}
// CreateUser creates user with transaction isolation to prevent race conditions
// CreateUser creates an administratively managed user without applying the
// public-registration capacity policy.
func (s *Store) CreateUser(record UserRecord) error {
return s.createUser(record, nil, nil)
}
// CreateUserWithinLimits atomically applies registration limits, evicts the
// oldest temporary account when required, creates the user, and optionally
// creates its initial session. No account is evicted on a duplicate request,
// and a session failure rolls back the user and eviction together.
func (s *Store) CreateUserWithinLimits(
record UserRecord,
session *SessionRecord,
limits UserLimits,
) error {
return s.createUser(record, session, &limits)
}
func (s *Store) createUser(record UserRecord, session *SessionRecord, limits *UserLimits) error {
tx, err := s.db.Begin()
if err != nil {
return fmt.Errorf("failed to begin transaction: %w", err)
@@ -81,7 +108,37 @@ func (s *Store) CreateUser(record UserRecord) error {
return err
}
if exists {
return fmt.Errorf("username or email already exists")
return ErrUserAlreadyExists
}
if limits != nil {
var total, permanent int
if err := tx.QueryRow(`SELECT COUNT(*),
COUNT(CASE WHEN account_type = 'permanent' THEN 1 END)
FROM users`).Scan(&total, &permanent); err != nil {
return fmt.Errorf("count users: %w", err)
}
if record.AccountType == "permanent" && permanent >= limits.PermanentSlots {
return ErrPermanentCapacity
}
if total >= limits.MaxUsers {
result, err := tx.Exec(`DELETE FROM users WHERE user_id = (
SELECT user_id FROM users
WHERE account_type = 'temp'
ORDER BY created_at ASC
LIMIT 1
)`)
if err != nil {
return fmt.Errorf("evict oldest temporary user: %w", err)
}
deleted, err := result.RowsAffected()
if err != nil {
return fmt.Errorf("inspect temporary user eviction: %w", err)
}
if deleted != 1 {
return ErrUserCapacity
}
}
}
// Insert user
@@ -96,14 +153,36 @@ func (s *Store) CreateUser(record UserRecord) error {
if err != nil {
return err
}
if session != nil {
if session.UserID != record.UserID {
return errors.New("initial session user does not match new user")
}
if _, err := tx.Exec(
`INSERT INTO sessions (session_id, user_id, created_at, expires_at) VALUES (?, ?, ?, ?)`,
session.SessionID, session.UserID, session.CreatedAt, session.ExpiresAt,
); err != nil {
return fmt.Errorf("create initial session: %w", err)
}
}
return tx.Commit()
if err := tx.Commit(); err != nil {
return err
}
slog.Debug("storage user created",
"user_id", record.UserID,
"account_type", record.AccountType,
"initial_session", session != nil,
)
return nil
}
// DeleteUserByID removes a user by ID (synchronous, for replacement logic)
func (s *Store) DeleteUserByID(userID string) error {
query := `DELETE FROM users WHERE user_id = ?`
_, err := s.db.Exec(query, userID)
if err == nil {
slog.Debug("storage user deleted", "user_id", userID)
}
return err
}
@@ -121,7 +200,9 @@ func (s *Store) userExists(tx *sql.Tx, username, email string) (bool, error) {
args := []any{username}
if email != "" {
query = `SELECT COUNT(*) FROM users WHERE username = ? COLLATE NOCASE OR email = ? COLLATE NOCASE`
query = `SELECT COUNT(*) FROM users
WHERE username = ? COLLATE NOCASE
OR (email = ? COLLATE NOCASE AND email IS NOT NULL AND email != '')`
args = append(args, email)
}
@@ -217,7 +298,7 @@ func (s *Store) GetUserByEmail(email string) (*UserRecord, error) {
var user UserRecord
var emailNull sql.NullString
query := `SELECT user_id, username, email, password_hash, account_type, created_at, expires_at, last_login_at
FROM users WHERE email = ? COLLATE NOCASE`
FROM users WHERE email = ? COLLATE NOCASE AND email IS NOT NULL AND email != ''`
err := s.db.QueryRow(query, email).Scan(
&user.UserID, &user.Username, &emailNull,
@@ -250,21 +331,8 @@ func (s *Store) GetUserByID(userID string) (*UserRecord, error) {
return &user, nil
}
// DeleteUser removes a user from the database (async)
// DeleteUser removes a user synchronously. Account operations are consistency
// sensitive and should not be reported successful before SQLite commits them.
func (s *Store) DeleteUser(userID string) error {
if !s.healthStatus.Load() {
return nil
}
select {
case s.writeChan <- func(tx *sql.Tx) error {
query := `DELETE FROM users WHERE user_id = ?`
_, err := tx.Exec(query, userID)
return err
}:
return nil
default:
log.Printf("Storage write queue full, dropping user deletion")
return nil
}
return s.DeleteUserByID(userID)
}