- Job kind=selection:executor 分支用 SelectionService(spec=SelectionQuery),结果 落 result_json、自动归档 Experiment(summary:as_of + 选出 N);/api/jobs 与 /api/experiments 的 result 解码支持 SelectionResult - POST /api/selections/jobs:提交 SelectionQuery 为异步 Job(BackgroundTasks/subprocess) - Web 选股页新增「异步(全市场)」按钮:提交 Job → waitJob 轮询结果(解决同步 60s+) - tests/test_selection_job.py:executor 执行归档(SelectionResult/Experiment kind)、 API 提交→轮询→结果与实验列表;全量 pytest + tsc 通过
165 lines
6.5 KiB
Python
165 lines
6.5 KiB
Python
"""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
|