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:
Simon
2026-09-09 07:18:36 +08:00
parent 5bde8f9f5f
commit 37510c1b89
6 changed files with 327 additions and 0 deletions
+9
View File
@@ -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)]
+33
View File
@@ -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
+2
View File
@@ -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)
+34
View File
@@ -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
+164
View File
@@ -0,0 +1,164 @@
"""M9-6 Bar Replay 测试:线性重放与回测调仓意图一致、范围约束、未来函数。"""
from __future__ import annotations
from datetime import date
import pandas as pd
import pytest
from app.api import deps
from app.application.services.replay_service import ReplayService
from app.domain.entities.market import Stock
from app.domain.entities.research import ResearchSpec
from app.domain.entities.selection import SelectionQuery
from app.domain.entities.signal import SignalRules
from app.infrastructure.persistence.sqlalchemy.base import Base
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
SqlAlchemyDailyBarRepository,
SqlAlchemyStockRepository,
)
from app.main import app
from app.quant.engine import LocalEngine
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from conftest_quant import bars_dataframe_to_daily_bars, synthetic_daily
_SYMS = ["600000.SH", "600001.SH", "600002.SH", "600003.SH", "600004.SH"]
def _query(symbols=_SYMS, as_of=None) -> SelectionQuery:
return SelectionQuery(
universe={"exclude_st": False, "min_listing_days": 0, "symbols": list(symbols)},
factors=[{"name": "momentum_60", "weight": 1.0}],
top_n=2,
as_of=as_of,
)
class _MemStockRepo:
def __init__(self, stocks):
self._stocks = stocks
def list(self):
return self._stocks
def get_by_symbol(self, symbol):
return next((s for s in self._stocks if s.symbol == symbol), None)
class _MemDailyRepo:
def __init__(self, df: pd.DataFrame) -> None:
self._bars = bars_dataframe_to_daily_bars(df)
def get_range(self, symbol, start, end):
return [b for b in self._bars if b.symbol == symbol and start <= b.trade_date <= end]
def get_range_many(self, symbols, start, end, adjust="none"):
syms = set(symbols)
return [b for b in self._bars if b.symbol in syms and start <= b.trade_date <= end]
@pytest.fixture()
def df() -> pd.DataFrame:
return synthetic_daily({s: 0.006 - 0.0015 * i for i, s in enumerate(_SYMS)}, n=320)
class TestReplayConsistency:
def test_replay_matches_backtest_selection(self, df) -> None:
stocks = [Stock(symbol=s, name=f"测试{i}", list_date=date(1999, 1, 1))
for i, s in enumerate(_SYMS)]
svc = ReplayService(_MemStockRepo(stocks), _MemDailyRepo(df))
start, end = date(2024, 5, 1), date(2024, 8, 31)
res = svc.replay(_query(), SignalRules(), start, end, top_n=2)
spec = ResearchSpec(
type="backtest",
universe={"exclude_st": False, "min_listing_days": 0, "symbols": _SYMS},
factors=[{"name": "momentum_60", "weight": 1.0}],
selection={"top_n": 2},
rebalance="monthly",
period=(date(2024, 5, 1), date(2024, 12, 31)),
)
bt = LocalEngine().run_backtest(df, spec)
by_date: dict[date, set[str]] = {}
for p in bt.selection_history:
if start <= p.date <= end:
by_date.setdefault(p.date, set()).add(p.symbol)
assert res.days and by_date
replay_by_date = {d.as_of: {t.symbol for t in d.top} for d in res.days}
for d, syms in by_date.items():
assert replay_by_date[d] == syms, f"as_of={d} 重放 {replay_by_date.get(d)} ≠ 回测 {syms}"
def test_bounds(self, df) -> None:
stocks = [Stock(symbol=s, name=f"测试{i}", list_date=date(1999, 1, 1))
for i, s in enumerate(_SYMS)]
svc = ReplayService(_MemStockRepo(stocks), _MemDailyRepo(df))
with pytest.raises(ValueError):
svc.replay(_query(symbols=[]), SignalRules(), date(2024, 5, 1), date(2024, 8, 31))
with pytest.raises(ValueError):
svc.replay(_query(), SignalRules(), date(2024, 1, 1), date(2024, 12, 31)) # >90 交易日
def test_future_does_not_leak_into_early_frames(self, df) -> None:
"""把 B 股在区间后段设为暴涨:前段重放帧不应出现 B(as_of 截断)。"""
dates = sorted(df["trade_date"].unique())
mid = dates[len(dates) // 2]
b = _SYMS[1]
for i, d in enumerate(dates):
if d > mid:
mask = (df["symbol"] == b) & (df["trade_date"] == d)
df.loc[mask, "close"] = df.loc[mask, "close"] * 1.03 # 后段持续暴涨
stocks = [Stock(symbol=s, name=f"测试{i}", list_date=date(1999, 1, 1))
for i, s in enumerate(_SYMS)]
svc = ReplayService(_MemStockRepo(stocks), _MemDailyRepo(df))
res = svc.replay(_query(), SignalRules(), date(2024, 2, 1), date(2024, 4, 30), top_n=2)
assert res.days
for d in res.days[: 5]: # 早段帧
assert all(t.symbol != b for t in d.top)
class TestReplayApi:
@pytest.fixture()
def client(self, tmp_path, df):
engine = create_engine(f"sqlite:///{tmp_path / 'replay.db'}", future=True)
Base.metadata.create_all(engine)
Session = sessionmaker(bind=engine, expire_on_commit=False)
with Session() as session:
SqlAlchemyStockRepository(session).upsert_many(
[Stock(symbol=s, name=f"测试{i}", list_date=date(1999, 1, 1))
for i, s in enumerate(_SYMS)]
)
SqlAlchemyDailyBarRepository(session).upsert_many(bars_dataframe_to_daily_bars(df))
session.commit()
def _override():
with Session() as s:
yield s
app.dependency_overrides[deps.get_session] = _override
with TestClient(app) as c:
yield c
app.dependency_overrides.clear()
def test_post_replay(self, client) -> None:
body = {
"query": {
"universe": {"exclude_st": False, "min_listing_days": 0, "symbols": _SYMS},
"factors": [{"name": "momentum_60", "weight": 1}],
"top_n": 2,
},
"rules": {"buy_rank_threshold": 1, "max_output_rank": 3},
"start": "2024-05-01",
"end": "2024-06-15",
"top_n": 2,
}
resp = client.post("/api/replays", json=body)
assert resp.status_code == 200
data = resp.json()
assert len(data["days"]) > 10
assert data["days"][0]["top"] and data["days"][0]["counts"]["sell"] >= 0
# 无白名单 → 400
bad = {**body, "query": {**body["query"], "universe": {"exclude_st": False, "min_listing_days": 0}}}
assert client.post("/api/replays", json=bad).status_code == 400