"""PostgreSQL reader contract tests using a fake connection.""" from datetime import date from typing import cast import psycopg import pytest from zhixing_server.modules.selection.infrastructure.postgres_reader import ( PostgresMarketDataReader, ) class FakeConnection: def __init__(self, rows: list[tuple[object, ...]]) -> None: self.rows = rows self.query: str | None = None self.parameters: tuple[object, ...] | None = None def __enter__(self) -> "FakeConnection": return self def __exit__(self, *args: object) -> None: return None def execute(self, query: str, parameters: tuple[object, ...]) -> "FakeResult": self.query = query self.parameters = parameters return FakeResult(self.rows) class FakeResult: def __init__(self, rows: list[tuple[object, ...]]) -> None: self.rows = rows def fetchall(self) -> list[tuple[object, ...]]: return self.rows def test_reader_parameterizes_target_and_maps_left_join(monkeypatch: pytest.MonkeyPatch) -> None: connection = FakeConnection( [ ( "000001.SZ", "平安银行", date(2024, 1, 2), "10", "11", "9", "10.5", "1000", None, None, ), ( "000001.SZ", "平安银行", date(2024, 1, 3), "10.5", "11", "10", "10.8", "1200", "1.2", "100000", ), ] ) def connect(database_url: str) -> FakeConnection: assert database_url == "postgresql://test" return connection monkeypatch.setattr(psycopg, "connect", connect) history = PostgresMarketDataReader("postgresql://test").load_history( "000001.SZ", date(2024, 1, 3) ) assert [bar.trade_date for bar in history.bars] == [date(2024, 1, 2), date(2024, 1, 3)] assert history.daily_basic[date(2024, 1, 2)].turnover_rate is None assert history.daily_basic[date(2024, 1, 3)].total_mv == 100000.0 assert connection.parameters == ("000001.SZ", date(2024, 1, 3)) assert "source_adj = 'qfq'" in cast(str, connection.query) assert "trade_date <= %s" in cast(str, connection.query)