Files
qlib/backend/tests/test_selection_condition.py
Simon ef09d5b419 feat(quant): M7.3 研究行情口径显式化(默认不复权 none,可切 qfq)
- DailyBarRepository.get_range_many / stream_range_many_columns 增加 adjust 参数
  (默认 'none')→ SQL 层过滤口径,消除 stock_daily 混 source/adjust 污染因子的风险
- ResearchSpec / SelectionQuery 增加 price_adjustment(none|qfq),随 config_snapshot
  落库可溯源;ResearchService._load_daily 与 SelectionService 装配按口径取数
- tests/test_price_adjustment.py:repo 读取按 adjust 过滤(none/qfq 各自命中)、
  spec 默认与字段记录;全量 pytest 通过
2026-09-09 00:32:55 +08:00

229 lines
8.4 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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 → 全部不通过