de04a8904b
三策略(聚宽py2→BulletTrade 0.9.2,BrokerFacade注入跨版本兼容): - momentum_timing 动量择时(牛熊分界+行业RPS+均线,切回10中证行业指数) - value_selection 价值精选(6条基本面过滤,切回沪深300) - small_cap 小市值(去IC对冲,切回000985中证全指) 框架: - runner_backtest 加 --strategy 分发(原硬编码all_weather) - provider 加 get_value_metrics(价值精选6条基本面,NOTICE_DATE治前视偏差) - 72单测全过(21+27+24) 修8个回测实测发现的真bug: - 01第⑥条EPS绝对值0.08~0.5与①大盘矛盾→6条交集恒空致全程空仓,按注释本意改净利润同比8~50% - 03原帖calRPS取数区间错(get_price start=end只取1天)→涨跌幅恒0 RPS失效;date.today()取真实今天非回测日 - 02 universe 000985不在constituent_unified→候选池空 VPS实测(短区间验证逻辑,非长期表现): 01价值+23%/03行业轮动+48%/02选出20只小盘 数据缺口(详见docs/research/joinquant_strategies/SUMMARY.md + data_gaps_fix_plan.md): - 三表"1/3损坏"误报已撤回(全扫5530文件/表0损坏,沪深95%+健康,仅北交所920xxx空,不做北交所) - 真实缺口: 行业成份股(G1已补)/000985(G2已补)/IC期货(02对冲去掉)/provider批量接口(G5待做,解锁长回测)
496 lines
22 KiB
Python
496 lines
22 KiB
Python
"""MomentumTimingStrategy 单元测试(mock provider + mock broker)。
|
|
|
|
策略层只测**逻辑分支正确**(RPS / 均线 / 牛熊信号 / 调仓),不测真实数据。
|
|
真实数据回测在 VPS 跑,这里只保证策略翻译等价 + 两个原始 bug 已修复。
|
|
"""
|
|
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,
|
|
price_df_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]
|
|
- price_df_map: get_price 按 (security, fields, count) 或 (security, 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_price 按 key 缓存(支持 count 模式 + start/end 模式)
|
|
# 规范化:把 key 第一项(list)转 tuple 以保证可 hash
|
|
def _normalize_key(k: Any) -> Any:
|
|
if isinstance(k, tuple) and k and isinstance(k[0], (list, tuple)):
|
|
return (tuple(k[0]),) + tuple(k[1:])
|
|
return k
|
|
|
|
price_df_map = {_normalize_key(k): v for k, v in (price_df_map or {}).items()}
|
|
|
|
def _get_price(security, **kwargs):
|
|
# 构造 cache key:两种取数模式
|
|
# 1) count 模式:(sec_key, fields, count)
|
|
# 2) start/end 模式:(sec_key, fields, start_date, end_date)
|
|
# 注意:list 不可 hash → 转 tuple
|
|
sec_key = tuple(security) if isinstance(security, list) else security
|
|
fields = tuple(kwargs.get("fields") or [])
|
|
if kwargs.get("count") is not None:
|
|
key = (sec_key, fields, kwargs.get("count"))
|
|
else:
|
|
key = (sec_key, fields, kwargs.get("start_date"), kwargs.get("end_date"))
|
|
return price_df_map.get(key, pd.DataFrame())
|
|
|
|
provider.get_price.side_effect = _get_price
|
|
|
|
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_panel(
|
|
codes: List[str],
|
|
closes: List[List[float]],
|
|
end_date: str = "2024-09-30",
|
|
days: int = 30,
|
|
) -> pd.DataFrame:
|
|
"""构造 panel=False 风格的 close DataFrame。
|
|
|
|
Args:
|
|
codes: 股票代码列表
|
|
closes: 每只股票的 close 序列(长度 <= days, 不足重复首值)
|
|
end_date: 最后一根 K 线日期
|
|
days: 总 K 线根数(默认 30)
|
|
"""
|
|
end_dt = datetime.strptime(end_date, "%Y-%m-%d")
|
|
dates = [(end_dt - timedelta(days=days - 1 - i)).strftime("%Y-%m-%d") for i in range(days)]
|
|
rows = []
|
|
for code, close_list in zip(codes, closes):
|
|
# 不足 days 的补首值
|
|
full = list(close_list) + [close_list[-1]] * (days - len(close_list))
|
|
for d, c in zip(dates, full):
|
|
rows.append({"time": pd.Timestamp(d), "code": code, "close": float(c)})
|
|
return pd.DataFrame(rows)
|
|
|
|
|
|
# =================== 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"]
|
|
df = pd.DataFrame([
|
|
{"time": pd.Timestamp("2024-09-01"), "code": "A.XSHG", "close": 10.0},
|
|
{"time": pd.Timestamp("2024-09-30"), "code": "A.XSHG", "close": 20.0}, # +100%
|
|
{"time": pd.Timestamp("2024-09-01"), "code": "B.XSHG", "close": 10.0},
|
|
{"time": pd.Timestamp("2024-09-30"), "code": "B.XSHG", "close": 15.0}, # +50%
|
|
{"time": pd.Timestamp("2024-09-01"), "code": "C.XSHG", "close": 10.0},
|
|
{"time": pd.Timestamp("2024-09-30"), "code": "C.XSHG", "close": 10.0}, # 0%
|
|
])
|
|
s = make_strategy(price_df_map={
|
|
# 按 start_date/end_date 取数,确认 _cal_rps 走的是区间查询
|
|
((str(codes),) if False else (tuple(codes), ("close",), "2024-09-01", "2024-09-30")): df,
|
|
})
|
|
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"]
|
|
df = pd.DataFrame([
|
|
{"time": pd.Timestamp("2024-09-01"), "code": "X.XSHG", "close": 10.0},
|
|
{"time": pd.Timestamp("2024-09-30"), "code": "X.XSHG", "close": 12.0}, # +20%
|
|
{"time": pd.Timestamp("2024-09-01"), "code": "Y.XSHG", "close": 10.0},
|
|
{"time": pd.Timestamp("2024-09-30"), "code": "Y.XSHG", "close": 15.0}, # +50%
|
|
])
|
|
s = make_strategy(price_df_map={
|
|
(tuple(codes), ("close",), "2024-09-01", "2024-09-30"): df,
|
|
})
|
|
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"]
|
|
df = pd.DataFrame([
|
|
{"time": pd.Timestamp("2024-09-01"), "code": "GOOD.XSHG", "close": 10.0},
|
|
{"time": pd.Timestamp("2024-09-30"), "code": "GOOD.XSHG", "close": 20.0},
|
|
{"time": pd.Timestamp("2024-09-01"), "code": "ZERO.XSHG", "close": 0.0},
|
|
{"time": pd.Timestamp("2024-09-30"), "code": "ZERO.XSHG", "close": 10.0},
|
|
{"time": pd.Timestamp("2024-09-01"), "code": "NAN.XSHG", "close": np.nan},
|
|
{"time": pd.Timestamp("2024-09-30"), "code": "NAN.XSHG", "close": 10.0},
|
|
])
|
|
s = make_strategy(price_df_map={
|
|
(tuple(codes), ("close",), "2024-09-01", "2024-09-30"): df,
|
|
})
|
|
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
|
|
df = _make_close_panel(["UP.XSHG"], [rising], end_date="2024-09-30", days=15)
|
|
s = make_strategy(price_df_map={
|
|
# 注意:_select_stocks 传 list,get_price 内部转 tuple → key 第一项必须是 tuple
|
|
(("UP.XSHG",), ("close",), 15): df,
|
|
})
|
|
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
|
|
df = _make_close_panel(["DOWN.XSHG"], [falling], end_date="2024-09-30", days=15)
|
|
s = make_strategy(price_df_map={
|
|
(("DOWN.XSHG"), ("close",), 15): df,
|
|
})
|
|
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
|
|
df = _make_close_panel(["FLAT.XSHG"], [series], end_date="2024-09-30", days=15)
|
|
s = make_strategy(price_df_map={
|
|
(("FLAT.XSHG"), ("close",), 15): df,
|
|
})
|
|
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 根 → 跳过。"""
|
|
short_df = _make_close_panel(["NEW.XSHG"], [[10, 11, 12]], end_date="2024-09-30", days=15)
|
|
s = make_strategy(price_df_map={
|
|
(("NEW.XSHG"), ("close",), 15): short_df,
|
|
})
|
|
# 序列被 _make_close_panel 补齐到 15,这里改为真短数据
|
|
s.provider.get_price.side_effect = None
|
|
s.provider.get_price.return_value = pd.DataFrame([
|
|
{"time": pd.Timestamp("2024-09-28"), "code": "NEW.XSHG", "close": 10.0},
|
|
{"time": pd.Timestamp("2024-09-29"), "code": "NEW.XSHG", "close": 11.0},
|
|
{"time": pd.Timestamp("2024-09-30"), "code": "NEW.XSHG", "close": 12.0},
|
|
])
|
|
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
|
|
df = _make_close_panel(idx_list, [rising, rising], end_date="2024-09-30", days=30)
|
|
s = make_strategy(price_df_map={
|
|
(tuple(idx_list), ("close",), 30): df,
|
|
})
|
|
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
|
|
df = _make_close_panel(idx_list, [falling, falling], end_date="2024-09-30", days=30)
|
|
s = make_strategy(price_df_map={
|
|
(tuple(idx_list), ("close",), 30): df,
|
|
})
|
|
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
|
|
df = _make_close_panel(idx_list, series_list, end_date="2024-09-30", days=30)
|
|
s = make_strategy(price_df_map={
|
|
(tuple(idx_list), ("close",), 30): df,
|
|
})
|
|
# 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
|
|
df_1 = _make_close_panel(idx_list, series_list_1, end_date="2024-09-30", days=30)
|
|
s.provider.get_price.side_effect = None
|
|
s.provider.get_price.return_value = df_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_price 返回下行 close
|
|
s.provider.get_price.side_effect = None
|
|
s.provider.get_price.return_value = _make_close_panel(
|
|
["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)] # 均线筛选用
|
|
|
|
# 提供所有可能查询路径的 price 数据
|
|
idx_codes = ["IDX.XSHG"]
|
|
stock_codes = ["CAND.XSHG"]
|
|
|
|
def _gp(security, **kwargs):
|
|
fields = tuple(kwargs.get("fields") or [])
|
|
# 1) _cal_buy_sign: idx_list, count=30
|
|
if security == idx_codes and kwargs.get("count") == 30:
|
|
return _make_close_panel(idx_codes, [rising_30], days=30)
|
|
# 2) _cal_rps for index 股池:股票, start/end 模式
|
|
if (
|
|
isinstance(security, list)
|
|
and security == stock_codes
|
|
and kwargs.get("start_date")
|
|
):
|
|
return pd.DataFrame([
|
|
{"time": pd.Timestamp("2024-09-01"), "code": "CAND.XSHG", "close": 10.0},
|
|
{"time": pd.Timestamp("2024-10-08"), "code": "CAND.XSHG", "close": 39.0},
|
|
])
|
|
# 3) _select_stocks: count=15
|
|
if (
|
|
isinstance(security, list)
|
|
and security == stock_codes
|
|
and kwargs.get("count") == cfg.ma_long
|
|
):
|
|
return _make_close_panel(stock_codes, [rising_15], days=15)
|
|
return pd.DataFrame()
|
|
|
|
s.provider.get_price.side_effect = _gp
|
|
|
|
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_price 的 end_date 跟随它
|
|
cfg = MomentumTimingConfig(index_list=["IDX.XSHG"])
|
|
s = make_strategy(config=cfg)
|
|
captured_end_dates: List[Any] = []
|
|
|
|
def _gp(security, **kwargs):
|
|
# 记录 end_date 用于断言
|
|
if kwargs.get("end_date"):
|
|
captured_end_dates.append(str(kwargs["end_date"]))
|
|
# 下行 → 熊市(快速 return,不查其他)
|
|
return _make_close_panel(
|
|
["IDX.XSHG"], [[40.0 - i for i in range(30)]],
|
|
end_date=str(kwargs.get("end_date", "2024-10-08"))[:10],
|
|
days=30,
|
|
)
|
|
|
|
s.provider.get_price.side_effect = _gp
|
|
ctx = FakeContext(current_dt=datetime(2024, 10, 8, 9, 30))
|
|
s.handle_data(ctx)
|
|
# 至少一次 get_price 的 end_date 是 "2024-10-08"(来自 current_dt),非今天
|
|
assert any("2024-10-08" in d for d in captured_end_dates)
|
|
|
|
|
|
# =================== _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_df_1 = pd.DataFrame([
|
|
{"time": pd.Timestamp("2024-09-01"), "code": "A.XSHG", "close": 10.0},
|
|
{"time": pd.Timestamp("2024-09-30"), "code": "A.XSHG", "close": 20.0},
|
|
{"time": pd.Timestamp("2024-09-01"), "code": "B.XSHG", "close": 10.0},
|
|
{"time": pd.Timestamp("2024-09-30"), "code": "B.XSHG", "close": 15.0},
|
|
{"time": pd.Timestamp("2024-09-01"), "code": "C.XSHG", "close": 10.0},
|
|
{"time": pd.Timestamp("2024-09-30"), "code": "C.XSHG", "close": 10.0},
|
|
])
|
|
rps_df_2 = pd.DataFrame([
|
|
{"time": pd.Timestamp("2024-09-01"), "code": "D.XSHG", "close": 10.0},
|
|
{"time": pd.Timestamp("2024-09-30"), "code": "D.XSHG", "close": 13.0},
|
|
{"time": pd.Timestamp("2024-09-01"), "code": "E.XSHG", "close": 10.0},
|
|
{"time": pd.Timestamp("2024-09-30"), "code": "E.XSHG", "close": 9.0},
|
|
])
|
|
|
|
def _gp(security, **kwargs):
|
|
if isinstance(security, list):
|
|
if "A.XSHG" in security:
|
|
return rps_df_1
|
|
if "D.XSHG" in security:
|
|
return rps_df_2
|
|
return pd.DataFrame()
|
|
|
|
s.provider.get_price.side_effect = _gp
|
|
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)
|