from __future__ import annotations from dataclasses import dataclass from pathlib import Path from typing import Any, Iterator import pytest from trading.domain.database import ( PostgresConnectionAdapter, PostgreSQLBackend, _postgres_sql, ) from trading.domain.schema import ( REQUIRED_TABLES, apply_migrations, migration_plan, validate_schema, ) @dataclass class FakeCursor: rows: list[Any] rowcount: int = 0 description: tuple[Any, ...] = () def fetchone(self) -> Any: return self.rows[0] if self.rows else None def fetchall(self) -> list[Any]: return list(self.rows) def __iter__(self) -> Iterator[Any]: return iter(self.rows) class FakePostgresConnection: dialect = "postgresql" def __init__(self, *, fail_script: str | None = None) -> None: self.applied: dict[str, str] = {} self.statements: list[tuple[str, tuple[Any, ...]]] = [] self.scripts: list[str] = [] self.commits = 0 self.rollbacks = 0 self.closed = False self.fail_script = fail_script self.tables = set(REQUIRED_TABLES) def execute(self, sql: str, params: tuple[Any, ...] = ()) -> FakeCursor: normalized = " ".join(sql.split()) self.statements.append((normalized, tuple(params))) upper = normalized.upper() if "PG_TRY_ADVISORY_LOCK" in upper or "PG_ADVISORY_UNLOCK" in upper: return FakeCursor([(True,)], rowcount=1) if "INFORMATION_SCHEMA.COLUMNS" in upper: return FakeCursor([("version",), ("applied_at",), ("checksum",)]) if upper.startswith("SELECT VERSION, CHECKSUM FROM SCHEMA_MIGRATIONS"): return FakeCursor([(version, checksum) for version, checksum in sorted(self.applied.items())]) if upper.startswith("SELECT VERSION FROM SCHEMA_MIGRATIONS"): return FakeCursor([(version,) for version in sorted(self.applied)]) if upper.startswith("INSERT INTO SCHEMA_MIGRATIONS"): version, _applied_at, checksum = params self.applied[str(version)] = str(checksum) return FakeCursor([], rowcount=1) if upper.startswith("UPDATE SCHEMA_MIGRATIONS SET CHECKSUM"): checksum, version = params self.applied[str(version)] = str(checksum) return FakeCursor([], rowcount=1) if "INFORMATION_SCHEMA.TABLES" in upper: return FakeCursor([(name,) for name in sorted(self.tables)]) if "FROM PG_CONSTRAINT" in upper: return FakeCursor([]) return FakeCursor([]) def executescript(self, sql: str) -> None: self.scripts.append(sql) if self.fail_script and self.fail_script in sql: raise RuntimeError("injected migration failure") def commit(self) -> None: self.commits += 1 def rollback(self) -> None: self.rollbacks += 1 def close(self) -> None: self.closed = True class RawCursor: def __init__(self) -> None: self.executed: list[tuple[str, tuple[Any, ...], dict[str, Any]]] = [] self.rowcount = 1 self.description = (("value",),) def execute(self, sql: str, params: tuple[Any, ...] = (), **kwargs: Any) -> None: self.executed.append((sql, tuple(params), kwargs)) def fetchone(self) -> tuple[int]: return (7,) def fetchall(self) -> list[tuple[int]]: return [(7,)] def __iter__(self) -> Iterator[tuple[int]]: return iter([(7,)]) class RawConnection: def __init__(self) -> None: self.cursors: list[RawCursor] = [] self.commits = 0 self.rollbacks = 0 self.closed = False def cursor(self) -> RawCursor: cursor = RawCursor() self.cursors.append(cursor) return cursor def commit(self) -> None: self.commits += 1 def rollback(self) -> None: self.rollbacks += 1 def close(self) -> None: self.closed = True def test_postgresql_migration_plan_is_explicit_and_sqlite_free() -> None: plan = migration_plan("postgresql") assert [item.version for item in plan] == [ "001_baseline", "002_operational_indexes", "003_scoped_lifecycle_metadata", "004_position_events", "005_execution_audit_session", "006_rate_limit_audit_integrity", "007_audit_append_only", "008_advisory_outcomes", ] text = "\n".join(item.path.read_text(encoding="utf-8") for item in plan).upper() for forbidden in ("AUTOINCREMENT", "PRAGMA ", "BEGIN IMMEDIATE", "SQLITE_MASTER", "RANDOMBLOB("): assert forbidden not in text for table in REQUIRED_TABLES - {"schema_migrations"}: assert f"TABLE IF NOT EXISTS {table.upper()}" in text def test_empty_postgresql_database_applies_all_migrations_under_advisory_lock() -> None: conn = FakePostgresConnection() ok, messages = apply_migrations(conn, dialect="postgresql") assert ok, messages assert list(conn.applied) == [item.version for item in migration_plan("postgresql")] assert len(conn.scripts) == len(migration_plan("postgresql")) rendered = "\n".join(sql for sql, _params in conn.statements) assert "pg_try_advisory_lock" in rendered assert "pg_advisory_unlock" in rendered def test_previous_postgresql_schema_receives_only_forward_migration() -> None: plan = migration_plan("postgresql") conn = FakePostgresConnection() conn.applied.update({item.version: item.sha256 for item in plan[:-1]}) ok, messages = apply_migrations(conn, dialect="postgresql") assert ok, messages assert messages == [f"applied {plan[-1].version}"] assert len(conn.scripts) == 1 assert "advisory_outcomes" in conn.scripts[0] def test_postgresql_migration_failure_rolls_back_and_preserves_prior_version() -> None: plan = migration_plan("postgresql") conn = FakePostgresConnection(fail_script="idx_execution_history_account_time") conn.applied[plan[0].version] = plan[0].sha256 ok, messages = apply_migrations(conn, dialect="postgresql") assert not ok assert conn.applied == {plan[0].version: plan[0].sha256} assert conn.rollbacks == 1 assert any("injected migration failure" in message for message in messages) assert any("pg_advisory_unlock" in sql for sql, _ in conn.statements) @pytest.mark.parametrize( "applied", [ {"999_unknown": "abc"}, {"001_baseline": "tampered"}, ], ) def test_invalid_postgresql_migration_state_fails_startup(applied: dict[str, str]) -> None: conn = FakePostgresConnection() conn.applied.update(applied) ok, messages = validate_schema(conn, dialect="postgresql") assert not ok assert any("unknown migration" in message or "checksum mismatch" in message for message in messages) def test_postgresql_schema_validation_checks_required_tables() -> None: conn = FakePostgresConnection() conn.applied.update({item.version: item.sha256 for item in migration_plan("postgresql")}) conn.tables.remove("fills") ok, messages = validate_schema(conn, dialect="postgresql") assert not ok assert "fills" in messages[0] def test_repository_sql_translation_is_bounded() -> None: assert _postgres_sql("BEGIN IMMEDIATE") == "BEGIN" assert _postgres_sql("SELECT * FROM orders WHERE order_id=?") == ( "SELECT * FROM orders WHERE order_id=%s" ) assert _postgres_sql("INSERT OR IGNORE INTO roles(role_id) VALUES (?)") == ( "INSERT INTO roles(role_id) VALUES (%s) ON CONFLICT DO NOTHING" ) def test_postgresql_connector_is_injectable_and_dsn_is_not_exposed(monkeypatch: pytest.MonkeyPatch) -> None: raw = RawConnection() captured: dict[str, Any] = {} def connector(dsn: str, **kwargs: Any) -> RawConnection: captured["dsn"] = dsn captured.update(kwargs) return raw monkeypatch.setenv("HERMES_DB_CONNECT_TIMEOUT_SECONDS", "4") backend = PostgreSQLBackend("postgresql://user:secret@db.invalid/hermes", connector=connector) conn = backend.connect() cursor = conn.execute("SELECT ?", (7,)) assert cursor.fetchone()[0] == 7 assert captured["connect_timeout"] == 4 assert captured["autocommit"] is False assert "secret" not in repr(backend) assert raw.cursors[-1].executed[0][0] == "SELECT %s" def test_postgresql_connection_failure_redacts_driver_details() -> None: def connector(_dsn: str, **_kwargs: Any) -> Any: raise RuntimeError("password=secret-value") backend = PostgreSQLBackend("postgresql://user:secret-value@db.invalid/hermes", connector=connector) with pytest.raises(RuntimeError, match="PostgreSQL connection failed") as exc_info: backend.connect() assert "secret-value" not in str(exc_info.value)