Files
zhixing-system/zhixing-server/src/zhixing_server/modules/selection/application/run.py
T

445 lines
16 KiB
Python
Raw Normal View History

"""Application orchestration for persisted whole-universe B1 runs."""
from __future__ import annotations
import logging
2026-08-12 09:45:16 +08:00
import time
from collections.abc import Callable, Sequence
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass
from datetime import date
2026-08-12 09:45:16 +08:00
from typing import Literal, Protocol, cast
2026-08-12 09:45:16 +08:00
from ..domain.models import SelectionEvaluation, StockHistory
from ..domain.pattern_scoring import (
PatternCase,
PatternCaseLibraryLoader,
PatternScore,
PatternScorer,
)
from ..domain.runs import (
2026-08-12 09:45:16 +08:00
BatchSelectionRunStore,
BatchSelectionUniverseReader,
SelectionExecutionSource,
SelectionRerunRequired,
SelectionResultQuery,
SelectionRun,
SelectionRunInProgress,
SelectionRunItem,
SelectionRunStatus,
SelectionRunStore,
2026-08-12 09:45:16 +08:00
SelectionStock,
SelectionUniverseReader,
)
from .evaluate import EvaluateZhixingB1
logger = logging.getLogger(__name__)
StrategyName = Literal["zhixing_b1"]
_FAILURE_STATUSES = {"insufficient_history", "missing_target_bar", "data_error"}
class SelectionEvaluator(Protocol):
"""Minimal single-stock evaluator required by the batch orchestrator."""
def execute(self, ts_code: str, target_trade_date: date) -> SelectionEvaluation: ...
@dataclass(frozen=True, slots=True)
class PreparedSelectionRun:
"""A claimed run and its immutable market-data source snapshot."""
run: SelectionRun
source: SelectionExecutionSource
class RunZhixingB1:
"""Prepare, execute, and query persisted Zhixing B1 result batches."""
def __init__(
self,
reader: SelectionUniverseReader,
store: SelectionRunStore,
evaluator: SelectionEvaluator | None = None,
pattern_case_loader: PatternCaseLibraryLoader | None = None,
pattern_scorer: PatternScorer | None = None,
2026-08-12 09:45:16 +08:00
*,
pattern_scoring_enabled: bool = False,
2026-08-12 09:45:16 +08:00
max_workers: int = 4,
batch_size: int = 200,
) -> None:
2026-08-12 09:45:16 +08:00
"""Inject storage ports and configure bounded chunk execution."""
2026-08-12 09:45:16 +08:00
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.pattern_case_loader = pattern_case_loader
self.pattern_scorer = pattern_scorer
self.pattern_scoring_enabled = pattern_scoring_enabled
2026-08-12 09:45:16 +08:00
self.max_workers = max_workers
self.batch_size = batch_size
def prepare(
self,
strategy: StrategyName,
target_trade_date: date,
*,
rerun: bool,
) -> PreparedSelectionRun:
"""Validate source eligibility before claiming the rerunnable key."""
source = self.reader.load_execution_source(strategy, target_trade_date)
run = self.store.prepare_run(
strategy,
target_trade_date,
source,
rerun=rerun,
)
return PreparedSelectionRun(run=run, source=source)
def execute(self, prepared: PreparedSelectionRun) -> None:
"""Evaluate every eligible stock and converge the persisted run status.
This method is the boundary used by FastAPI's in-process background
task. An unexpected batch-level error is recorded before the worker
returns so the UI never mistakes a lost worker exception for success.
"""
2026-08-12 09:45:16 +08:00
stocks = _unique_stocks(prepared.source.stocks)
evaluated_count = 0
selected_stock_count = 0
signal_count = 0
failed_count = 0
2026-08-12 09:45:16 +08:00
history_rows = 0
batch_count = _chunk_count(len(stocks), self.batch_size)
read_seconds = 0.0
evaluate_seconds = 0.0
persist_seconds = 0.0
scoring_seconds = 0.0
try:
pattern_cases, pattern_library_error = self._prepare_pattern_cases(prepared.run.id)
2026-08-12 09:45:16 +08:00
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,
)
2026-08-12 09:45:16 +08:00
read_seconds += time.perf_counter() - read_started
history_rows += sum(
len(history.bars) for history in histories if history is not None
)
2026-08-12 09:45:16 +08:00
evaluate_started = time.perf_counter()
evaluations = tuple(
executor.map(
self._evaluate_stock,
batch_stocks,
histories,
[prepared.source.target_trade_date] * len(batch_stocks),
)
)
evaluate_seconds += time.perf_counter() - evaluate_started
scoring_started = time.perf_counter()
2026-08-12 09:45:16 +08:00
items = tuple(
_to_item(
stock.ts_code,
stock.name,
evaluation,
pattern_score=self._score_stock(
prepared.run.id,
stock,
history,
evaluation,
pattern_cases,
pattern_library_error,
),
2026-08-12 09:45:16 +08:00
)
for stock, history, evaluation in zip(
2026-08-12 09:45:16 +08:00
batch_stocks,
histories,
evaluations,
2026-08-12 09:45:16 +08:00
strict=True,
)
)
scoring_seconds += time.perf_counter() - scoring_started
2026-08-12 09:45:16 +08:00
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(
prepared.run.id,
status,
evaluated_count=evaluated_count,
selected_stock_count=selected_stock_count,
signal_count=signal_count,
failed_count=failed_count,
)
except Exception as exc: # noqa: BLE001 - worker boundary must persist failure state
2026-08-12 09:45:16 +08:00
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,
"failed",
evaluated_count=evaluated_count,
selected_stock_count=selected_stock_count,
signal_count=signal_count,
failed_count=max(failed_count, 1),
error_type="batch_error",
error_message=str(exc),
)
except Exception: # noqa: BLE001 - preserve the original worker failure
2026-08-12 09:45:16 +08:00
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 scoring_seconds=%.3f persist_seconds=%.3f",
2026-08-12 09:45:16 +08:00
prepared.run.id,
len(stocks),
history_rows,
batch_count,
self.max_workers,
read_seconds,
evaluate_seconds,
scoring_seconds,
2026-08-12 09:45:16 +08:00
persist_seconds,
)
def _prepare_pattern_cases(
self,
run_id: str,
) -> tuple[tuple[PatternCase, ...] | None, str | None]:
"""Load the complete case library once without failing selection."""
if not self.pattern_scoring_enabled:
return None, None
if self.pattern_case_loader is None or self.pattern_scorer is None:
reason = "pattern scoring is enabled but not configured"
logger.error("selection_pattern_library_failed run_id=%s reason=%s", run_id, reason)
return None, reason
try:
return self.pattern_case_loader.load(), None
except Exception as exc: # noqa: BLE001 - scoring enrichment must not fail selection
reason = _safe_item_error(exc)
logger.warning(
"selection_pattern_library_failed run_id=%s error_type=%s reason=%s",
run_id,
exc.__class__.__name__,
reason,
)
return None, reason
def _score_stock(
self,
run_id: str,
stock: SelectionStock,
history: StockHistory | None,
evaluation: SelectionEvaluation,
cases: tuple[PatternCase, ...] | None,
library_error: str | None,
) -> PatternScore:
"""Score one selected stock once and isolate enrichment failures."""
if not self.pattern_scoring_enabled or evaluation.status != "selected":
return PatternScore()
if library_error is not None:
return PatternScore.failed(library_error)
if history is None or cases is None or self.pattern_scorer is None:
return PatternScore.failed("pattern scoring history or case library is unavailable")
try:
return self.pattern_scorer.score(history, cases)
except Exception as exc: # noqa: BLE001 - one score must not fail the selection run
reason = _safe_item_error(exc)
logger.warning(
"selection_pattern_score_failed run_id=%s ts_code=%s error_type=%s reason=%s",
run_id,
stock.ts_code,
exc.__class__.__name__,
reason,
)
return PatternScore.failed(reason)
2026-08-12 09:45:16 +08:00
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,
run_id: str,
*,
query: SelectionResultQuery | None = None,
) -> SelectionRun | None:
"""Read one persisted run for polling."""
return self.store.get_run(run_id, query=query)
def get_latest(
self,
strategy: StrategyName,
target_trade_date: date | None = None,
*,
query: SelectionResultQuery | None = None,
) -> SelectionRun | None:
"""Read the current result by date or the latest result for a strategy."""
return self.store.get_latest_run(strategy, target_trade_date, query=query)
def _to_item(
ts_code: str,
name: str,
evaluation: SelectionEvaluation,
*,
pattern_score: PatternScore | None = None,
) -> SelectionRunItem:
"""Translate a single-stock domain result into a stored item."""
return SelectionRunItem(
ts_code=ts_code,
name=name or (evaluation.signals[0].name if evaluation.signals else ""),
status=evaluation.status,
signal_count=len(evaluation.signals),
reason=evaluation.reason,
pattern_score=pattern_score or PatternScore(),
signals=evaluation.signals,
)
def _run_status(evaluated_count: int, failed_count: int) -> SelectionRunStatus:
"""Map per-stock outcomes into a visible batch status."""
if failed_count == 0:
return "success"
if evaluated_count == 0 or failed_count >= evaluated_count:
return "failed"
return "partial_success"
def _safe_item_error(error: Exception) -> str:
"""Keep per-stock failure context readable without persisting tracebacks."""
return " ".join(str(error).split())[:500] or error.__class__.__name__
2026-08-12 09:45:16 +08:00
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",
"SelectionRerunRequired",
"SelectionRunInProgress",
]