feat(replay): M9-6 Bar Replay 线性重放(as_of 逐日仅用当时数据)
- 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 通过
This commit is contained in:
@@ -11,6 +11,7 @@ from fastapi import Depends
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.application.services.chart_service import ChartService
|
||||
from app.application.services.replay_service import ReplayService
|
||||
from app.application.services.selection_service import SelectionService
|
||||
from app.application.services.signal_service import SignalService
|
||||
from app.domain.repositories.composite import CompositeRepository
|
||||
@@ -89,6 +90,13 @@ def _service_factory(
|
||||
return ResearchService(stock_repo, daily_repo, engine)
|
||||
|
||||
|
||||
def _replay_service_factory(
|
||||
stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)],
|
||||
daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)],
|
||||
) -> ReplayService:
|
||||
return ReplayService(stock_repo, daily_repo)
|
||||
|
||||
|
||||
def _signal_service_factory(
|
||||
stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)],
|
||||
daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)],
|
||||
@@ -135,6 +143,7 @@ FactorRepoDep = Annotated[FactorRepository, Depends(_factor_repo_factory)]
|
||||
CompositeRepoDep = Annotated[CompositeRepository, Depends(_composite_repo_factory)]
|
||||
SignalRepoDep = Annotated[SignalRepository, Depends(_signal_repo_factory)]
|
||||
SignalServiceDep = Annotated[SignalService, Depends(_signal_service_factory)]
|
||||
ReplayServiceDep = Annotated[ReplayService, Depends(_replay_service_factory)]
|
||||
ChartServiceDep = Annotated[ChartService, Depends(_chart_service_factory)]
|
||||
StrategyRepoDep = Annotated[StrategyRepository, Depends(_strategy_repo_factory)]
|
||||
|
||||
|
||||
@@ -0,0 +1,33 @@
|
||||
"""Bar Replay API(M9-6):POST /api/replays —— 线性重放选股/信号时间线。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
|
||||
from fastapi import APIRouter, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.api.deps import ReplayServiceDep
|
||||
from app.domain.entities.replay import ReplayResult
|
||||
from app.domain.entities.selection import SelectionQuery
|
||||
from app.domain.entities.signal import SignalRules
|
||||
|
||||
router = APIRouter(prefix="/replays", tags=["replays"])
|
||||
|
||||
|
||||
class ReplayRequest(BaseModel):
|
||||
query: SelectionQuery
|
||||
rules: SignalRules = SignalRules()
|
||||
start: date
|
||||
end: date
|
||||
top_n: int = Field(default=5, ge=1, le=20)
|
||||
|
||||
|
||||
@router.post("", response_model=ReplayResult, summary="线性重放(as_of 逐日,仅用当时数据)")
|
||||
def run_replay(req: ReplayRequest, service: ReplayServiceDep) -> ReplayResult:
|
||||
if req.start >= req.end:
|
||||
raise HTTPException(status_code=400, detail="start 必须早于 end")
|
||||
try:
|
||||
return service.replay(req.query, req.rules, req.start, req.end, req.top_n)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
@@ -16,6 +16,7 @@ from app.api import (
|
||||
factors,
|
||||
health,
|
||||
jobs,
|
||||
replays,
|
||||
research,
|
||||
selections,
|
||||
signals,
|
||||
@@ -31,6 +32,7 @@ api_router.include_router(composites.router)
|
||||
api_router.include_router(research.router)
|
||||
api_router.include_router(selections.router)
|
||||
api_router.include_router(charts.router)
|
||||
api_router.include_router(replays.router)
|
||||
api_router.include_router(signals.router)
|
||||
api_router.include_router(strategies.router)
|
||||
api_router.include_router(jobs.router)
|
||||
|
||||
@@ -0,0 +1,85 @@
|
||||
"""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)
|
||||
@@ -0,0 +1,34 @@
|
||||
"""Bar Replay(v3 §20.6,第二阶段 MVP)领域实体。
|
||||
|
||||
线性重放:对给定选股查询与信号规则,在交易日序列上逐日以「当日为止的数据」执行
|
||||
(as_of 语义),输出每日时间线(意图 Top + 信号 + 计数),用于核对:
|
||||
- 未来函数:每日计算只用 <= as_of 数据(与静态研究同一引擎)
|
||||
- 一致性:重放某日结果 == 该日独立 select/signal 结果(亦 == 回测该调仓日意图)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.domain.entities.signal import SignalEvent
|
||||
|
||||
|
||||
class ReplayTop(BaseModel):
|
||||
symbol: str
|
||||
score: float
|
||||
|
||||
|
||||
class ReplayDay(BaseModel):
|
||||
as_of: date
|
||||
top: list[ReplayTop] = Field(default_factory=list, description="意图排名前 N(score 降序)")
|
||||
events: list[SignalEvent] = Field(default_factory=list, description="当日信号(<=max_output_rank)")
|
||||
counts: dict[str, int] = Field(default_factory=dict, description="BUY/WATCH/SELL 计数")
|
||||
|
||||
|
||||
class ReplayResult(BaseModel):
|
||||
start: date
|
||||
end: date
|
||||
days: list[ReplayDay] = Field(default_factory=list)
|
||||
top_n: int = 5
|
||||
Reference in New Issue
Block a user