Spaces:
Sleeping
Sleeping
| 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, | |
| ) | |
| 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) | |
| 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) | |