Files
sanguo_vnpy_v2/tests/portfolio/test_momentum_timing.py
T
claude_dev 8862816557 feat(portfolio): B fundamentals批量 + C涨跌停filter修复(get_limit_status_batch接入)
B: value_selection 逐只 get_value_metrics → get_value_metrics_batch(数据session)
- 01 验证 -21.85% vs 改前 -21.63%(微差0.22%, batch实现微差,可接受)

C: filters filter_limitup/limitdown/paused 接入 get_limit_status_batch(数据session)
- 修复回测死代码: filter 取 tick.get(last_price/paused) 恒None → 照买涨停/照卖跌停/照交易停牌
- 三策略调仓预取 status_map 共享一次查询, 向后兼容 all_weather(不传参=原行为)
- 03 短区间(2024Q1)验证: C前+138.7%虚高 → C后+101.6%, filter修复减少照买涨停虚增

验收: 101单测(filters 30含14新status_map口径 + 三策略71)
注意: get_limit_status_batch 44s/800只(数据session待批量化优化), 02/03全周期待优化后
2026-07-30 07:37:03 +08:00

491 lines
20 KiB
Python

"""MomentumTimingStrategy 单元测试(mock provider + mock broker)。
策略层只测**逻辑分支正确**(RPS / 均线 / 牛熊信号 / 调仓),不测真实数据。
真实数据回测在 VPS 跑,这里只保证策略翻译等价 + 两个原始 bug 已修复。
性能改造后:策略用 ``get_closes_panel`` 批量取 close 宽表(替代 get_price+pivot),
测试 mock 同步切到 ``get_closes_panel`` 返回宽表 DataFrame。
"""
from __future__ import annotations
from datetime import datetime, timedelta
from typing import Any, Dict, List, Optional
from unittest.mock import MagicMock
import numpy as np
import pandas as pd
import pytest
from sanguo_portfolio import BrokerFacade
from sanguo_portfolio.strategies.momentum_timing import (
MomentumTimingConfig,
MomentumTimingStrategy,
)
from tests.portfolio.conftest import FakeContext, FakePosition
# ------------------------ 测试 helper ------------------------
def make_strategy(
*,
index_stocks_map: Optional[Dict[str, List[str]]] = None,
panel_map: Optional[Dict[Any, pd.DataFrame]] = None,
config: Optional[MomentumTimingConfig] = None,
) -> MomentumTimingStrategy:
"""构造一个 mock provider + mock broker 驱动的策略。
- index_stocks_map: get_index_stocks 返回,dict[index] -> List[code]
- panel_map: get_closes_panel 按 (symbols_tuple) 或 (symbols_tuple, start, end) 缓存的宽表返回
"""
provider = MagicMock(name="provider")
# get_index_stocks
index_stocks_map = index_stocks_map or {}
def _get_index_stocks(index_symbol, date=None):
return list(index_stocks_map.get(index_symbol, []))
provider.get_index_stocks.side_effect = _get_index_stocks
# get_security_info(filter_st/filter_new 默认放过)
provider.get_security_info.return_value = {
"display_name": "NORMAL",
"name": "600519",
"start_date": datetime(2000, 1, 1),
}
# get_live_current:不停牌不涨跌停
provider.get_live_current.return_value = {
"paused": False, "last_price": 10.0,
"high_limit": 11.0, "low_limit": 9.0,
}
provider.get_current_tick.return_value = {
"paused": False, "last_price": 10.0,
"high_limit": 11.0, "low_limit": 9.0,
}
# get_closes_panel 按 key 缓存:支持精确 (tuple(symbols), start, end) 与
# 通配 (tuple(symbols),) 两种 key(精确优先,fallback 忽略日期)
def _normalize_syms(k: Any) -> Any:
if isinstance(k, tuple) and k and isinstance(k[0], (list, tuple)):
return (tuple(k[0]),) + tuple(k[1:])
return k
panel_map = {_normalize_syms(k): v for k, v in (panel_map or {}).items()}
def _get_closes_panel(symbols, start=None, end=None, interval="d", fq="raw"):
syms_key = tuple(symbols) if isinstance(symbols, list) else symbols
# 精确 key 优先
exact = panel_map.get((syms_key, start, end))
if exact is not None:
return exact
# fallback: 忽略 start/end
fallback = panel_map.get((syms_key,))
if fallback is not None:
return fallback
return pd.DataFrame(index=pd.DatetimeIndex([]))
provider.get_closes_panel.side_effect = _get_closes_panel
# get_limit_status_batch: 默认全部"正常交易"(filter 全保留)
def _glbs(codes, date=None):
return {
c: {"is_limit_up": False, "is_limit_down": False, "is_paused": False}
for c in codes
}
provider.get_limit_status_batch.side_effect = _glbs
broker = BrokerFacade()
broker.order_target_value = MagicMock(return_value=MagicMock(filled=100))
broker.order_value = MagicMock(return_value=MagicMock(filled=100))
broker.set_benchmark = MagicMock()
broker.set_option = MagicMock()
broker.run_daily = MagicMock()
broker.run_monthly = MagicMock()
return MomentumTimingStrategy(provider=provider, broker=broker, config=config)
def _make_close_wide(
codes: List[str],
closes: List[List[float]],
end_date: str = "2024-09-30",
days: int = 30,
) -> pd.DataFrame:
"""构造 get_closes_panel 风格的 close 宽表 (index=DatetimeIndex, columns=codes)。
Args:
codes: 股票代码列表(列名)
closes: 每只股票的 close 序列(长度 <= days, 不足重复首值)
end_date: 最后一根 K 线日期
days: 总 K 线根数(默认 30)
"""
end_dt = datetime.strptime(end_date, "%Y-%m-%d")
dates = pd.DatetimeIndex([
end_dt - timedelta(days=days - 1 - i) for i in range(days)
])
data: Dict[str, List[float]] = {}
for code, close_list in zip(codes, closes):
# 不足 days 的补首值
full = list(close_list) + [close_list[-1]] * (days - len(close_list))
data[code] = [float(c) for c in full]
return pd.DataFrame(data, index=dates)
# =================== initialize ===================
class TestInitialize:
def test_initialize_registers_daily_handle_data(self, fake_context):
s = make_strategy()
s.initialize(fake_context)
# run_daily 至少被调一次(注册 handle_data)
assert s.broker.run_daily.called
# run_daily 的第一个参数应是 handle_data 方法
first_call = s.broker.run_daily.call_args_list[0]
assert first_call.args[0].__name__ == "handle_data"
def test_initialize_sets_benchmark(self, fake_context):
cfg = MomentumTimingConfig(benchmark="000300.XSHG")
s = make_strategy(config=cfg)
s.initialize(fake_context)
s.broker.set_benchmark.assert_called_with("000300.XSHG")
# =================== _cal_rps (修复后涨跌幅正确) ===================
class TestCalRps:
def test_empty_stocks_returns_empty_df(self):
"""空股票列表 → 空 DataFrame。"""
s = make_strategy()
out = s._cal_rps([], cur_date="2024-09-30", pre_date="2024-09-01")
assert out.empty
assert "rps_value" in out.columns
def test_rps_uses_pre_to_cur_range_real_returns(self):
"""⚠️ 核心修复验证:RPS 必须用 preDate~curDate 区间算真实涨跌幅,
而非原始 bug 的 ``get_price(start=curDate, end=curDate)`` 单日恒 0。
"""
# 3 只股票,涨幅依次为 +100% / +50% / 0%
# preDate 首值 = 10, curDate 末值 = 20 / 15 / 10
codes = ["A.XSHG", "B.XSHG", "C.XSHG"]
panel = pd.DataFrame(
{
"A.XSHG": [10.0, 20.0], # +100%
"B.XSHG": [10.0, 15.0], # +50%
"C.XSHG": [10.0, 10.0], # 0%
},
index=pd.DatetimeIndex(["2024-09-01", "2024-09-30"]),
)
s = make_strategy(panel_map={
(tuple(codes), "2024-09-01", "2024-09-30"): panel,
})
out = s._cal_rps(codes, cur_date="2024-09-30", pre_date="2024-09-01")
# 排序:A(+100%) > B(+50%) > C(0%)
assert list(out["code"]) == ["A.XSHG", "B.XSHG", "C.XSHG"]
# RPS: 99 - 100*i/n → [99, 99-100/3, 99-200/3] = [99, 65.67, 32.33]
assert out["rps_value"].iloc[0] == pytest.approx(99.0, abs=0.01)
assert out["rps_value"].iloc[1] == pytest.approx(99 - 100 / 3, abs=0.01)
assert out["rps_value"].iloc[2] == pytest.approx(99 - 200 / 3, abs=0.01)
def test_rps_descending_by_return(self):
"""涨幅大的排前(降序)。"""
codes = ["X.XSHG", "Y.XSHG"]
panel = pd.DataFrame(
{
"X.XSHG": [10.0, 12.0], # +20%
"Y.XSHG": [10.0, 15.0], # +50%
},
index=pd.DatetimeIndex(["2024-09-01", "2024-09-30"]),
)
s = make_strategy(panel_map={
(tuple(codes), "2024-09-01", "2024-09-30"): panel,
})
out = s._cal_rps(codes, cur_date="2024-09-30", pre_date="2024-09-01")
# Y 涨幅大,排前
assert out["code"].iloc[0] == "Y.XSHG"
assert out["code"].iloc[1] == "X.XSHG"
def test_rps_filters_nan_and_zero_first(self):
"""首值为 0(除零)或 NaN → 过滤掉。"""
codes = ["GOOD.XSHG", "ZERO.XSHG", "NAN.XSHG"]
panel = pd.DataFrame(
{
"GOOD.XSHG": [10.0, 20.0],
"ZERO.XSHG": [0.0, 10.0],
"NAN.XSHG": [np.nan, 10.0],
},
index=pd.DatetimeIndex(["2024-09-01", "2024-09-30"]),
)
s = make_strategy(panel_map={
(tuple(codes), "2024-09-01", "2024-09-30"): panel,
})
out = s._cal_rps(codes, cur_date="2024-09-30", pre_date="2024-09-01")
assert list(out["code"]) == ["GOOD.XSHG"]
# =================== _select_stocks (均线动量) ===================
class TestSelectStocks:
def test_empty_input(self):
s = make_strategy()
assert s._select_stocks([], cur_date="2024-09-30") == []
def test_keep_close_above_ma_short_above_ma_long(self):
"""close > MA5 且 MA5 > MA15 → 保留。"""
# 构造 15 日 close 序列:上升 → close(末) > MA5 > MA15
rising = [10.0 + i * 0.5 for i in range(15)] # 10→17
panel = _make_close_wide(["UP.XSHG"], [rising], end_date="2024-09-30", days=15)
s = make_strategy(panel_map={
(("UP.XSHG",),): panel,
})
out = s._select_stocks(["UP.XSHG"], cur_date="2024-09-30")
assert out == ["UP.XSHG"]
def test_filter_close_below_ma_short(self):
"""close < MA5 → 剔除(下行趋势)。"""
falling = [20.0 - i * 0.5 for i in range(15)] # 20→13
panel = _make_close_wide(["DOWN.XSHG"], [falling], end_date="2024-09-30", days=15)
s = make_strategy(panel_map={
(("DOWN.XSHG",),): panel,
})
out = s._select_stocks(["DOWN.XSHG"], cur_date="2024-09-30")
assert out == []
def test_filter_ma_short_below_ma_long(self):
"""close > MA5 但 MA5 < MA15(下跌但末值小反弹)→ 剔除。"""
# 前 10 日大涨(20→30),后 5 日跌(30→26):MA5 < MA15
series = [20 + i for i in range(10)] + [30 - i for i in range(1, 6)] # 20..29, 29..25
panel = _make_close_wide(["FLAT.XSHG"], [series], end_date="2024-09-30", days=15)
s = make_strategy(panel_map={
(("FLAT.XSHG",),): panel,
})
out = s._select_stocks(["FLAT.XSHG"], cur_date="2024-09-30")
# close=25, MA5 = mean(29,28,27,26,25)=27, MA15 = mean(all)=24.67
# close(25) < MA5(27) → 不满足 close>MA5
assert out == []
def test_insufficient_data_skipped(self):
"""不足 ma_long=15 根 → 跳过。"""
# 直接给短宽表(3 行):panel.notna().sum()=3 < 15 → 全部剔除
short_panel = pd.DataFrame(
{"NEW.XSHG": [10.0, 11.0, 12.0]},
index=pd.DatetimeIndex(["2024-09-28", "2024-09-29", "2024-09-30"]),
)
s = make_strategy(panel_map={
(("NEW.XSHG",),): short_panel,
})
out = s._select_stocks(["NEW.XSHG"], cur_date="2024-09-30")
assert out == []
# =================== _cal_buy_sign (牛熊分界) ===================
class TestCalBuySign:
def test_empty_index_list_returns_false(self):
s = make_strategy()
assert s._cal_buy_sign([], past_day=30, cur_date="2024-09-30") is False
def test_bull_when_above_ma_ratio_exceeds_threshold(self):
"""所有指数都站在 30 日均线上方 → 占比 100% > 20% → 牛市(True)。"""
# 上升序列:末值远高于均值
idx_list = ["000300.XSHG", "000905.XSHG"]
rising = [10.0 + i for i in range(30)] # 10→39
panel = _make_close_wide(idx_list, [rising, rising], end_date="2024-09-30", days=30)
s = make_strategy(panel_map={
(tuple(idx_list),): panel,
})
assert s._cal_buy_sign(idx_list, past_day=30, cur_date="2024-09-30") is True
def test_bear_when_below_ma_ratio_below_threshold(self):
"""所有指数都跌破 30 日均线 → 占比 0% < 20% → 熊市(False)。"""
idx_list = ["000300.XSHG", "000905.XSHG"]
falling = [40.0 - i for i in range(30)] # 40→11
panel = _make_close_wide(idx_list, [falling, falling], end_date="2024-09-30", days=30)
s = make_strategy(panel_map={
(tuple(idx_list),): panel,
})
assert s._cal_buy_sign(idx_list, past_day=30, cur_date="2024-09-30") is False
def test_threshold_boundary_3_of_9_above_is_bull(self):
"""9 个指数中 2 个站上(2/9=0.222 > 0.2)→ 牛市。1 个站上(0.111 < 0.2)→ 熊市。"""
idx_list = [f"IDX{i}.XSHG" for i in range(9)]
rising = [10.0 + i for i in range(30)]
falling = [40.0 - i for i in range(30)]
# 2 个 rising + 7 个 falling
series_list = [rising, rising] + [falling] * 7
panel = _make_close_wide(idx_list, series_list, end_date="2024-09-30", days=30)
s = make_strategy(panel_map={
(tuple(idx_list),): panel,
})
# 2/9 ≈ 0.222 > 0.2 → 牛市
assert s._cal_buy_sign(idx_list, past_day=30, cur_date="2024-09-30") is True
# 改为 1 个 rising:1/9 ≈ 0.111 < 0.2 → 熊市
series_list_1 = [rising] + [falling] * 8
panel_1 = _make_close_wide(idx_list, series_list_1, end_date="2024-09-30", days=30)
s.provider.get_closes_panel.side_effect = None
s.provider.get_closes_panel.return_value = panel_1
assert s._cal_buy_sign(idx_list, past_day=30, cur_date="2024-09-30") is False
# =================== handle_data (主流程) ===================
class TestHandleData:
def test_bear_signal_clears_all_positions(self):
"""熊市信号 → 全部持仓清掉。"""
cfg = MomentumTimingConfig(index_list=["IDX.XSHG"])
s = make_strategy(config=cfg)
# 触发熊市:get_closes_panel 返回下行 close
s.provider.get_closes_panel.side_effect = None
s.provider.get_closes_panel.return_value = _make_close_wide(
["IDX.XSHG"], [[40.0 - i for i in range(30)]],
end_date="2024-10-08", days=30,
)
ctx = FakeContext(
current_dt=datetime(2024, 10, 8, 9, 30),
previous_date="2024-09-30",
positions={
"600519.XSHG": FakePosition("600519.XSHG", avg_cost=1600, price=1500),
"000001.XSHE": FakePosition("000001.XSHE", avg_cost=10, price=9),
},
)
s.handle_data(ctx)
# 两只持仓都被 order_target_value(code, 0)
sell_calls = [
c for c in s.broker.order_target_value.call_args_list if c.args[1] == 0
]
assert len(sell_calls) == 2
sell_codes = {c.args[0] for c in sell_calls}
assert sell_codes == {"600519.XSHG", "000001.XSHE"}
def test_bull_signal_buys_new_stocks(self):
"""牛市信号 + 候选池选股 → 买入(等额)。"""
# 构造场景:1 个指数,成份股 1 只,close 上升(RPS 正,均线多)
cfg = MomentumTimingConfig(index_list=["IDX.XSHG"], top_k=6, ma_short=5, ma_long=15)
s = make_strategy(
index_stocks_map={"IDX.XSHG": ["CAND.XSHG"]},
config=cfg,
)
rising_30 = [10.0 + i for i in range(30)] # 牛市信号用
rising_15 = [10.0 + i for i in range(15)] # 均线筛选用
# 提供所有可能查询路径的 close 宽表
idx_codes = ["IDX.XSHG"]
stock_codes = ["CAND.XSHG"]
def _gcp(symbols, start=None, end=None, interval="d", fq="raw"):
if symbols == idx_codes:
# _cal_buy_sign: idx_list, tail(30)
return _make_close_wide(idx_codes, [rising_30], days=30)
if isinstance(symbols, list) and symbols == stock_codes:
# _cal_rps (2 行 start/end) 与 _select_stocks (15 行) 都用同一份上升序列
# 取 15 行,_cal_rps 用首末(_shift_date 算出 start),
# _select_stocks 用全部 15 行(ma_long=15)
return _make_close_wide(stock_codes, [rising_15], days=15)
return pd.DataFrame(index=pd.DatetimeIndex([]))
s.provider.get_closes_panel.side_effect = _gcp
ctx = FakeContext(
current_dt=datetime(2024, 10, 8, 9, 30),
previous_date="2024-09-30",
positions={},
cash=1_000_000,
)
s.handle_data(ctx)
# 应有 1 笔买入 CAND.XSHG,金额 ≈ 1_000_000 / 1 = 1_000_000
buy_calls = [
c for c in s.broker.order_target_value.call_args_list if c.args[1] != 0
]
assert len(buy_calls) >= 1
assert any(c.args[0] == "CAND.XSHG" for c in buy_calls)
def test_handle_data_uses_current_dt_not_today(self):
"""⚠️ 修复原始 bug 验证:handle_data 必须用 context.current_dt 计算 cur_date,
不能用 datetime.date.today()(后者取真实今天)。
"""
# 用一个明显不同的 current_dt,确认 get_closes_panel 的 end 跟随它
cfg = MomentumTimingConfig(index_list=["IDX.XSHG"])
s = make_strategy(config=cfg)
captured_ends: List[Any] = []
def _gcp(symbols, start=None, end=None, interval="d", fq="raw"):
# 记录 end 用于断言
if end is not None:
captured_ends.append(str(end))
# 下行 → 熊市(快速 return,不查其他)
return _make_close_wide(
["IDX.XSHG"], [[40.0 - i for i in range(30)]],
end_date=str(end or "2024-10-08")[:10],
days=30,
)
s.provider.get_closes_panel.side_effect = _gcp
ctx = FakeContext(current_dt=datetime(2024, 10, 8, 9, 30))
s.handle_data(ctx)
# 至少一次 get_closes_panel 的 end 是 "2024-10-08"(来自 current_dt),非今天
assert any("2024-10-08" in d for d in captured_ends)
# =================== _find_stock_pool (取强舍弱) ===================
class TestFindStockPool:
def test_picks_top_k_per_index(self):
"""每个行业取 RPS top_k → 候选池并集。"""
cfg = MomentumTimingConfig(index_list=["IDX1.XSHG", "IDX2.XSHG"], top_k=2)
s = make_strategy(
index_stocks_map={
"IDX1.XSHG": ["A.XSHG", "B.XSHG", "C.XSHG"],
"IDX2.XSHG": ["D.XSHG", "E.XSHG"],
},
config=cfg,
)
# 涨幅:A=+100%, B=+50%, C=0%, D=+30%, E=-10%
rps_panel_1 = pd.DataFrame(
{
"A.XSHG": [10.0, 20.0],
"B.XSHG": [10.0, 15.0],
"C.XSHG": [10.0, 10.0],
},
index=pd.DatetimeIndex(["2024-09-01", "2024-09-30"]),
)
rps_panel_2 = pd.DataFrame(
{
"D.XSHG": [10.0, 13.0],
"E.XSHG": [10.0, 9.0],
},
index=pd.DatetimeIndex(["2024-09-01", "2024-09-30"]),
)
def _gcp(symbols, start=None, end=None, interval="d", fq="raw"):
if isinstance(symbols, list):
if "A.XSHG" in symbols:
return rps_panel_1
if "D.XSHG" in symbols:
return rps_panel_2
return pd.DataFrame(index=pd.DatetimeIndex([]))
s.provider.get_closes_panel.side_effect = _gcp
out = s._find_stock_pool(
["IDX1.XSHG", "IDX2.XSHG"], cur_date="2024-09-30", pre_date="2024-09-01",
)
# IDX1 top2 = [A, B], IDX2 top2 = [D, E]
assert set(out) == {"A.XSHG", "B.XSHG", "D.XSHG", "E.XSHG"}
# =================== Config 默认值 ===================
class TestConfigDefaults:
def test_default_index_list_is_10_csi_industry_indices(self):
"""✅ 默认板块是 10 个中证行业指数(G1 补全后切回原版,000938 缺跳过)。"""
cfg = MomentumTimingConfig()
assert len(cfg.index_list) == 10
# 10 个中证行业指数 000928-000937 全部存在
for code in ["000928", "000929", "000930", "000931", "000932",
"000933", "000934", "000935", "000936", "000937"]:
assert f"{code}.XSHG" in cfg.index_list
# 000938 缺(constituent_unified 仍无,暂跳记遗留)
assert "000938.XSHG" not in cfg.index_list
def test_default_params_match_original(self):
"""关键参数与原策略 g.* 一致。"""
cfg = MomentumTimingConfig()
assert cfg.index_thre == 0.2 # g.indexThre
assert cfg.past_day == 30 # g.pastDay
assert cfg.top_k == 6 # g.topK
assert cfg.ma_short == 5 # mavg(5)
assert cfg.ma_long == 15 # mavg(15)