perf(selection): 优化选股执行性能
This commit is contained in:
+119
-56
@@ -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."""
|
||||
|
||||
|
||||
Reference in New Issue
Block a user