"""研究服务:把 Research Specification 编排为数据获取 + 引擎执行。 本层是业务入口:API / Agent 只能调用这里的用例(AGENT.md §16/§17), 禁止直接拼接引擎配置。数据一律经 Repository 获取(防未来函数由查询层保证)。 内存优化:大数据面板优先走 Repository 的流式列裁剪查询 (stream_range_many_columns,SQL 侧转 REAL、分批拉取),避免 ORM 对象 / Decimal 全量物化;老实现回退到 get_range_many 逐实体路径。 """ from __future__ import annotations from collections.abc import Iterable from datetime import date, timedelta import pandas as pd from app.domain.entities.research import ( BacktestResult, FactorCorrelationReport, FactorTestReport, ResearchSpec, ) from app.domain.repositories.market import ( DailyBarRepository, StockRepository, ) from app.quant.composite import build_factor_panels from app.quant.engine import QuantEngine from app.quant.evaluation import factor_correlation_report from app.quant.universe import filter_stocks, resolve_members # noqa: F401 —— 范围过滤 # 流式路径每攒多少行落一个 DataFrame 分片(控制 concat 峰值) _FRAME_CHUNK_ROWS = 50_000 def bars_to_daily_df(bars) -> pd.DataFrame: """DailyBar 列表 → 引擎长表 DataFrame(symbol/trade_date/ohlc/volume/amount)。 领域实体中的 Decimal 在此转 float,供 pandas 数值运算(保持 DataFrame 全数值列)。 """ df = pd.DataFrame([b.model_dump() for b in bars]) if not df.empty: for col in ("open", "high", "low", "close", "volume", "amount"): if col in df.columns: df[col] = df[col].astype(float) return df def _frame_from_stream(rows: Iterable[tuple], columns: list[str]) -> pd.DataFrame: """把流式 (symbol, trade_date_iso, *float_cols) 分批拼成 float 长表。 全程只保留分批 DataFrame + 最终一份结果,避免整批 tuple/Decimal 同时驻留。 返回列:symbol, trade_date(datetime64), *columns(float64)。 """ cols = ["symbol", "trade_date", *columns] buf: list[tuple] = [] pieces: list[pd.DataFrame] = [] for row in rows: buf.append(row) if len(buf) >= _FRAME_CHUNK_ROWS: pieces.append(pd.DataFrame(buf, columns=cols)) buf = [] if buf: pieces.append(pd.DataFrame(buf, columns=cols)) if not pieces: return pd.DataFrame() df = pd.concat(pieces, ignore_index=True) df["trade_date"] = pd.to_datetime(df["trade_date"]) for col in columns: # NULL → NaN,统一 float64 df[col] = pd.to_numeric(df[col], errors="coerce") return df def load_daily_df( daily_repo, symbols: list[str], start: date, end: date, columns: list[str], adjust: str = "none", ) -> pd.DataFrame: """从 Repository 装配行情长表(供研究/选股共用)。 优先走流式列裁剪(stream_range_many_columns,SQL 侧转 REAL、分批), 失败或实现缺失时回退 get_range_many / 逐只 get_range。 """ if not symbols: return pd.DataFrame() streamer = getattr(daily_repo, "stream_range_many_columns", None) if streamer is not None: try: return _frame_from_stream( streamer(symbols, start, end, sorted(columns), adjust=adjust), sorted(columns) ) except Exception: # noqa: BLE001 —— 流式路径失败回退旧路径(兼容非 SQL 实现) pass get_many = getattr(daily_repo, "get_range_many", None) if get_many is not None: bars = list(get_many(symbols, start, end, adjust=adjust)) else: # 兜底:逐只查询 bars = [] for sym in symbols: bars.extend(daily_repo.get_range(sym, start, end)) return bars_to_daily_df(bars) class ResearchService: """研究用例入口(因子测试 / 回测)。依赖注入 Repository 与引擎。""" def __init__( self, stock_repo: StockRepository, daily_repo: DailyBarRepository, engine: QuantEngine, index_repo=None, ) -> None: self._stock_repo = stock_repo self._daily_repo = daily_repo self._engine = engine self._index_repo = index_repo def run_factor_test(self, spec: ResearchSpec, horizon_days: int = 21) -> FactorTestReport: if spec.type != "factor_test": raise ValueError("factor_test 用例需要 spec.type=factor_test") daily = self._load_daily(spec) return self._engine.run_factor_test(daily, spec, horizon_days=horizon_days) def run_backtest(self, spec: ResearchSpec) -> BacktestResult: if spec.type != "backtest": raise ValueError("backtest 用例需要 spec.type=backtest") daily = self._load_daily(spec) return self._engine.run_backtest(daily, spec) def run_factor_correlation(self, spec: ResearchSpec) -> FactorCorrelationReport: """多因子两两相关(v3 §12):同 universe/period 装配 → 横截面相关矩阵。""" daily = self._load_daily(spec) panels = {fs.name: build_factor_panels(daily, [fs])[0][1] for fs in spec.factors} return factor_correlation_report(panels) # ---- 数据装配 ---- def _load_daily(self, spec: ResearchSpec) -> pd.DataFrame: start, end = spec.period # 回测前预留因子 warmup(lookback≤120 交易日,取 300 自然日余量) data_start = start - timedelta(days=300) stocks = filter_stocks( self._stock_repo.list(), spec.universe, as_of=start, members=resolve_members(self._index_repo, spec.universe, start), ) # 引擎所需列裁剪(LocalEngine 只取 close + 因子字段;Qlib 回测取全 OHLCV) required = self._engine.required_columns(spec) return load_daily_df( self._daily_repo, [s.symbol for s in stocks], data_start, end, sorted(required), adjust=spec.price_adjustment, )