"""Application orchestration for persisted whole-universe B1 runs.""" 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, cast from ..domain.models import SelectionEvaluation, StockHistory from ..domain.runs import ( BatchSelectionRunStore, BatchSelectionUniverseReader, SelectionExecutionSource, SelectionRerunRequired, SelectionResultQuery, SelectionRun, SelectionRunInProgress, SelectionRunItem, SelectionRunStatus, SelectionRunStore, 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, *, max_workers: int = 4, batch_size: int = 200, ) -> None: """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, 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. """ 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: 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, ) read_seconds += time.perf_counter() - read_started history_rows += sum( len(history.bars) for history in histories if history is not None ) 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, ) ) 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( 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 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 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, 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) -> 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, 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__ 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", ]