Files
qlib/backend/tests/test_selection.py
T
Simon ef09d5b419 feat(quant): M7.3 研究行情口径显式化(默认不复权 none,可切 qfq)
- DailyBarRepository.get_range_many / stream_range_many_columns 增加 adjust 参数
  (默认 'none')→ SQL 层过滤口径,消除 stock_daily 混 source/adjust 污染因子的风险
- ResearchSpec / SelectionQuery 增加 price_adjustment(none|qfq),随 config_snapshot
  落库可溯源;ResearchService._load_daily 与 SelectionService 装配按口径取数
- tests/test_price_adjustment.py:repo 读取按 adjust 过滤(none/qfq 各自命中)、
  spec 默认与字段记录;全量 pytest 通过
2026-09-09 00:32:55 +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, adjust="none") -> 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