8862816557
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全周期待优化后
491 lines
20 KiB
Python
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)
|