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

281 lines
8.4 KiB
Python
Raw Normal View History

"""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)