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:
@@ -2,10 +2,11 @@
|
|||||||
|
|
||||||
- 输入:SelectionQuery(universe + method + factors/conditions + top_n/pct + as_of)
|
- 输入:SelectionQuery(universe + method + factors/conditions + top_n/pct + as_of)
|
||||||
- 装配:股票池(universe 过滤)→ 行情长表(含预热窗口)→ Selection Engine
|
- 装配:股票池(universe 过滤)→ 行情长表(含预热窗口)→ Selection Engine
|
||||||
- 输出:SelectionResult(可解释:factor_values / selection_reason)
|
(method=score 因子评分 / method=condition 结构化条件)
|
||||||
- 未来函数红线:全部数据只取到 <= as_of(v2 §9);财务条件(后续)只取已公告值
|
- 输出:SelectionResult(可解释:factor_values / filter_status / selection_reason)
|
||||||
|
- 未来函数红线:行情只取 <= as_of;财务条件只取 announce_date <= as_of 的已公告值(v2 §9)
|
||||||
|
|
||||||
MVP 为同步执行(单日全市场因子计算量轻);如需异步可复用 Job 链路(M6.3 决策)。
|
MVP 为同步执行(单日全市场因子/条件计算量轻);如需异步可复用 Job 链路。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
@@ -14,11 +15,23 @@ from datetime import date, timedelta
|
|||||||
|
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
|
|
||||||
|
from app.domain.entities.market import FinancialIndicator
|
||||||
from app.domain.entities.selection import SelectionQuery, SelectionResult
|
from app.domain.entities.selection import SelectionQuery, SelectionResult
|
||||||
from app.domain.repositories.market import DailyBarRepository, StockRepository
|
from app.domain.repositories.market import (
|
||||||
from app.quant.selection import factor_columns, run_score_selection
|
DailyBarRepository,
|
||||||
|
FinancialRepository,
|
||||||
|
StockRepository,
|
||||||
|
)
|
||||||
|
from app.quant.selection import (
|
||||||
|
condition_needed_columns,
|
||||||
|
factor_columns,
|
||||||
|
run_condition_selection,
|
||||||
|
run_score_selection,
|
||||||
|
)
|
||||||
from app.quant.service import filter_stocks, load_daily_df
|
from app.quant.service import filter_stocks, load_daily_df
|
||||||
|
|
||||||
|
_FUNDAMENTAL_PREFIX = "fundamental."
|
||||||
|
|
||||||
|
|
||||||
class SelectionService:
|
class SelectionService:
|
||||||
"""选股用例入口:select(query) → SelectionResult(当前或历史 as_of)。"""
|
"""选股用例入口:select(query) → SelectionResult(当前或历史 as_of)。"""
|
||||||
@@ -27,17 +40,22 @@ class SelectionService:
|
|||||||
self,
|
self,
|
||||||
stock_repo: StockRepository,
|
stock_repo: StockRepository,
|
||||||
daily_repo: DailyBarRepository,
|
daily_repo: DailyBarRepository,
|
||||||
|
financial_repo: FinancialRepository | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
self._stock_repo = stock_repo
|
self._stock_repo = stock_repo
|
||||||
self._daily_repo = daily_repo
|
self._daily_repo = daily_repo
|
||||||
|
self._financial_repo = financial_repo
|
||||||
|
|
||||||
def select(self, query: SelectionQuery) -> SelectionResult:
|
def select(self, query: SelectionQuery) -> SelectionResult:
|
||||||
as_of = query.as_of or date.today()
|
as_of = query.as_of or date.today()
|
||||||
stocks = filter_stocks(self._stock_repo.list(), query.universe, as_of=as_of)
|
stocks = filter_stocks(self._stock_repo.list(), query.universe, as_of=as_of)
|
||||||
if not stocks:
|
if not stocks:
|
||||||
return run_score_selection(pd.DataFrame(), query, as_of)
|
return self._run(query, pd.DataFrame(), stocks, as_of, financial={})
|
||||||
symbols = [s.symbol for s in stocks]
|
symbols = [s.symbol for s in stocks]
|
||||||
columns = sorted(factor_columns(query)) if query.method == "score" else ["close"]
|
if query.method == "score":
|
||||||
|
columns = sorted(factor_columns(query))
|
||||||
|
else:
|
||||||
|
columns = sorted(condition_needed_columns(query))
|
||||||
daily = load_daily_df(
|
daily = load_daily_df(
|
||||||
self._daily_repo,
|
self._daily_repo,
|
||||||
symbols,
|
symbols,
|
||||||
@@ -45,6 +63,53 @@ class SelectionService:
|
|||||||
as_of,
|
as_of,
|
||||||
columns,
|
columns,
|
||||||
)
|
)
|
||||||
if daily.empty:
|
financial: dict[str, FinancialIndicator] = {}
|
||||||
|
if query.method == "condition" and self._uses_fundamental(query):
|
||||||
|
financial = self._load_financial(symbols, as_of)
|
||||||
|
return self._run(query, daily, stocks, as_of, financial)
|
||||||
|
|
||||||
|
# ---- 内部 ----
|
||||||
|
|
||||||
|
def _run(
|
||||||
|
self,
|
||||||
|
query: SelectionQuery,
|
||||||
|
daily: pd.DataFrame,
|
||||||
|
stocks: list,
|
||||||
|
as_of: date,
|
||||||
|
financial: dict[str, FinancialIndicator],
|
||||||
|
) -> SelectionResult:
|
||||||
|
if query.method == "score":
|
||||||
return run_score_selection(daily, query, as_of)
|
return run_score_selection(daily, query, as_of)
|
||||||
return run_score_selection(daily, query, as_of)
|
return run_condition_selection(daily, stocks, query, as_of, financial)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _uses_fundamental(query: SelectionQuery) -> bool:
|
||||||
|
for c in query.conditions:
|
||||||
|
if c.field.startswith(_FUNDAMENTAL_PREFIX) or (
|
||||||
|
c.ref is not None and c.ref.startswith(_FUNDAMENTAL_PREFIX)
|
||||||
|
):
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
def _load_financial(
|
||||||
|
self, symbols: list[str], as_of: date
|
||||||
|
) -> dict[str, FinancialIndicator]:
|
||||||
|
"""按 announce_date <= as_of 批量取财务,每 symbol 保留最新一版。"""
|
||||||
|
if self._financial_repo is None:
|
||||||
|
raise ValueError("condition 引用了 fundamental.* 字段,但未注入 FinancialRepository")
|
||||||
|
getter = getattr(self._financial_repo, "list_announced_many", None)
|
||||||
|
if getter is not None:
|
||||||
|
rows = list(getter(symbols, as_of))
|
||||||
|
else: # 回退逐只
|
||||||
|
rows = []
|
||||||
|
for sym in symbols:
|
||||||
|
rows.extend(self._financial_repo.list_announced(sym, as_of))
|
||||||
|
by_symbol: dict[str, FinancialIndicator] = {}
|
||||||
|
for row in rows:
|
||||||
|
cur = by_symbol.get(row.symbol)
|
||||||
|
if cur is None or (row.announce_date, row.report_date) > (
|
||||||
|
cur.announce_date,
|
||||||
|
cur.report_date,
|
||||||
|
):
|
||||||
|
by_symbol[row.symbol] = row
|
||||||
|
return by_symbol
|
||||||
|
|||||||
@@ -43,12 +43,13 @@ class SelectionQuery(BaseModel):
|
|||||||
|
|
||||||
@model_validator(mode="after")
|
@model_validator(mode="after")
|
||||||
def _check_method_args(self) -> SelectionQuery:
|
def _check_method_args(self) -> SelectionQuery:
|
||||||
if self.method == "score" and not self.factors:
|
if self.method == "score":
|
||||||
raise ValueError("method=score 需要至少一个 factors")
|
if not self.factors:
|
||||||
|
raise ValueError("method=score 需要至少一个 factors")
|
||||||
|
if self.top_n is None and self.top_pct is None:
|
||||||
|
raise ValueError("method=score 需要 top_n 与 top_pct 至少提供一个")
|
||||||
if self.method == "condition" and not self.conditions:
|
if self.method == "condition" and not self.conditions:
|
||||||
raise ValueError("method=condition 需要至少一个 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
|
return self
|
||||||
|
|
||||||
@model_validator(mode="after")
|
@model_validator(mode="after")
|
||||||
|
|||||||
@@ -86,6 +86,17 @@ class FinancialRepository(Protocol):
|
|||||||
) -> list[FinancialIndicator]:
|
) -> list[FinancialIndicator]:
|
||||||
"""只返回 announce_date <= as_of_date 的记录 —— 未来函数红线。"""
|
"""只返回 announce_date <= as_of_date 的记录 —— 未来函数红线。"""
|
||||||
|
|
||||||
|
def list_announced_many(
|
||||||
|
self,
|
||||||
|
symbols: Sequence[str],
|
||||||
|
as_of_date: date,
|
||||||
|
) -> list[FinancialIndicator]:
|
||||||
|
"""批量版:返回这些股票 announce_date <= as_of_date 的全部记录。
|
||||||
|
|
||||||
|
供选股/截面研究一次性取财务字段(调用方按需取每 symbol 最新一版)。
|
||||||
|
实现可选 —— 未提供时 SelectionService 回退逐只 list_announced。
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
class SyncLogRepository(Protocol):
|
class SyncLogRepository(Protocol):
|
||||||
def add(self, log: SyncLog) -> SyncLog: ...
|
def add(self, log: SyncLog) -> SyncLog: ...
|
||||||
|
|||||||
@@ -280,6 +280,28 @@ class SqlAlchemyFinancialRepository:
|
|||||||
rows = self._session.scalars(stmt).all()
|
rows = self._session.scalars(stmt).all()
|
||||||
return [FinancialIndicator.model_validate(r, from_attributes=True) for r in rows]
|
return [FinancialIndicator.model_validate(r, from_attributes=True) for r in rows]
|
||||||
|
|
||||||
|
def list_announced_many(
|
||||||
|
self,
|
||||||
|
symbols: Sequence[str],
|
||||||
|
as_of_date: date,
|
||||||
|
) -> list[FinancialIndicator]:
|
||||||
|
"""批量:这些股票 announce_date <= as_of_date 的全部记录(防未来函数)。"""
|
||||||
|
if not symbols:
|
||||||
|
return []
|
||||||
|
rows = self._session.scalars(
|
||||||
|
select(FinancialIndicatorModel)
|
||||||
|
.where(
|
||||||
|
FinancialIndicatorModel.symbol.in_(list(symbols)),
|
||||||
|
FinancialIndicatorModel.announce_date <= as_of_date,
|
||||||
|
)
|
||||||
|
.order_by(
|
||||||
|
FinancialIndicatorModel.symbol,
|
||||||
|
FinancialIndicatorModel.announce_date,
|
||||||
|
FinancialIndicatorModel.report_date,
|
||||||
|
)
|
||||||
|
).all()
|
||||||
|
return [FinancialIndicator.model_validate(r, from_attributes=True) for r in rows]
|
||||||
|
|
||||||
|
|
||||||
class SqlAlchemySyncLogRepository:
|
class SqlAlchemySyncLogRepository:
|
||||||
def __init__(self, session: Session) -> None:
|
def __init__(self, session: Session) -> None:
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ from datetime import date
|
|||||||
|
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
|
|
||||||
|
from app.domain.entities.market import FinancialIndicator
|
||||||
from app.domain.entities.selection import (
|
from app.domain.entities.selection import (
|
||||||
SelectionCandidate,
|
SelectionCandidate,
|
||||||
SelectionQuery,
|
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
|
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:
|
def _to_float(v) -> float | None:
|
||||||
if v is None:
|
if v is None:
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -0,0 +1,228 @@
|
|||||||
|
"""M6.2 条件选股测试:结构化条件(static.*/tech 字段与因子/fundamental.*)与未来函数防护。
|
||||||
|
|
||||||
|
Fake(内存)Repository;财务行显式携带 announce_date 验证 as_of 可见性。
|
||||||
|
"""
|
||||||
|
|
||||||
|
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 FinancialIndicator, Stock
|
||||||
|
from app.domain.entities.research import UniverseSpec
|
||||||
|
from app.domain.entities.selection import SelectionQuery
|
||||||
|
|
||||||
|
from conftest_quant import bars_dataframe_to_daily_bars, synthetic_daily
|
||||||
|
|
||||||
|
_SYMS = ["600000.SH", "600001.SH", "600002.SH"]
|
||||||
|
|
||||||
|
|
||||||
|
def _stocks(industries: dict[str, str] | None = None) -> list[Stock]:
|
||||||
|
industries = industries or {}
|
||||||
|
return [
|
||||||
|
Stock(
|
||||||
|
symbol=s, name=f"股{i}", industry=industries.get(s),
|
||||||
|
list_date=date(1999, 1, 1),
|
||||||
|
)
|
||||||
|
for i, s in enumerate(_SYMS)
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
class _MemDailyRepo:
|
||||||
|
def __init__(self, df: pd.DataFrame) -> None:
|
||||||
|
self._bars_all = bars_dataframe_to_daily_bars(df)
|
||||||
|
|
||||||
|
def get_range(self, symbol, start, end):
|
||||||
|
return [b for b in self._bars_all if b.symbol == symbol and start <= b.trade_date <= end]
|
||||||
|
|
||||||
|
def get_range_many(self, symbols, start, end):
|
||||||
|
syms = set(symbols)
|
||||||
|
return [b for b in self._bars_all if b.symbol in syms and start <= b.trade_date <= end]
|
||||||
|
|
||||||
|
def latest_date(self, symbol):
|
||||||
|
rows = [b.trade_date for b in self._bars_all if b.symbol == symbol]
|
||||||
|
return max(rows) if rows else None
|
||||||
|
|
||||||
|
|
||||||
|
class _MemFinancialRepo:
|
||||||
|
def __init__(self, rows: list[FinancialIndicator]) -> None:
|
||||||
|
self._rows = rows
|
||||||
|
|
||||||
|
def list_announced(self, symbol, as_of_date, report_start=None):
|
||||||
|
return [
|
||||||
|
r for r in self._rows
|
||||||
|
if r.symbol == symbol and r.announce_date <= as_of_date
|
||||||
|
and (report_start is None or r.report_date >= report_start)
|
||||||
|
]
|
||||||
|
|
||||||
|
def list_announced_many(self, symbols, as_of_date):
|
||||||
|
syms = set(symbols)
|
||||||
|
return [
|
||||||
|
r for r in self._rows
|
||||||
|
if r.symbol in syms and r.announce_date <= as_of_date
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def _svc(df, stocks=None, fin_rows=None) -> SelectionService:
|
||||||
|
return SelectionService(
|
||||||
|
_MemStockRepo(stocks or _stocks()),
|
||||||
|
_MemDailyRepo(df),
|
||||||
|
_MemFinancialRepo(fin_rows or []),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class _MemStockRepo:
|
||||||
|
def __init__(self, stocks) -> None:
|
||||||
|
self._stocks = stocks
|
||||||
|
|
||||||
|
def list(self):
|
||||||
|
return self._stocks
|
||||||
|
|
||||||
|
def get_by_symbol(self, symbol):
|
||||||
|
return next((s for s in self._stocks if s.symbol == symbol), None)
|
||||||
|
|
||||||
|
|
||||||
|
def _cond_query(conditions, **kw) -> SelectionQuery:
|
||||||
|
base = dict(method="condition", conditions=conditions, as_of=date(2024, 12, 31))
|
||||||
|
base.update(kw)
|
||||||
|
return SelectionQuery(**base)
|
||||||
|
|
||||||
|
|
||||||
|
class TestStaticCondition:
|
||||||
|
def test_industry_filter(self) -> None:
|
||||||
|
df = synthetic_daily({s: 0.002 for s in _SYMS}, n=320)
|
||||||
|
industries = {_SYMS[0]: "白酒", _SYMS[1]: "银行", _SYMS[2]: "白酒"}
|
||||||
|
svc = _svc(df, stocks=_stocks(industries))
|
||||||
|
res = svc.select(
|
||||||
|
_cond_query(
|
||||||
|
[{"field": "static.industry", "op": "in", "value": ["白酒"]}],
|
||||||
|
universe=UniverseSpec(min_listing_days=0),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
got = {c.symbol for c in res.candidates}
|
||||||
|
assert got == {_SYMS[0], _SYMS[2]}
|
||||||
|
assert res.statistics.selected == 2
|
||||||
|
# 每候选带条件状态与理由
|
||||||
|
assert all(c.filter_status for c in res.candidates)
|
||||||
|
assert all(c.selection_reason for c in res.candidates)
|
||||||
|
|
||||||
|
def test_static_ne(self) -> None:
|
||||||
|
df = synthetic_daily({s: 0.002 for s in _SYMS}, n=320)
|
||||||
|
industries = {_SYMS[0]: "白酒", _SYMS[1]: "银行", _SYMS[2]: "白酒"}
|
||||||
|
svc = _svc(df, stocks=_stocks(industries))
|
||||||
|
res = svc.select(
|
||||||
|
_cond_query(
|
||||||
|
[{"field": "static.industry", "op": "ne", "value": "白酒"}],
|
||||||
|
universe=UniverseSpec(min_listing_days=0),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
assert {c.symbol for c in res.candidates} == {_SYMS[1]}
|
||||||
|
|
||||||
|
|
||||||
|
class TestTechCondition:
|
||||||
|
def test_momentum_gt_zero(self) -> None:
|
||||||
|
df = synthetic_daily({_SYMS[0]: 0.008, _SYMS[1]: -0.008, _SYMS[2]: 0.002}, n=320)
|
||||||
|
svc = _svc(df)
|
||||||
|
res = svc.select(
|
||||||
|
_cond_query(
|
||||||
|
[{"field": "momentum_60", "op": "gt", "value": 0}],
|
||||||
|
universe=UniverseSpec(min_listing_days=0),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
got = {c.symbol for c in res.candidates}
|
||||||
|
assert _SYMS[1] not in got # 下跌股 60 日动量为负
|
||||||
|
assert _SYMS[0] in got and _SYMS[2] in got
|
||||||
|
|
||||||
|
def test_close_above_ma60_ref(self) -> None:
|
||||||
|
df = synthetic_daily({_SYMS[0]: 0.006, _SYMS[1]: -0.006, _SYMS[2]: 0.0005}, n=320)
|
||||||
|
svc = _svc(df)
|
||||||
|
res = svc.select(
|
||||||
|
_cond_query(
|
||||||
|
[{"field": "close", "op": "gt", "ref": "ma60"}],
|
||||||
|
universe=UniverseSpec(min_listing_days=0),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
got = {c.symbol for c in res.candidates}
|
||||||
|
assert _SYMS[1] not in got # 下跌股收盘在 MA60 之下
|
||||||
|
assert _SYMS[0] in got
|
||||||
|
|
||||||
|
def test_lte_threshold(self) -> None:
|
||||||
|
df = synthetic_daily({_SYMS[0]: 0.01, _SYMS[1]: -0.01, _SYMS[2]: 0.0001}, n=320)
|
||||||
|
svc = _svc(df)
|
||||||
|
res = svc.select(
|
||||||
|
_cond_query(
|
||||||
|
[{"field": "momentum_60", "op": "lte", "value": 0}],
|
||||||
|
universe=UniverseSpec(min_listing_days=0),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
got = {c.symbol for c in res.candidates}
|
||||||
|
assert _SYMS[1] in got
|
||||||
|
assert _SYMS[0] not in got
|
||||||
|
|
||||||
|
|
||||||
|
class TestFundamentalCondition:
|
||||||
|
def _fin_rows(self) -> list[FinancialIndicator]:
|
||||||
|
return [
|
||||||
|
FinancialIndicator(
|
||||||
|
symbol=_SYMS[0], report_date=date(2024, 9, 30), announce_date=date(2024, 10, 25),
|
||||||
|
eps=Decimal("3.5"), roe=Decimal("20.0"), source="tushare",
|
||||||
|
),
|
||||||
|
# B:只在 as_of 之后才公告(未来数据)→ as_of 时不可见
|
||||||
|
FinancialIndicator(
|
||||||
|
symbol=_SYMS[1], report_date=date(2024, 9, 30), announce_date=date(2025, 3, 30),
|
||||||
|
roe=Decimal("99.0"), source="tushare",
|
||||||
|
),
|
||||||
|
# C:roe 低于阈值
|
||||||
|
FinancialIndicator(
|
||||||
|
symbol=_SYMS[2], report_date=date(2024, 6, 30), announce_date=date(2024, 8, 20),
|
||||||
|
roe=Decimal("5.0"), source="tushare",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
def test_roe_filter_no_future_leak(self) -> None:
|
||||||
|
df = synthetic_daily({s: 0.002 for s in _SYMS}, n=320)
|
||||||
|
svc = _svc(df, fin_rows=self._fin_rows())
|
||||||
|
res = svc.select(
|
||||||
|
_cond_query(
|
||||||
|
[{"field": "fundamental.roe", "op": "gte", "value": 15}],
|
||||||
|
universe=UniverseSpec(min_listing_days=0),
|
||||||
|
),
|
||||||
|
# 上面 helper 已带 as_of
|
||||||
|
)
|
||||||
|
got = {c.symbol for c in res.candidates}
|
||||||
|
# A 可见且 roe=20 → 入选;B 公告在未来(防未来函数)→ 不入选;C roe=5 → 不入选
|
||||||
|
assert got == {_SYMS[0]}
|
||||||
|
|
||||||
|
def test_missing_financial_repo_raises(self) -> None:
|
||||||
|
df = synthetic_daily({s: 0.002 for s in _SYMS}, n=320)
|
||||||
|
svc = SelectionService(_MemStockRepo(_stocks()), _MemDailyRepo(df)) # 无财务 repo
|
||||||
|
with pytest.raises(ValueError):
|
||||||
|
svc.select(
|
||||||
|
_cond_query(
|
||||||
|
[{"field": "fundamental.roe", "op": "gte", "value": 15}],
|
||||||
|
universe=UniverseSpec(min_listing_days=0),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_as_of_earlier_excludes_announced_after(self) -> None:
|
||||||
|
"""更早 as_of:A 的 roe=20 若在 as_of 之后才公告也不可见。"""
|
||||||
|
df = synthetic_daily({s: 0.002 for s in _SYMS}, n=320)
|
||||||
|
rows = [
|
||||||
|
FinancialIndicator(
|
||||||
|
symbol=_SYMS[0], report_date=date(2024, 6, 30),
|
||||||
|
announce_date=date(2024, 10, 1), roe=Decimal("99.0"), source="tushare",
|
||||||
|
)
|
||||||
|
]
|
||||||
|
svc = _svc(df, fin_rows=rows)
|
||||||
|
res = svc.select(
|
||||||
|
SelectionQuery(
|
||||||
|
method="condition",
|
||||||
|
conditions=[{"field": "fundamental.roe", "op": "gte", "value": 50}],
|
||||||
|
as_of=date(2024, 9, 1), # announce(10-01) 尚未来
|
||||||
|
universe=UniverseSpec(min_listing_days=0),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
assert {c.symbol for c in res.candidates} == set() # 无人可见 roe → 全部不通过
|
||||||
Reference in New Issue
Block a user