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