feat(selection): M6.2 条件选股(method=condition + 财务可见性防护)
- quant/selection.run_condition_selection:结构化条件 AND 求值 —— 字段域 static.*(行业/市场…)、技术列与派生量(close/volume/ma20/ma60)、已注册因子 (momentum_60 等)、fundamental.*(announce_date<=as_of 的最新已公告财务值); 条件支持 value 字面量与 ref 字段比较(如 close > ma60);结果带 filter_status/reason - SelectionQuery 校验调整:condition 模式为纯过滤(不再强制 top_n/top_pct) - FinancialRepository 新增 list_announced_many(批量防未来函数读取)+ SQLAlchemy 实现; SelectionService 注入 financial_repo 并按 announce_date 取每股最新一版 - tests/test_selection_condition.py:8 例(行业 in/ne、动量>0、close>ma60 ref、阈值、 ROE 过滤且未来公告不可见、缺财务 repo 报错、更早 as_of 排除);全量 pytest 通过
This commit is contained in:
@@ -14,6 +14,7 @@ from datetime import date
|
||||
|
||||
import pandas as pd
|
||||
|
||||
from app.domain.entities.market import FinancialIndicator
|
||||
from app.domain.entities.selection import (
|
||||
SelectionCandidate,
|
||||
SelectionQuery,
|
||||
@@ -163,6 +164,195 @@ 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: SelectionQuery) -> set[str]:
|
||||
"""条件引用的行情列(fundamental/static 走元数据与财务表,不需要行情列)。"""
|
||||
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
|
||||
try: # 其余按已注册因子处理
|
||||
defn, _fn = get_factor(f)
|
||||
except FactorError:
|
||||
raise ValueError(
|
||||
f"条件字段未知:{f}(可用: 行情列/ma20/ma60/已注册因子/static.*/fundamental.*)"
|
||||
) from None
|
||||
needed.update(defn.requires)
|
||||
return needed
|
||||
|
||||
|
||||
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"),
|
||||
)
|
||||
|
||||
view = daily[pd.to_datetime(daily["trade_date"]) <= obs]
|
||||
close = view.pivot(index="trade_date", columns="symbol", values="close").sort_index()
|
||||
close.index = pd.to_datetime(close.index)
|
||||
|
||||
# 技术字段面板(obs 行)
|
||||
tech: dict[str, pd.Series] = {}
|
||||
for col in ("close", "open", "high", "low", "volume", "amount"):
|
||||
if col in view.columns and col != "close":
|
||||
panel = view.pivot(index="trade_date", columns="symbol", values=col).sort_index()
|
||||
panel.index = pd.to_datetime(panel.index)
|
||||
tech[col] = panel.loc[obs]
|
||||
tech["close"] = close.loc[obs]
|
||||
tech["ma20"] = close.rolling(20).mean().loc[obs]
|
||||
tech["ma60"] = close.rolling(60).mean().loc[obs]
|
||||
# 因子字段按需计算
|
||||
for cond in query.conditions:
|
||||
for f in (cond.field, cond.ref):
|
||||
if f is None or f.startswith((_STATIC_PREFIX, _FUNDAMENTAL_PREFIX)) or f in tech:
|
||||
continue
|
||||
if f in _TECH_DERIVED or f in ("close", "open", "high", "low", "volume", "amount"):
|
||||
continue
|
||||
try:
|
||||
_defn, panel = compute_factor(f, view)
|
||||
except FactorError:
|
||||
continue # 已在 condition_needed_columns 报错;此处防御
|
||||
if obs in panel.index:
|
||||
tech[f] = panel.loc[obs]
|
||||
|
||||
statics = {s.symbol: s.model_dump() for s in stocks}
|
||||
candidates: list[SelectionCandidate] = []
|
||||
passed_symbols: list[str] = []
|
||||
for sym in sorted(statics):
|
||||
statuses: list[str] = []
|
||||
all_ok = True
|
||||
for cond in query.conditions:
|
||||
ok = _eval_condition(cond, sym, statics, tech, financial or {})
|
||||
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_symbols.append(sym)
|
||||
candidates.append(
|
||||
SelectionCandidate(
|
||||
symbol=sym,
|
||||
rank=0, # 占位,末尾统一编号
|
||||
score=1.0,
|
||||
filter_status=statuses,
|
||||
selection_reason=[f"通过全部 {len(query.conditions)} 条条件"],
|
||||
)
|
||||
)
|
||||
for rank, c in enumerate(candidates, start=1):
|
||||
c.rank = rank
|
||||
|
||||
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 外不通过;数值/字符串分别处理。"""
|
||||
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":
|
||||
return left in right
|
||||
if op == "not_in":
|
||||
return left not in right
|
||||
if op in ("gt", "gte", "lt", "lte"):
|
||||
return _num_cmp(left, right, op) # 尝试数值,字符串会 ValueError → 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
|
||||
|
||||
Reference in New Issue
Block a user