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)