87 lines
3.4 KiB
Python
87 lines
3.4 KiB
Python
"""通路测试策略(ChannelTestStrategy)单元测试。
|
|
|
|
验证轮换目标集 + 卖旧买新调度 + T+1 探针,用 mock broker 记录下单调用。
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
from sanguo_portfolio.strategies import ChannelTestConfig, ChannelTestStrategy
|
|
from sanguo_portfolio.strategies.all_weather import BrokerFacade
|
|
|
|
|
|
class _MockBroker(BrokerFacade):
|
|
def __init__(self) -> None:
|
|
# 注意:BrokerFacade 是 dataclass,父类 __init__ 会用字段默认值覆盖同名
|
|
# 实例属性,子类方法重写会被遮蔽 → 必须在 super().__init__() 之后
|
|
# 用实例属性注入记录函数。
|
|
self.calls: list[tuple[str, str, float]] = [] # (method, code, value)
|
|
super().__init__()
|
|
self.order_target_value = self._record_otv
|
|
|
|
def _record_otv(self, code: str, value: float):
|
|
self.calls.append(("otv", code, value))
|
|
return None
|
|
|
|
|
|
class _Pos:
|
|
def __init__(self, value: float) -> None:
|
|
self.value = value
|
|
|
|
|
|
class _Ctx:
|
|
def __init__(self, positions: dict, cash: float) -> None:
|
|
self.portfolio = type("P", (), {"positions": positions,
|
|
"available_cash": cash,
|
|
"total_value": cash + sum(p.value for p in positions.values())})()
|
|
|
|
|
|
def test_target_set_rotates_with_day():
|
|
s = ChannelTestStrategy(provider=None, config=ChannelTestConfig(
|
|
universe=["A", "B", "C", "D"], hold_n=2, period=1))
|
|
s._day = 0
|
|
assert set(s._target_set()) <= {"A", "B", "C", "D"}
|
|
s._day = 1
|
|
t1 = s._target_set()
|
|
s._day = 2
|
|
t2 = s._target_set()
|
|
assert len(t1) == 2 and len(t2) == 2
|
|
assert t1 != t2 # 不同周期目标偏移
|
|
|
|
|
|
def test_rotate_sells_non_target_and_buys_target():
|
|
broker = _MockBroker()
|
|
s = ChannelTestStrategy(provider=None, broker=broker, config=ChannelTestConfig(
|
|
universe=["A", "B", "C", "D"], hold_n=2, period=1, probe_t1=False))
|
|
# rotate#1 → day=1 → offset=1 → target=[B,C];当前持仓 C,D
|
|
# → 卖出 D(C 在目标内保留),对 B,C 调仓(买入)
|
|
ctx = _Ctx({"C": _Pos(1000), "D": _Pos(1000)}, cash=8000)
|
|
s.rotate(ctx)
|
|
zero_calls = {c for m, c, v in broker.calls if m == "otv" and v == 0}
|
|
assert zero_calls == {"D"} # 只卖非目标的 D
|
|
buys = {c for m, c, v in broker.calls if m == "otv" and v > 0}
|
|
assert buys == {"B", "C"} # 买新 B + 调仓 C
|
|
|
|
|
|
def test_rotate_t1_probe_fires():
|
|
broker = _MockBroker()
|
|
s = ChannelTestStrategy(provider=None, broker=broker, config=ChannelTestConfig(
|
|
universe=["A", "B"], hold_n=1, period=1, probe_t1=True))
|
|
ctx = _Ctx({}, cash=10000)
|
|
s.rotate(ctx)
|
|
# rotate#1 → target=[B];探针对 target[0]=B 再次 order_target_value(0)(当日卖,预期被 T+1 拒)
|
|
otv_calls = [c for m, c, v in broker.calls if m == "otv"]
|
|
assert otv_calls.count("B") >= 2 # 一次买入调仓 + 一次 T+1 探针
|
|
|
|
|
|
def test_period_skips_off_cycle_days():
|
|
broker = _MockBroker()
|
|
s = ChannelTestStrategy(provider=None, broker=broker, config=ChannelTestConfig(
|
|
universe=["A", "B"], hold_n=1, period=3, probe_t1=False))
|
|
ctx = _Ctx({}, cash=10000)
|
|
s.rotate(ctx) # day1 → 触发
|
|
n1 = len(broker.calls)
|
|
s.rotate(ctx) # day2 → 跳过
|
|
s.rotate(ctx) # day3 → 跳过
|
|
assert len(broker.calls) == n1
|
|
s.rotate(ctx) # day4 → 触发
|
|
assert len(broker.calls) > n1
|