Files
zhixing-system/zhixing-server/tests/unit/selection/test_postgres_runs.py
T
2026-08-31 16:14:16 +08:00

502 lines
16 KiB
Python

"""Persistence transaction tests for selection runs."""
from datetime import date
from decimal import Decimal
import psycopg
import pytest
from psycopg.types.json import Jsonb
from zhixing_server.modules.selection.domain.models import SelectionSignal, ZhixingB1Category
from zhixing_server.modules.selection.domain.pattern_scoring import (
PATTERN_SCORING_VERSION,
ZHIXING_B1_PATTERN_CASES,
PatternScore,
PatternScoreBreakdown,
)
from zhixing_server.modules.selection.domain.runs import (
SelectionExecutionSource,
SelectionRerunRequired,
SelectionResultQuery,
SelectionRunInProgress,
SelectionRunItem,
SelectionStock,
)
from zhixing_server.modules.selection.domain.zhixing_b1 import ZHIXING_B1_SIGNAL_ORDER
from zhixing_server.modules.selection.infrastructure.postgres_runs import (
PostgresSelectionRunRepository,
)
TARGET = date(2026, 8, 8)
class FakeResult:
def __init__(self, row: tuple[object, ...] | None = None) -> None:
self.row = row
def fetchone(self) -> tuple[object, ...] | None:
return self.row
class FakeTransaction:
def __enter__(self) -> "FakeTransaction":
return self
def __exit__(self, *args: object) -> None:
return None
class FakeConnection:
def __init__(self, existing: tuple[object, ...] | None) -> None:
self.existing = existing
self.statements: list[tuple[str, tuple[object, ...]]] = []
def __enter__(self) -> "FakeConnection":
return self
def __exit__(self, *args: object) -> None:
return None
def transaction(self) -> FakeTransaction:
return FakeTransaction()
def execute(self, query: str, parameters: tuple[object, ...]) -> FakeResult:
self.statements.append((query, parameters))
if "SELECT id, status" in query:
return FakeResult(self.existing)
return FakeResult()
class BatchConnection(FakeConnection):
def __init__(self) -> None:
super().__init__(None)
self.executemany_calls: list[tuple[str, tuple[tuple[object, ...], ...]]] = []
def cursor(self) -> "BatchCursor":
return BatchCursor(self.executemany_calls)
class BatchCursor:
def __init__(self, calls: list[tuple[str, tuple[tuple[object, ...], ...]]]) -> None:
self.calls = calls
def __enter__(self) -> "BatchCursor":
return self
def __exit__(self, *args: object) -> None:
return None
def executemany(
self,
query: str,
parameters: tuple[tuple[object, ...], ...],
) -> None:
self.calls.append((query, parameters))
def _source() -> SelectionExecutionSource:
return SelectionExecutionSource(
market_sync_batch_id="market-run-1",
target_trade_date=TARGET,
target_count=1,
valid_count=1,
coverage=Decimal("1"),
stocks=(SelectionStock("000001.SZ", "平安银行"),),
)
def _repository(
monkeypatch: pytest.MonkeyPatch,
connection: FakeConnection,
) -> PostgresSelectionRunRepository:
def connect(database_url: str) -> FakeConnection:
assert database_url == "postgresql://test"
return connection
monkeypatch.setattr(psycopg, "connect", connect)
return PostgresSelectionRunRepository("postgresql://test")
def test_prepare_claims_new_business_key(monkeypatch: pytest.MonkeyPatch) -> None:
connection = FakeConnection(None)
run = _repository(monkeypatch, connection).prepare_run(
"zhixing_b1",
TARGET,
_source(),
rerun=False,
)
assert run.status == "running"
assert run.target_trade_date == TARGET
assert any("INSERT INTO selection_run" in query for query, _ in connection.statements)
assert not any("DELETE FROM selection_run" in query for query, _ in connection.statements)
def test_prepare_requires_confirmation_for_terminal_run(monkeypatch: pytest.MonkeyPatch) -> None:
connection = FakeConnection(("old-run", "success"))
repository = _repository(monkeypatch, connection)
with pytest.raises(SelectionRerunRequired):
repository.prepare_run("zhixing_b1", TARGET, _source(), rerun=False)
assert not any("DELETE FROM selection_run" in query for query, _ in connection.statements)
def test_prepare_rejects_duplicate_running_run(monkeypatch: pytest.MonkeyPatch) -> None:
connection = FakeConnection(("old-run", "running"))
repository = _repository(monkeypatch, connection)
with pytest.raises(SelectionRunInProgress):
repository.prepare_run("zhixing_b1", TARGET, _source(), rerun=True)
def test_prepare_rerun_deletes_old_result_before_insert(monkeypatch: pytest.MonkeyPatch) -> None:
connection = FakeConnection(("old-run", "failed"))
run = _repository(monkeypatch, connection).prepare_run(
"zhixing_b1",
TARGET,
_source(),
rerun=True,
)
statements = [query for query, _ in connection.statements]
assert "DELETE FROM selection_run WHERE strategy = %s AND target_trade_date = %s" in statements
assert run.id != "old-run"
def test_record_item_persists_independent_signals_as_jsonb(monkeypatch: pytest.MonkeyPatch) -> None:
connection = FakeConnection(None)
repository = _repository(monkeypatch, connection)
signal = SelectionSignal(
ts_code="000001.SZ",
name="平安银行",
target_trade_date=TARGET,
strategy="zhixing_b1",
category=ZhixingB1Category.ORIGINAL_B1,
close=10.5,
details={"j": 12.0},
)
repository.record_item(
"run-1",
SelectionRunItem(
ts_code="000001.SZ",
name="平安银行",
status="selected",
signal_count=1,
signals=(signal,),
),
)
signal_insert = next(
parameters
for query, parameters in connection.statements
if "INSERT INTO selection_signal" in query
)
assert isinstance(signal_insert[-1], Jsonb)
def test_record_items_uses_one_delete_and_two_batch_upserts(
monkeypatch: pytest.MonkeyPatch,
) -> None:
connection = BatchConnection()
repository = _repository(monkeypatch, connection)
first = SelectionSignal(
ts_code="000001.SZ",
name="平安银行",
target_trade_date=TARGET,
strategy="zhixing_b1",
category=ZHIXING_B1_SIGNAL_ORDER[0],
close=10.5,
details={"j": 12.0},
)
second = SelectionSignal(
ts_code="000001.SZ",
name="平安银行",
target_trade_date=TARGET,
strategy="zhixing_b1",
category=ZHIXING_B1_SIGNAL_ORDER[-1],
close=10.5,
details={"j": 13.0},
)
repository.record_items(
"run-1",
(
SelectionRunItem(
ts_code="000001.SZ",
name="平安银行",
status="selected",
signal_count=2,
pattern_score=PatternScore(
status="matched",
value=86.4,
threshold=60.0,
version=PATTERN_SCORING_VERSION,
case=ZHIXING_B1_PATTERN_CASES[0],
breakdown=PatternScoreBreakdown(71.2, 83.0, 88.0, 90.1),
),
signals=(first, second),
),
SelectionRunItem(
ts_code="600000.SH",
name="浦发银行",
status="no_signal",
),
),
)
delete_query, delete_parameters = connection.statements[0]
assert "DELETE FROM selection_signal" in delete_query
assert "ANY(%s)" in delete_query
assert delete_parameters == ("run-1", ["000001.SZ", "600000.SH"])
assert len(connection.executemany_calls) == 2
assert "INSERT INTO selection_run_item" in connection.executemany_calls[0][0]
item_parameters = connection.executemany_calls[0][1]
assert item_parameters[0][6:13] == (
"matched",
86.4,
60.0,
PATTERN_SCORING_VERSION,
"case_001",
"华纳药厂",
date(2025, 5, 12),
)
assert isinstance(item_parameters[0][13], Jsonb)
assert item_parameters[1][6:] == (
"not_executed",
None,
None,
None,
None,
None,
None,
None,
None,
)
assert "INSERT INTO selection_signal" in connection.executemany_calls[1][0]
signal_parameters = connection.executemany_calls[1][1]
assert len(signal_parameters) == 2
assert {values[5] for values in signal_parameters} == {
ZHIXING_B1_SIGNAL_ORDER[0].value,
ZHIXING_B1_SIGNAL_ORDER[-1].value,
}
assert all(isinstance(values[-1], Jsonb) for values in signal_parameters)
class LoadConnection:
def __init__(self) -> None:
self.statements: list[tuple[str, tuple[object, ...]]] = []
def __enter__(self) -> "LoadConnection":
return self
def __exit__(self, *args: object) -> None:
return None
def execute(self, query: str, parameters: tuple[object, ...]) -> "LoadResult":
self.statements.append((query, parameters))
if "FROM selection_run\n" in query:
return LoadResult(
row=(
"run-1",
"zhixing_b1",
TARGET,
"market-run-1",
"success",
1,
1,
1,
1,
2,
0,
Decimal("1"),
None,
None,
None,
None,
)
)
if "FROM selection_run_item\n" in query:
return LoadResult(
rows=[
(
"000001.SZ",
"平安银行",
"selected",
2,
None,
"matched",
Decimal("86.40"),
Decimal("60.00"),
PATTERN_SCORING_VERSION,
"case_001",
"华纳药厂",
date(2025, 5, 12),
{
"trend_structure": 71.2,
"kdj_state": 83.0,
"volume_pattern": 88.0,
"price_shape": 90.1,
},
None,
)
]
)
if "SELECT COUNT(*) FROM selection_run_item AS item" in query:
return LoadResult(row=(2,))
if "SELECT item.ts_code" in query:
return LoadResult(rows=[("000001.SZ",)])
return LoadResult(
rows=[
(
"000001.SZ",
"平安银行",
TARGET,
"zhixing_b1",
ZHIXING_B1_SIGNAL_ORDER[-1].value,
Decimal("10.5"),
{},
),
(
"000001.SZ",
"平安银行",
TARGET,
"zhixing_b1",
ZhixingB1Category.ORIGINAL_B1.value,
Decimal("10.5"),
{},
),
]
)
class EmptyStockPageConnection(LoadConnection):
"""Return a non-zero filtered total with no stocks on the requested page."""
def execute(self, query: str, parameters: tuple[object, ...]) -> "LoadResult":
if "SELECT item.ts_code" in query:
self.statements.append((query, parameters))
return LoadResult(rows=[])
return super().execute(query, parameters)
class LoadResult:
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_get_run_pages_stocks_and_loads_all_signals_for_category_matches(
monkeypatch: pytest.MonkeyPatch,
) -> None:
connection = LoadConnection()
def connect(database_url: str) -> LoadConnection:
assert database_url == "postgresql://test"
return connection
monkeypatch.setattr(psycopg, "connect", connect)
run = PostgresSelectionRunRepository("postgresql://test").get_run(
"run-1",
query=SelectionResultQuery(
page=1,
page_size=1,
search="100%",
category="pullback",
sort="score_desc",
),
)
assert run is not None
assert [signal.category for signal in run.signals] == [
ZhixingB1Category.ORIGINAL_B1,
ZHIXING_B1_SIGNAL_ORDER[-1],
]
assert run.stocks_total == 2
assert run.items[0].pattern_score.status == "matched"
assert run.items[0].pattern_score.value == 86.4
count_query, count_parameters = next(
(query, parameters)
for query, parameters in connection.statements
if "SELECT COUNT(*) FROM selection_run_item AS item" in query
)
assert "name ILIKE %s ESCAPE" in count_query
assert count_parameters == ("run-1", "%100\\%%", "%100\\%%", "zhixing_b1_pullback_%")
stock_page_query, page_parameters = next(
(query, parameters)
for query, parameters in connection.statements
if "SELECT item.ts_code" in query
)
assert "ORDER BY item.score_value DESC NULLS LAST, item.ts_code ASC" in stock_page_query
assert page_parameters[-2:] == (1, 0)
signal_query, signal_parameters = next(
(query, parameters)
for query, parameters in connection.statements
if "ts_code = ANY(%s)" in query and "SELECT\n" in query
)
assert "ORDER BY ts_code, CASE category" in signal_query
assert "category LIKE" not in signal_query
assert signal_parameters == ("run-1", ["000001.SZ"])
def test_get_run_does_not_load_signals_for_an_empty_stock_page(
monkeypatch: pytest.MonkeyPatch,
) -> None:
connection = EmptyStockPageConnection()
def connect(database_url: str) -> EmptyStockPageConnection:
assert database_url == "postgresql://test"
return connection
monkeypatch.setattr(psycopg, "connect", connect)
run = PostgresSelectionRunRepository("postgresql://test").get_run(
"run-1",
query=SelectionResultQuery(page=3, page_size=1),
)
assert run is not None
assert run.stocks_total == 2
assert run.signals == ()
stock_page_query, stock_page_parameters = next(
(query, parameters)
for query, parameters in connection.statements
if "SELECT item.ts_code" in query
)
assert "ORDER BY item.ts_code ASC" in stock_page_query
assert stock_page_parameters[-2:] == (1, 2)
assert not any("ts_code = ANY(%s)" in query for query, _ in connection.statements)
def test_get_run_sorts_scores_ascending_with_nulls_last_and_code_tiebreak(
monkeypatch: pytest.MonkeyPatch,
) -> None:
connection = LoadConnection()
def connect(database_url: str) -> LoadConnection:
assert database_url == "postgresql://test"
return connection
monkeypatch.setattr(psycopg, "connect", connect)
run = PostgresSelectionRunRepository("postgresql://test").get_run(
"run-1",
query=SelectionResultQuery(sort="score_asc"),
)
assert run is not None
stock_page_query = next(
query for query, _ in connection.statements if "SELECT item.ts_code" in query
)
assert "ORDER BY item.score_value ASC NULLS LAST, item.ts_code ASC" in stock_page_query