From f3586adb250a722d57bdeb46af6d073c86b55e3b Mon Sep 17 00:00:00 2001 From: Simon Date: Wed, 9 Sep 2026 00:12:28 +0800 Subject: [PATCH] =?UTF-8?q?feat(selection):=20M6.0=20=E9=80=89=E8=82=A1?= =?UTF-8?q?=E5=A5=91=E7=BA=A6=E4=B8=8E=E8=AF=84=E5=88=86=E5=BC=95=E6=93=8E?= =?UTF-8?q?=EF=BC=88SelectionQuery/Result=20+=20select(as=5Fof)=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 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 通过 --- .../application/services/selection_service.py | 50 ++++ backend/app/domain/entities/selection.py | 125 ++++++++++ backend/app/quant/selection.py | 175 ++++++++++++++ backend/app/quant/service.py | 54 +++-- backend/tests/test_selection.py | 223 ++++++++++++++++++ 5 files changed, 606 insertions(+), 21 deletions(-) create mode 100644 backend/app/application/services/selection_service.py create mode 100644 backend/app/domain/entities/selection.py create mode 100644 backend/app/quant/selection.py create mode 100644 backend/tests/test_selection.py diff --git a/backend/app/application/services/selection_service.py b/backend/app/application/services/selection_service.py new file mode 100644 index 0000000..dde9052 --- /dev/null +++ b/backend/app/application/services/selection_service.py @@ -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) diff --git a/backend/app/domain/entities/selection.py b/backend/app/domain/entities/selection.py new file mode 100644 index 0000000..a6c4583 --- /dev/null +++ b/backend/app/domain/entities/selection.py @@ -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() diff --git a/backend/app/quant/selection.py b/backend/app/quant/selection.py new file mode 100644 index 0000000..e232d9f --- /dev/null +++ b/backend/app/quant/selection.py @@ -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 diff --git a/backend/app/quant/service.py b/backend/app/quant/service.py index e5d3a03..738b0da 100644 --- a/backend/app/quant/service.py +++ b/backend/app/quant/service.py @@ -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) + ) diff --git a/backend/tests/test_selection.py b/backend/tests/test_selection.py new file mode 100644 index 0000000..a3fc283 --- /dev/null +++ b/backend/tests/test_selection.py @@ -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