""" Orchestrator for task coordination and execution Manages backtesting tasks with lazy imports """ import asyncio import os import uuid 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 def _record_submit(self, task_id: str) -> None: """提交即记提交时间(真创建时间)。失败静默——缺失只影响耗时展示。""" if not self.db_path: return try: from sanguo_backtest.result_store import record_submission record_submission(task_id, self.db_path) except Exception: pass 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", capital: float = 1_000_000, position_pct: float = 0.95, interval: str = "d") -> str: """Submit a CTA backtesting task asynchronously""" # Stable uuid up front → reused as the persisted DB task_id, so runner-id == # DB task_id (durable across restarts; previously used id(params) memory addr). task_id = f"cta_{uuid.uuid4().hex[:8]}" self._record_submit(task_id) 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, capital=capital, position_pct=position_pct, interval=interval, ) 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_id, spec["capital"], spec["position_pct"], spec["interval"] ) 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_{uuid.uuid4().hex[:8]}" self._record_submit(task_id) 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_id ) 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_{uuid.uuid4().hex[:8]}" self._record_submit(task_id) 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 submit_portfolio(self, start: str, end: str, cash: float, benchmark: str, strategy: str = "all_weather", max_pool: int = 30, provider_config=None, commission_rate: float = 0.0003, stamp_duty_rate: float = 0.001, min_commission: float = 5.0, slippage: float = 0.0, interval: str = "d") -> str: """Submit a portfolio backtest task asynchronously. Runs runner_backtest as a subprocess (3600s hard cap) inside the process pool, reusing the same _wait_future/_on_done bridge as CTA/optimize/factor. """ task_id = f"portfolio_{uuid.uuid4().hex[:8]}" self._record_submit(task_id) self.pool.submit(task_id, "portfolio") file_dir = os.path.dirname(os.path.abspath(self.db_path)) if self.db_path else None self._pending[task_id] = dict( task_id=task_id, start=start, end=end, cash=cash, benchmark=benchmark, strategy=strategy, max_pool=max_pool, provider_config=provider_config, commission_rate=commission_rate, stamp_duty_rate=stamp_duty_rate, min_commission=min_commission, slippage=slippage, interval=interval, db_path=self.db_path, file_dir=file_dir, ) await self._notify_stage(task_id, "排队中") spec = self._pending[task_id] fut: Future = self.pool.submit_work(task_id, _portfolio_worker, spec) 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: import logging logging.getLogger(__name__).error( "task %s failed in pool: %s: %s", task_id, type(e).__name__, 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 # S2: persist factor result so it appears in task list & survives restart if getattr(result, "ic_summary", None) and getattr(result, "factor_names", None) is not None: self._persist_factor(task_id, result) await self._notify_stage(task_id, "完成") def _persist_factor(self, task_id: str, fr) -> None: """Persist FactorReport to backtest_results.db (type=factor) so it shows in task list and survives API restart. Mirrors cta/optimize persistence.""" from sanguo_backtest.result_store import save_result, BacktestResult try: save_result(BacktestResult( task_id=task_id, type="factor", status="done", strategy=",".join(fr.factor_names), symbol=",".join(getattr(fr, "symbols", []) or []), params={"factor_names": fr.factor_names}, start=getattr(fr, "start", "") or "", end=getattr(fr, "end", "") or "", statistics={"ic_summary": fr.ic_summary, "report_paths": fr.report_paths}, equity_curve=None, trades=None, ), db_path=self.db_path) except Exception as e: import logging logging.getLogger(__name__).warning("persist factor %s failed: %s", task_id, e) 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, task_id: str, capital: float = 1_000_000, position_pct: float = 0.95, interval: str = "d") -> 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, task_id=task_id, capital=capital, position_pct=position_pct, interval=interval) def _opt_worker(strategy_class, symbol: str, grid: dict, start: str, end: str, cfg, db_path: str, task_id: 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, task_id=task_id) 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) def _portfolio_worker(spec: dict) -> any: """Worker for portfolio backtest (lazy import, spawn-friendly)""" from sanguo_orchestrator.portfolio_worker import run_portfolio_task return run_portfolio_task(spec)