- ResearchService.run_backtest/run_factor_test 支持 on_stage 回调(data_loading →
factor_calculation/backtesting → analysis);job_executor 经独立短会话把 stage 写库
(子进程模式同样走 DB),收尾回读最后阶段避免覆盖
- /api/jobs:GET 列表(kind/limit)、POST /{id}/cancel(queued/running → CANCELLED +
终止子进程 terminate_active;父进程兜底已跳过 CANCELLED)
- SSE /jobs/{id}/events 现会携带 stage
- tests/test_job_stages.py:成功 job 终态 stage=analysis;executor 取消 queued/幂等/
已完成不可取消;API 列表+终态不可取消+404;全量 pytest 通过
178 lines
6.6 KiB
Python
178 lines
6.6 KiB
Python
"""研究服务:把 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, on_stage=None
|
||
) -> FactorTestReport:
|
||
"""on_stage(str):执行阶段回调(data_loading / factor_calculation / analysis),
|
||
供 Job 状态机上报 stage(v3 §23)。"""
|
||
if spec.type != "factor_test":
|
||
raise ValueError("factor_test 用例需要 spec.type=factor_test")
|
||
_stage(on_stage, "data_loading")
|
||
daily = self._load_daily(spec)
|
||
_stage(on_stage, "factor_calculation")
|
||
report = self._engine.run_factor_test(daily, spec, horizon_days=horizon_days)
|
||
_stage(on_stage, "analysis")
|
||
return report
|
||
|
||
def run_backtest(self, spec: ResearchSpec, on_stage=None) -> BacktestResult:
|
||
if spec.type != "backtest":
|
||
raise ValueError("backtest 用例需要 spec.type=backtest")
|
||
_stage(on_stage, "data_loading")
|
||
daily = self._load_daily(spec)
|
||
_stage(on_stage, "backtesting")
|
||
result = self._engine.run_backtest(daily, spec)
|
||
_stage(on_stage, "analysis")
|
||
return result
|
||
|
||
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,
|
||
)
|
||
|
||
|
||
def _stage(cb, name: str) -> None:
|
||
if cb is not None:
|
||
cb(name)
|