252 lines
7.4 KiB
Python
252 lines
7.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,
|
|
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 __enter__(self) -> "LoadConnection":
|
|
return self
|
|
|
|
def __exit__(self, *args: object) -> None:
|
|
return None
|
|
|
|
def execute(self, query: str, parameters: tuple[object, ...]) -> "LoadResult":
|
|
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)])
|
|
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")
|
|
|
|
assert run is not None
|
|
assert [signal.category for signal in run.signals] == [
|
|
ZHIXING_B1_SIGNAL_ORDER[0],
|
|
ZHIXING_B1_SIGNAL_ORDER[-1],
|
|
]
|