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
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user