"""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.market import DAILY_BASIC_NUMERIC_FIELDS, FinancialIndicator from app.domain.entities.selection import ( SelectionCandidate, SelectionQuery, SelectionResult, SelectionStatistics, ) from app.quant.composite import build_score_panel from app.quant.factors import FactorError, compute_factor, get_factor _UNIMPLEMENTED_DEFAULT = [ "exclude_suspended 依赖停牌数据,当前未建模(结果可能包含停牌股)", ] def score_panel_for_factors(daily: pd.DataFrame, factor_specs) -> pd.DataFrame: """因子加权复合分面板(index=trade_date, columns=symbol)。 回测(LocalEngine)与选股(run_score_selection)共用同一构建 —— 保证 v2 §25/§27「历史回测与当前选股使用同一套引擎」的一致性。 """ return build_score_panel(daily, factor_specs) # 未知因子在此抛 FactorError 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"), ) score = score_panel_for_factors(view, query.factors).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 # ---------- method=condition:结构化条件选股(M6.2) ---------- # 技术字段:预计算派生量 + 行情原列(原列需在装配列中才可用) _TECH_DERIVED = ("ma20", "ma60") _STATIC_PREFIX = "static." _FUNDAMENTAL_PREFIX = "fundamental." def condition_needed_columns(query) -> set[str]: """条件引用的行情/指标列(fundamental/static 走元数据与财务表,不需要面板列)。 query 可以是 SelectionQuery,也可以是任何带 `conditions` 的对象 (ResearchSpec 亦然)—— 回测与选股共用本函数,保证列裁剪一致。 """ needed = {"close"} names = [c.field for c in query.conditions] + [ c.ref for c in query.conditions if c.ref and not c.ref.startswith(_FUNDAMENTAL_PREFIX) ] for f in names: if not f or f.startswith((_STATIC_PREFIX, _FUNDAMENTAL_PREFIX)): continue if f in {"open", "high", "low", "close", "volume", "amount", *_TECH_DERIVED}: if f not in _TECH_DERIVED: needed.add(f) continue if f in DAILY_BASIC_NUMERIC_FIELDS: needed.add(f) # 每日指标列(dv_ratio / pe / pb / total_mv …) continue try: # 其余按已注册因子处理 defn, _fn = get_factor(f) except FactorError: raise ValueError( f"条件字段未知:{f}(可用: 行情列/ma20/ma60/已注册因子/" "参数化因子键(如 momentum(window=90,direction=higher_is_better))/" "每日指标列(dv_ratio 等)/static.*/fundamental.*)" ) from None needed.update(defn.requires) return needed def build_condition_fields( daily: pd.DataFrame, conditions, obs: pd.Timestamp, ) -> dict[str, pd.Series]: """条件各字段在 obs(<= as_of 的最近交易日)的截面值。 返回 {字段名: Series(index=symbol)},覆盖: - 行情原列:close / open / high / low / volume / amount - 技术派生:ma20 / ma60 - 每日指标列:dv_ratio / dv_ttm / pe / pb / total_mv …(由 Service 并入 daily) - 已注册因子:momentum_60 等(在 <=obs 的截断数据上计算,无未来函数) 回测与选股共用本函数(v2 §25 一致性)。 """ view = daily[pd.to_datetime(daily["trade_date"]) <= obs] if view.empty: return {} close = view.pivot(index="trade_date", columns="symbol", values="close").sort_index() close.index = pd.to_datetime(close.index) fields: dict[str, pd.Series] = {} for col in ("open", "high", "low", "volume", "amount"): if col in view.columns: panel = view.pivot(index="trade_date", columns="symbol", values=col).sort_index() panel.index = pd.to_datetime(panel.index) if obs in panel.index: fields[col] = panel.loc[obs] if obs in close.index: fields["close"] = close.loc[obs] ma20 = close.rolling(20).mean() ma60 = close.rolling(60).mean() if obs in ma20.index: fields["ma20"] = ma20.loc[obs] if obs in ma60.index: fields["ma60"] = ma60.loc[obs] wanted: set[str] = set() for cond in conditions: for f in (cond.field, cond.ref): if not f or f.startswith((_STATIC_PREFIX, _FUNDAMENTAL_PREFIX)): continue if f in _TECH_DERIVED or f in fields: continue wanted.add(f) for f in sorted(wanted): if f in DAILY_BASIC_NUMERIC_FIELDS: if f in view.columns: panel = view.pivot(index="trade_date", columns="symbol", values=f).sort_index() panel.index = pd.to_datetime(panel.index) if obs in panel.index: fields[f] = panel.loc[obs] continue try: _defn, panel = compute_factor(f, view) except FactorError: continue # 已在 condition_needed_columns 报错;此处防御 if obs in panel.index: fields[f] = panel.loc[obs] return fields def eligible_symbols( candidates, conditions, statics: dict[str, dict], fields: dict[str, pd.Series], financial: dict[str, FinancialIndicator], ) -> dict[str, list[str]]: """逐股求值全部条件(AND),返回 {symbol: [各条件通过情况文案]}(仅通过者)。 `candidates` 限定参与求值的股票(通常 = universe 过滤后的 symbol 列表)。 回测(ResearchSpec.conditions)与选股(SelectionQuery.conditions)共用, 确保「历史某日 Selection == 回测当日 Selection」(v2 §25 / v3 §28)。 """ passed: dict[str, list[str]] = {} for sym in candidates: statuses: list[str] = [] all_ok = True for cond in conditions: ok = _eval_condition(cond, sym, statics, fields, financial) statuses.append( f"{cond.field} {cond.op} {cond.ref or cond.value}: {'通过' if ok else '未通过'}" ) all_ok = all_ok and ok if all_ok: passed[sym] = statuses return passed def run_condition_selection( daily: pd.DataFrame, stocks: list, query: SelectionQuery, as_of: date | None, financial: dict[str, FinancialIndicator] | None = None, ) -> SelectionResult: """条件选股(v2 §14.1A):全部条件 AND 通过者入选(无排序;truncation 不适用)。 fields 域:static.*(股票基础)、close/volume/amount/ma20/ma60/已注册因子(行情)、 fundamental.*(announce_date <= as_of 的最新已公告财务值 —— 防未来函数由 Service 取数保证)。 """ if query.method != "condition": raise ValueError(f"run_condition_selection 需要 method=condition,当前 {query.method}") obs = resolve_observation_date(daily, as_of) resolved = (obs.date() if obs is not None else as_of) or date.today() if obs is None or daily.empty: return SelectionResult( as_of_date=resolved, method=query.method, statistics=SelectionStatistics(), candidates=[], unimplemented=list(_UNIMPLEMENTED_DEFAULT), config_snapshot=query.model_dump(mode="json"), ) # 共享字段面板 + 求值器(回测与选股同一实现,v2 §25 一致性) tech = build_condition_fields(daily, query.conditions, obs) statics = {s.symbol: s.model_dump() for s in stocks} passed = eligible_symbols(sorted(statics), query.conditions, statics, tech, financial or {}) candidates: list[SelectionCandidate] = [] for rank, sym in enumerate(sorted(passed), start=1): candidates.append( SelectionCandidate( symbol=sym, rank=rank, score=1.0, filter_status=passed[sym], selection_reason=[f"通过全部 {len(query.conditions)} 条条件"], ) ) return SelectionResult( as_of_date=resolved, method=query.method, statistics=SelectionStatistics( universe_size=len(statics), evaluated=len(statics), selected=len(candidates), ), candidates=candidates, unimplemented=list(_UNIMPLEMENTED_DEFAULT) + [ "条件选股为纯过滤(AND),未排序/未截断;如需排序请在 factors 中提供评分", ], config_snapshot=query.model_dump(mode="json"), ) def _eval_condition( cond, sym: str, statics: dict, tech: dict[str, pd.Series], financial: dict, ) -> bool: """求值单条条件:value 与 ref 二选一;left 与 right 同为 field 或 field vs 字面量。""" left = _field_value(cond.field, sym, statics, tech, financial) if cond.ref is not None: right = _field_value(cond.ref, sym, statics, tech, financial) else: right = cond.value return _compare(left, right, cond.op) def _field_value(field, sym, statics, tech, financial): if field.startswith(_STATIC_PREFIX): return statics.get(sym, {}).get(field[len(_STATIC_PREFIX):]) if field.startswith(_FUNDAMENTAL_PREFIX): fin = financial.get(sym) return getattr(fin, field[len(_FUNDAMENTAL_PREFIX):], None) if fin else None series = tech.get(field) if series is None: return None v = series.get(sym) return None if v is None or (isinstance(v, float) and v != v) else v # NaN → None def _compare(left, right, op: str) -> bool: """混合比较:None 视为不可用 → 除 ne 外不通过;数值/字符串分别处理。 任何类型不匹配(拿日期字段去比大小、in 的右侧不是列表…)一律返回 False, **绝不抛异常**:条件是用户可以随手改的输入,一个手滑的字段名不该把选股/回测 打成 500。字段库(quant.condition_fields)会在源头上拒绝不可比较的字段, 这里只是最后一道防线。 """ if op == "ne": return left != right if left is None or right is None: return False try: if isinstance(left, (int, float)) or isinstance(right, (int, float)): return _num_cmp(float(left), float(right), op) except (TypeError, ValueError): pass # 字符串/其它:支持 eq/ne/in/not_in if op == "eq": return left == right if op in ("in", "not_in"): try: return left in right if op == "in" else left not in right except TypeError: # 右侧不是容器 → 该条件无法求值 return False if op in ("gt", "gte", "lt", "lte"): try: return _num_cmp(left, right, op) # 尝试数值,字符串会 ValueError → False except (TypeError, ValueError): return False return False def _num_cmp(a: float, b: float, op: str) -> bool: if op == "gt": return a > b if op == "gte": return a >= b if op == "lt": return a < b if op == "lte": return a <= b if op == "eq": return a == b return a != b 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