2026-08-09 09:34:46 +08:00
|
|
|
"""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,
|
2026-08-10 11:09:23 +08:00
|
|
|
SelectionResultQuery,
|
2026-08-09 09:34:46 +08:00
|
|
|
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()
|
|
|
|
|
|
|
|
|
|
|
2026-08-12 09:45:16 +08:00
|
|
|
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))
|
|
|
|
|
|
|
|
|
|
|
2026-08-09 09:34:46 +08:00
|
|
|
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)
|
|
|
|
|
|
|
|
|
|
|
2026-08-12 09:45:16 +08:00
|
|
|
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)
|
|
|
|
|
|
|
|
|
|
|
2026-08-09 09:34:46 +08:00
|
|
|
class LoadConnection:
|
2026-08-10 11:09:23 +08:00
|
|
|
def __init__(self) -> None:
|
|
|
|
|
self.statements: list[tuple[str, tuple[object, ...]]] = []
|
|
|
|
|
|
2026-08-09 09:34:46 +08:00
|
|
|
def __enter__(self) -> "LoadConnection":
|
|
|
|
|
return self
|
|
|
|
|
|
|
|
|
|
def __exit__(self, *args: object) -> None:
|
|
|
|
|
return None
|
|
|
|
|
|
|
|
|
|
def execute(self, query: str, parameters: tuple[object, ...]) -> "LoadResult":
|
2026-08-10 11:09:23 +08:00
|
|
|
self.statements.append((query, parameters))
|
2026-08-09 09:34:46 +08:00
|
|
|
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)])
|
2026-08-28 11:39:05 +08:00
|
|
|
if "COUNT(DISTINCT ts_code) FROM selection_signal" in query:
|
2026-08-10 11:09:23 +08:00
|
|
|
return LoadResult(row=(2,))
|
2026-08-28 11:39:05 +08:00
|
|
|
if "SELECT DISTINCT ts_code" in query:
|
|
|
|
|
return LoadResult(rows=[("000001.SZ",)])
|
2026-08-09 09:34:46 +08:00
|
|
|
return LoadResult(
|
|
|
|
|
rows=[
|
|
|
|
|
(
|
|
|
|
|
"000001.SZ",
|
|
|
|
|
"平安银行",
|
|
|
|
|
TARGET,
|
|
|
|
|
"zhixing_b1",
|
|
|
|
|
ZHIXING_B1_SIGNAL_ORDER[-1].value,
|
|
|
|
|
Decimal("10.5"),
|
|
|
|
|
{},
|
|
|
|
|
),
|
|
|
|
|
(
|
|
|
|
|
"000001.SZ",
|
|
|
|
|
"平安银行",
|
|
|
|
|
TARGET,
|
|
|
|
|
"zhixing_b1",
|
2026-08-28 11:39:05 +08:00
|
|
|
ZhixingB1Category.ORIGINAL_B1.value,
|
2026-08-09 09:34:46 +08:00
|
|
|
Decimal("10.5"),
|
|
|
|
|
{},
|
|
|
|
|
),
|
|
|
|
|
]
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
2026-08-28 11:39:05 +08:00
|
|
|
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 DISTINCT ts_code" in query:
|
|
|
|
|
self.statements.append((query, parameters))
|
|
|
|
|
return LoadResult(rows=[])
|
|
|
|
|
return super().execute(query, parameters)
|
|
|
|
|
|
|
|
|
|
|
2026-08-09 09:34:46 +08:00
|
|
|
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
|
|
|
|
|
|
|
|
|
|
|
2026-08-28 11:39:05 +08:00
|
|
|
def test_get_run_pages_stocks_and_loads_all_signals_for_category_matches(
|
|
|
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
|
|
|
) -> None:
|
2026-08-09 09:34:46 +08:00
|
|
|
connection = LoadConnection()
|
|
|
|
|
|
|
|
|
|
def connect(database_url: str) -> LoadConnection:
|
|
|
|
|
assert database_url == "postgresql://test"
|
|
|
|
|
return connection
|
|
|
|
|
|
|
|
|
|
monkeypatch.setattr(psycopg, "connect", connect)
|
2026-08-10 11:09:23 +08:00
|
|
|
run = PostgresSelectionRunRepository("postgresql://test").get_run(
|
|
|
|
|
"run-1",
|
|
|
|
|
query=SelectionResultQuery(
|
2026-08-28 11:39:05 +08:00
|
|
|
page=1,
|
2026-08-10 11:09:23 +08:00
|
|
|
page_size=1,
|
|
|
|
|
search="100%",
|
|
|
|
|
category="pullback",
|
|
|
|
|
),
|
|
|
|
|
)
|
2026-08-09 09:34:46 +08:00
|
|
|
|
|
|
|
|
assert run is not None
|
|
|
|
|
assert [signal.category for signal in run.signals] == [
|
2026-08-28 11:39:05 +08:00
|
|
|
ZhixingB1Category.ORIGINAL_B1,
|
2026-08-09 09:34:46 +08:00
|
|
|
ZHIXING_B1_SIGNAL_ORDER[-1],
|
|
|
|
|
]
|
2026-08-28 11:39:05 +08:00
|
|
|
assert run.stocks_total == 2
|
2026-08-10 11:09:23 +08:00
|
|
|
count_query, count_parameters = next(
|
|
|
|
|
(query, parameters)
|
|
|
|
|
for query, parameters in connection.statements
|
2026-08-28 11:39:05 +08:00
|
|
|
if "COUNT(DISTINCT ts_code) FROM selection_signal" in query
|
2026-08-10 11:09:23 +08:00
|
|
|
)
|
|
|
|
|
assert "name ILIKE %s ESCAPE" in count_query
|
|
|
|
|
assert count_parameters == ("run-1", "%100\\%%", "%100\\%%", "zhixing_b1_pullback_%")
|
2026-08-28 11:39:05 +08:00
|
|
|
stock_page_query, page_parameters = next(
|
|
|
|
|
(query, parameters)
|
|
|
|
|
for query, parameters in connection.statements
|
|
|
|
|
if "SELECT DISTINCT ts_code" in query
|
|
|
|
|
)
|
|
|
|
|
assert "ORDER BY ts_code" 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(
|
2026-08-10 11:09:23 +08:00
|
|
|
(query, parameters)
|
|
|
|
|
for query, parameters in connection.statements
|
2026-08-28 11:39:05 +08:00
|
|
|
if "SELECT DISTINCT ts_code" in query
|
2026-08-10 11:09:23 +08:00
|
|
|
)
|
2026-08-28 11:39:05 +08:00
|
|
|
assert "ORDER BY ts_code" 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)
|