"""PostgreSQL reader contract tests using a fake connection.""" from collections.abc import Generator from contextlib import contextmanager from datetime import date from decimal import Decimal from typing import cast import psycopg import pytest from zhixing_server.modules.selection.domain.runs import SelectionStock from zhixing_server.modules.selection.infrastructure.postgres_pool import SelectionPostgresPool 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 class Pool: def __init__(self, connection: FakeConnection) -> None: self.connection_value = connection self.opened = 0 def open(self, *, wait: bool = True) -> None: assert wait is True self.opened += 1 def close(self) -> None: return None @contextmanager def connection(self) -> Generator[FakeConnection, None, None]: yield self.connection_value 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) def test_reader_batches_qfq_rows_by_stock_without_historical_basic_join() -> None: connection = FakeConnection( [ ( "600000.SH", "浦发银行", date(2024, 1, 2), "8", "8.5", "7.8", "8.2", "900", ), ( "000001.SZ", "平安银行", date(2024, 1, 3), "10.5", "11", "10", "10.8", "1200", ), ( "000001.SZ", "平安银行", date(2024, 1, 2), "10", "11", "9", "10.5", "1000", ), ] ) pool = Pool(connection) owner = SelectionPostgresPool("postgresql://test", max_connections=6, pool=pool) histories = PostgresMarketDataReader("postgresql://test", pool=owner).load_histories( ( SelectionStock("000001.SZ", "平安银行"), SelectionStock("600000.SH", "浦发银行"), ), date(2024, 1, 3), ) assert [history.ts_code for history in histories] == ["000001.SZ", "600000.SH"] assert [bar.trade_date for bar in histories[0].bars] == [ date(2024, 1, 2), date(2024, 1, 3), ] assert histories[0].daily_basic == {} assert connection.parameters == (["000001.SZ", "600000.SH"], date(2024, 1, 3)) assert "bar.ts_code = ANY(%s)" in cast(str, connection.query) assert "market_daily_basic" not in cast(str, connection.query) assert pool.opened == 1 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), )