Files
qlib/backend/tests/test_selection.py
T
Simon f3586adb25 feat(selection): M6.0 选股契约与评分引擎(SelectionQuery/Result + select(as_of))
- domain/entities/selection.py:SelectionQuery(universe+method+factors+top_n/top_pct/
  min_score+as_of+预热)与 SelectionResult/Candidate/Statistics(v2 §14.2/§21.1 DTO);
  ConditionSpec 字段就位供 M6.2 条件选股
- quant/selection.py:Selection Engine method=score —— 复合分(zscore×权重×方向)
  → TopN/Top% 截断;observation_date=<=as_of 最近交易日(防未来函数,v2 §9);
  候选带 factor_values 与 selection_reason(可解释)
- application/services/selection_service.py:选股用例(universe 过滤 → 装配 → 引擎)
- quant/service.py:抽取公共 load_daily_df 供研究/选股共用(行为不变)
- tests/test_selection.py:11 例 —— TopN/排序/理由、as_of 防未来函数、ST/上市天数/
  退市过滤、top_pct/min_score、空数据与查询校验;全量 pytest 通过
2026-09-09 00:12:28 +08:00

224 lines
8.8 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""M6.0 选股契约与服务测试:因子评分 TopN 选股、as_of 未来函数防护、universe 过滤。
使用 Fake(内存)Repository + 合成/自定义行情,不触数据库。
"""
from __future__ import annotations
from datetime import date
from decimal import Decimal
import pandas as pd
import pytest
from app.application.services.selection_service import SelectionService
from app.domain.entities.market import DailyBar, Stock
from app.domain.entities.selection import SelectionQuery
from pydantic import ValidationError
from conftest_quant import synthetic_daily
_SYMS = ["600000.SH", "600001.SH", "600002.SH", "600003.SH", "600004.SH"]
def _mem_stocks() -> list[Stock]:
return [
Stock(symbol=s, name=f"测试股份{i}", list_date=date(1999, 1, 1))
for i, s in enumerate(_SYMS)
]
class _MemStockRepo:
def __init__(self, stocks: list[Stock]) -> None:
self._stocks = stocks
def list(self) -> list[Stock]:
return self._stocks
def get_by_symbol(self, symbol: str) -> Stock | None:
return next((s for s in self._stocks if s.symbol == symbol), None)
class _MemDailyRepo:
"""内存日线仓库:支持 get_range / get_range_many(无流式 → 走回退路径)。"""
def __init__(self, df: pd.DataFrame) -> None:
self._df = df
def _bars(self, symbols, start, end) -> list[DailyBar]:
sub = self._df[
self._df["symbol"].isin(symbols)
& (self._df["trade_date"] >= start)
& (self._df["trade_date"] <= end)
]
out: list[DailyBar] = []
for r in sub.itertuples():
out.append(
DailyBar(
symbol=r.symbol,
trade_date=r.trade_date,
open=Decimal(str(r.open)),
high=Decimal(str(r.high)),
low=Decimal(str(r.low)),
close=Decimal(str(r.close)),
volume=Decimal(str(r.volume)),
amount=Decimal(str(r.amount)),
)
)
return out
def get_range(self, symbol, start, end) -> list[DailyBar]:
return self._bars([symbol], start, end)
def get_range_many(self, symbols, start, end) -> list[DailyBar]:
return self._bars(list(symbols), start, end)
def latest_date(self, symbol: str) -> date | None:
sub = self._df[self._df["symbol"] == symbol]
return None if sub.empty else sub["trade_date"].max()
@pytest.fixture()
def svc():
drifts = {sym: 0.006 - 0.0015 * i for i, sym in enumerate(_SYMS)}
df = synthetic_daily(drifts, n=320)
return SelectionService(_MemStockRepo(_mem_stocks()), _MemDailyRepo(df))
def _score_query(**kw) -> SelectionQuery:
base = dict(
factors=[{"name": "momentum_60", "weight": 1.0}],
top_n=3,
as_of=date(2024, 12, 31),
)
base.update(kw)
return SelectionQuery(**base)
class TestScoreSelection:
def test_returns_top_n_with_rank_reason(self, svc) -> None:
res = svc.select(_score_query())
assert res.method == "score"
assert len(res.candidates) == 3
assert [c.rank for c in res.candidates] == [1, 2, 3]
# 分数降序
scores = [c.score for c in res.candidates]
assert scores == sorted(scores, reverse=True)
# 每个候选带因子值与理由(可解释)
for c in res.candidates:
assert c.symbol in _SYMS
assert "momentum_60" in c.factor_values
assert any("momentum_60" in r for r in c.selection_reason)
# 统计:评估数 > 0
assert res.statistics.evaluated >= 3
assert res.statistics.selected == 3
assert res.as_of_date <= date(2024, 12, 31)
def test_highest_drift_ranked_high(self, svc) -> None:
res = svc.select(_score_query(top_n=1))
# 最高漂移股票(600000.SH)应位列前三(动量因子对强趋势敏感)
assert res.candidates[0].symbol in _SYMS[:3]
def test_top_pct(self, svc) -> None:
res = svc.select(_score_query(top_n=None, top_pct=0.4))
assert 1 <= len(res.candidates) <= 3 # 5 只的 40% ≈ 2
assert res.candidates[0].rank == 1
def test_min_score_filters(self, svc) -> None:
res = svc.select(_score_query(min_score=10_000))
# 合成数据最高动量约 +60%,min_score=10000 应无候选
assert res.candidates == []
class TestAsOfNoFutureFunction:
def _two_stage_df(self) -> tuple[pd.DataFrame, date, date]:
"""X 前段强、Y 前段平后段暴涨:as_of=split 时 Y 的 60 日动量应为 0/无。"""
dates = pd.bdate_range("2024-01-01", periods=200)
split = dates[100].date() # d101:Y 从这天开始暴涨
rows: list[dict] = []
for d in dates:
# X:恒定 +0.8%/日
rows.append({"symbol": "600000.SH", "trade_date": d.date(),
"close": 100 * 1.008 ** ((d - dates[0]).days)})
y_price = 100.0
for d in dates:
if d.date() > split:
y_price = y_price * 1.05
rows.append({"symbol": "600001.SH", "trade_date": d.date(), "close": y_price})
df = pd.DataFrame(rows)
for col in ("open", "high", "low", "volume", "amount"):
df[col] = df["close"] * 1.001 if col != "volume" else 1_000_000
return df, split, dates[-1].date()
def test_as_of_cutoff_excludes_future(self) -> None:
df, split, _end = self._two_stage_df()
stocks = [
Stock(symbol="600000.SH", name="A", list_date=date(1999, 1, 1)),
Stock(symbol="600001.SH", name="B", list_date=date(1999, 1, 1)),
]
svc = SelectionService(_MemStockRepo(stocks), _MemDailyRepo(df))
# 历史选股:as_of=split(Y 尚未暴涨)→ Y 不应排第一(其 60 日动量≈0/缺失)
res = svc.select(_score_query(factors=[{"name": "momentum_60", "weight": 1.0}], as_of=split))
assert res.candidates
assert res.candidates[0].symbol == "600000.SH"
# Y 若进入候选,其理由值应接近 0(而非泄漏未来暴涨的巨幅动量)
y = next((c for c in res.candidates if c.symbol == "600001.SH"), None)
if y is not None:
assert y.factor_values["momentum_60"] < 0.1
class TestUniverseFilter:
def test_exclude_st(self, svc) -> None:
stocks = _mem_stocks()[:2]
stocks[0] = stocks[0].model_copy(update={"name": "ST 风险股份"})
s = SelectionService(_MemStockRepo(stocks), svc._daily_repo)
res = s.select(_score_query())
assert all(c.symbol != "600000.SH" for c in res.candidates)
def test_min_listing_days(self) -> None:
stocks = [
Stock(symbol=_SYMS[0], name="新上市", list_date=date(2024, 12, 1)),
Stock(symbol=_SYMS[1], name="老股", list_date=date(1999, 1, 1)),
]
df = synthetic_daily({_SYMS[0]: 0.01, _SYMS[1]: 0.001}, n=320)
s = SelectionService(_MemStockRepo(stocks), _MemDailyRepo(df))
res = s.select(_score_query(as_of=date(2024, 12, 31)))
assert all(c.symbol != _SYMS[0] for c in res.candidates) # 上市不足 250 自然日被滤除
def test_delisted_before_as_of(self) -> None:
stocks = [
Stock(symbol=_SYMS[0], name="已退市", list_date=date(1990, 1, 1),
delist_date=date(2023, 6, 30)),
Stock(symbol=_SYMS[1], name="正常", list_date=date(1999, 1, 1)),
]
df = synthetic_daily({_SYMS[0]: 0.01, _SYMS[1]: 0.001}, n=320)
s = SelectionService(_MemStockRepo(stocks), _MemDailyRepo(df))
res = s.select(_score_query(as_of=date(2024, 6, 30)))
assert all(c.symbol != _SYMS[0] for c in res.candidates)
class TestEmptyAndValidation:
def test_no_stocks_returns_empty(self) -> None:
s = SelectionService(_MemStockRepo([]), _MemDailyRepo(pd.DataFrame()))
res = s.select(_score_query())
assert res.candidates == []
assert res.statistics.selected == 0
def test_query_validation(self) -> None:
with pytest.raises(ValidationError):
SelectionQuery(method="score", top_n=5) # score 缺 factors
with pytest.raises(ValidationError):
SelectionQuery(method="condition", conditions=[], top_n=5) # condition 缺 conditions
with pytest.raises(ValidationError):
SelectionQuery(factors=[{"name": "a", "weight": 1}], top_n=None, top_pct=None)
with pytest.raises(ValidationError):
SelectionQuery(
factors=[{"name": "a", "weight": 1}, {"name": "a", "weight": 2}], top_n=5
)
# 合法
SelectionQuery(factors=[{"name": "a", "weight": 1}], top_n=5)
def test_query_top_pct_valid(self) -> None:
q = SelectionQuery(factors=[{"name": "momentum_60", "weight": 1}], top_pct=0.5)
assert q.top_n is None and q.top_pct == 0.5