perf(selection): 优化选股执行性能
This commit is contained in:
@@ -61,6 +61,33 @@ class FakeConnection:
|
||||
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",
|
||||
@@ -163,6 +190,64 @@ def test_record_item_persists_independent_signals_as_jsonb(monkeypatch: pytest.M
|
||||
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, ...]]] = []
|
||||
|
||||
Reference in New Issue
Block a user