Files
dzentra_bot/app/tests/unit/storage/test_migrations.py

296 lines
8.8 KiB
Python

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 (
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"),
)
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) <> '')") == 3
assert sql.count("CHECK (BTRIM(symbol) <> '')") == 3
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
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)
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)
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)
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)
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]