Files
zhixing-system/zhixing-server/tests/unit/selection/test_postgres_runs.py
T
2026-08-10 11:09:23 +08:00

281 lines
8.4 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.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()
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)
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" in query:
return LoadResult(rows=[("000001.SZ", "平安银行", "selected", 2, None)])
if "COUNT(*) FROM selection_signal" in query:
return LoadResult(row=(2,))
return LoadResult(
rows=[
(
"000001.SZ",
"平安银行",
TARGET,
"zhixing_b1",
ZHIXING_B1_SIGNAL_ORDER[-1].value,
Decimal("10.5"),
{},
),
(
"000001.SZ",
"平安银行",
TARGET,
"zhixing_b1",
ZHIXING_B1_SIGNAL_ORDER[0].value,
Decimal("10.5"),
{},
),
]
)
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_orders_signals_by_formula_priority(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=2,
page_size=1,
search="100%",
category="pullback",
),
)
assert run is not None
assert [signal.category for signal in run.signals] == [
ZHIXING_B1_SIGNAL_ORDER[0],
ZHIXING_B1_SIGNAL_ORDER[-1],
]
count_query, count_parameters = next(
(query, parameters)
for query, parameters in connection.statements
if "COUNT(*) FROM selection_signal" in query
)
assert "name ILIKE %s ESCAPE" in count_query
assert count_parameters == ("run-1", "%100\\%%", "%100\\%%", "zhixing_b1_pullback_%")
page_query, page_parameters = next(
(query, parameters)
for query, parameters in connection.statements
if "LIMIT %s OFFSET %s" in query
)
assert "ORDER BY ts_code, CASE category" in page_query
assert page_parameters[-2:] == (1, 1)