- index_weight 表(migration f5e0d1c2b3a4,MySQL 已应用;index_code+date+symbol 唯一) + IndexWeight 实体 + IndexConstituentRepository(members_at:取 <=as_of 最近一期快照, Survivorship-free / 无未来成分;latest_date) - UniverseSpec.index_code + universe.filter_stocks members 交集 + resolve_members; Research/Selection/Signal/Replay 服务注入 index repo(历史成分过滤,选股/回测共用) - tests/test_index_universe.py:快照历史成分(成分变更不入早期结果)、幂等、 空快照期空集、index_code 过滤下 as_of 一致性;全量 pytest 通过
152 lines
5.5 KiB
Python
152 lines
5.5 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,
|
||
FactorTestReport,
|
||
ResearchSpec,
|
||
)
|
||
from app.domain.repositories.market import (
|
||
DailyBarRepository,
|
||
StockRepository,
|
||
)
|
||
from app.quant.engine import QuantEngine
|
||
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 _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,
|
||
)
|