Files
easy-tdx/tests/unit/test_portfolio_rebalance.py
T
Justin Gu 5fc398255d fix(types): 修复 CI mypy strict + ruff format 失败
mypy (13 errors → 0):
- portfolio/optimizer: register_optimizer 返回类型改为 Callable 装饰器签名
  (原标注 type[WeightOptimizer] 导致 4 个子类 Too many arguments)
- factor/engine: _datetime_to_int 用 isinstance 收窄替代 object→int 强转
- factor/analysis: 删多余 type:ignore(改由 mypy override 统一处理 scipy)
- backtest/orders, execution: np.sqrt 表达式用 float() 包裹消除 no-any-return
- MyTT.pyi: MACD 签名删除错误的 LOW/HIGH 参数(与 MyTT.py 实际签名对齐)
- pyproject: 新增 scipy mypy override (ignore_missing_imports)

ruff format: 8 个 test 文件格式化

验证: 564 passed, mypy 192 文件零错误, ruff check/format 全绿
2026-06-13 21:21:33 +08:00

77 lines
2.6 KiB
Python

"""Test RebalanceEngine."""
from __future__ import annotations
import numpy as np
import pandas as pd
from easy_tdx.portfolio.optimizer import EqualWeightOptimizer, FactorWeightedOptimizer
from easy_tdx.portfolio.rebalance import RebalanceEngine
def _make_market(n_stocks: int = 10, n_days: int = 120, seed: int = 42) -> dict[str, pd.DataFrame]:
rng = np.random.default_rng(seed)
data = {}
for i in range(n_stocks):
close = 10.0 + np.cumsum(rng.normal(0.01, 0.5, n_days))
close = np.maximum(close, 1.0)
data[f"{i:06d}"] = pd.DataFrame(
{
"datetime": pd.date_range("2024-01-01", periods=n_days, freq="D"),
"open": close,
"high": close + 0.3,
"low": close - 0.3,
"close": close,
"vol": rng.integers(1e5, 1e7, n_days).astype(float),
"amount": close * 1e6,
}
)
return data
class TestRebalanceEngine:
def test_basic_run(self):
engine = RebalanceEngine(
optimizer=EqualWeightOptimizer(),
factor_name="momentum_20d",
n_stocks=5,
rebalance_freq="M",
cash=1_000_000,
)
result = engine.run(_make_market(), start_date=20240101, end_date=20240430)
assert len(result.states) > 0
assert len(result.rebalance_dates) > 0
assert len(result.equity_curve) > 0
assert "total_return" in result.performance
def test_with_factor_weighted(self):
engine = RebalanceEngine(
optimizer=FactorWeightedOptimizer(),
factor_name="momentum_20d",
n_stocks=5,
)
result = engine.run(_make_market(), start_date=20240101, end_date=20240430)
assert len(result.states) > 0
def test_empty_data(self):
result = RebalanceEngine(optimizer=EqualWeightOptimizer()).run({})
assert result.performance["total_return"] == 0.0
def test_equity_curve_dates_sorted(self):
result = RebalanceEngine(
optimizer=EqualWeightOptimizer(),
rebalance_freq="M",
).run(_make_market(), start_date=20240101, end_date=20240430)
dates = result.equity_curve["datetime"].tolist()
assert dates == sorted(dates)
def test_trades_recorded(self):
engine = RebalanceEngine(
optimizer=EqualWeightOptimizer(),
n_stocks=3,
rebalance_freq="M",
)
result = engine.run(_make_market(), start_date=20240101, end_date=20240430)
assert len(result.trades) > 0
assert "BUY" in result.trades["direction"].values