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
@@ -2,6 +2,7 @@
from __future__ import annotations
from collections.abc import Sequence
from datetime import date
from ..domain.models import SelectionEvaluation, StockHistory
@@ -44,3 +45,12 @@ class EvaluateZhixingB1:
"""Evaluate an already loaded history for deterministic unit tests."""
return self.strategy.evaluate(history, target_trade_date)
def execute_histories(
self,
histories: Sequence[StockHistory],
target_trade_date: date,
) -> tuple[SelectionEvaluation, ...]:
"""Evaluate loaded histories without issuing one read per stock."""
return tuple(self.execute_history(history, target_trade_date) for history in histories)
@@ -3,12 +3,17 @@
from __future__ import annotations
import logging
import time
from collections.abc import Callable, Sequence
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass
from datetime import date
from typing import Literal, Protocol
from typing import Literal, Protocol, cast
from ..domain.models import SelectionEvaluation
from ..domain.models import SelectionEvaluation, StockHistory
from ..domain.runs import (
BatchSelectionRunStore,
BatchSelectionUniverseReader,
SelectionExecutionSource,
SelectionRerunRequired,
SelectionResultQuery,
@@ -17,6 +22,7 @@ from ..domain.runs import (
SelectionRunItem,
SelectionRunStatus,
SelectionRunStore,
SelectionStock,
SelectionUniverseReader,
)
from .evaluate import EvaluateZhixingB1
@@ -48,12 +54,21 @@ class RunZhixingB1:
reader: SelectionUniverseReader,
store: SelectionRunStore,
evaluator: SelectionEvaluator | None = None,
*,
max_workers: int = 4,
batch_size: int = 200,
) -> None:
"""Inject storage ports and optionally a test evaluator."""
"""Inject storage ports and configure bounded chunk execution."""
if max_workers < 1:
raise ValueError("max_workers must be at least 1")
if batch_size < 1:
raise ValueError("batch_size must be at least 1")
self.reader = reader
self.store = store
self.evaluator = evaluator or EvaluateZhixingB1(reader)
self.max_workers = max_workers
self.batch_size = batch_size
def prepare(
self,
@@ -81,35 +96,57 @@ class RunZhixingB1:
returns so the UI never mistakes a lost worker exception for success.
"""
stocks = _unique_stocks(prepared.source.stocks)
evaluated_count = 0
selected_stock_count = 0
signal_count = 0
failed_count = 0
history_rows = 0
batch_count = _chunk_count(len(stocks), self.batch_size)
read_seconds = 0.0
evaluate_seconds = 0.0
persist_seconds = 0.0
try:
for stock in prepared.source.stocks:
try:
evaluation = self.evaluator.execute(
stock.ts_code,
with ThreadPoolExecutor(max_workers=self.max_workers) as executor:
for batch_stocks in _chunks(stocks, self.batch_size):
read_started = time.perf_counter()
histories = self._load_histories(
batch_stocks,
prepared.source.target_trade_date,
)
except Exception as exc: # noqa: BLE001 - isolate one stock from the batch
logger.exception(
"selection_item_failed run_id=%s ts_code=%s",
prepared.run.id,
stock.ts_code,
read_seconds += time.perf_counter() - read_started
history_rows += sum(
len(history.bars) for history in histories if history is not None
)
evaluation = SelectionEvaluation(
ts_code=stock.ts_code,
target_trade_date=prepared.source.target_trade_date,
status="data_error",
reason=_safe_item_error(exc),
evaluate_started = time.perf_counter()
items = tuple(
_to_item(
stock.ts_code,
stock.name,
evaluation,
)
for stock, evaluation in zip(
batch_stocks,
executor.map(
self._evaluate_stock,
batch_stocks,
histories,
[prepared.source.target_trade_date] * len(batch_stocks),
),
strict=True,
)
)
item = _to_item(stock.ts_code, stock.name, evaluation)
self.store.record_item(prepared.run.id, item)
evaluated_count += 1
selected_stock_count += evaluation.status == "selected"
signal_count += len(evaluation.signals)
failed_count += evaluation.status in _FAILURE_STATUSES
evaluate_seconds += time.perf_counter() - evaluate_started
evaluated_count += len(items)
selected_stock_count += sum(item.status == "selected" for item in items)
signal_count += sum(item.signal_count for item in items)
failed_count += sum(item.status in _FAILURE_STATUSES for item in items)
persist_started = time.perf_counter()
self._record_items(prepared.run.id, items)
persist_seconds += time.perf_counter() - persist_started
status = _run_status(evaluated_count, failed_count)
self.store.finish_run(
@@ -121,7 +158,12 @@ class RunZhixingB1:
failed_count=failed_count,
)
except Exception as exc: # noqa: BLE001 - worker boundary must persist failure state
logger.exception("selection_run_failed run_id=%s", prepared.run.id)
logger.error(
"selection_run_failed run_id=%s error_type=%s reason=%s",
prepared.run.id,
exc.__class__.__name__,
_safe_item_error(exc),
)
try:
self.store.finish_run(
prepared.run.id,
@@ -134,7 +176,94 @@ class RunZhixingB1:
error_message=str(exc),
)
except Exception: # noqa: BLE001 - preserve the original worker failure
logger.exception("selection_run_failure_persist_failed run_id=%s", prepared.run.id)
logger.error(
"selection_run_failure_persist_failed run_id=%s",
prepared.run.id,
)
finally:
logger.info(
"selection_run_summary run_id=%s stock_count=%d history_rows=%d "
"batch_count=%d worker_count=%d read_seconds=%.3f "
"evaluate_seconds=%.3f persist_seconds=%.3f",
prepared.run.id,
len(stocks),
history_rows,
batch_count,
self.max_workers,
read_seconds,
evaluate_seconds,
persist_seconds,
)
def _load_histories(
self,
stocks: Sequence[SelectionStock],
target_trade_date: date,
) -> tuple[StockHistory | None, ...]:
"""Load one chunk when the reader supports it, with old-path fallback."""
typed_stocks = tuple(stocks)
loader = getattr(self.reader, "load_histories", None)
if callable(loader):
batch_reader = cast(BatchSelectionUniverseReader, self.reader)
loaded = batch_reader.load_histories(typed_stocks, target_trade_date)
histories_by_code = {history.ts_code: history for history in loaded}
return tuple(
histories_by_code.get(
stock.ts_code,
StockHistory(ts_code=stock.ts_code, name=stock.name),
)
for stock in typed_stocks
)
if isinstance(self.evaluator, EvaluateZhixingB1):
return tuple(
self.reader.load_history(stock.ts_code, target_trade_date) for stock in typed_stocks
)
return (None,) * len(typed_stocks)
def _evaluate_stock(
self,
stock: SelectionStock,
history: StockHistory | None,
target_trade_date: date,
) -> SelectionEvaluation:
"""Evaluate one stock inside a worker and isolate its exception."""
ts_code = stock.ts_code
try:
execute_history: Callable[[StockHistory, date], SelectionEvaluation] | None = getattr(
self.evaluator,
"execute_history",
None,
)
if history is not None and execute_history is not None:
return execute_history(history, target_trade_date)
return self.evaluator.execute(ts_code, target_trade_date)
except Exception as exc: # noqa: BLE001 - isolate one stock from the batch
logger.warning(
"selection_item_failed ts_code=%s error_type=%s reason=%s",
ts_code,
exc.__class__.__name__,
_safe_item_error(exc),
)
return SelectionEvaluation(
ts_code=ts_code,
target_trade_date=target_trade_date,
status="data_error",
reason=_safe_item_error(exc),
)
def _record_items(self, run_id: str, items: Sequence[SelectionRunItem]) -> None:
"""Use batch persistence while retaining the old single-item seam."""
record_items = getattr(self.store, "record_items", None)
if callable(record_items):
batch_store = cast(BatchSelectionRunStore, self.store)
batch_store.record_items(run_id, tuple(items))
return
for item in items:
self.store.record_item(run_id, item)
def get_run(
self,
@@ -187,6 +316,35 @@ def _safe_item_error(error: Exception) -> str:
return " ".join(str(error).split())[:500] or error.__class__.__name__
def _chunks(
values: Sequence[SelectionStock],
size: int,
) -> tuple[tuple[SelectionStock, ...], ...]:
"""Split a stable stock sequence into bounded immutable chunks."""
return tuple(tuple(values[index : index + size]) for index in range(0, len(values), size))
def _chunk_count(value_count: int, size: int) -> int:
"""Return the number of chunks without materializing empty chunks."""
return (value_count + size - 1) // size
def _unique_stocks(stocks: Sequence[SelectionStock]) -> tuple[SelectionStock, ...]:
"""Keep the first source row for each stock so it is evaluated once."""
seen: set[str] = set()
unique: list[SelectionStock] = []
for stock in stocks:
ts_code = stock.ts_code
if ts_code in seen:
continue
seen.add(ts_code)
unique.append(stock)
return tuple(unique)
__all__ = [
"PreparedSelectionRun",
"RunZhixingB1",
@@ -2,6 +2,7 @@
from __future__ import annotations
from collections.abc import Sequence
from dataclasses import dataclass, field
from datetime import date, datetime
from decimal import Decimal
@@ -150,3 +151,19 @@ class SelectionUniverseReader(Protocol):
) -> SelectionExecutionSource: ...
def load_history(self, ts_code: str, target_trade_date: date) -> StockHistory: ...
class BatchSelectionRunStore(Protocol):
"""Optional batch-write extension for selection stores."""
def record_items(self, run_id: str, items: Sequence[SelectionRunItem]) -> None: ...
class BatchSelectionUniverseReader(Protocol):
"""Optional bounded batch-history extension for selection readers."""
def load_histories(
self,
stocks: Sequence[SelectionStock],
target_trade_date: date,
) -> tuple[StockHistory, ...]: ...
@@ -0,0 +1,96 @@
"""Bounded PostgreSQL pool ownership for selection adapters."""
from __future__ import annotations
import threading
from collections.abc import Generator
from contextlib import contextmanager
from typing import Any, Protocol, cast
from psycopg_pool import ConnectionPool
class SelectionConnectionPool(Protocol):
"""Small pool surface shared by psycopg and unit-test fakes."""
def open(self, *, wait: bool = True) -> None: ...
def close(self) -> None: ...
def connection(self) -> Any: ...
class SelectionPostgresPool:
"""Own one bounded PostgreSQL pool for the selection read/write adapters.
The owner opens lazily on the first borrowed connection, which keeps app
construction cheap while still making the pool lifetime process-scoped in
the HTTP composition layer. A fake pool can be injected for unit tests.
"""
def __init__(
self,
database_url: str,
*,
max_connections: int,
pool: SelectionConnectionPool | None = None,
) -> None:
"""Create a bounded pool owner with an optional injected pool."""
if max_connections < 1:
raise ValueError("max_connections must be at least 1")
self.database_url = database_url
self.max_connections = max_connections
self.pool: SelectionConnectionPool = pool or cast(
SelectionConnectionPool,
ConnectionPool(
conninfo=database_url,
min_size=1,
max_size=max_connections,
open=False,
),
)
self._pool_open = False
self._pool_state_lock = threading.Lock()
def open(self) -> None:
"""Open the underlying pool once and wait for its minimum connection."""
with self._pool_state_lock:
if self._pool_open:
return
if bool(getattr(self.pool, "_opened", False)):
self._pool_open = True
return
self.pool.open(wait=True)
self._pool_open = True
def close(self) -> None:
"""Close the pool after borrowed connections have been returned."""
with self._pool_state_lock:
if self._pool_open or bool(getattr(self.pool, "_opened", False)):
self.pool.close()
self._pool_open = False
@contextmanager
def connection(self) -> Generator[Any, None, None]:
"""Borrow one connection and return it to the bounded pool."""
self.open()
with self.pool.connection() as connection:
yield connection
def __enter__(self) -> SelectionPostgresPool:
"""Open and return this resource owner."""
self.open()
return self
def __exit__(self, exc_type: object, exc_value: object, traceback: object) -> None:
"""Release the owned pool at the end of a context."""
self.close()
__all__ = ["SelectionConnectionPool", "SelectionPostgresPool"]
@@ -2,9 +2,11 @@
from __future__ import annotations
from collections.abc import Generator, Sequence
from contextlib import contextmanager
from datetime import date, datetime
from decimal import Decimal, InvalidOperation
from typing import cast
from typing import Any, cast
import psycopg
@@ -12,6 +14,7 @@ from ....bootstrap.config import Settings
from ..domain.models import SelectionBar, SelectionDailyBasic, StockHistory
from ..domain.ports import MarketDataReaderError
from ..domain.runs import SelectionExecutionSource, SelectionStock
from .postgres_pool import SelectionConnectionPool, SelectionPostgresPool
class SelectionReaderError(MarketDataReaderError):
@@ -23,6 +26,25 @@ class SelectionMarketDataNotReady(MarketDataReaderError):
_HISTORY_QUERY = """
SELECT
bar.ts_code,
stock.name,
bar.trade_date,
bar.open,
bar.high,
bar.low,
bar.close,
bar.vol
FROM market_daily_bar AS bar
LEFT JOIN market_stock AS stock
ON stock.ts_code = bar.ts_code
WHERE bar.ts_code = ANY(%s)
AND bar.source_adj = 'qfq'
AND bar.trade_date <= %s
ORDER BY bar.ts_code ASC, bar.trade_date ASC
"""
_SINGLE_HISTORY_QUERY = """
SELECT
bar.ts_code,
stock.name,
@@ -111,10 +133,30 @@ def _as_float(value: object) -> float | None:
class PostgresMarketDataReader:
"""Load qfq bars and same-day basic facts without writing market data."""
def __init__(self, settings: Settings | str) -> None:
"""Create a reader from injected settings or a compatible URL string."""
def __init__(
self,
settings: Settings | str,
*,
pool: SelectionPostgresPool | SelectionConnectionPool | None = None,
) -> None:
"""Create a reader from settings/URL and an optional shared pool.
Omitting ``pool`` intentionally retains the direct ``psycopg.connect``
path used by one-shot callers and existing adapter tests. The HTTP
composition layer always supplies the process-scoped selection pool.
"""
self.database_url = settings.database_url if isinstance(settings, Settings) else settings
if isinstance(pool, SelectionPostgresPool):
self.pool: SelectionPostgresPool | None = pool
elif pool is not None:
self.pool = SelectionPostgresPool(
self.database_url,
max_connections=1,
pool=pool,
)
else:
self.pool = None
def load_history(self, ts_code: str, target_trade_date: date) -> StockHistory:
"""Read all retained qfq rows through the explicit target date.
@@ -133,34 +175,57 @@ class PostgresMarketDataReader:
"""
try:
with psycopg.connect(self.database_url) as connection:
with self._connection() as connection:
rows = connection.execute(
_HISTORY_QUERY,
_SINGLE_HISTORY_QUERY,
(ts_code, target_trade_date),
).fetchall()
except psycopg.Error as exc:
except SelectionReaderError:
raise
except Exception as exc: # noqa: BLE001 - redact driver/pool details at the port boundary
raise SelectionReaderError(
f"failed to load market history for {ts_code} at {target_trade_date.isoformat()}"
) from exc
bars: dict[date, SelectionBar] = {}
daily_basic: dict[date, SelectionDailyBasic] = {}
name = ""
for raw_row in rows:
row = cast(tuple[object, ...], raw_row)
row_code, row_name, bar, basic = self._map_row(row, ts_code)
if row_code != ts_code:
raise ValueError(f"reader returned unexpected stock code: {row_code}")
name = row_name or name
if bar.trade_date <= target_trade_date:
bars[bar.trade_date] = bar
daily_basic[bar.trade_date] = basic
return StockHistory(
ts_code=ts_code,
name=name,
bars=tuple(bars[trade_date] for trade_date in sorted(bars)),
daily_basic={trade_date: daily_basic[trade_date] for trade_date in sorted(daily_basic)},
return self._histories_from_rows(
rows,
(SelectionStock(ts_code, ""),),
target_trade_date,
)[0]
def load_histories(
self,
stocks: Sequence[SelectionStock] | Sequence[str],
target_trade_date: date,
) -> tuple[StockHistory, ...]:
"""Read one bounded stock chunk with one parameterized qfq query.
Historical daily-basic values are deliberately not joined here: B1
only needs OHLCV for its historical formula. The execution-source
query still requires a complete target-day basic row before a stock is
admitted to a run.
"""
normalized = tuple(
stock if isinstance(stock, SelectionStock) else SelectionStock(stock, "")
for stock in stocks
)
if not normalized:
return ()
codes = [stock.ts_code for stock in normalized]
try:
with self._connection() as connection:
rows = connection.execute(
_HISTORY_QUERY,
(codes, target_trade_date),
).fetchall()
except SelectionReaderError:
raise
except Exception as exc: # noqa: BLE001 - redact driver/pool details at the port boundary
raise SelectionReaderError(
f"failed to load market history batch at {target_trade_date.isoformat()}"
) from exc
return self._histories_from_rows(rows, normalized, target_trade_date)
def load_execution_source(
self,
@@ -187,7 +252,7 @@ class PostgresMarketDataReader:
if strategy != "zhixing_b1":
raise SelectionMarketDataNotReady(f"unsupported selection strategy: {strategy}")
try:
with psycopg.connect(self.database_url) as connection:
with self._connection() as connection:
source_row = connection.execute(_SOURCE_QUERY, (target_trade_date,)).fetchone()
if source_row is None:
raise SelectionMarketDataNotReady(
@@ -197,9 +262,9 @@ class PostgresMarketDataReader:
_ELIGIBLE_STOCKS_QUERY,
(target_trade_date, target_trade_date),
).fetchall()
except SelectionMarketDataNotReady:
except (SelectionMarketDataNotReady, SelectionReaderError):
raise
except psycopg.Error as exc:
except Exception as exc: # noqa: BLE001 - redact driver/pool details at the port boundary
raise SelectionReaderError(
f"failed to load selection source at {target_trade_date.isoformat()}"
) from exc
@@ -224,8 +289,8 @@ class PostgresMarketDataReader:
def _map_row(
row: tuple[object, ...],
expected_code: str,
) -> tuple[str, str, SelectionBar, SelectionDailyBasic]:
"""Map the current query row, tolerating a legacy test row without name."""
) -> tuple[str, str, SelectionBar, SelectionDailyBasic | None]:
"""Map qfq OHLCV rows and tolerate the legacy basic-join test shape."""
if len(row) >= 10:
code, raw_name, raw_date = row[0], row[1], row[2]
@@ -233,13 +298,19 @@ class PostgresMarketDataReader:
elif len(row) >= 9:
code, raw_name, raw_date = row[0], "", row[1]
values = row[2:]
elif len(row) >= 8:
code, raw_name, raw_date = row[0], row[1], row[2]
values = row[3:]
elif len(row) >= 7:
code, raw_name, raw_date = row[0], "", row[1]
values = row[2:]
else:
raise ValueError("market history row has too few columns")
row_code = str(code or expected_code)
name = str(raw_name or "")
trade_date = _as_date(raw_date)
if len(values) < 7:
raise ValueError("market history row is missing OHLCV/basic columns")
if len(values) < 5:
raise ValueError("market history row is missing OHLCV columns")
bar = SelectionBar(
trade_date=trade_date,
open=_as_float(values[0]),
@@ -248,9 +319,70 @@ class PostgresMarketDataReader:
close=_as_float(values[3]),
volume=_as_float(values[4]),
)
basic = SelectionDailyBasic(
trade_date=trade_date,
turnover_rate=_as_float(values[5]),
total_mv=_as_float(values[6]),
basic = (
SelectionDailyBasic(
trade_date=trade_date,
turnover_rate=_as_float(values[5]),
total_mv=_as_float(values[6]),
)
if len(values) >= 7
else None
)
return row_code, name, bar, basic
@classmethod
def _histories_from_rows(
cls,
rows: Sequence[object],
stocks: Sequence[SelectionStock],
target_trade_date: date,
) -> tuple[StockHistory, ...]:
"""Group sorted/possibly duplicated database rows by requested stock."""
requested = {stock.ts_code: stock for stock in stocks}
bars_by_code: dict[str, dict[date, SelectionBar]] = {code: {} for code in requested}
basics_by_code: dict[str, dict[date, SelectionDailyBasic]] = {
code: {} for code in requested
}
names = {stock.ts_code: stock.name for stock in stocks}
for raw_row in rows:
row = cast(tuple[object, ...], raw_row)
row_hint = str(row[0]) if row and row[0] is not None else ""
row_code, row_name, bar, basic = cls._map_row(row, row_hint)
if row_code not in requested:
raise ValueError(f"reader returned unexpected stock code: {row_code}")
names[row_code] = row_name or names[row_code]
if bar.trade_date <= target_trade_date:
bars_by_code[row_code][bar.trade_date] = bar
if basic is not None:
basics_by_code[row_code][bar.trade_date] = basic
histories: list[StockHistory] = []
for stock in stocks:
code = stock.ts_code
bars = bars_by_code[code]
basics = basics_by_code[code]
histories.append(
StockHistory(
ts_code=code,
name=names[code],
bars=tuple(bars[trade_date] for trade_date in sorted(bars)),
daily_basic={trade_date: basics[trade_date] for trade_date in sorted(basics)},
)
)
return tuple(histories)
@contextmanager
def _connection(self) -> Generator[Any, None, None]:
"""Borrow from the shared pool or use the legacy direct connection."""
try:
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 (SelectionMarketDataNotReady, SelectionReaderError, ValueError):
raise
except Exception as exc: # noqa: BLE001 - normalize pool/driver failures
raise SelectionReaderError("selection database operation failed") from exc
@@ -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."""
@@ -1,5 +1,7 @@
"""HTTP presentation for persisted strategy execution results."""
import atexit
import threading
from datetime import date, datetime
from typing import Annotated, Literal
@@ -17,6 +19,7 @@ from zhixing_server.modules.selection.domain.runs import (
SelectionRunInProgress,
SelectionRunStoreError,
)
from zhixing_server.modules.selection.infrastructure.postgres_pool import SelectionPostgresPool
from zhixing_server.modules.selection.infrastructure.postgres_reader import (
PostgresMarketDataReader,
SelectionMarketDataNotReady,
@@ -27,6 +30,8 @@ from zhixing_server.modules.selection.infrastructure.postgres_runs import (
)
selection_router = APIRouter()
_SELECTION_POOL_CACHE_LOCK = threading.Lock()
_SELECTION_POOL_CACHE: dict[tuple[str, int], SelectionPostgresPool] = {}
StrategyValue = Literal["zhixing_b1"]
SelectionStatusValue = Literal[
@@ -117,11 +122,45 @@ class SelectionResultsResponse(BaseModel):
def get_selection_service(
settings: Annotated[Settings, Depends(get_settings)],
) -> RunZhixingB1:
"""Build one request-scoped selection application service."""
"""Build the selection service on top of process-scoped shared resources."""
reader = PostgresMarketDataReader(settings)
store = PostgresSelectionRunRepository(settings.database_url)
return RunZhixingB1(reader, store)
pool = get_selection_postgres_pool(settings)
reader = PostgresMarketDataReader(settings, pool=pool)
store = PostgresSelectionRunRepository(settings.database_url, pool=pool)
return RunZhixingB1(
reader,
store,
max_workers=settings.selection_max_workers,
batch_size=settings.selection_batch_size,
)
def get_selection_postgres_pool(settings: Settings) -> SelectionPostgresPool:
"""Return the cached bounded pool shared by selection adapters."""
key = (settings.database_url, settings.selection_max_workers + 2)
with _SELECTION_POOL_CACHE_LOCK:
pool = _SELECTION_POOL_CACHE.get(key)
if pool is None:
pool = SelectionPostgresPool(
settings.database_url,
max_connections=key[1],
)
_SELECTION_POOL_CACHE[key] = pool
return pool
def _close_cached_selection_pools() -> None:
"""Close all process-cached selection pools during interpreter shutdown."""
with _SELECTION_POOL_CACHE_LOCK:
pools = tuple(_SELECTION_POOL_CACHE.values())
_SELECTION_POOL_CACHE.clear()
for pool in pools:
pool.close()
atexit.register(_close_cached_selection_pools)
@selection_router.post(
@@ -292,5 +331,6 @@ __all__ = [
"SelectionRunAcceptedResponse",
"SelectionRunRequest",
"get_selection_service",
"get_selection_postgres_pool",
"selection_router",
]