Files
sanguo_vnpy_v2/tests/portfolio/test_factor_topn.py
T

197 lines
8.0 KiB
Python

# tests/portfolio/test_factor_topn.py
"""FactorTopNStrategy 单测(mock provider + mock broker + tmp 截面 parquet).
只测逻辑分支: 截面 asof(严格早于调仓日)/TopN 选取/过滤链(留创业板)/
交易日计数器/等权买入/截面不足跳过. 真实数据回测在 NAS 跑(③ 选股组合回测).
"""
from __future__ import annotations
import os
import sys
from datetime import datetime
from typing import Any, Dict
from unittest.mock import MagicMock
import pandas as pd
import pytest
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..")))
from sanguo_portfolio import BrokerFacade
from sanguo_portfolio.strategies.factor_topn import (
FactorTopNConfig,
FactorTopNStrategy,
_filter_kcbj_keep_chinext,
)
from tests.portfolio.conftest import FakeContext, FakePosition
# ------------------------ helper ------------------------
SYMS = [ # 8 只(vnpy 后缀,截面真实格式): 6 主板/创业板可交易 + 科创 68 + 北交 920(应被剔)
"600000.SSE", "600001.SSE", "000001.SZSE", "000002.SZSE",
"300001.SZSE", "300002.SZSE", "688001.SSE", "920001.SZSE",
]
DATES = ["2024-09-27", "2024-09-30", "2024-10-08"]
@pytest.fixture
def factor_dir(tmp_path) -> str:
"""3 日 × 8 股截面宽表;值=股票序号(越大越"") + 日期偏移."""
frame = pd.DataFrame(
{s: [float(i) + d for d in range(3)] for i, s in enumerate(SYMS)},
index=pd.DatetimeIndex(pd.to_datetime(DATES)),
)
d = tmp_path / "cross_section"
d.mkdir()
frame.to_parquet(d / "composite_test.parquet", index=True)
return str(d)
def _ok_status(codes):
return {c: {"is_paused": False, "is_limit_up": False, "is_limit_down": False}
for c in codes}
def make_strategy(factor_dir: str, top_n: int = 3, **cfg_over) -> tuple[FactorTopNStrategy, MagicMock]:
"""mock provider(get_security_info_batch={}→ST/次新全保留) + 记录型 broker."""
provider = MagicMock(name="provider")
provider.get_security_info_batch = lambda stocks: {}
provider.get_limit_status_batch = lambda stocks, date: _ok_status(stocks)
broker_otv = MagicMock(name="order_target_value")
broker = BrokerFacade(order_target_value=broker_otv)
cfg = FactorTopNConfig(factor_name="composite_test", factor_dir=factor_dir,
top_n=top_n, rebalance_days=5, **cfg_over)
return FactorTopNStrategy(provider=provider, broker=broker, config=cfg), broker_otv
def _ctx(day: str, positions: Dict[str, Any] | None = None) -> FakeContext:
return FakeContext(
current_dt=datetime.fromisoformat(f"{day}T09:30:00"),
previous_date="2024-09-30",
positions=positions,
)
# ------------------------ 宇宙过滤 ------------------------
def test_filter_kcbj_keep_chinext():
"""剔 68/4/8/920,保留主板+创业板 300/301(filter_kcbj_stock 会误杀创业板)."""
got = _filter_kcbj_keep_chinext(list(SYMS))
assert got == ["600000.SSE", "600001.SSE", "000001.SZSE",
"000002.SZSE", "300001.SZSE", "300002.SZSE"]
assert _filter_kcbj_keep_chinext(["430001.BJ", "830001.BJ", "870001.BJ"]) == []
# ------------------------ 截面 asof(防前视核心) ------------------------
def test_factor_scores_strictly_before_rebalance_day(factor_dir):
"""调仓日=截面最后一日(10-08)时取前一日(09-30)列——当日截面不可用(前视);
vt_symbol 后缀同时转聚宽格式(SSE→XSHG/SZSE→XSHE)."""
st, _ = make_strategy(factor_dir)
scores = st._factor_scores("2024-10-08")
assert scores is not None
# 09-30 列: 值 = 序号 + 1 → 920001 最高 8.0;索引已转聚宽后缀
assert scores["920001.XSHE"] == 8.0
assert scores["600000.XSHG"] == 1.0
assert ".SSE" not in "".join(scores.index) and ".SZSE" not in "".join(scores.index)
def test_factor_scores_exact_boundary(factor_dir):
"""调仓日恰在截面日序列中间:取严格小于的最近截面日(09-30 → 取 09-27 列)."""
st, _ = make_strategy(factor_dir)
scores = st._factor_scores("2024-09-30")
assert scores is not None
assert scores["920001.XSHE"] == 7.0 # 09-27 列: 序号 + 0
def test_factor_scores_before_all_returns_none(factor_dir):
"""调仓日早于全部截面日(预热期) → None,跳过调仓."""
st, _ = make_strategy(factor_dir)
assert st._factor_scores("2024-09-26") is None
# ------------------------ 调仓主流程 ------------------------
def test_rebalance_topn_and_universe(factor_dir):
"""TopN=3 取值最大 3 只(截面 asof 前一日)且 688/920 被剔,创业板可入;
下单代码为聚宽格式(撮合层能查到行情)."""
st, otv = make_strategy(factor_dir, top_n=3)
st.rebalance(_ctx("2024-10-08"))
# 09-30 截面(剔科创北交后): 300002=6.0 > 300001=5.0 > 000002=4.0 > ...
bought = [c.args[0] for c in otv.call_args_list]
assert bought == ["000002.XSHE", "300001.XSHE", "300002.XSHE"]
# 等权: total 1e6 / 3
assert all(c.args[1] == pytest.approx(1e6 / 3) for c in otv.call_args_list)
def test_rebalance_counter_only_every_k_days(factor_dir):
"""rebalance_days=5: 第 1/6 日触发,第 2-5 日不动."""
st, otv = make_strategy(factor_dir)
for _ in range(5):
st.rebalance(_ctx("2024-10-08"))
assert otv.call_count == 3 # 仅第 1 日建仓 3 笔,第 2-5 日 0 笔
st.rebalance(_ctx("2024-10-08")) # 第 6 日再触发
assert otv.call_count == 6
def test_rebalance_sells_stale_positions(factor_dir):
"""旧持仓不在新名单 → 先卖后买."""
st, otv = make_strategy(factor_dir)
ctx = _ctx("2024-10-08", positions={
"600000.XSHG": FakePosition("600000.XSHG", avg_cost=10.0, price=10.0)})
st.rebalance(ctx)
sells = [c for c in otv.call_args_list if c.args[1] == 0]
assert len(sells) == 1 and sells[0].args[0] == "600000.XSHG"
def test_to_jq_symbol():
"""后缀转换: SSE→XSHG / SZSE→XSHE / 其他原样."""
from sanguo_portfolio.strategies.factor_topn import _to_jq_symbol
assert _to_jq_symbol("600000.SSE") == "600000.XSHG"
assert _to_jq_symbol("300001.SZSE") == "300001.XSHE"
assert _to_jq_symbol("000300.XSHG") == "000300.XSHG"
def test_rebalance_skips_when_scores_insufficient(factor_dir):
"""截面有效值 < top_n → 整档跳过(不下单)."""
st, otv = make_strategy(factor_dir, top_n=100) # 截面只有 8 只
st.rebalance(_ctx("2024-10-08"))
otv.assert_not_called()
def test_rebalance_no_factor_dir_noop():
"""factor_dir 未配置 → 不炸不下单."""
st, otv = make_strategy("nonexistent_dir")
st.rebalance(_ctx("2024-10-08"))
otv.assert_not_called()
def test_initialize_registers_daily_schedule(factor_dir):
"""initialize 经 facade.run_daily 挂 9:30(回测真正注册在 runner _register_schedule)."""
st, _ = make_strategy(factor_dir)
calls: list[tuple] = []
st.broker.run_daily = lambda fn, t: calls.append((fn, t))
st.initialize(_ctx("2024-10-08"))
assert calls == [(st.rebalance, "9:30")]
def test_run_backtest_json_passes_factor_topn_params(monkeypatch):
"""CLI --json 路径(run_backtest_json 手工拼 Namespace)必须透传 factor 四参——
2026-09-12 矩阵三跑全空仓根因: 漏透传→factor_dir=""→策略「未配置」直接空仓,
而进程内 run_backtest(parse_args()) 路径完整(诊断一直成交掩盖此 bug)."""
from sanguo_portfolio import runner_backtest
captured: dict = {}
def _fake_run_backtest(args):
captured.update({
k: getattr(args, k, "<MISSING>")
for k in ("factor_name", "factor_dir", "top_n", "rebalance_days")
})
return {}
monkeypatch.setattr(runner_backtest, "run_backtest", _fake_run_backtest)
runner_backtest.run_backtest_json({
"strategy": "factor_topn",
"factor_name": "composite_test", "factor_dir": "/x/y",
"top_n": 50, "rebalance_days": 10,
})
assert captured == {"factor_name": "composite_test", "factor_dir": "/x/y",
"top_n": 50, "rebalance_days": 10}