280 lines
8.8 KiB
Python
280 lines
8.8 KiB
Python
"""PostgreSQL reader contract tests using a fake connection."""
|
|
|
|
from collections.abc import Generator
|
|
from contextlib import contextmanager
|
|
from datetime import date, timedelta
|
|
from decimal import Decimal
|
|
from typing import cast
|
|
|
|
import psycopg
|
|
import pytest
|
|
|
|
from zhixing_server.modules.selection.domain.pattern_scoring import ZHIXING_B1_PATTERN_CASES
|
|
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,
|
|
PostgresPatternCaseLibraryLoader,
|
|
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_with_target_day_turnover_only() -> 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",
|
|
"1.2",
|
|
"100000",
|
|
),
|
|
(
|
|
"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),
|
|
)
|
|
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)
|
|
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]
|