- 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 通过
224 lines
8.8 KiB
Python
224 lines
8.8 KiB
Python
"""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
|