from __future__ import annotations import sqlite3 import pytest from backend.database import MIGRATIONS, Database, Migration, MigrationError, MigrationRunner from backend.database.repositories import DatabaseStatusRepository def create_example(connection: sqlite3.Connection) -> None: connection.execute("CREATE TABLE example (id INTEGER PRIMARY KEY, value TEXT NOT NULL)") def drop_example(connection: sqlite3.Connection) -> None: connection.execute("DROP TABLE example") def example_migration(signature: str = "example:v1") -> Migration: return Migration(1, "create_example", signature, create_example, drop_example) def table_names(database: Database) -> set[str]: with database.read() as connection: return { str(row["name"]) for row in connection.execute("SELECT name FROM sqlite_master WHERE type='table'") } def test_connection_enables_wal_foreign_keys_and_busy_timeout(tmp_path) -> None: database = Database(tmp_path / "app.db") with database.read() as connection: assert connection.execute("PRAGMA journal_mode").fetchone()[0] == "wal" assert connection.execute("PRAGMA foreign_keys").fetchone()[0] == 1 assert connection.execute("PRAGMA busy_timeout").fetchone()[0] == 20_000 def test_migration_can_upgrade_idempotently_and_downgrade(tmp_path) -> None: database = Database(tmp_path / "app.db") runner = MigrationRunner(database) migration = example_migration() assert runner.upgrade((migration,)) == (1,) assert runner.upgrade((migration,)) == () assert "example" in table_names(database) assert DatabaseStatusRepository(database).get().schema_version == 1 assert runner.downgrade((migration,), target_version=0) == (1,) assert "example" not in table_names(database) assert DatabaseStatusRepository(database).get().schema_version == 0 def test_failed_migration_is_atomic_and_not_recorded(tmp_path) -> None: database = Database(tmp_path / "app.db") def fail(connection: sqlite3.Connection) -> None: connection.execute("CREATE TABLE should_rollback (id INTEGER)") raise RuntimeError("stop") migration = Migration(1, "failure", "failure:v1", fail, lambda connection: None) with pytest.raises(MigrationError, match="upgrade failed"): MigrationRunner(database).upgrade((migration,)) assert "should_rollback" not in table_names(database) assert DatabaseStatusRepository(database).get().schema_version == 0 def test_applied_migration_checksum_cannot_change(tmp_path) -> None: database = Database(tmp_path / "app.db") runner = MigrationRunner(database) runner.upgrade((example_migration(),)) with pytest.raises(MigrationError, match="checksum changed"): runner.upgrade((example_migration("example:v2"),)) with pytest.raises(MigrationError, match="checksum changed"): runner.downgrade((example_migration("example:v2"),), target_version=0) def test_unknown_database_migration_is_rejected(tmp_path) -> None: database = Database(tmp_path / "app.db") runner = MigrationRunner(database) runner.upgrade((example_migration(),)) with pytest.raises(MigrationError, match="unknown migrations"): runner.upgrade(()) def test_non_contiguous_database_history_is_rejected(tmp_path) -> None: database = Database(tmp_path / "app.db") first = example_migration() second = Migration( 2, "second", "second:v1", lambda connection: None, lambda connection: None, ) runner = MigrationRunner(database) runner.upgrade((first, second)) with database.transaction() as connection: connection.execute("DELETE FROM schema_migrations WHERE version = 1") with pytest.raises(MigrationError, match="not contiguous"): runner.upgrade((first, second)) def test_real_account_schema_can_upgrade_and_rollback(tmp_path) -> None: database = Database(tmp_path / "app.db") runner = MigrationRunner(database) assert runner.upgrade(MIGRATIONS) == (1, 2, 3, 4, 5, 6, 7, 8, 9) assert { "users", "memberships", "sessions", "birth_profiles", "llm_usage_daily", "system_credentials", "llm_models", "llm_configuration", "trading_days", "market_entities", "market_summaries", "chart_series", "sector_member_snapshots", "market_insight_snapshots", "seat_aliases", "watchlist_entries", "screener_factor_snapshots", "screener_factor_values", "screener_runs", "custom_screener_strategies", "strategy_tracks", "strategy_track_bars", "strategy_track_events", "mentor_preferences", "mentor_messages", "llm_requests", "llm_attempts", "heaven_readings", } <= table_names(database) assert runner.downgrade(MIGRATIONS, target_version=0) == (9, 8, 7, 6, 5, 4, 3, 2, 1) assert "users" not in table_names(database) assert "llm_models" not in table_names(database)