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") }