v0.11.0 harden persistence and prepare game replays
This commit is contained in:
+426
-86
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
`
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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(¤tVersion); 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, ¬Null, &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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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, ¬Null, &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, ¬Used, &detail); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
details = append(details, detail)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return strings.Join(details, "\n")
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user