Files
zhixing-system/zhixing-server/tests/unit/selection/test_postgres_reader.py
T

280 lines
8.8 KiB
Python
Raw Normal View History

2026-08-08 22:41:45 +08:00
"""PostgreSQL reader contract tests using a fake connection."""
2026-08-12 09:45:16 +08:00
from collections.abc import Generator
from contextlib import contextmanager
from datetime import date, timedelta
from decimal import Decimal
2026-08-08 22:41:45 +08:00
from typing import cast
import psycopg
import pytest
from zhixing_server.modules.selection.domain.pattern_scoring import ZHIXING_B1_PATTERN_CASES
2026-08-12 09:45:16 +08:00
from zhixing_server.modules.selection.domain.runs import SelectionStock
from zhixing_server.modules.selection.infrastructure.postgres_pool import SelectionPostgresPool
2026-08-08 22:41:45 +08:00
from zhixing_server.modules.selection.infrastructure.postgres_reader import (
PostgresMarketDataReader,
PostgresPatternCaseLibraryLoader,
SelectionMarketDataNotReady,
2026-08-08 22:41:45 +08:00
)
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
2026-08-12 09:45:16 +08:00
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
2026-08-08 22:41:45 +08:00
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_with_target_day_turnover_only() -> None:
2026-08-12 09:45:16 +08:00
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",
"1.2",
"100000",
2026-08-12 09:45:16 +08:00
),
(
"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[date(2024, 1, 3)].turnover_rate == 1.2
assert connection.parameters == (
date(2024, 1, 3),
["000001.SZ", "600000.SH"],
date(2024, 1, 3),
)
2026-08-12 09:45:16 +08:00
assert "bar.ts_code = ANY(%s)" in cast(str, connection.query)
assert "market_daily_basic" in cast(str, connection.query)
assert "basic.trade_date = %s" in cast(str, connection.query)
2026-08-12 09:45:16 +08:00
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),
)
def test_pattern_case_loader_reads_one_complete_exclusive_qfq_library() -> None:
rows: list[tuple[object, ...]] = []
for definition in ZHIXING_B1_PATTERN_CASES:
for offset in range(definition.lookback_days, 0, -1):
rows.append(
(
definition.id,
definition.ts_code,
definition.breakout_date - timedelta(days=offset),
"10",
"11",
"9",
str(10 + offset / 100),
str(1000 + offset),
)
)
connection = FakeConnection(rows)
pool = Pool(connection)
owner = SelectionPostgresPool("postgresql://test", max_connections=2, pool=pool)
cases = PostgresPatternCaseLibraryLoader("postgresql://test", pool=owner).load()
assert tuple(case.definition for case in cases) == ZHIXING_B1_PATTERN_CASES
assert all(len(case.history.bars) == 25 for case in cases)
assert all(case.history.bars[-1].trade_date < case.definition.breakout_date for case in cases)
assert "bar.trade_date < definition.breakout_date" in cast(str, connection.query)
assert "bar.source_adj = 'qfq'" in cast(str, connection.query)
assert connection.parameters is not None
assert connection.parameters[0] == [definition.id for definition in ZHIXING_B1_PATTERN_CASES]