perf(selection): 优化选股执行性能
This commit is contained in:
@@ -3,9 +3,12 @@
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
import zhixing_server.modules.selection.presentation.http as selection_http
|
||||
from zhixing_server.bootstrap.app import create_app
|
||||
from zhixing_server.bootstrap.config import Settings
|
||||
from zhixing_server.modules.selection.application.run import PreparedSelectionRun
|
||||
from zhixing_server.modules.selection.domain.models import SelectionSignal, ZhixingB1Category
|
||||
from zhixing_server.modules.selection.domain.runs import (
|
||||
@@ -19,7 +22,10 @@ from zhixing_server.modules.selection.domain.runs import (
|
||||
from zhixing_server.modules.selection.infrastructure.postgres_reader import (
|
||||
SelectionMarketDataNotReady,
|
||||
)
|
||||
from zhixing_server.modules.selection.presentation.http import get_selection_service
|
||||
from zhixing_server.modules.selection.presentation.http import (
|
||||
get_selection_postgres_pool,
|
||||
get_selection_service,
|
||||
)
|
||||
|
||||
TARGET = date(2026, 8, 8)
|
||||
|
||||
@@ -128,6 +134,17 @@ def _client(service: FakeSelectionService) -> TestClient:
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
def test_selection_dependency_reuses_one_bounded_pool(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(selection_http, "_SELECTION_POOL_CACHE", {})
|
||||
settings = Settings(database_url="postgresql://test", selection_max_workers=4)
|
||||
|
||||
first = get_selection_postgres_pool(settings)
|
||||
second = get_selection_postgres_pool(settings)
|
||||
|
||||
assert first is second
|
||||
assert first.max_connections == 6
|
||||
|
||||
|
||||
def test_trigger_returns_accepted_run_and_schedules_execution() -> None:
|
||||
service = FakeSelectionService()
|
||||
|
||||
|
||||
@@ -0,0 +1,45 @@
|
||||
"""Lifecycle tests for the shared selection PostgreSQL pool owner."""
|
||||
|
||||
from collections.abc import Generator
|
||||
from contextlib import contextmanager
|
||||
|
||||
from zhixing_server.modules.selection.infrastructure.postgres_pool import SelectionPostgresPool
|
||||
|
||||
|
||||
class FakePool:
|
||||
def __init__(self) -> None:
|
||||
self.open_calls = 0
|
||||
self.close_calls = 0
|
||||
self.connection_calls = 0
|
||||
|
||||
def open(self, *, wait: bool = True) -> None:
|
||||
assert wait is True
|
||||
self.open_calls += 1
|
||||
|
||||
def close(self) -> None:
|
||||
self.close_calls += 1
|
||||
|
||||
@contextmanager
|
||||
def connection(self) -> Generator[str, None, None]:
|
||||
self.connection_calls += 1
|
||||
yield "connection"
|
||||
|
||||
|
||||
def test_selection_pool_opens_once_borrows_and_closes_injected_pool() -> None:
|
||||
fake = FakePool()
|
||||
owner = SelectionPostgresPool(
|
||||
"postgresql://test",
|
||||
max_connections=6,
|
||||
pool=fake,
|
||||
)
|
||||
|
||||
with owner.connection() as connection:
|
||||
assert connection == "connection"
|
||||
with owner.connection() as connection:
|
||||
assert connection == "connection"
|
||||
|
||||
assert owner.max_connections == 6
|
||||
assert fake.open_calls == 1
|
||||
assert fake.connection_calls == 2
|
||||
owner.close()
|
||||
assert fake.close_calls == 1
|
||||
@@ -1,5 +1,7 @@
|
||||
"""PostgreSQL reader contract tests using a fake connection."""
|
||||
|
||||
from collections.abc import Generator
|
||||
from contextlib import contextmanager
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
from typing import cast
|
||||
@@ -7,6 +9,8 @@ from typing import cast
|
||||
import psycopg
|
||||
import pytest
|
||||
|
||||
from zhixing_server.modules.selection.domain.runs import SelectionStock
|
||||
from zhixing_server.modules.selection.infrastructure.postgres_pool import SelectionPostgresPool
|
||||
from zhixing_server.modules.selection.infrastructure.postgres_reader import (
|
||||
PostgresMarketDataReader,
|
||||
SelectionMarketDataNotReady,
|
||||
@@ -39,6 +43,23 @@ class FakeResult:
|
||||
return self.rows
|
||||
|
||||
|
||||
class Pool:
|
||||
def __init__(self, connection: FakeConnection) -> None:
|
||||
self.connection_value = connection
|
||||
self.opened = 0
|
||||
|
||||
def open(self, *, wait: bool = True) -> None:
|
||||
assert wait is True
|
||||
self.opened += 1
|
||||
|
||||
def close(self) -> None:
|
||||
return None
|
||||
|
||||
@contextmanager
|
||||
def connection(self) -> Generator[FakeConnection, None, None]:
|
||||
yield self.connection_value
|
||||
|
||||
|
||||
def test_reader_parameterizes_target_and_maps_left_join(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
connection = FakeConnection(
|
||||
[
|
||||
@@ -86,6 +107,64 @@ def test_reader_parameterizes_target_and_maps_left_join(monkeypatch: pytest.Monk
|
||||
assert "trade_date <= %s" in cast(str, connection.query)
|
||||
|
||||
|
||||
def test_reader_batches_qfq_rows_by_stock_without_historical_basic_join() -> None:
|
||||
connection = FakeConnection(
|
||||
[
|
||||
(
|
||||
"600000.SH",
|
||||
"浦发银行",
|
||||
date(2024, 1, 2),
|
||||
"8",
|
||||
"8.5",
|
||||
"7.8",
|
||||
"8.2",
|
||||
"900",
|
||||
),
|
||||
(
|
||||
"000001.SZ",
|
||||
"平安银行",
|
||||
date(2024, 1, 3),
|
||||
"10.5",
|
||||
"11",
|
||||
"10",
|
||||
"10.8",
|
||||
"1200",
|
||||
),
|
||||
(
|
||||
"000001.SZ",
|
||||
"平安银行",
|
||||
date(2024, 1, 2),
|
||||
"10",
|
||||
"11",
|
||||
"9",
|
||||
"10.5",
|
||||
"1000",
|
||||
),
|
||||
]
|
||||
)
|
||||
pool = Pool(connection)
|
||||
owner = SelectionPostgresPool("postgresql://test", max_connections=6, pool=pool)
|
||||
|
||||
histories = PostgresMarketDataReader("postgresql://test", pool=owner).load_histories(
|
||||
(
|
||||
SelectionStock("000001.SZ", "平安银行"),
|
||||
SelectionStock("600000.SH", "浦发银行"),
|
||||
),
|
||||
date(2024, 1, 3),
|
||||
)
|
||||
|
||||
assert [history.ts_code for history in histories] == ["000001.SZ", "600000.SH"]
|
||||
assert [bar.trade_date for bar in histories[0].bars] == [
|
||||
date(2024, 1, 2),
|
||||
date(2024, 1, 3),
|
||||
]
|
||||
assert histories[0].daily_basic == {}
|
||||
assert connection.parameters == (["000001.SZ", "600000.SH"], date(2024, 1, 3))
|
||||
assert "bar.ts_code = ANY(%s)" in cast(str, connection.query)
|
||||
assert "market_daily_basic" not in cast(str, connection.query)
|
||||
assert pool.opened == 1
|
||||
|
||||
|
||||
class SourceConnection:
|
||||
def __init__(self, source_row: tuple[object, ...] | None) -> None:
|
||||
self.source_row = source_row
|
||||
|
||||
@@ -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, ...]]] = []
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
"""Application tests for persisted whole-universe selection runs."""
|
||||
|
||||
import threading
|
||||
import time
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
from typing import Literal
|
||||
@@ -42,6 +44,30 @@ class FakeReader:
|
||||
raise AssertionError("the fake evaluator should be used")
|
||||
|
||||
|
||||
class BatchReader(FakeReader):
|
||||
def __init__(self, source: SelectionExecutionSource) -> None:
|
||||
super().__init__(source)
|
||||
self.batch_calls: list[tuple[str, ...]] = []
|
||||
|
||||
def load_histories(
|
||||
self,
|
||||
stocks: tuple[SelectionStock, ...],
|
||||
target_trade_date: date,
|
||||
) -> tuple[StockHistory, ...]:
|
||||
self.batch_calls.append(tuple(stock.ts_code for stock in stocks))
|
||||
return tuple(StockHistory(ts_code=stock.ts_code, name=stock.name) for stock in stocks)
|
||||
|
||||
|
||||
class PartialBatchReader(BatchReader):
|
||||
def load_histories(
|
||||
self,
|
||||
stocks: tuple[SelectionStock, ...],
|
||||
target_trade_date: date,
|
||||
) -> tuple[StockHistory, ...]:
|
||||
self.batch_calls.append(tuple(stock.ts_code for stock in stocks))
|
||||
return ()
|
||||
|
||||
|
||||
class FakeStore:
|
||||
def __init__(self) -> None:
|
||||
self.items: list[SelectionRunItem] = []
|
||||
@@ -114,6 +140,26 @@ class FakeStore:
|
||||
return None
|
||||
|
||||
|
||||
class BatchStore(FakeStore):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.batches: list[tuple[SelectionRunItem, ...]] = []
|
||||
|
||||
def record_item(self, run_id: str, item: SelectionRunItem) -> None:
|
||||
raise AssertionError("the batch path should use record_items")
|
||||
|
||||
def record_items(self, run_id: str, items: tuple[SelectionRunItem, ...]) -> None:
|
||||
assert run_id == "run-1"
|
||||
batch = tuple(items)
|
||||
self.batches.append(batch)
|
||||
self.items.extend(batch)
|
||||
|
||||
|
||||
class FailingBatchStore(BatchStore):
|
||||
def record_items(self, run_id: str, items: tuple[SelectionRunItem, ...]) -> None:
|
||||
raise RuntimeError("batch write unavailable")
|
||||
|
||||
|
||||
class FakeEvaluator:
|
||||
def __init__(self, results: dict[str, SelectionEvaluation]) -> None:
|
||||
self.results = results
|
||||
@@ -129,6 +175,29 @@ class RaisingEvaluator:
|
||||
return SelectionEvaluation(ts_code, target_trade_date, "no_signal")
|
||||
|
||||
|
||||
class ConcurrentHistoryEvaluator:
|
||||
def __init__(self) -> None:
|
||||
self.active = 0
|
||||
self.peak = 0
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def execute(self, ts_code: str, target_trade_date: date) -> SelectionEvaluation:
|
||||
raise AssertionError("the batch evaluator path should be used")
|
||||
|
||||
def execute_history(
|
||||
self,
|
||||
history: StockHistory,
|
||||
target_trade_date: date,
|
||||
) -> SelectionEvaluation:
|
||||
with self._lock:
|
||||
self.active += 1
|
||||
self.peak = max(self.peak, self.active)
|
||||
time.sleep(0.02)
|
||||
with self._lock:
|
||||
self.active -= 1
|
||||
return SelectionEvaluation(history.ts_code, target_trade_date, "no_signal")
|
||||
|
||||
|
||||
def _source() -> SelectionExecutionSource:
|
||||
return SelectionExecutionSource(
|
||||
market_sync_batch_id="market-run-1",
|
||||
@@ -233,3 +302,65 @@ def test_execute_isolates_unexpected_single_stock_failure() -> None:
|
||||
assert store.finished[0:2] == ("run-1", "partial_success")
|
||||
assert store.finished[2]["evaluated_count"] == 2
|
||||
assert store.finished[2]["failed_count"] == 1
|
||||
|
||||
|
||||
def test_execute_batches_history_reads_writes_and_limits_evaluation_workers() -> None:
|
||||
source = SelectionExecutionSource(
|
||||
market_sync_batch_id="market-run-1",
|
||||
target_trade_date=TARGET,
|
||||
target_count=8,
|
||||
valid_count=8,
|
||||
coverage=Decimal("1"),
|
||||
stocks=tuple(SelectionStock(f"{index:06d}.SZ", f"stock-{index}") for index in range(8)),
|
||||
)
|
||||
reader = BatchReader(source)
|
||||
store = BatchStore()
|
||||
evaluator = ConcurrentHistoryEvaluator()
|
||||
|
||||
service = RunZhixingB1(
|
||||
reader,
|
||||
store,
|
||||
evaluator,
|
||||
max_workers=4,
|
||||
batch_size=4,
|
||||
)
|
||||
service.execute(service.prepare("zhixing_b1", TARGET, rerun=False))
|
||||
|
||||
assert reader.batch_calls == [
|
||||
("000000.SZ", "000001.SZ", "000002.SZ", "000003.SZ"),
|
||||
("000004.SZ", "000005.SZ", "000006.SZ", "000007.SZ"),
|
||||
]
|
||||
assert [len(batch) for batch in store.batches] == [4, 4]
|
||||
assert evaluator.peak <= 4
|
||||
assert evaluator.peak >= 2
|
||||
assert store.finished is not None
|
||||
assert store.finished[0:2] == ("run-1", "success")
|
||||
|
||||
|
||||
def test_execute_maps_missing_batch_history_without_a_single_stock_read() -> None:
|
||||
source = _source()
|
||||
reader = PartialBatchReader(source)
|
||||
store = FakeStore()
|
||||
|
||||
service = RunZhixingB1(reader, store)
|
||||
service.execute(service.prepare("zhixing_b1", TARGET, rerun=False))
|
||||
|
||||
assert reader.batch_calls == [("000001.SZ", "600000.SH")]
|
||||
assert [item.status for item in store.items] == ["missing_target_bar", "missing_target_bar"]
|
||||
assert store.finished is not None
|
||||
assert store.finished[0:2] == ("run-1", "failed")
|
||||
|
||||
|
||||
def test_execute_marks_batch_write_failure_as_failed() -> None:
|
||||
source = _source()
|
||||
reader = BatchReader(source)
|
||||
store = FailingBatchStore()
|
||||
evaluator = ConcurrentHistoryEvaluator()
|
||||
|
||||
service = RunZhixingB1(reader, store, evaluator)
|
||||
service.execute(service.prepare("zhixing_b1", TARGET, rerun=False))
|
||||
|
||||
assert store.finished is not None
|
||||
assert store.finished[0:2] == ("run-1", "failed")
|
||||
assert store.finished[2]["error_type"] == "batch_error"
|
||||
assert store.finished[2]["failed_count"] == 1
|
||||
|
||||
Reference in New Issue
Block a user