from __future__ import annotations from dataclasses import dataclass, field from typing import Any import pytest from src.storage.exceptions import StorageMigrationError from src.storage.migrations import ( MARKET_DATA_PARTITION_ADVISORY_LOCK_ID, STORAGE_MIGRATION_ADVISORY_LOCK_ID, STORAGE_MIGRATIONS, StorageMigration, StorageMigrationRunner, ) @dataclass class RecordingCursor: applied_rows: list[tuple[int, str]] = field(default_factory=list) fail_on: str | None = None calls: list[tuple[str, object | None]] = field(default_factory=list) def __enter__(self) -> RecordingCursor: return self def __exit__(self, *args: object) -> None: return None def execute( self, statement: str, parameters: object | None = None, ) -> None: normalized = " ".join(statement.split()) self.calls.append((normalized, parameters)) if self.fail_on is not None and self.fail_on in normalized: raise RuntimeError("database failed") def fetchall(self) -> list[tuple[int, str]]: return list(self.applied_rows) @dataclass class RecordingConnection: cursor_value: RecordingCursor entered: int = 0 exited: int = 0 def __enter__(self) -> RecordingConnection: self.entered += 1 return self def __exit__(self, *args: object) -> None: self.exited += 1 return None def cursor(self) -> RecordingCursor: return self.cursor_value @dataclass class RecordingProvider: connection: RecordingConnection calls: int = 0 def __call__(self) -> RecordingConnection: self.calls += 1 return self.connection def _runner( *, applied_rows: list[tuple[int, str]] | None = None, fail_on: str | None = None, migrations: tuple[StorageMigration, ...] = STORAGE_MIGRATIONS, ) -> tuple[ StorageMigrationRunner, RecordingCursor, RecordingConnection, RecordingProvider, ]: cursor = RecordingCursor( applied_rows=applied_rows or [], fail_on=fail_on, ) connection = RecordingConnection(cursor) provider = RecordingProvider(connection) runner = StorageMigrationRunner( connection_provider=provider, migrations=migrations, ) return runner, cursor, connection, provider def test_default_migrations_have_stable_order_and_names() -> None: assert tuple( (migration.version, migration.name) for migration in STORAGE_MIGRATIONS ) == ( (1, "create_market_data_schema"), (2, "create_canonical_trades"), (3, "create_canonical_quotes"), (4, "create_canonical_candle_revisions"), (5, "add_trade_observation_sources"), (6, "add_quote_and_candle_observation_sources"), (7, "create_market_data_partition_registry"), (8, "create_trade_stream_checkpoints"), (9, "add_global_market_data_replay_sequence"), ) def test_default_schema_defines_partitions_identities_and_constraints() -> None: sql = "\n".join( statement for migration in STORAGE_MIGRATIONS for statement in migration.statements ) assert "CREATE SCHEMA IF NOT EXISTS market_data" in sql assert "CREATE TABLE market_data.trades" in sql assert "PRIMARY KEY (venue, symbol, trade_id, executed_at)" in sql assert "trade_id BETWEEN -2147483648 AND 2147483647" in sql assert sql.count("CHECK (BTRIM(venue) <> '')") == 4 assert sql.count("CHECK (BTRIM(symbol) <> '')") == 4 assert sql.count("CHECK (BTRIM(source) <> '')") == 3 assert "PARTITION BY RANGE (executed_at)" in sql assert "CREATE TABLE market_data.quotes" in sql assert "PRIMARY KEY (venue, symbol, received_at)" in sql assert "PARTITION BY RANGE (received_at)" in sql assert "CREATE TABLE market_data.candle_revisions" in sql assert "open_time,\n observed_at" in sql assert "PARTITION BY RANGE (open_time)" in sql assert sql.count("PARTITION OF") == 3 assert "ADD COLUMN observation_sources TEXT[]" in sql assert "SET observation_sources = ARRAY[source]" in sql assert "CARDINALITY(observation_sources) > 0" in sql assert "ALTER TABLE market_data.quotes" in sql assert "ALTER TABLE market_data.candle_revisions" in sql assert "quotes_observation_sources_not_empty" in sql assert "candle_revisions_observation_sources_not_empty" in sql assert "CREATE TABLE market_data.partition_registry" in sql assert "PRIMARY KEY (data_type, range_start)" in sql assert "UNIQUE (partition_name)" in sql assert "partition_bound TEXT NOT NULL" in sql assert "BTRIM(partition_bound) <> ''" in sql assert "range_end > range_start" in sql assert "CREATE TABLE market_data.trade_stream_checkpoints" in sql assert "PRIMARY KEY (venue, symbol)" in sql assert "revision BIGINT NOT NULL" in sql assert "checkpoint_schema_version INTEGER NOT NULL DEFAULT 1" in sql assert "CONSTRAINT trade_stream_checkpoints_trade_fk" in sql assert "FOREIGN KEY (" in sql assert ") REFERENCES market_data.trades (" in sql assert "ON UPDATE NO ACTION" in sql assert "ON DELETE NO ACTION" in sql assert "DEFERRABLE INITIALLY DEFERRED" in sql assert "CHECK (revision > 0)" in sql assert "CHECK (checkpoint_schema_version > 0)" in sql def test_replay_sequence_migration_has_atomic_global_order_contract() -> None: migration = STORAGE_MIGRATIONS[-1] statements = tuple( " ".join(statement.split()) for statement in migration.statements ) sql = "\n".join(statements) assert migration.version == 9 assert migration.name == "add_global_market_data_replay_sequence" assert statements[:4] == ( ( "SELECT pg_advisory_xact_lock(" f"{MARKET_DATA_PARTITION_ADVISORY_LOCK_ID}" ")" ), "LOCK TABLE market_data.trades IN ACCESS EXCLUSIVE MODE", "LOCK TABLE market_data.quotes IN ACCESS EXCLUSIVE MODE", ( "LOCK TABLE market_data.candle_revisions " "IN ACCESS EXCLUSIVE MODE" ), ) assert "CREATE SEQUENCE market_data.replay_sequence AS BIGINT" in sql assert "MINVALUE 1" in sql assert "CACHE 1" in sql assert "NO CYCLE" in sql assert "OWNED BY NONE" in sql assert sql.count("ADD COLUMN replay_sequence BIGINT") == 3 assert "CREATE TABLE market_data.replay_sequence" not in sql assert "CREATE TEMPORARY TABLE market_data_replay_sequence_backfill" in sql assert "ON COMMIT DROP" in sql assert "ROW_NUMBER() OVER" in sql assert "UNION ALL" in sql assert "event_time, data_type_rank, venue COLLATE \"C\"" in sql assert "symbol COLLATE \"C\"" in sql assert "candle_interval COLLATE \"C\"" in sql assert "ctid" not in sql.lower() assert sql.count("SET replay_sequence = backfill.replay_sequence") == 3 assert "target.trade_id = backfill.trade_id" in sql assert "target.received_at = backfill.received_at" in sql assert "target.interval = backfill.interval" in sql assert "target.open_time = backfill.open_time" in sql assert "target.observed_at = backfill.observed_at" in sql assert "SELECT pg_catalog.setval(" in sql assert "EXISTS ( SELECT 1 FROM market_data_replay_sequence_backfill )" in sql assert sql.count("SET DEFAULT nextval(") == 3 assert sql.count("ALTER COLUMN replay_sequence SET NOT NULL") == 3 assert sql.count("CHECK (replay_sequence > 0)") == 3 assert "CREATE FUNCTION market_data.reject_replay_sequence_change()" in sql assert sql.count("BEFORE UPDATE OF replay_sequence") == 3 assert "CREATE INDEX trades_history_keyset_idx" in sql assert "CREATE INDEX quotes_history_keyset_idx" in sql assert "CREATE INDEX candle_revisions_history_keyset_idx" in sql assert "CREATE INDEX candle_revisions_replay_keyset_idx" in sql assert "executed_at, replay_sequence" in sql assert "received_at, replay_sequence" in sql assert "interval, open_time, replay_sequence" in sql assert "interval, observed_at, replay_sequence" in sql def test_run_locks_and_applies_every_pending_migration_in_order() -> None: runner, cursor, connection, provider = _runner() result = runner.run() assert result == (1, 2, 3, 4, 5, 6, 7, 8, 9) assert provider.calls == 1 assert connection.entered == 1 assert connection.exited == 1 assert cursor.calls[0] == ( "SELECT pg_advisory_xact_lock(%s)", (STORAGE_MIGRATION_ADVISORY_LOCK_ID,), ) assert "public.storage_schema_migrations" in cursor.calls[1][0] inserted_versions = tuple( parameters[0] for statement, parameters in cursor.calls if statement.startswith( "INSERT INTO public.storage_schema_migrations" ) and isinstance(parameters, tuple) ) assert inserted_versions == (1, 2, 3, 4, 5, 6, 7, 8, 9) def test_run_skips_already_applied_migrations() -> None: applied = [ (migration.version, migration.name) for migration in STORAGE_MIGRATIONS ] runner, cursor, _, _ = _runner(applied_rows=applied) result = runner.run() assert result == () assert not any( statement.startswith("CREATE SCHEMA") or statement.startswith("CREATE TABLE market_data") for statement, _ in cursor.calls ) def test_run_applies_only_migrations_after_existing_prefix() -> None: applied = [ (migration.version, migration.name) for migration in STORAGE_MIGRATIONS[:2] ] runner, cursor, _, _ = _runner(applied_rows=applied) result = runner.run() assert result == (3, 4, 5, 6, 7, 8, 9) inserted_versions = tuple( parameters[0] for statement, parameters in cursor.calls if statement.startswith( "INSERT INTO public.storage_schema_migrations" ) and isinstance(parameters, tuple) ) assert inserted_versions == (3, 4, 5, 6, 7, 8, 9) def test_run_rejects_unknown_applied_version() -> None: runner, _, _, _ = _runner(applied_rows=[(99, "future")]) with pytest.raises(StorageMigrationError, match="unknown.*99"): runner.run() def test_run_rejects_changed_applied_name() -> None: runner, _, _, _ = _runner(applied_rows=[(1, "renamed")]) with pytest.raises(StorageMigrationError, match="name mismatch"): runner.run() def test_run_rejects_non_prefix_applied_history() -> None: second = STORAGE_MIGRATIONS[1] runner, _, _, _ = _runner( applied_rows=[(second.version, second.name)], ) with pytest.raises(StorageMigrationError, match="ordered prefix"): runner.run() def test_database_error_is_wrapped() -> None: runner, _, _, _ = _runner(fail_on="CREATE SCHEMA") with pytest.raises(StorageMigrationError) as error_info: runner.run() assert isinstance(error_info.value.__cause__, RuntimeError) def test_keyboard_interrupt_is_not_wrapped() -> None: def interrupted_provider() -> Any: raise KeyboardInterrupt runner = StorageMigrationRunner( connection_provider=interrupted_provider, ) with pytest.raises(KeyboardInterrupt): runner.run() def test_runner_rejects_duplicate_versions() -> None: migration = StorageMigration( version=1, name="one", statements=("SELECT 1",), ) with pytest.raises(ValueError, match="unique"): StorageMigrationRunner( connection_provider=lambda: None, # type: ignore[arg-type] migrations=(migration, migration), ) def test_runner_rejects_out_of_order_versions() -> None: first = StorageMigration(1, "one", ("SELECT 1",)) second = StorageMigration(2, "two", ("SELECT 2",)) with pytest.raises(ValueError, match="ordered"): StorageMigrationRunner( connection_provider=lambda: None, # type: ignore[arg-type] migrations=(second, first), ) @pytest.mark.parametrize( "arguments", ( {"version": 0, "name": "zero", "statements": ("SELECT 0",)}, {"version": True, "name": "one", "statements": ("SELECT 1",)}, {"version": 1, "name": " ", "statements": ("SELECT 1",)}, {"version": 1, "name": "one", "statements": ()}, {"version": 1, "name": "one", "statements": (" ",)}, ), ) def test_migration_rejects_invalid_definition( arguments: dict[str, object], ) -> None: with pytest.raises(ValueError): StorageMigration(**arguments) # type: ignore[arg-type]