"""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, adjust="none"): 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 → 全部不通过