perf(selection): 优化选股执行性能

This commit is contained in:
yuxuanhui
2026-08-12 09:45:16 +08:00
parent dd04933d63
commit 8963c067b3
22 changed files with 1333 additions and 120 deletions
+18 -1
View File
@@ -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