- 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 通过
229 lines
8.4 KiB
Python
229 lines
8.4 KiB
Python
"""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 → 全部不通过
|