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:
@@ -0,0 +1,50 @@
|
||||
"""选股用例入口(ARCHITECTURE_v2 §14 Selection Engine · 业务层)。
|
||||
|
||||
- 输入:SelectionQuery(universe + method + factors/conditions + top_n/pct + as_of)
|
||||
- 装配:股票池(universe 过滤)→ 行情长表(含预热窗口)→ Selection Engine
|
||||
- 输出:SelectionResult(可解释:factor_values / selection_reason)
|
||||
- 未来函数红线:全部数据只取到 <= as_of(v2 §9);财务条件(后续)只取已公告值
|
||||
|
||||
MVP 为同步执行(单日全市场因子计算量轻);如需异步可复用 Job 链路(M6.3 决策)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, timedelta
|
||||
|
||||
import pandas as pd
|
||||
|
||||
from app.domain.entities.selection import SelectionQuery, SelectionResult
|
||||
from app.domain.repositories.market import DailyBarRepository, StockRepository
|
||||
from app.quant.selection import factor_columns, run_score_selection
|
||||
from app.quant.service import filter_stocks, load_daily_df
|
||||
|
||||
|
||||
class SelectionService:
|
||||
"""选股用例入口:select(query) → SelectionResult(当前或历史 as_of)。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
stock_repo: StockRepository,
|
||||
daily_repo: DailyBarRepository,
|
||||
) -> None:
|
||||
self._stock_repo = stock_repo
|
||||
self._daily_repo = daily_repo
|
||||
|
||||
def select(self, query: SelectionQuery) -> SelectionResult:
|
||||
as_of = query.as_of or date.today()
|
||||
stocks = filter_stocks(self._stock_repo.list(), query.universe, as_of=as_of)
|
||||
if not stocks:
|
||||
return run_score_selection(pd.DataFrame(), query, as_of)
|
||||
symbols = [s.symbol for s in stocks]
|
||||
columns = sorted(factor_columns(query)) if query.method == "score" else ["close"]
|
||||
daily = load_daily_df(
|
||||
self._daily_repo,
|
||||
symbols,
|
||||
as_of - timedelta(days=query.warmup_days),
|
||||
as_of,
|
||||
columns,
|
||||
)
|
||||
if daily.empty:
|
||||
return run_score_selection(daily, query, as_of)
|
||||
return run_score_selection(daily, query, as_of)
|
||||
@@ -0,0 +1,125 @@
|
||||
"""选股系统领域对象(ARCHITECTURE_v2 §14 Selection Engine)。
|
||||
|
||||
回答两个核心问题(v2 §8):
|
||||
- 「某历史日(as_of)为什么选出这些股票?」→ SelectionResult 带 factor_values / selection_reason
|
||||
- 「当前(as_of)有哪些股票满足策略?」→ 同一条查询对当前日期执行
|
||||
|
||||
设计:
|
||||
- SelectionQuery = v2 §14.2 的 Selection 输入(universe 范围 + 评分因子 + TopN 截断 + as_of)。
|
||||
- method=score:按因子加权复合分取 TopN(复用现有 9 个内置因子);
|
||||
method=condition:结构化条件选股(M6.2 加入 ConditionSpec)。
|
||||
- 结果不落库由本实体负责(落库表在 M6.3);本实体是前后端/Agent 的统一 DTO(v2 §21.1)。
|
||||
- 所有查询天然带 as_of 语义:只允许使用 <= as_of 的数据(v2 §9 防未来函数)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
|
||||
from pydantic import BaseModel, Field, field_validator, model_validator
|
||||
|
||||
from app.domain.entities.research import FactorSpec, UniverseSpec
|
||||
|
||||
|
||||
class SelectionQuery(BaseModel):
|
||||
"""一次选股查询(v2 §14.2 Selection 输入)。"""
|
||||
|
||||
universe: UniverseSpec = UniverseSpec()
|
||||
# 研究时点:None → 引擎用 <= 今天最近可用交易日;显式给历史日期即做历史选股
|
||||
as_of: date | None = Field(
|
||||
default=None, description="选股时点;历史回测/解释用具体日期,当前选股可留空"
|
||||
)
|
||||
method: str = Field(default="score", pattern="^(score|condition)$")
|
||||
# method=score:因子 + 权重(至少 1 个;方向由因子元数据决定)
|
||||
factors: list[FactorSpec] = Field(default_factory=list)
|
||||
# method=condition:结构化条件(M6.2 引入 ConditionSpec 后启用)
|
||||
conditions: list[ConditionSpec] = Field(default_factory=list)
|
||||
# 截断:top_n(绝对数量)与 top_pct(占可评分股票比例)二选一;可选 min_score 下限
|
||||
top_n: int | None = Field(default=None, ge=1, le=2000)
|
||||
top_pct: float | None = Field(default=None, gt=0, le=1)
|
||||
min_score: float | None = None
|
||||
# 因子预热窗口(自然日):覆盖 lookback 前导数据,Lookback 放大时需同步加大
|
||||
warmup_days: int = Field(default=300, ge=0)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _check_method_args(self) -> SelectionQuery:
|
||||
if self.method == "score" and not self.factors:
|
||||
raise ValueError("method=score 需要至少一个 factors")
|
||||
if self.method == "condition" and not self.conditions:
|
||||
raise ValueError("method=condition 需要至少一个 conditions")
|
||||
if self.top_n is None and self.top_pct is None:
|
||||
raise ValueError("top_n 与 top_pct 至少提供一个")
|
||||
return self
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _no_duplicate_factors(self) -> SelectionQuery:
|
||||
names = [f.name for f in self.factors]
|
||||
if len(set(names)) != len(names):
|
||||
raise ValueError("factors 存在重复因子名")
|
||||
return self
|
||||
|
||||
|
||||
class ConditionSpec(BaseModel):
|
||||
"""结构化选股条件(M6.2 使用)。
|
||||
|
||||
field 域:
|
||||
- static.*:股票基础字段(industry / market / area / exchange / status…)
|
||||
- 行情/技术字段:close / ma20 / ma60 / volume 及全部已注册因子名(momentum_60 等)
|
||||
- fundamental.*:财务字段(eps / roe / total_revenue / net_profit / gross_margin),
|
||||
仅取 announce_date <= as_of 的最新已公告值(防未来函数)
|
||||
右操作数取 value(字面量)或 ref(另一字段名),二者二选一。
|
||||
"""
|
||||
|
||||
field: str
|
||||
op: str = Field(pattern="^(gt|gte|lt|lte|eq|ne|in|not_in)$")
|
||||
value: float | int | str | list | None = None
|
||||
ref: str | None = None # 与另一字段比较(如 close vs ma60)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _require_operand(self) -> ConditionSpec:
|
||||
if self.value is None and self.ref is None:
|
||||
raise ValueError("value 与 ref 必须提供一个")
|
||||
if self.value is not None and self.ref is not None:
|
||||
raise ValueError("value 与 ref 只能提供一个")
|
||||
if self.op in ("in", "not_in") and not isinstance(self.value, list):
|
||||
raise ValueError("in/not_in 的 value 必须是列表")
|
||||
return self
|
||||
|
||||
|
||||
class SelectionCandidate(BaseModel):
|
||||
"""单只候选股(v2 §14.3/§21.1)。"""
|
||||
|
||||
symbol: str
|
||||
rank: int
|
||||
score: float
|
||||
factor_values: dict[str, float] = Field(default_factory=dict)
|
||||
filter_status: list[str] = Field(default_factory=list, description="各条件通过/未通过")
|
||||
selection_reason: list[str] = Field(default_factory=list, description="为什么选它(可解释)")
|
||||
|
||||
|
||||
class SelectionStatistics(BaseModel):
|
||||
universe_size: int = 0 # 股票池过滤后数量
|
||||
evaluated: int = 0 # 有有效分数的股票数量
|
||||
selected: int = 0 # 最终选出数量
|
||||
|
||||
|
||||
class SelectionResult(BaseModel):
|
||||
"""选股结果(v2 §21.1)。前端 / Agent 只依赖该结构。"""
|
||||
|
||||
as_of_date: date
|
||||
method: str
|
||||
statistics: SelectionStatistics
|
||||
candidates: list[SelectionCandidate] = Field(default_factory=list)
|
||||
unimplemented: list[str] = Field(
|
||||
default_factory=list,
|
||||
description="本结果中未建模的约束(如 exclude_suspended 依赖停牌数据未实现)",
|
||||
)
|
||||
config_snapshot: dict = Field(default_factory=dict, description="复现用查询快照")
|
||||
|
||||
@field_validator("candidates")
|
||||
@classmethod
|
||||
def _rank_sorted(cls, candidates: list[SelectionCandidate]) -> list[SelectionCandidate]:
|
||||
return sorted(candidates, key=lambda c: c.rank)
|
||||
|
||||
|
||||
SelectionQuery.model_rebuild()
|
||||
@@ -0,0 +1,175 @@
|
||||
"""Selection Engine(ARCHITECTURE_v2 §14)—— 纯 pandas 执行层。
|
||||
|
||||
当前实现 method=score:因子加权复合分 → TopN/Top% 截断,输出 SelectionResult。
|
||||
M6.2 在同一模块加入 method=condition(结构化条件选股)。
|
||||
|
||||
未来函数纪律:面板只在 <= observation_date 的数据上计算;observation_date 是
|
||||
<= as_of 的最近可用交易日(as_of 显式传入即历史选股,None 则到数据最新)。
|
||||
data 长表由 Service 装配(已按 universe 过滤 symbol、含预热窗口)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
|
||||
import pandas as pd
|
||||
|
||||
from app.domain.entities.selection import (
|
||||
SelectionCandidate,
|
||||
SelectionQuery,
|
||||
SelectionResult,
|
||||
SelectionStatistics,
|
||||
)
|
||||
from app.quant.factors import FactorError, compute_factor, get_factor
|
||||
from app.quant.local_engine import build_factor_panels, composite_score
|
||||
|
||||
_UNIMPLEMENTED_DEFAULT = [
|
||||
"exclude_suspended 依赖停牌数据,当前未建模(结果可能包含停牌股)",
|
||||
]
|
||||
|
||||
|
||||
def resolve_observation_date(daily: pd.DataFrame, as_of: date | None) -> pd.Timestamp | None:
|
||||
"""<= as_of 的最近可用交易日;as_of=None 取数据最新一日。"""
|
||||
if daily.empty:
|
||||
return None
|
||||
dates = pd.to_datetime(daily["trade_date"])
|
||||
if as_of is None:
|
||||
return dates.max()
|
||||
avail = dates[dates <= pd.Timestamp(as_of)]
|
||||
return avail.max() if len(avail) else None
|
||||
|
||||
|
||||
def factor_columns(query: SelectionQuery) -> set[str]:
|
||||
"""score 模式所需行情数值列(数据装配裁剪用)。"""
|
||||
needed = {"close"}
|
||||
for fs in query.factors:
|
||||
try:
|
||||
defn, _fn = get_factor(fs.name)
|
||||
except FactorError:
|
||||
continue # 未知因子由执行期统一报错(score_selection 中 build_factor_panels)
|
||||
needed.update(defn.requires)
|
||||
return needed
|
||||
|
||||
|
||||
def run_score_selection(
|
||||
daily: pd.DataFrame,
|
||||
query: SelectionQuery,
|
||||
as_of: date | None,
|
||||
) -> SelectionResult:
|
||||
"""因子评分选股(v2 §14.1B):复合分 → 排序 → TopN/Top%。"""
|
||||
if query.method != "score":
|
||||
raise ValueError(f"run_score_selection 需要 method=score,当前 {query.method}")
|
||||
obs = resolve_observation_date(daily, as_of)
|
||||
if obs is None:
|
||||
resolved = as_of or date.today()
|
||||
return SelectionResult(
|
||||
as_of_date=resolved,
|
||||
method=query.method,
|
||||
statistics=SelectionStatistics(),
|
||||
candidates=[],
|
||||
unimplemented=list(_UNIMPLEMENTED_DEFAULT),
|
||||
config_snapshot=query.model_dump(mode="json"),
|
||||
)
|
||||
resolved = obs.date()
|
||||
# 只允许使用 <= obs 的数据(面板计算在截断后数据上进行)
|
||||
view = daily[pd.to_datetime(daily["trade_date"]) <= obs]
|
||||
if view.empty:
|
||||
return SelectionResult(
|
||||
as_of_date=resolved,
|
||||
method=query.method,
|
||||
statistics=SelectionStatistics(),
|
||||
candidates=[],
|
||||
unimplemented=list(_UNIMPLEMENTED_DEFAULT),
|
||||
config_snapshot=query.model_dump(mode="json"),
|
||||
)
|
||||
|
||||
panels = build_factor_panels(view, query.factors) # 未知因子在此抛 FactorError
|
||||
score = composite_score(panels).loc[obs].dropna().sort_values(ascending=False)
|
||||
|
||||
# 每因子在 obs 行的原始值(factor_values 供展示与解释;与 build_factor_panels 同数据)
|
||||
raw: dict[str, pd.Series] = {}
|
||||
for fs in query.factors:
|
||||
_defn, panel = compute_factor(fs.name, view)
|
||||
if obs in panel.index:
|
||||
raw[fs.name] = panel.loc[obs]
|
||||
|
||||
candidates_df = _truncate(score, query)
|
||||
evaluated = int(len(score)) # score 已 dropna,长度即有分股票数
|
||||
candidates: list[SelectionCandidate] = []
|
||||
for rank, (sym, sc) in enumerate(candidates_df.items(), start=1):
|
||||
factor_values = {
|
||||
name: _to_float(series.get(sym))
|
||||
for name, series in raw.items()
|
||||
if isinstance(series, pd.Series)
|
||||
}
|
||||
factor_values = {k: v for k, v in factor_values.items() if v is not None}
|
||||
candidates.append(
|
||||
SelectionCandidate(
|
||||
symbol=sym,
|
||||
rank=rank,
|
||||
score=round(float(sc), 6),
|
||||
factor_values=factor_values,
|
||||
selection_reason=_score_reason(query, sym, raw),
|
||||
)
|
||||
)
|
||||
|
||||
return SelectionResult(
|
||||
as_of_date=resolved,
|
||||
method=query.method,
|
||||
statistics=SelectionStatistics(
|
||||
universe_size=_symbol_count(view),
|
||||
evaluated=evaluated,
|
||||
selected=len(candidates),
|
||||
),
|
||||
candidates=candidates,
|
||||
unimplemented=list(_UNIMPLEMENTED_DEFAULT),
|
||||
config_snapshot=query.model_dump(mode="json"),
|
||||
)
|
||||
|
||||
|
||||
def _truncate(score: pd.Series, query: SelectionQuery) -> pd.Series:
|
||||
"""按 top_n / top_pct / min_score 截断(入参已按分数降序)。"""
|
||||
s = score
|
||||
if query.min_score is not None:
|
||||
s = s[s >= query.min_score]
|
||||
if query.top_pct is not None:
|
||||
n = max(int(round(len(s) * query.top_pct)), 1)
|
||||
s = s.head(n)
|
||||
elif query.top_n is not None:
|
||||
s = s.head(query.top_n)
|
||||
return s
|
||||
|
||||
|
||||
def _score_reason(query: SelectionQuery, symbol: str, raw: dict[str, pd.Series]) -> list[str]:
|
||||
"""生成可读的入选理由:列每个因子的观测值与权重。"""
|
||||
reasons: list[str] = []
|
||||
for fs in query.factors:
|
||||
try:
|
||||
defn, _fn = get_factor(fs.name)
|
||||
except FactorError:
|
||||
continue
|
||||
series = raw.get(fs.name)
|
||||
val = _to_float(series.get(symbol)) if isinstance(series, pd.Series) else None
|
||||
if val is None:
|
||||
continue
|
||||
good = defn.direction == "higher_is_better"
|
||||
reasons.append(
|
||||
f"{fs.name}={val:.4f}(权重 {fs.weight},{'越高越好' if good else '越低越好'})"
|
||||
)
|
||||
return reasons
|
||||
|
||||
|
||||
def _symbol_count(daily: pd.DataFrame) -> int:
|
||||
return int(daily["symbol"].nunique()) if not daily.empty and "symbol" in daily else 0
|
||||
|
||||
|
||||
def _to_float(v) -> float | None:
|
||||
if v is None:
|
||||
return None
|
||||
try:
|
||||
f = float(v)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
if f != f: # NaN
|
||||
return None
|
||||
return f
|
||||
@@ -88,6 +88,36 @@ def _frame_from_stream(rows: Iterable[tuple], columns: list[str]) -> pd.DataFram
|
||||
return df
|
||||
|
||||
|
||||
def load_daily_df(
|
||||
daily_repo,
|
||||
symbols: list[str],
|
||||
start: date,
|
||||
end: date,
|
||||
columns: list[str],
|
||||
) -> pd.DataFrame:
|
||||
"""从 Repository 装配行情长表(供研究/选股共用)。
|
||||
|
||||
优先走流式列裁剪(stream_range_many_columns,SQL 侧转 REAL、分批),
|
||||
失败或实现缺失时回退 get_range_many / 逐只 get_range。
|
||||
"""
|
||||
if not symbols:
|
||||
return pd.DataFrame()
|
||||
streamer = getattr(daily_repo, "stream_range_many_columns", None)
|
||||
if streamer is not None:
|
||||
try:
|
||||
return _frame_from_stream(streamer(symbols, start, end, sorted(columns)), sorted(columns))
|
||||
except Exception: # noqa: BLE001 —— 流式路径失败回退旧路径(兼容非 SQL 实现)
|
||||
pass
|
||||
get_many = getattr(daily_repo, "get_range_many", None)
|
||||
if get_many is not None:
|
||||
bars = list(get_many(symbols, start, end))
|
||||
else: # 兜底:逐只查询
|
||||
bars = []
|
||||
for sym in symbols:
|
||||
bars.extend(daily_repo.get_range(sym, start, end))
|
||||
return bars_to_daily_df(bars)
|
||||
|
||||
|
||||
class ResearchService:
|
||||
"""研究用例入口(因子测试 / 回测)。依赖注入 Repository 与引擎。"""
|
||||
|
||||
@@ -120,26 +150,8 @@ class ResearchService:
|
||||
# 回测前预留因子 warmup(lookback≤120 交易日,取 300 自然日余量)
|
||||
data_start = start - timedelta(days=300)
|
||||
stocks = filter_stocks(self._stock_repo.list(), spec.universe, as_of=start)
|
||||
if not stocks:
|
||||
return pd.DataFrame()
|
||||
symbols = [s.symbol for s in stocks]
|
||||
|
||||
# 引擎所需列裁剪(LocalEngine 只取 close + 因子字段;Qlib 回测取全 OHLCV)
|
||||
required = self._engine.required_columns(spec)
|
||||
streamer = getattr(self._daily_repo, "stream_range_many_columns", None)
|
||||
if streamer is not None:
|
||||
try:
|
||||
return _frame_from_stream(
|
||||
streamer(symbols, data_start, end, sorted(required)), sorted(required)
|
||||
)
|
||||
except Exception: # noqa: BLE001 —— 流式路径失败回退旧路径(兼容非 SQL 实现)
|
||||
pass
|
||||
# 旧路径:逐实体(供内存 / Fake 仓储等实现使用)
|
||||
get_many = getattr(self._daily_repo, "get_range_many", None)
|
||||
if get_many is not None:
|
||||
bars = list(get_many(symbols, data_start, end))
|
||||
else: # 兜底:逐只查询
|
||||
bars = []
|
||||
for s in stocks:
|
||||
bars.extend(self._daily_repo.get_range(s.symbol, data_start, end))
|
||||
return bars_to_daily_df(bars)
|
||||
return load_daily_df(
|
||||
self._daily_repo, [s.symbol for s in stocks], data_start, end, sorted(required)
|
||||
)
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user