Spaces:
Sleeping
Sleeping
File size: 8,797 Bytes
2e658e7 1f8cf56 2e658e7 1f8cf56 2e658e7 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 | 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)
|