"""PostgreSQL reader contract tests using a fake connection.""" from datetime import date from decimal import Decimal from typing import cast import psycopg import pytest from zhixing_server.modules.selection.infrastructure.postgres_reader import ( PostgresMarketDataReader, SelectionMarketDataNotReady, ) 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) class SourceConnection: def __init__(self, source_row: tuple[object, ...] | None) -> None: self.source_row = source_row self.queries: list[tuple[str, tuple[object, ...]]] = [] def __enter__(self) -> "SourceConnection": return self def __exit__(self, *args: object) -> None: return None def execute(self, query: str, parameters: tuple[object, ...]) -> "SourceResult": self.queries.append((query, parameters)) if "FROM market_sync_batch" in query: return SourceResult(row=self.source_row) return SourceResult(rows=[("000001.SZ", "平安银行")]) class SourceResult: def __init__( self, row: tuple[object, ...] | None = None, rows: list[tuple[object, ...]] | None = None, ) -> None: self.row = row self.rows = rows or [] def fetchone(self) -> tuple[object, ...] | None: return self.row def fetchall(self) -> list[tuple[object, ...]]: return self.rows def test_reader_loads_only_eligible_stocks_from_finished_market_batch( monkeypatch: pytest.MonkeyPatch, ) -> None: connection = SourceConnection(("market-run-1", 2, 2, Decimal("1"))) def connect(database_url: str) -> SourceConnection: assert database_url == "postgresql://test" return connection monkeypatch.setattr(psycopg, "connect", connect) source = PostgresMarketDataReader("postgresql://test").load_execution_source( "zhixing_b1", date(2026, 8, 8), ) assert source.market_sync_batch_id == "market-run-1" assert source.target_count == 2 assert source.coverage == Decimal("1") assert source.stocks[0].ts_code == "000001.SZ" assert "strategy_eligible = true" in connection.queries[0][0] assert "source_adj = 'qfq'" in connection.queries[1][0] def test_reader_rejects_date_without_eligible_market_batch(monkeypatch: pytest.MonkeyPatch) -> None: connection = SourceConnection(None) def connect(database_url: str) -> SourceConnection: assert database_url == "postgresql://test" return connection monkeypatch.setattr(psycopg, "connect", connect) with pytest.raises(SelectionMarketDataNotReady): PostgresMarketDataReader("postgresql://test").load_execution_source( "zhixing_b1", date(2026, 8, 8), )