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 通过
This commit is contained in:
Simon
2026-09-09 00:12:28 +08:00
parent 697ffc767b
commit f3586adb25
5 changed files with 606 additions and 21 deletions
+223
View File
@@ -0,0 +1,223 @@
"""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