- domain/entities/replay.py:ReplayDay{top/events/counts}/ReplayResult 时间线
- ReplayService:universe.symbols 白名单必填(≤40)且重放交易日 ≤90(防全市场长任务);
每个交易日以 <=as_of 数据经同一 signal/score 引擎生成帧
- POST /api/replays(边界校验)→ 时间线;供前端 Bar Replay 控件(v3 §20.6 阶段二)
- tests/test_replays.py:重放帧 == 回测 selection_history 逐调仓日一致;范围约束;
后段暴涨股不泄漏进早段帧(未来函数);API 400/200;全量 pytest 通过
86 lines
3.2 KiB
Python
86 lines
3.2 KiB
Python
"""Bar Replay 服务(M9-6):线性逐交易日重放选股+信号(as_of 语义)。
|
|
|
|
- 范围约束:universe.symbols 必填(≤ 40 只)、重放交易日 ≤ 90 —— 避免全市场长任务
|
|
- 每日计算只使用 <= as_of 数据(与 select/signal/回测同一引擎与口径)
|
|
- ReplayDay.events 按 rank 升序;top 取前 N(意图排名,与回测 selection_history 对齐)
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
from datetime import date, timedelta
|
|
|
|
import pandas as pd
|
|
|
|
from app.domain.entities.replay import ReplayDay, ReplayResult, ReplayTop
|
|
from app.domain.entities.selection import SelectionQuery
|
|
from app.domain.entities.signal import SignalRules
|
|
from app.domain.repositories.market import DailyBarRepository, StockRepository
|
|
from app.quant.selection import factor_columns
|
|
from app.quant.service import filter_stocks, load_daily_df
|
|
from app.quant.signal import generate_signals
|
|
|
|
MAX_SYMBOLS = 40
|
|
MAX_DAYS = 90
|
|
|
|
|
|
class ReplayService:
|
|
def __init__(self, stock_repo: StockRepository, daily_repo: DailyBarRepository) -> None:
|
|
self._stock_repo = stock_repo
|
|
self._daily_repo = daily_repo
|
|
|
|
def replay(
|
|
self,
|
|
query: SelectionQuery,
|
|
rules: SignalRules,
|
|
start: date,
|
|
end: date,
|
|
top_n: int = 5,
|
|
) -> ReplayResult:
|
|
symbols = list(query.universe.symbols or [])
|
|
if not symbols:
|
|
raise ValueError("Bar Replay 需要 universe.symbols 白名单(≤40 只),避免全市场长任务")
|
|
if len(symbols) > MAX_SYMBOLS:
|
|
raise ValueError(f"Bar Replay 白名单最多 {MAX_SYMBOLS} 只,当前 {len(symbols)}")
|
|
stocks = filter_stocks(self._stock_repo.list(), query.universe, as_of=start)
|
|
if not stocks:
|
|
return ReplayResult(start=start, end=end, top_n=top_n)
|
|
|
|
columns = sorted(factor_columns(query))
|
|
daily = load_daily_df(
|
|
self._daily_repo,
|
|
symbols,
|
|
start - timedelta(days=query.warmup_days),
|
|
end,
|
|
columns,
|
|
adjust=query.price_adjustment,
|
|
)
|
|
if daily.empty:
|
|
return ReplayResult(start=start, end=end, top_n=top_n)
|
|
trading_days = sorted(
|
|
pd.to_datetime(daily["trade_date"].unique())
|
|
)
|
|
days = [d for d in trading_days if start <= d.date() <= end]
|
|
if len(days) > MAX_DAYS:
|
|
raise ValueError(f"重放区间交易日 {len(days)} > 上限 {MAX_DAYS},请缩短区间")
|
|
|
|
out_days: list[ReplayDay] = []
|
|
for d in days:
|
|
res = generate_signals(daily, query, rules, as_of=d.date())
|
|
top = [
|
|
ReplayTop(symbol=e.symbol, score=e.score or 0.0)
|
|
for e in res.events[:top_n]
|
|
]
|
|
out_days.append(
|
|
ReplayDay(
|
|
as_of=d.date(),
|
|
top=top,
|
|
events=res.events,
|
|
counts={
|
|
"buy": res.statistics.buy,
|
|
"watch": res.statistics.watch,
|
|
"sell": res.statistics.sell,
|
|
},
|
|
)
|
|
)
|
|
return ReplayResult(start=start, end=end, days=out_days, top_n=top_n)
|