526 lines
16 KiB
Go
526 lines
16 KiB
Go
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")
|
|
}
|