"""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() 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, 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] 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" 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)