""" Orchestrator for task coordination and execution Manages backtesting tasks with lazy imports """ import asyncio from concurrent.futures import Future from .pool import TaskPool from .task import TaskState class Orchestrator: """Task coordinator for backtesting operations""" def __init__(self, db_path: str, file_dir=None, max_workers: int = 2): """Initialize orchestrator with database path and worker limits""" self.db_path = db_path self.file_dir = file_dir self.pool = TaskPool(max_workers=max_workers) self._pending: dict[str, dict] = {} self._on_stage = None # async callback(task_id, stage) def set_on_stage(self, cb): """Set callback for stage updates (async callable)""" self._on_stage = cb async def _notify_stage(self, task_id: str, stage: str) -> None: """Update task stage and fire callback if set""" self.pool.update_stage(task_id, stage) if self._on_stage: await self._on_stage(task_id, stage) async def submit_cta(self, strategy_class, symbol: str, params: dict, start: str, end: str, cfg, benchmark: str = "hs300") -> str: """Submit a CTA backtesting task asynchronously""" task_id = f"cta_{symbol}_{id(params)}" self.pool.submit(task_id, "cta") self._pending[task_id] = dict( strategy_class=strategy_class, symbol=symbol, params=params, start=start, end=end, cfg=cfg, benchmark=benchmark ) await self._notify_stage(task_id, "排队中") spec = self._pending[task_id] fut: Future = self.pool.submit_work( task_id, _cta_worker, spec["strategy_class"], spec["symbol"], spec["params"], spec["start"], spec["end"], spec["cfg"], spec["benchmark"], self.db_path ) task = self.pool.get_task(task_id) task.start() await self._notify_stage(task_id, "回测中") asyncio.ensure_future(self._wait_future(task_id, fut)) return task_id async def submit_optimize(self, strategy_class, symbol: str, grid: dict, start: str, end: str, cfg) -> str: """Submit a CTA optimization task asynchronously""" task_id = f"opt_{symbol}_{id(grid)}" self.pool.submit(task_id, "optimize") self._pending[task_id] = dict( strategy_class=strategy_class, symbol=symbol, grid=grid, start=start, end=end, cfg=cfg ) await self._notify_stage(task_id, "参数优化中") spec = self._pending[task_id] fut: Future = self.pool.submit_work( task_id, _opt_worker, spec["strategy_class"], spec["symbol"], spec["grid"], spec["start"], spec["end"], spec["cfg"], self.db_path ) task = self.pool.get_task(task_id) task.start() await self._notify_stage(task_id, "参数优化中") asyncio.ensure_future(self._wait_future(task_id, fut)) return task_id async def submit_factor(self, symbols: list, factor_names: list, start: str, end: str, cfg, output_dir: str) -> str: """Submit a factor analysis task asynchronously""" task_id = f"factor_{id(factor_names)}" self.pool.submit(task_id, "factor") self._pending[task_id] = dict( symbols=symbols, factor_names=factor_names, start=start, end=end, cfg=cfg, output_dir=output_dir ) await self._notify_stage(task_id, "因子分析中") spec = self._pending[task_id] fut: Future = self.pool.submit_work( task_id, _factor_worker, spec["symbols"], spec["factor_names"], spec["start"], spec["end"], spec["cfg"], spec["output_dir"] ) task = self.pool.get_task(task_id) task.start() await self._notify_stage(task_id, "分析中") asyncio.ensure_future(self._wait_future(task_id, fut)) return task_id async def _wait_future(self, task_id: str, fut: Future) -> None: """Wait for Future to complete and handle result/exception Bridges concurrent.futures.Future (from ProcessPoolExecutor) to asyncio coroutine. """ try: result = await asyncio.wrap_future(fut) await self._on_done(task_id, result) except Exception as e: task = self.pool.get_task(task_id) if task: task.fail(f"{type(e).__name__}: {e}") await self._notify_stage(task_id, "失败") async def _on_done(self, task_id: str, result) -> None: """Handle task completion (with None-guard for unknown tasks)""" task = self.pool.get_task(task_id) if task is None: # Unknown task - fire callback but don't crash await self._notify_stage(task_id, "完成") return # S1.1: use the persisted DB row id (BacktestResult.id) so get_result can # load_result(result.id). FactorReport (no .id) falls back to None until S2. task.complete(result_id=getattr(result, "id", None)) task.raw_result = result # S2: keep in-memory result (FactorReport) for ic-summary/report await self._notify_stage(task_id, "完成") def get_status(self, task_id: str) -> TaskState | None: """Get task status by ID""" return self.pool.get_status(task_id) def get_result(self, task_id: str): """Get task result by ID. Tries in-memory (current run) then DB (history).""" task = self.pool.get_task(task_id) if task and task.status == TaskState.DONE and task.result_id: # Lazy import to avoid vnpy dependency issues from sanguo_backtest.result_store import load_result return load_result(task.result_id, self.db_path) # Fallback: historical task persisted in DB (e.g. after restart) from sanguo_backtest.result_store import load_result_by_task_id return load_result_by_task_id(task_id, self.db_path) def get_raw_result(self, task_id: str): """Get the raw in-memory result object (e.g. FactorReport) by task ID. Used by factor endpoints (ic-summary, tears report) where the result isn't a BacktestResult persisted to the DB. """ task = self.pool.get_task(task_id) return task.raw_result if task else None # Module-level worker functions (must be top-level for ProcessPoolExecutor pickle) def _cta_worker(strategy_class, symbol: str, params: dict, start: str, end: str, cfg, benchmark: str, db_path: str) -> any: """Worker for CTA backtest (lazy import, spawn-friendly)""" from sanguo_backtest.cta_engine import run_cta_backtest return run_cta_backtest(strategy_class, symbol, params, start, end, cfg, db_path, benchmark=benchmark) def _opt_worker(strategy_class, symbol: str, grid: dict, start: str, end: str, cfg, db_path: str) -> any: """Worker for CTA optimization (lazy import, spawn-friendly)""" from sanguo_backtest.cta_optimizer import run_cta_optimization return run_cta_optimization(strategy_class, symbol, grid, start, end, cfg, db_path) def _factor_worker(symbols: list, factor_names: list, start: str, end: str, cfg, output_dir: str) -> any: """Worker for factor analysis (lazy import, spawn-friendly)""" from sanguo_factor.analyzer import run_factor_analysis return run_factor_analysis(symbols, factor_names, start, end, cfg, output_dir)