From 37510c1b896dfbce6ef1939a6feae0e10f48367a Mon Sep 17 00:00:00 2001 From: Simon Date: Wed, 9 Sep 2026 07:18:36 +0800 Subject: [PATCH] =?UTF-8?q?feat(replay):=20M9-6=20Bar=20Replay=20=E7=BA=BF?= =?UTF-8?q?=E6=80=A7=E9=87=8D=E6=94=BE=EF=BC=88as=5Fof=20=E9=80=90?= =?UTF-8?q?=E6=97=A5=E4=BB=85=E7=94=A8=E5=BD=93=E6=97=B6=E6=95=B0=E6=8D=AE?= =?UTF-8?q?=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 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 通过 --- backend/app/api/deps.py | 9 + backend/app/api/replays.py | 33 ++++ backend/app/api/router.py | 2 + .../application/services/replay_service.py | 85 +++++++++ backend/app/domain/entities/replay.py | 34 ++++ backend/tests/test_replays.py | 164 ++++++++++++++++++ 6 files changed, 327 insertions(+) create mode 100644 backend/app/api/replays.py create mode 100644 backend/app/application/services/replay_service.py create mode 100644 backend/app/domain/entities/replay.py create mode 100644 backend/tests/test_replays.py diff --git a/backend/app/api/deps.py b/backend/app/api/deps.py index a61402f..d5d22cd 100644 --- a/backend/app/api/deps.py +++ b/backend/app/api/deps.py @@ -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)] diff --git a/backend/app/api/replays.py b/backend/app/api/replays.py new file mode 100644 index 0000000..7356835 --- /dev/null +++ b/backend/app/api/replays.py @@ -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 diff --git a/backend/app/api/router.py b/backend/app/api/router.py index 32d3f97..6a7e19e 100644 --- a/backend/app/api/router.py +++ b/backend/app/api/router.py @@ -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) diff --git a/backend/app/application/services/replay_service.py b/backend/app/application/services/replay_service.py new file mode 100644 index 0000000..ab3c55d --- /dev/null +++ b/backend/app/application/services/replay_service.py @@ -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) diff --git a/backend/app/domain/entities/replay.py b/backend/app/domain/entities/replay.py new file mode 100644 index 0000000..462945b --- /dev/null +++ b/backend/app/domain/entities/replay.py @@ -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 diff --git a/backend/tests/test_replays.py b/backend/tests/test_replays.py new file mode 100644 index 0000000..223b365 --- /dev/null +++ b/backend/tests/test_replays.py @@ -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