Files
Simon 861a4051ca feat(job): D2 全市场选股 Job 化(kind=selection 异步)
- 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 通过
2026-09-09 07:43:35 +08:00

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