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
@@ -4,7 +4,7 @@ from __future__ import annotations
import json
from collections import defaultdict
from collections.abc import Generator, Mapping
from collections.abc import Generator, Mapping, Sequence
from contextlib import contextmanager
from datetime import date, datetime
from decimal import Decimal
@@ -28,6 +28,7 @@ from ..domain.runs import (
SelectionRunStoreError,
)
from ..domain.zhixing_b1 import ZHIXING_B1_SIGNAL_ORDER
from .postgres_pool import SelectionConnectionPool, SelectionPostgresPool
_SIGNAL_PRIORITY = {category: index for index, category in enumerate(ZHIXING_B1_SIGNAL_ORDER)}
_CATEGORY_PREFIXES = {
@@ -45,13 +46,52 @@ _SIGNAL_ORDER_SQL = (
)
_ITEM_UPSERT = """
INSERT INTO selection_run_item
(run_id, ts_code, name, status, signal_count, reason)
VALUES (%s, %s, %s, %s, %s, %s)
ON CONFLICT (run_id, ts_code) DO UPDATE SET
name = EXCLUDED.name,
status = EXCLUDED.status,
signal_count = EXCLUDED.signal_count,
reason = EXCLUDED.reason
"""
_SIGNAL_UPSERT = """
INSERT INTO selection_signal
(
run_id, ts_code, name, target_trade_date, strategy,
category, close, details
)
VALUES (%s, %s, %s, %s, %s, %s, %s, %s)
ON CONFLICT (run_id, ts_code, category) DO UPDATE SET
name = EXCLUDED.name,
close = EXCLUDED.close,
details = EXCLUDED.details
"""
class PostgresSelectionRunRepository(SelectionRunStore):
"""Persist one current result attempt per strategy and target date."""
def __init__(self, database_url: str) -> None:
"""Create the adapter with an injected PostgreSQL URL."""
def __init__(
self,
database_url: str,
*,
pool: SelectionPostgresPool | SelectionConnectionPool | None = None,
) -> None:
"""Create the adapter with a URL and optional shared PostgreSQL pool."""
self.database_url = database_url
if isinstance(pool, SelectionPostgresPool):
self.pool: SelectionPostgresPool | None = pool
elif pool is not None:
self.pool = SelectionPostgresPool(
database_url,
max_connections=1,
pool=pool,
)
else:
self.pool = None
def prepare_run(
self,
@@ -137,62 +177,62 @@ class PostgresSelectionRunRepository(SelectionRunStore):
)
def record_item(self, run_id: str, item: SelectionRunItem) -> None:
"""Upsert one stock outcome and all of its independent signal rows."""
"""Persist one item through the batch path for compatibility."""
self.record_items(run_id, (item,))
def record_items(self, run_id: str, items: Sequence[SelectionRunItem]) -> None:
"""Persist one chunk in one transaction with set-based driver calls.
Existing signal rows are removed before the upserts so retrying a
chunk cannot retain a category that disappeared from a recalculation.
``executemany`` is used for both materialized tables; the small
fallback keeps the direct fake connections used by older tests usable.
"""
if not items:
return
item_values = tuple(
(
run_id,
item.ts_code,
item.name,
item.status,
item.signal_count,
item.reason,
)
for item in items
)
signal_values = tuple(
(
run_id,
signal.ts_code,
signal.name,
signal.target_trade_date,
signal.strategy,
signal.category.value,
signal.close,
Jsonb(dict(signal.details)),
)
for item in items
for signal in item.signals
)
codes = [item.ts_code for item in items]
try:
with self._connection() as connection, connection.transaction():
connection.execute(
"""
INSERT INTO selection_run_item
(run_id, ts_code, name, status, signal_count, reason)
VALUES (%s, %s, %s, %s, %s, %s)
ON CONFLICT (run_id, ts_code) DO UPDATE SET
name = EXCLUDED.name,
status = EXCLUDED.status,
signal_count = EXCLUDED.signal_count,
reason = EXCLUDED.reason
""",
(
run_id,
item.ts_code,
item.name,
item.status,
item.signal_count,
item.reason,
),
"DELETE FROM selection_signal WHERE run_id = %s AND ts_code = ANY(%s)",
(run_id, codes),
)
connection.execute(
"DELETE FROM selection_signal WHERE run_id = %s AND ts_code = %s",
(run_id, item.ts_code),
)
for signal in item.signals:
connection.execute(
"""
INSERT INTO selection_signal
(
run_id, ts_code, name, target_trade_date, strategy,
category, close, details
)
VALUES (%s, %s, %s, %s, %s, %s, %s, %s)
ON CONFLICT (run_id, ts_code, category) DO UPDATE SET
name = EXCLUDED.name,
close = EXCLUDED.close,
details = EXCLUDED.details
""",
(
run_id,
signal.ts_code,
signal.name,
signal.target_trade_date,
signal.strategy,
signal.category.value,
signal.close,
Jsonb(dict(signal.details)),
),
)
except psycopg.Error as exc:
_executemany(connection, _ITEM_UPSERT, item_values)
if signal_values:
_executemany(connection, _SIGNAL_UPSERT, signal_values)
except SelectionRunError:
raise
except Exception as exc: # noqa: BLE001 - redact driver/pool details
code_context = items[0].ts_code if len(items) == 1 else f"{len(items)} items"
raise SelectionRunStoreError(
f"failed to persist selection item {item.ts_code}"
f"failed to persist selection item {code_context}"
) from exc
def finish_run(
@@ -398,12 +438,35 @@ class PostgresSelectionRunRepository(SelectionRunStore):
"""Translate psycopg failures without exposing driver details."""
try:
with psycopg.connect(self.database_url) as connection:
yield connection
except psycopg.Error as exc:
if self.pool is None:
with psycopg.connect(self.database_url) as connection:
yield connection
else:
with self.pool.connection() as connection:
yield connection
except SelectionRunError:
raise
except Exception as exc: # noqa: BLE001 - normalize driver/pool errors
raise SelectionRunStoreError("selection database operation failed") from exc
def _executemany(connection: Any, query: str, parameters: Sequence[tuple[object, ...]]) -> None:
"""Use psycopg's batch API while retaining a minimal fake connection seam."""
executemany = getattr(connection, "executemany", None)
if callable(executemany):
executemany(query, parameters)
return
cursor_factory = getattr(connection, "cursor", None)
if callable(cursor_factory):
cursor_context = cast(Any, cursor_factory())
with cursor_context as cursor:
cursor.executemany(query, parameters)
return
for values in parameters:
connection.execute(query, values)
def _signal_from_row(row: tuple[object, ...]) -> SelectionSignal:
"""Map a persisted signal row back to the domain signal model."""