"""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 d in 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