618 lines
26 KiB
Python
618 lines
26 KiB
Python
"""AllWeatherStrategy 单元测试(mock provider + mock broker)。
|
|
|
|
策略层只测**逻辑分支正确**(选股 / 轮动决策 / 调仓),不测真实数据。
|
|
真实数据回测在 VPS 跑,这里只保证策略翻译等价。
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
from datetime import datetime
|
|
from typing import Any, Dict, List
|
|
from unittest.mock import MagicMock
|
|
|
|
import numpy as np
|
|
import pandas as pd
|
|
import pytest
|
|
|
|
from sanguo_portfolio import AllWeatherConfig, AllWeatherStrategy, BrokerFacade
|
|
|
|
|
|
# ------------------------ 测试 helper:构造策略实例 ------------------------
|
|
def make_strategy(
|
|
*,
|
|
fund_df: pd.DataFrame | None = None,
|
|
index_stocks_map: Dict[str, List[str]] | None = None,
|
|
price_df_map: Dict[str, pd.DataFrame] | None = None,
|
|
) -> AllWeatherStrategy:
|
|
"""构造一个 mock provider + mock broker 驱动的策略。
|
|
|
|
- fund_df: 默认 get_fundamentals_df 返回
|
|
- index_stocks_map: get_index_stocks 返回,dict[index] -> List[code]
|
|
- price_df_map: get_price 按 (code, fields) 缓存的返回
|
|
"""
|
|
provider = MagicMock(name="provider")
|
|
|
|
# 默认 fundamentals:空表,测试里覆盖
|
|
if fund_df is None:
|
|
fund_df = pd.DataFrame(columns=["code"])
|
|
provider.get_fundamentals_df.return_value = fund_df
|
|
|
|
# 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,
|
|
}
|
|
|
|
# get_price 按 key 缓存
|
|
price_df_map = price_df_map or {}
|
|
|
|
def _get_price(security, **kwargs):
|
|
# 构造 cache key:不严格,按 security+fields+count 取
|
|
fields = tuple(kwargs.get("fields") or [])
|
|
count = kwargs.get("count", 1)
|
|
key = (str(security), fields, count)
|
|
return price_df_map.get(key, pd.DataFrame())
|
|
|
|
provider.get_price.side_effect = _get_price
|
|
|
|
# get_closes_panel:从 price_df_map 同源构造 close 宽表(忽略 count 差异)
|
|
def _get_closes_panel(symbols, start=None, end=None, interval="d", fq="raw"):
|
|
frames = [
|
|
df[["time", "code", "close"]]
|
|
for (sec, fields, _cnt), df in price_df_map.items()
|
|
if sec == str(symbols) and tuple(fields) == ("close",)
|
|
and df is not None and len(df)
|
|
]
|
|
if not frames:
|
|
return pd.DataFrame()
|
|
long_df = pd.concat(frames)
|
|
return long_df.pivot(index="time", columns="code", values="close")
|
|
|
|
provider.get_closes_panel.side_effect = _get_closes_panel
|
|
|
|
# 涨跌停/停牌批量:默认全正常(filter_paused 无bar=剔除语义,MagicMock/{} 会误杀)
|
|
def _get_limit_status_batch(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 = _get_limit_status_batch
|
|
|
|
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 AllWeatherStrategy(provider=provider, broker=broker)
|
|
|
|
|
|
def make_fund_df(rows: List[Dict[str, Any]]) -> pd.DataFrame:
|
|
"""构造 fundamentals DataFrame(带 index = code)。"""
|
|
if not rows:
|
|
return pd.DataFrame(columns=["code"])
|
|
df = pd.DataFrame(rows)
|
|
df["code"] = df.get("code", df.index.astype(str))
|
|
df = df.set_index("code", drop=False)
|
|
return df
|
|
|
|
|
|
# =================== initialize ===================
|
|
class TestInitialize:
|
|
def test_initialize_registers_scheduled_tasks(self, fake_context):
|
|
# Arrange
|
|
s = make_strategy()
|
|
# Act
|
|
s.initialize(fake_context)
|
|
# Assert:run_daily / run_monthly 各被调一次(至少)
|
|
assert s.broker.run_daily.called
|
|
assert s.broker.run_monthly.called
|
|
assert s.broker.set_benchmark.called
|
|
|
|
def test_initialize_sets_benchmark_from_config(self, fake_context):
|
|
cfg = AllWeatherConfig(benchmark="000300.XSHG")
|
|
s = make_strategy()
|
|
s.config = cfg
|
|
s.initialize(fake_context)
|
|
s.broker.set_benchmark.assert_called_with("000300.XSHG")
|
|
|
|
|
|
# =================== prepare_stock_list ===================
|
|
class TestPrepareStockList:
|
|
def test_empty_positions_clears_lists(self):
|
|
# Arrange
|
|
s = make_strategy()
|
|
ctx = MagicMock()
|
|
ctx.portfolio.positions = {}
|
|
ctx.previous_date = "2024-09-30"
|
|
# Act
|
|
s.prepare_stock_list(ctx)
|
|
# Assert
|
|
assert s.hold_list == []
|
|
assert s.yesterday_hl_list == []
|
|
|
|
def test_populates_hold_list_from_positions(self):
|
|
s = make_strategy()
|
|
pos = MagicMock(); pos.security = "600519.XSHG"
|
|
ctx = MagicMock()
|
|
ctx.portfolio.positions = {"600519.XSHG": pos}
|
|
ctx.previous_date = "2024-09-30"
|
|
# get_price 返回空(不报错即可)
|
|
s.provider.get_price.return_value = pd.DataFrame()
|
|
s.prepare_stock_list(ctx)
|
|
assert s.hold_list == ["600519.XSHG"]
|
|
|
|
def test_records_yesterday_limit_up(self):
|
|
s = make_strategy()
|
|
pos = MagicMock(); pos.security = "600519.XSHG"
|
|
ctx = MagicMock()
|
|
ctx.portfolio.positions = {"600519.XSHG": pos}
|
|
ctx.previous_date = "2024-09-30"
|
|
# get_limit_status_batch 返回昨日涨停(G5-P2:原 get_price close==high_limit 改批量口径)
|
|
s.provider.get_limit_status_batch.side_effect = \
|
|
lambda codes, date=None: {
|
|
c: {"is_limit_up": True, "is_limit_down": False, "is_paused": False}
|
|
for c in codes
|
|
}
|
|
s.prepare_stock_list(ctx)
|
|
assert "600519.XSHG" in s.yesterday_hl_list
|
|
|
|
def test_uses_batch_endpoint_once(self):
|
|
"""G5-P2:原 get_price(close+high_limit) → get_limit_status_batch 一次批量。"""
|
|
s = make_strategy()
|
|
pos = MagicMock(); pos.security = "600519.XSHG"
|
|
ctx = MagicMock()
|
|
ctx.portfolio.positions = {"600519.XSHG": pos, "000001.XSHE": MagicMock(security="000001.XSHE")}
|
|
ctx.previous_date = "2024-09-30"
|
|
s.prepare_stock_list(ctx)
|
|
assert s.provider.get_limit_status_batch.call_count == 1
|
|
args, _kwargs = s.provider.get_limit_status_batch.call_args
|
|
assert sorted(args[0]) == ["000001.XSHE", "600519.XSHG"]
|
|
s.provider.get_price.assert_not_called()
|
|
|
|
|
|
# =================== trend_mean ===================
|
|
class TestTrendMean:
|
|
def _panel_seed(self):
|
|
return pd.DataFrame({
|
|
"time": pd.to_datetime(["2024-09-20", "2024-09-30"]),
|
|
"code": ["600519.XSHG"] * 2,
|
|
"close": [10.0, 15.0], # 涨 50%
|
|
})
|
|
|
|
def test_uses_closes_panel_not_get_price(self):
|
|
"""G5-P2:原 get_price 长表+pivot → get_closes_panel 宽表。"""
|
|
s = make_strategy(price_df_map={
|
|
("['600519.XSHG']", ("close",), 10): self._panel_seed(),
|
|
})
|
|
out = s._trend_mean(["600519.XSHG"], "2024-09-30", 10)
|
|
assert out == 50.0
|
|
s.provider.get_closes_panel.assert_called_once()
|
|
s.provider.get_price.assert_not_called()
|
|
|
|
def test_empty_or_short_panel_returns_zero(self):
|
|
"""panel 缺失/不足 2 行 → 0.0 不崩(降级同原 get_price 空df语义)。"""
|
|
s = make_strategy()
|
|
assert s._trend_mean([], "2024-09-30", 10) == 0.0
|
|
assert s._trend_mean(["600519.XSHG"], "2024-09-30", 10) == 0.0 # 无种子数据
|
|
|
|
def test_nan_only_change_not_propagated(self):
|
|
"""某列 NaN(close 缺失)→ nan_to_num 归 0,不产出 NaN。"""
|
|
wide = pd.DataFrame(
|
|
{"600519.XSHG": [np.nan, 15.0], "000001.XSHE": [10.0, 11.0]},
|
|
index=pd.to_datetime(["2024-09-20", "2024-09-30"]),
|
|
)
|
|
s = make_strategy()
|
|
s.provider.get_closes_panel.side_effect = None
|
|
s.provider.get_closes_panel.return_value = wide
|
|
# 600519 NaN→0%, 000001 +10% → 均值 5(浮点容差)
|
|
assert s._trend_mean(["600519.XSHG", "000001.XSHE"], "2024-09-30", 10) == pytest.approx(5.0)
|
|
class TestStopLoss:
|
|
def test_stop_loss_triggers_when_price_drops_8pct(self):
|
|
"""avg_cost=100, price=91 (< 100*0.92=92) → 止损。"""
|
|
from tests.portfolio.conftest import FakePosition, FakeContext
|
|
|
|
s = make_strategy()
|
|
pos = FakePosition("600519.XSHG", avg_cost=100.0, price=91.0)
|
|
ctx = FakeContext(positions={"600519.XSHG": pos})
|
|
s.yesterday_hl_list = [] # 跳过昨日涨停分支
|
|
s.stop_loss(ctx)
|
|
s.broker.order_target_value.assert_called_with("600519.XSHG", 0)
|
|
|
|
def test_stop_loss_skipped_when_price_above_threshold(self):
|
|
from tests.portfolio.conftest import FakePosition, FakeContext
|
|
|
|
s = make_strategy()
|
|
pos = FakePosition("600519.XSHG", avg_cost=100.0, price=95.0) # > 92
|
|
ctx = FakeContext(positions={"600519.XSHG": pos})
|
|
s.yesterday_hl_list = []
|
|
s.stop_loss(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 sell_calls == []
|
|
|
|
# ---- 昨日涨停分支:P1.2 去 1m 依赖,改 get_limit_status_batch 日线口径 ----
|
|
|
|
@staticmethod
|
|
def _limit_status_side_effect(is_limit_up):
|
|
def _glbs(codes, date):
|
|
return {c: {"is_limit_up": is_limit_up, "is_limit_down": False,
|
|
"is_paused": False} for c in codes}
|
|
return _glbs
|
|
|
|
def test_yesterday_limitup_sold_when_today_not_limitup(self):
|
|
"""昨日涨停 + 当日(日线)未涨停 → 涨停打开卖出。"""
|
|
from tests.portfolio.conftest import FakePosition, FakeContext
|
|
|
|
s = make_strategy()
|
|
pos = FakePosition("600519.XSHG", avg_cost=100.0, price=95.0) # 不触发 -8%
|
|
ctx = FakeContext(positions={"600519.XSHG": pos})
|
|
s.yesterday_hl_list = ["600519.XSHG"]
|
|
s.provider.get_limit_status_batch.side_effect = \
|
|
self._limit_status_side_effect(is_limit_up=False)
|
|
s.stop_loss(ctx)
|
|
s.broker.order_target_value.assert_called_with("600519.XSHG", 0)
|
|
|
|
def test_yesterday_limitup_kept_when_still_limitup(self):
|
|
"""昨日涨停 + 当日仍涨停 → 继续持有,不卖。"""
|
|
from tests.portfolio.conftest import FakePosition, FakeContext
|
|
|
|
s = make_strategy()
|
|
pos = FakePosition("600519.XSHG", avg_cost=100.0, price=95.0)
|
|
ctx = FakeContext(positions={"600519.XSHG": pos})
|
|
s.yesterday_hl_list = ["600519.XSHG"]
|
|
s.provider.get_limit_status_batch.side_effect = \
|
|
self._limit_status_side_effect(is_limit_up=True)
|
|
s.stop_loss(ctx)
|
|
sell_calls = [c for c in s.broker.order_target_value.call_args_list if c.args[1] == 0]
|
|
assert sell_calls == []
|
|
|
|
def test_yesterday_limitup_provider_failure_skips_branch(self):
|
|
"""provider 无 get_limit_status_batch / 查询异常 → 跳过该分支不崩(降级)。"""
|
|
from tests.portfolio.conftest import FakePosition, FakeContext
|
|
|
|
s = make_strategy()
|
|
pos = FakePosition("600519.XSHG", avg_cost=100.0, price=95.0)
|
|
ctx = FakeContext(positions={"600519.XSHG": pos})
|
|
s.yesterday_hl_list = ["600519.XSHG"]
|
|
s.provider.get_limit_status_batch.side_effect = RuntimeError("boom")
|
|
s.stop_loss(ctx) # 不抛异常
|
|
sell_calls = [c for c in s.broker.order_target_value.call_args_list if c.args[1] == 0]
|
|
assert sell_calls == []
|
|
|
|
|
|
# =================== monthly_adjustment:轮动决策分支 ===================
|
|
class TestMonthlyAdjustmentDecision:
|
|
def test_foreign_etf_branch_when_both_trends_negative(self):
|
|
"""b_mean < 0 且 s_mean < 0 → 开外盘(海外 ETF)。"""
|
|
# Arrange
|
|
s = make_strategy(
|
|
index_stocks_map={
|
|
"000300.XSHG": ["600519.XSHG"],
|
|
"399101.XSHE": ["000001.XSHE"],
|
|
},
|
|
price_df_map={
|
|
# trend window = 10, 但 close 都跌
|
|
("['600519.XSHG']", ("close",), 10): pd.DataFrame({
|
|
"time": pd.to_datetime(["2024-09-20", "2024-09-30"]),
|
|
"code": ["600519.XSHG"] * 2,
|
|
"close": [15.0, 10.0], # 跌
|
|
}),
|
|
("['000001.XSHE']", ("close",), 10): pd.DataFrame({
|
|
"time": pd.to_datetime(["2024-09-20", "2024-09-30"]),
|
|
"code": ["000001.XSHE"] * 2,
|
|
"close": [15.0, 10.0],
|
|
}),
|
|
},
|
|
)
|
|
# 流通市值 top/bottom 的 fund_df:让 _market_cap_top 仍能跑
|
|
s.provider.get_fundamentals_df.return_value = make_fund_df([
|
|
{"code": "600519.XSHG", "circulating_market_cap": 20000, "market_cap": 20000},
|
|
{"code": "000001.XSHE", "circulating_market_cap": 500, "market_cap": 500},
|
|
])
|
|
ctx = MagicMock()
|
|
ctx.current_dt = datetime(2024, 10, 8, 9, 30)
|
|
ctx.previous_date = "2024-09-30"
|
|
ctx.portfolio.positions = {}
|
|
ctx.portfolio.available_cash = 1_000_000
|
|
|
|
# Act
|
|
s.monthly_adjustment(ctx)
|
|
|
|
# Assert:海外 ETF 在 order_target_value 入参里
|
|
called_codes = [c.args[0] for c in s.broker.order_target_value.call_args_list]
|
|
for etf in s.config.foreign_etf:
|
|
assert etf in called_codes, f"未触发海外 ETF 下单: {etf}"
|
|
|
|
def test_rotation_buys_new_after_selling_old(self):
|
|
"""换仓月回归:卖出旧仓后须重取持仓再决定买入。
|
|
|
|
2025-09-01/12-01 实况:持仓数≥目标数时,step5 卖完全部旧仓,
|
|
step6 却用卖出【前】的陈旧持仓计数(len≥target)→ 一股不买 →
|
|
空仓躺到下月。聚宽原版卖出后持仓同步更新,翻译时快照没刷新。
|
|
"""
|
|
from tests.portfolio.conftest import FakePosition
|
|
|
|
s = make_strategy(
|
|
index_stocks_map={
|
|
"000300.XSHG": ["600519.XSHG"],
|
|
"399101.XSHE": ["000001.XSHE"],
|
|
},
|
|
price_df_map={
|
|
("['600519.XSHG']", ("close",), 10): pd.DataFrame({
|
|
"time": pd.to_datetime(["2024-09-20", "2024-09-30"]),
|
|
"code": ["600519.XSHG"] * 2, "close": [15.0, 10.0],
|
|
}),
|
|
("['000001.XSHE']", ("close",), 10): pd.DataFrame({
|
|
"time": pd.to_datetime(["2024-09-20", "2024-09-30"]),
|
|
"code": ["000001.XSHE"] * 2, "close": [15.0, 10.0],
|
|
}),
|
|
},
|
|
)
|
|
s.provider.get_fundamentals_df.return_value = make_fund_df([
|
|
{"code": "600519.XSHG", "circulating_market_cap": 20000, "market_cap": 20000},
|
|
{"code": "000001.XSHE", "circulating_market_cap": 500, "market_cap": 500},
|
|
])
|
|
|
|
# 持仓 5 只(= foreign_etf 数量),卖出后动态清空(模拟引擎持仓属性)
|
|
held = {f"60000{i}.XSHG": FakePosition(f"60000{i}.XSHG", 10.0, 10.0)
|
|
for i in range(5)}
|
|
state = dict(held)
|
|
|
|
class DynPortfolio:
|
|
available_cash = 1_000_000
|
|
cash = 1_000_000
|
|
|
|
@property
|
|
def positions(self):
|
|
return dict(state)
|
|
|
|
ctx = MagicMock()
|
|
ctx.current_dt = datetime(2024, 10, 8, 9, 30)
|
|
ctx.previous_date = "2024-09-30"
|
|
ctx.portfolio = DynPortfolio()
|
|
|
|
def _order(code, value):
|
|
if value == 0:
|
|
state.pop(code, None) # 卖出→持仓减少(引擎语义)
|
|
return MagicMock(filled=100)
|
|
s.broker.order_target_value.side_effect = _order
|
|
|
|
s.monthly_adjustment(ctx)
|
|
|
|
# 旧仓 5 只全卖
|
|
sell_codes = [c.args[0] for c in s.broker.order_target_value.call_args_list
|
|
if c.args[1] == 0]
|
|
assert sorted(sell_codes) == sorted(held.keys())
|
|
# 新仓(5 只 ETF)要买进来——陈旧计数会让 5>5=False 一股不买
|
|
buy_codes = [c.args[0] for c in s.broker.order_target_value.call_args_list
|
|
if c.args[1] > 0]
|
|
assert sorted(buy_codes) == sorted(s.config.foreign_etf), \
|
|
f"换仓月未买入新目标: 实买={buy_codes}"
|
|
|
|
def test_foreign_etf_branch_skips_limitup_and_paused(self):
|
|
"""P1.3:涨停(未持有)与停牌的 ETF 不买入——filter 批量预取接线。"""
|
|
s = make_strategy(
|
|
index_stocks_map={
|
|
"000300.XSHG": ["600519.XSHG"],
|
|
"399101.XSHE": ["000001.XSHE"],
|
|
},
|
|
price_df_map={
|
|
("['600519.XSHG']", ("close",), 10): pd.DataFrame({
|
|
"time": pd.to_datetime(["2024-09-20", "2024-09-30"]),
|
|
"code": ["600519.XSHG"] * 2,
|
|
"close": [15.0, 10.0], # 跌
|
|
}),
|
|
("['000001.XSHE']", ("close",), 10): pd.DataFrame({
|
|
"time": pd.to_datetime(["2024-09-20", "2024-09-30"]),
|
|
"code": ["000001.XSHE"] * 2,
|
|
"close": [15.0, 10.0], # 跌
|
|
}),
|
|
},
|
|
)
|
|
s.provider.get_fundamentals_df.return_value = make_fund_df([
|
|
{"code": "600519.XSHG", "circulating_market_cap": 20000, "market_cap": 20000},
|
|
{"code": "000001.XSHE", "circulating_market_cap": 500, "market_cap": 500},
|
|
])
|
|
# 518880 涨停(未持有不买)、513030 停牌(不交易),其余正常
|
|
def _glbs(codes, date):
|
|
out = {}
|
|
for c in codes:
|
|
if c == "518880.XSHG":
|
|
out[c] = {"is_limit_up": True, "is_limit_down": False, "is_paused": False}
|
|
elif c == "513030.XSHG":
|
|
out[c] = {"is_limit_up": False, "is_limit_down": False, "is_paused": True}
|
|
else:
|
|
out[c] = {"is_limit_up": False, "is_limit_down": False, "is_paused": False}
|
|
return out
|
|
s.provider.get_limit_status_batch.side_effect = _glbs
|
|
|
|
ctx = MagicMock()
|
|
ctx.current_dt = datetime(2024, 10, 8, 9, 30)
|
|
ctx.previous_date = "2024-09-30"
|
|
ctx.portfolio.positions = {}
|
|
ctx.portfolio.available_cash = 1_000_000
|
|
|
|
s.monthly_adjustment(ctx)
|
|
|
|
called_codes = [c.args[0] for c in s.broker.order_target_value.call_args_list]
|
|
assert "518880.XSHG" not in called_codes, "涨停 ETF 不应买入"
|
|
assert "513030.XSHG" not in called_codes, "停牌 ETF 不应交易"
|
|
for etf in ("513100.XSHG", "164824.XSHE", "159866.XSHE"):
|
|
assert etf in called_codes, f"正常 ETF 应下单: {etf}"
|
|
|
|
def test_big_market_branch_when_b_trend_dominant(self):
|
|
"""b_mean > s_mean 且 b_mean > 0 → 开大(选 B_stocks)。"""
|
|
s = make_strategy(
|
|
index_stocks_map={
|
|
"000300.XSHG": ["600519.XSHG"],
|
|
"399101.XSHE": ["000001.XSHE"],
|
|
},
|
|
price_df_map={
|
|
("['600519.XSHG']", ("close",), 10): pd.DataFrame({
|
|
"time": pd.to_datetime(["2024-09-20", "2024-09-30"]),
|
|
"code": ["600519.XSHG"] * 2,
|
|
"close": [10.0, 15.0], # 涨 50%
|
|
}),
|
|
("['000001.XSHE']", ("close",), 10): pd.DataFrame({
|
|
"time": pd.to_datetime(["2024-09-20", "2024-09-30"]),
|
|
"code": ["000001.XSHE"] * 2,
|
|
"close": [10.0, 11.0], # 涨 10%
|
|
}),
|
|
},
|
|
)
|
|
# 选股函数返回的 fund_df:让 big 路径选到 1 只
|
|
big_fund = make_fund_df([{
|
|
"code": "600519.XSHG", "market_cap": 20000,
|
|
"circulating_market_cap": 20000,
|
|
"pe_ratio": 10.0, "ps_ratio": 2.0, "pcf_ratio": 2.0,
|
|
"eps": 1.0, "roe": 0.2, "roa": 0.15,
|
|
"net_profit_margin": 0.2, "gross_profit_margin": 0.5,
|
|
"inc_revenue_year_on_year": 0.3,
|
|
"inc_operation_profit_year_on_year": 0.2,
|
|
"inc_total_revenue_year_on_year": 0.4,
|
|
"total_liability": 1e9, "total_sheet_owner_equities": 1e10,
|
|
"retained_profit": 5e9, "roic": 0.15, "pb_ratio": 2.0,
|
|
}])
|
|
s.provider.get_fundamentals_df.return_value = big_fund
|
|
ctx = MagicMock()
|
|
ctx.current_dt = datetime(2024, 10, 8, 9, 30)
|
|
ctx.previous_date = "2024-09-30"
|
|
ctx.portfolio.positions = {}
|
|
ctx.portfolio.available_cash = 1_000_000
|
|
|
|
s.monthly_adjustment(ctx)
|
|
|
|
# 600519 应被买入(开大 + 多个选股函数都会选它)
|
|
buy_calls = [
|
|
c.args[0] for c in s.broker.order_target_value.call_args_list
|
|
if c.args[1] != 0
|
|
]
|
|
assert "600519.XSHG" in buy_calls
|
|
|
|
|
|
# =================== 选股函数直接测试 ===================
|
|
class TestStockPickers:
|
|
def test_small_filters_by_roe_roa(self):
|
|
"""roe>0.05 & roa>0.02 → 仅保留合格股,按 market_cap asc。
|
|
|
|
阈值是 e807bed 验证用放宽口径(原 0.15/0.10 对中证1000 命中仅~5%),
|
|
最终业务决策再定——测试锚定当前实现。
|
|
"""
|
|
df = make_fund_df([
|
|
{"code": "A.XSHG", "roe": 0.20, "roa": 0.15, "market_cap": 500},
|
|
{"code": "B.XSHG", "roe": 0.03, "roa": 0.20, "market_cap": 300}, # roe 不够
|
|
{"code": "C.XSHG", "roe": 0.30, "roa": 0.01, "market_cap": 200}, # roa 不够
|
|
{"code": "D.XSHG", "roe": 0.25, "roa": 0.12, "market_cap": 100},
|
|
])
|
|
s = make_strategy()
|
|
s.provider.get_fundamentals_df.return_value = df
|
|
out = s.small(["A", "B", "C", "D"], current_dt=None, previous_date="2024-09-30")
|
|
# A 和 D 合格,D 市值小排前
|
|
assert out == ["D.XSHG", "A.XSHG"]
|
|
|
|
def test_big_applies_full_multi_factor_filter(self):
|
|
df = make_fund_df([{
|
|
# 全部满足
|
|
"code": "PASS.XSHG",
|
|
"market_cap": 500,
|
|
"pe_ratio": 15.0, "ps_ratio": 3.0, "pcf_ratio": 5.0,
|
|
"eps": 1.0, "roe": 0.2, "net_profit_margin": 0.2,
|
|
"gross_profit_margin": 0.5, "inc_revenue_year_on_year": 0.3,
|
|
}, {
|
|
"code": "FAIL.XSHG",
|
|
"market_cap": 800,
|
|
"pe_ratio": 50.0, # pe 不在 [0,30]
|
|
"ps_ratio": 3.0, "pcf_ratio": 5.0,
|
|
"eps": 1.0, "roe": 0.2, "net_profit_margin": 0.2,
|
|
"gross_profit_margin": 0.5, "inc_revenue_year_on_year": 0.3,
|
|
}])
|
|
s = make_strategy()
|
|
s.provider.get_fundamentals_df.return_value = df
|
|
out = s.big(["PASS", "FAIL"], current_dt=None, previous_date="2024-09-30")
|
|
assert out == ["PASS.XSHG"]
|
|
|
|
def test_roic_big_filters_by_roic_threshold(self):
|
|
"""ROIC > 0.08 才保留。"""
|
|
df = make_fund_df([
|
|
{"code": "HIGH.XSHG", "market_cap": 500, "pe_ratio": 20,
|
|
"eps": 0.5, "roa": 0.20, "total_liability": 1e8,
|
|
"total_sheet_owner_equities": 1e10, "retained_profit": 5e9,
|
|
"inc_total_revenue_year_on_year": 0.4,
|
|
"inc_revenue_year_on_year": 0.3, "roic": 0.15},
|
|
{"code": "LOW.XSHG", "market_cap": 500, "pe_ratio": 20,
|
|
"eps": 0.5, "roa": 0.20, "total_liability": 1e8,
|
|
"total_sheet_owner_equities": 1e10, "retained_profit": 5e9,
|
|
"inc_total_revenue_year_on_year": 0.4,
|
|
"inc_revenue_year_on_year": 0.3, "roic": 0.05}, # ROIC 不够
|
|
])
|
|
s = make_strategy()
|
|
s.provider.get_fundamentals_df.return_value = df
|
|
out = s.roic_big(["HIGH", "LOW"], current_dt=None, previous_date="2024-09-30")
|
|
assert "HIGH.XSHG" in out
|
|
assert "LOW.XSHG" not in out
|
|
|
|
def test_bm_uses_mid_cap_value_filters(self):
|
|
df = make_fund_df([{
|
|
"code": "GOOD.XSHG",
|
|
"market_cap": 500, "pb_ratio": 2.0, "pcf_ratio": 2.0,
|
|
"eps": 1.0, "roe": 0.3, "net_profit_margin": 0.2,
|
|
"inc_revenue_year_on_year": 0.3,
|
|
"inc_operation_profit_year_on_year": 0.2,
|
|
}, {
|
|
"code": "BIG.XSHG",
|
|
"market_cap": 1000, # 不在 [100, 900]
|
|
"pb_ratio": 2.0, "pcf_ratio": 2.0,
|
|
"eps": 1.0, "roe": 0.3, "net_profit_margin": 0.2,
|
|
"inc_revenue_year_on_year": 0.3,
|
|
"inc_operation_profit_year_on_year": 0.2,
|
|
}])
|
|
s = make_strategy()
|
|
s.provider.get_fundamentals_df.return_value = df
|
|
out = s.bm(["GOOD", "BIG"], current_dt=None, previous_date="2024-09-30")
|
|
assert out == ["GOOD.XSHG"]
|
|
|
|
|
|
# =================== filter_roic ===================
|
|
class TestFilterRoic:
|
|
def test_filters_below_threshold(self):
|
|
df = make_fund_df([{"code": "A.XSHG", "roic": 0.15}])
|
|
s = make_strategy()
|
|
s.provider.get_fundamentals_df.return_value = df
|
|
out = s.filter_roic(["A.XSHG", "B.XSHG"], previous_date="2024-09-30")
|
|
# 第 2 次调用 fund_df 也是同一个 mock,所以 B 也算 roic=0.15 → 都保留
|
|
assert "A.XSHG" in out
|
|
|
|
def test_empty_input_returns_empty(self):
|
|
s = make_strategy()
|
|
out = s.filter_roic([], previous_date="2024-09-30")
|
|
assert out == []
|
|
|
|
def test_batch_single_call_with_fields_shortcut(self):
|
|
"""G5-P2:原逐只循环 → 一次批量 + fields=['roic'] 短路。"""
|
|
df = make_fund_df([{"code": "A.XSHG", "roic": 0.15}])
|
|
s = make_strategy()
|
|
s.provider.get_fundamentals_df.return_value = df
|
|
s.filter_roic(["A.XSHG", "B.XSHG"], previous_date="2024-09-30")
|
|
# 一次批量调用(非逐只 N 次),且带 fields 短路
|
|
assert s.provider.get_fundamentals_df.call_count == 1
|
|
args, kwargs = s.provider.get_fundamentals_df.call_args
|
|
assert list(args[0]) == ["A.XSHG", "B.XSHG"]
|
|
assert kwargs.get("fields") == ["roic"]
|