初始提交:高股息策略研究与回测系统
从 Point-in-Time 股票筛选到统一 Web 前端的完整链路: 筛选 → 画像 → 策略 → 回测 → Walk-forward → 绩效分析 → 报告/前端。 架构 - 数据层与策略层分离;策略代码不写 SQL,只经 data/repo.py 取数 - 所有业务阈值集中在 config/*.yml,代码零硬编码(字段写错直接报错) - 报告只做「run_id → SQL → 渲染」,不做任何计算,数字可追溯 - 前后端分离:output/ 静态站点 + hdiv web 提供的 REST API 数据安全 - 只增不删:SQL 钩子拦截 DELETE/DROP/TRUNCATE,并有源码扫描测试守护 - qlib 原有表只读,本项目数据写入 hd_ 前缀表 - 回补使用 INSERT IGNORE,保证既有行零改动 - .env 存密钥且已 gitignore;output/、logs/、.venv/ 不入库 交付物 - 30 张 hd_* 表、7 个 YAML 配置、283 项自动化测试 - 统一 Web 前端(hash 路由 SPA)+ nginx 部署配置与 launchd 托管脚本 如实声明的限制 - 策略缺少稳定的样本外超额收益(Walk-forward 7 窗口均值 -0.95%, 基准 +2.29%);其价值体现在回撤控制,而非超额收益 - 涨跌停/停牌约束仅覆盖 2019 年起;index_weight 尚未填充 - AI Agent 层(plan.md 第四版 P8)未实现 详见 docs/user-guide.md 与 docs/implementation-status.md。
This commit is contained in:
@@ -0,0 +1,460 @@
|
||||
"""回测引擎、成本、分红、绩效与敏感性测试。
|
||||
|
||||
覆盖的都是「错了也不会报错、只会静默给出错误结论」的地方:
|
||||
- 总市值漏掉现金 → 净值曲线失真;
|
||||
- 会计恒等式不平 → 成本/分红有遗漏;
|
||||
- 分位阈值口径混用 → 未来函数;
|
||||
- 分批建仓与减仓互相冲突 → 高频无效交易;
|
||||
- 参数耦合未同步 → 扫出的差异来自形状畸变而非阈值本身。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
from hdiv.backtest.engine import (
|
||||
CostModel,
|
||||
Position,
|
||||
Signal,
|
||||
_months_between,
|
||||
_round_lot,
|
||||
reconcile,
|
||||
)
|
||||
from hdiv.backtest.walk_forward import WalkForwardRunner, _add_months, _add_years
|
||||
from hdiv.core.config import CostConfig, load_config
|
||||
from hdiv.strategy.registry import StrategyRegistry, apply_sweep, parse_sweep
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 成本模型
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def cost() -> CostModel:
|
||||
return CostModel(load_config("cost"))
|
||||
|
||||
|
||||
def test_commission_has_minimum(cost: CostModel) -> None:
|
||||
"""小额成交必须触发最低佣金 5 元。"""
|
||||
c, s, t = cost.fees(1000.0, "BUY")
|
||||
assert c == pytest.approx(5.0), "1000 × 0.025% = 0.25 元,应被最低佣金托底"
|
||||
assert s == 0.0, "买入不收印花税"
|
||||
assert t == pytest.approx(1000.0 * 0.00001)
|
||||
|
||||
|
||||
def test_large_commission_uses_rate(cost: CostModel) -> None:
|
||||
c, _, _ = cost.fees(1_000_000.0, "BUY")
|
||||
assert c == pytest.approx(250.0)
|
||||
|
||||
|
||||
def test_stamp_duty_sell_only(cost: CostModel) -> None:
|
||||
_, s_buy, _ = cost.fees(1_000_000.0, "BUY")
|
||||
_, s_sell, _ = cost.fees(1_000_000.0, "SELL")
|
||||
assert s_buy == 0.0
|
||||
assert s_sell == pytest.approx(500.0)
|
||||
|
||||
|
||||
def test_slippage_direction(cost: CostModel) -> None:
|
||||
"""买入价上滑、卖出价下滑 —— 方向反了会凭空产生收益。"""
|
||||
assert cost.slip(100.0, "BUY") > 100.0
|
||||
assert cost.slip(100.0, "SELL") < 100.0
|
||||
assert cost.slip(100.0, "BUY") == pytest.approx(100.1) # 10bps
|
||||
|
||||
|
||||
def test_slippage_modes() -> None:
|
||||
cfg = CostConfig.model_validate(
|
||||
{"version": 1, "slippage": {"mode": "fixed", "value": 0.02}}
|
||||
)
|
||||
assert CostModel(cfg).slip(100.0, "BUY") == pytest.approx(100.02)
|
||||
cfg2 = CostConfig.model_validate(
|
||||
{"version": 1, "slippage": {"mode": "tick", "value": 2}}
|
||||
)
|
||||
assert CostModel(cfg2).slip(100.0, "BUY") == pytest.approx(100.02)
|
||||
|
||||
|
||||
def test_dividend_tax_by_holding_period(cost: CostModel) -> None:
|
||||
"""plan.md §30:持股越久税率越低,超过 1 年免税。"""
|
||||
assert cost.dividend_tax_rate(10) == pytest.approx(0.20)
|
||||
assert cost.dividend_tax_rate(100) == pytest.approx(0.10)
|
||||
assert cost.dividend_tax_rate(400) == pytest.approx(0.00)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# A 股交易规则
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_round_lot_is_100_shares() -> None:
|
||||
assert _round_lot(150) == 100
|
||||
assert _round_lot(99) == 0
|
||||
assert _round_lot(1000) == 1000
|
||||
assert _round_lot(-5) == 0
|
||||
|
||||
|
||||
def test_months_between() -> None:
|
||||
assert _months_between(date(2024, 1, 15), date(2024, 1, 30)) == 0
|
||||
assert _months_between(date(2024, 1, 15), date(2024, 2, 1)) == 1
|
||||
assert _months_between(date(2023, 12, 1), date(2024, 12, 1)) == 12
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 资金对账(P4 硬验收)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _T:
|
||||
def __init__(self, side: str, amount: float, fee: float) -> None:
|
||||
self.side = side
|
||||
self.amount = amount
|
||||
self.total_cost = fee
|
||||
|
||||
|
||||
def test_reconcile_balanced() -> None:
|
||||
eq = pd.DataFrame({"cash": [1000.0, 300.0, 550.0]})
|
||||
trades = [_T("BUY", 700.0, 5.0), _T("SELL", 300.0, 3.0)]
|
||||
# 1000 - 700 - 5 - 3 + 300 + 分红 0 = 592?
|
||||
# 实际:1000 - 700 - 5(买佣) - 3(卖佣) + 300 = 592;现金应为 592
|
||||
eq = pd.DataFrame({"cash": [592.0]})
|
||||
rc = reconcile(eq, trades, total_dividend_net=0.0, initial_capital=1000.0)
|
||||
assert rc["balanced"] is True
|
||||
assert rc["residual"] == pytest.approx(0.0)
|
||||
|
||||
|
||||
def test_reconcile_detects_missing_dividend() -> None:
|
||||
eq = pd.DataFrame({"cash": [700.0]})
|
||||
trades = [_T("BUY", 300.0, 0.0)]
|
||||
rc = reconcile(eq, trades, total_dividend_net=0.0, initial_capital=1000.0)
|
||||
# 期望现金 700,实际 700 → 平衡
|
||||
assert rc["balanced"] is True
|
||||
# 若真实有 50 元分红但没记账,现金会多出 50 → 应被检出
|
||||
rc2 = reconcile(eq, trades, total_dividend_net=50.0, initial_capital=1000.0)
|
||||
assert rc2["balanced"] is False
|
||||
assert rc2["residual"] == pytest.approx(-50.0)
|
||||
|
||||
|
||||
def test_reconcile_does_not_include_position_value() -> None:
|
||||
"""买入的股票仍在账上,其市值不是现金口径的误差。"""
|
||||
eq = pd.DataFrame({"cash": [200.0]})
|
||||
trades = [_T("BUY", 800.0, 0.0)]
|
||||
rc = reconcile(eq, trades, 0.0, 1000.0)
|
||||
assert rc["balanced"] is True, "1000 − 800 = 200,持仓市值不应进入残差"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 目标仓位阶梯(防「分批建仓/减仓互相冲突」)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def engine():
|
||||
from hdiv.backtest.engine import BacktestEngine
|
||||
|
||||
return BacktestEngine.from_strategy("config/strategy/high_dividend_v1.yml")
|
||||
|
||||
|
||||
def test_target_weight_ladder(engine) -> None:
|
||||
"""plan.md §19 的建仓阶梯 + §18 的减仓阶梯,必须合成单一函数。"""
|
||||
assert engine._target_weight(95) == pytest.approx(1.00)
|
||||
assert engine._target_weight(88) == pytest.approx(0.75)
|
||||
assert engine._target_weight(82) == pytest.approx(0.50)
|
||||
assert engine._target_weight(76) == pytest.approx(0.25)
|
||||
assert engine._target_weight(45) == pytest.approx(0.50)
|
||||
assert engine._target_weight(30) == pytest.approx(0.50)
|
||||
assert engine._target_weight(10) == pytest.approx(0.00)
|
||||
|
||||
|
||||
def test_target_weight_has_dead_zone(engine) -> None:
|
||||
"""死区必须存在 —— 否则分位抖动会导致高频无效交易(原实现年换手 8.9)。"""
|
||||
for pct in (51, 60, 70, 74.9):
|
||||
assert engine._target_weight(pct) is None, f"分位 {pct} 应落在死区"
|
||||
|
||||
|
||||
def test_target_weight_is_monotonic_on_entry_side(engine) -> None:
|
||||
"""买入侧:分位越高仓位越重(单调不减)。"""
|
||||
vals = [engine._target_weight(p) for p in (75, 80, 85, 90)]
|
||||
assert all(v is not None for v in vals)
|
||||
assert vals == sorted(vals)
|
||||
|
||||
|
||||
def test_target_weight_is_monotonic_on_exit_side(engine) -> None:
|
||||
"""卖出侧:分位越低仓位越轻(单调不减)。"""
|
||||
vals = [engine._target_weight(p) for p in (10, 25, 40, 50)]
|
||||
assert all(v is not None for v in vals)
|
||||
assert vals == sorted(vals)
|
||||
|
||||
|
||||
def test_held_position_at_high_percentile_is_not_trimmed(engine) -> None:
|
||||
"""已持仓且分位很高时不应被误判为减仓 —— 这是原实现的真实 bug。"""
|
||||
assert engine._target_weight(92) == pytest.approx(1.0)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 参数扫描的耦合处理
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_parse_sweep() -> None:
|
||||
g = parse_sweep("entry.yield_percentile=70,75,80")
|
||||
assert g == {"entry.yield_percentile": [70, 75, 80]}
|
||||
g2 = parse_sweep("a.b=1,2;c.d=x,y")
|
||||
assert g2 == {"a.b": [1, 2], "c.d": ["x", "y"]}
|
||||
|
||||
|
||||
def test_parse_sweep_rejects_bad_format() -> None:
|
||||
from hdiv.core.errors import SchemaValidationError
|
||||
|
||||
with pytest.raises(SchemaValidationError):
|
||||
parse_sweep("entry.yield_percentile")
|
||||
|
||||
|
||||
def test_apply_sweep_shifts_ladder_shape() -> None:
|
||||
"""扫描 entry.yield_percentile 时必须整体平移 scale_in,而非只改首档。"""
|
||||
r = StrategyRegistry()
|
||||
s = r.load("config/strategy/high_dividend_v1.yml")
|
||||
original_weights = [st.weight for st in s.entry.scale_in]
|
||||
for p in (65, 70, 80, 90, 95):
|
||||
x = apply_sweep(s, {"entry.yield_percentile": p})
|
||||
ladder = [st.percentile for st in x.entry.scale_in]
|
||||
assert x.entry.yield_percentile == pytest.approx(ladder[0])
|
||||
assert ladder == sorted(set(ladder)), f"P{p} 的阶梯必须严格升序:{ladder}"
|
||||
assert max(ladder) <= 100, f"P{p} 的阶梯越界:{ladder}"
|
||||
# 权重形状必须保持
|
||||
assert [st.weight for st in x.entry.scale_in] == original_weights
|
||||
|
||||
|
||||
def test_apply_sweep_exit_coupling() -> None:
|
||||
r = StrategyRegistry()
|
||||
s = r.load("config/strategy/high_dividend_v1.yml")
|
||||
x = apply_sweep(s, {"exit.yield_percentile": 30})
|
||||
assert x.exit.yield_percentile == pytest.approx(30)
|
||||
assert x.exit.scale_out[-1].percentile == pytest.approx(30)
|
||||
|
||||
|
||||
def test_apply_sweep_unknown_path() -> None:
|
||||
from hdiv.core.errors import SchemaValidationError
|
||||
|
||||
r = StrategyRegistry()
|
||||
s = r.load("config/strategy/high_dividend_v1.yml")
|
||||
with pytest.raises(SchemaValidationError):
|
||||
apply_sweep(s, {"entry.not_a_field": 1})
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 敏感性判读
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_sensitivity_analysis_flags_spike() -> None:
|
||||
"""plan.md §27 的尖峰情形必须被识别为疑似过拟合。"""
|
||||
from hdiv.analysis.sensitivity import SensitivityRunner
|
||||
|
||||
points = [
|
||||
{"cagr": 0.13, "max_drawdown": -0.2},
|
||||
{"cagr": 0.135, "max_drawdown": -0.2},
|
||||
{"cagr": 0.20, "max_drawdown": -0.2},
|
||||
{"cagr": 0.132, "max_drawdown": -0.2},
|
||||
{"cagr": 0.128, "max_drawdown": -0.2},
|
||||
]
|
||||
a = SensitivityRunner._analyse(points, {"entry.yield_percentile": [75]})
|
||||
assert a["spikes"], "P80 的 20% 相对邻居是明显尖峰,必须被检出"
|
||||
assert a["robust"] is False
|
||||
assert "过拟合" in a["verdict"]
|
||||
|
||||
|
||||
def test_sensitivity_analysis_accepts_smooth_curve() -> None:
|
||||
"""plan.md §27 的平滑情形应被判定为对参数不敏感。"""
|
||||
from hdiv.analysis.sensitivity import SensitivityRunner
|
||||
|
||||
points = [{"cagr": c, "max_drawdown": -0.2} for c in (0.13, 0.133, 0.135, 0.132, 0.130)]
|
||||
a = SensitivityRunner._analyse(points, {"entry.yield_percentile": [75]})
|
||||
assert not a["spikes"]
|
||||
assert a["smoothness"] > 0.6
|
||||
assert a["robust"] is True
|
||||
assert "不敏感" in a["verdict"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Walk-forward 窗口
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_walk_forward_windows_are_disjoint_and_ordered() -> None:
|
||||
w = WalkForwardRunner()
|
||||
wins = w.windows()
|
||||
assert wins, "应至少切出一个窗口"
|
||||
for a, b in zip(wins, wins[1:], strict=False):
|
||||
assert a.test_end < b.test_start, "测试区间不得重叠"
|
||||
assert a.test_start > a.train_end, "测试必须晚于训练"
|
||||
for x in wins:
|
||||
assert x.train_start < x.train_end < x.test_start <= x.test_end
|
||||
|
||||
|
||||
def test_walk_forward_freeze_is_enforced_by_config() -> None:
|
||||
"""plan.md §25:测试阶段禁止重新调参 —— 配置层必须拒绝关闭该开关。"""
|
||||
cfg = load_config("backtest")
|
||||
assert cfg.walk_forward.freeze_params_in_test is True
|
||||
|
||||
|
||||
def test_add_months_and_years() -> None:
|
||||
assert _add_months(date(2024, 1, 31), 1) == date(2024, 2, 29), "闰年 2 月"
|
||||
assert _add_months(date(2023, 1, 31), 1) == date(2023, 2, 28)
|
||||
assert _add_months(date(2024, 12, 15), 1) == date(2025, 1, 15)
|
||||
assert _add_years(date(2020, 2, 29), 1) == date(2021, 2, 28)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 引擎端到端(小样本,含数据时执行)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.db
|
||||
def test_engine_end_to_end_reconciliation() -> None:
|
||||
"""完整跑一段回测并验证资金对账必须平衡。"""
|
||||
from hdiv.backtest.engine import BacktestEngine
|
||||
|
||||
db_ok = True
|
||||
try:
|
||||
from hdiv.data import db as _db
|
||||
|
||||
_db.load_dotenv_once()
|
||||
_db.list_tables(load_config("datasource"))
|
||||
except Exception:
|
||||
db_ok = False
|
||||
if not db_ok:
|
||||
pytest.skip("数据库不可用")
|
||||
|
||||
engine = BacktestEngine.from_strategy("config/strategy/high_dividend_v1.yml")
|
||||
res = engine.run(start=date(2023, 1, 1), end=date(2024, 6, 28), persist=False, verbose=False)
|
||||
rc = res["reconciliation"]
|
||||
assert rc["balanced"], f"资金对账不平,残差 {rc['residual']}"
|
||||
assert res["equity"]["nav"].iloc[0] == pytest.approx(1.0)
|
||||
|
||||
eq = res["equity"]
|
||||
# 净值必须等于 总市值 / 期初资金
|
||||
assert (eq["nav"] - eq["total_value"] / res["initial_capital"]).abs().max() < 1e-9
|
||||
# 总市值必须等于 现金 + 持仓
|
||||
assert (eq["total_value"] - (eq["cash"] + eq["position_value"])).abs().max() < 1e-6
|
||||
# 回撤不得为正
|
||||
assert eq["drawdown"].max() <= 1e-9
|
||||
|
||||
|
||||
@pytest.mark.db
|
||||
def test_engine_uses_next_open_no_lookahead() -> None:
|
||||
"""成交日必须晚于信号日(plan.md §6 无未来函数)。"""
|
||||
from hdiv.backtest.engine import BacktestEngine
|
||||
|
||||
try:
|
||||
from hdiv.data import db as _db
|
||||
|
||||
_db.load_dotenv_once()
|
||||
_db.list_tables(load_config("datasource"))
|
||||
except Exception:
|
||||
pytest.skip("数据库不可用")
|
||||
|
||||
engine = BacktestEngine.from_strategy("config/strategy/high_dividend_v1.yml")
|
||||
res = engine.run(start=date(2023, 1, 1), end=date(2024, 6, 28), persist=False, verbose=False)
|
||||
for t in res["trades"]:
|
||||
assert t.execution_date > t.signal_date, (
|
||||
f"{t.symbol} 成交日 {t.execution_date} 未晚于信号日 {t.signal_date}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.db
|
||||
def test_engine_dividends_are_creditable() -> None:
|
||||
"""持有期间应确实收到现金分红(高股息策略的核心收益来源)。"""
|
||||
from hdiv.backtest.engine import BacktestEngine
|
||||
|
||||
try:
|
||||
from hdiv.data import db as _db
|
||||
|
||||
_db.load_dotenv_once()
|
||||
_db.list_tables(load_config("datasource"))
|
||||
except Exception:
|
||||
pytest.skip("数据库不可用")
|
||||
|
||||
engine = BacktestEngine.from_strategy("config/strategy/high_dividend_v1.yml")
|
||||
res = engine.run(start=date(2021, 1, 1), end=date(2024, 6, 28), persist=False, verbose=False)
|
||||
assert res["total_dividend_net"] > 0, "高股息策略在 3.5 年里不可能没有现金分红"
|
||||
d = res["dividends"]
|
||||
assert not d.empty
|
||||
assert (d["net"] <= d["gross"] + 1e-9).all(), "税后不得大于税前"
|
||||
assert (d["tax"] >= 0).all()
|
||||
|
||||
|
||||
@pytest.mark.db
|
||||
def test_engine_pit_discipline_reference_window(engine) -> None:
|
||||
"""rolling 参照窗口必须完全落在评估日之前(无未来函数)。"""
|
||||
for day in (date(2022, 6, 30), date(2024, 1, 15)):
|
||||
ref = engine._reference_window(day)
|
||||
assert ref is not None
|
||||
assert ref[1] == day, "参照窗口右端必须是评估日本身"
|
||||
assert ref[0] < day
|
||||
assert engine._reference_mode() == "rolling"
|
||||
|
||||
|
||||
def test_frozen_reference_overrides_rolling() -> None:
|
||||
"""frozen 模式下必须使用冻结窗口,不得回退到滚动窗口。"""
|
||||
from hdiv.backtest.engine import BacktestEngine
|
||||
|
||||
s = StrategyRegistry().load("config/strategy/high_dividend_v1.yml")
|
||||
frozen = (date(2015, 1, 1), date(2019, 12, 31))
|
||||
e = BacktestEngine(s, frozen_reference=frozen)
|
||||
assert e._reference_window(date(2021, 6, 30)) == frozen
|
||||
assert e._reference_mode() == "frozen"
|
||||
|
||||
|
||||
def test_reference_window_end_is_evaluation_day() -> None:
|
||||
"""rolling 窗口的右边界必须是评估日 —— 否则会用未来数据。"""
|
||||
from hdiv.backtest.engine import BacktestEngine
|
||||
|
||||
s = StrategyRegistry().load("config/strategy/high_dividend_v1.yml")
|
||||
e = BacktestEngine(s)
|
||||
day = date(2023, 5, 10)
|
||||
lo, hi = e._reference_window(day)
|
||||
assert hi == day
|
||||
assert (day - lo).days == int(365.25 * e.bt_cfg.percentile_reference.lookback_years)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 绩效指标
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_compute_metrics_flags_insufficient_data() -> None:
|
||||
from hdiv.analysis.performance import compute_metrics
|
||||
|
||||
m = compute_metrics(pd.DataFrame(), [], load_config("backtest"), "r1")
|
||||
assert m == {}, "空曲线不应编造指标"
|
||||
|
||||
eq = pd.DataFrame({
|
||||
"trade_date": [date(2024, 1, 2)],
|
||||
"total_value": [1_000_000.0],
|
||||
"daily_return": [0.0],
|
||||
"drawdown": [0.0],
|
||||
"cash": [1_000_000.0],
|
||||
"position_value": [0.0],
|
||||
})
|
||||
m2 = compute_metrics(eq, [], load_config("backtest"), "r2")
|
||||
assert m2["total_return"] == pytest.approx(0.0)
|
||||
assert m2["max_drawdown"] == pytest.approx(0.0)
|
||||
|
||||
|
||||
def test_metrics_do_not_invent_values() -> None:
|
||||
"""样本不足时 Sharpe 必须为 None,而不是 0。"""
|
||||
from hdiv.analysis.performance import compute_metrics
|
||||
|
||||
eq = pd.DataFrame({
|
||||
"trade_date": [date(2024, 1, 2), date(2024, 1, 3)],
|
||||
"total_value": [1_000_000.0, 1_010_000.0],
|
||||
"daily_return": [0.0, 0.01],
|
||||
"drawdown": [0.0, 0.0],
|
||||
"cash": [0.0, 0.0],
|
||||
"position_value": [1_000_000.0, 1_010_000.0],
|
||||
})
|
||||
m = compute_metrics(eq, [], load_config("backtest"), "r3")
|
||||
assert m["sharpe"] is None, "1 个观测算不出波动率,Sharpe 必须是 None"
|
||||
@@ -0,0 +1,303 @@
|
||||
"""CLI 接口契约测试。
|
||||
|
||||
**为什么需要**:使用手册(``docs/user-guide.md``)里的命令就是用户看到的接口。
|
||||
文档与实现脱节是极易发生且很难察觉的问题 —— 实测发现手册记载的
|
||||
``hdiv sync financial --interleaved`` 在 CLI 中根本没有暴露,
|
||||
而 ``hdiv ddl verify`` 的输出格式也与手册不符。
|
||||
|
||||
本测试把「手册承诺的接口」变成可执行断言:
|
||||
只要有人删掉一个参数或改掉命令名,测试立刻失败。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from hdiv.cli import build_parser
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def parser():
|
||||
return build_parser()
|
||||
|
||||
|
||||
def _subparsers(p):
|
||||
"""取子命令表。同时支持两种声明方式:
|
||||
|
||||
- ``add_subparsers()`` → choices 是 dict(sync/universe/...)
|
||||
- ``add_argument("action", choices=[...])`` → choices 是 list(ddl/report/strategy)
|
||||
"""
|
||||
for action in p._actions:
|
||||
ch = getattr(action, "choices", None)
|
||||
if not ch:
|
||||
continue
|
||||
if isinstance(ch, dict):
|
||||
return ch
|
||||
raise AssertionError("未找到子命令")
|
||||
|
||||
|
||||
def _action_choices(p) -> set[str]:
|
||||
"""取位置参数(如 ddl 的 plan/apply/verify)的取值集合。"""
|
||||
for action in p._actions:
|
||||
ch = getattr(action, "choices", None)
|
||||
if isinstance(ch, (list, tuple, set)) and not action.option_strings:
|
||||
return set(ch)
|
||||
return set()
|
||||
|
||||
|
||||
def _opts(p) -> set[str]:
|
||||
out: set[str] = set()
|
||||
for a in p._actions:
|
||||
out.update(a.option_strings)
|
||||
return out
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 手册记载的顶层命令必须全部存在
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
DOCUMENTED_COMMANDS = [
|
||||
"ddl", "sync", "audit", "report", "universe",
|
||||
"profile", "strategy", "backtest", "sensitivity",
|
||||
"web", "site",
|
||||
]
|
||||
|
||||
|
||||
def test_all_documented_commands_exist(parser) -> None:
|
||||
subs = _subparsers(parser)
|
||||
missing = [c for c in DOCUMENTED_COMMANDS if c not in subs]
|
||||
assert not missing, f"手册记载但 CLI 不存在的命令:{missing}"
|
||||
|
||||
|
||||
def test_no_undocumented_commands(parser) -> None:
|
||||
"""反向检查:CLI 不应有手册未提及的命令。"""
|
||||
subs = set(_subparsers(parser))
|
||||
extra = subs - set(DOCUMENTED_COMMANDS)
|
||||
assert not extra, f"CLI 存在手册未记载的命令:{sorted(extra)}"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 手册记载的动作/参数必须存在
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"cmd,expected",
|
||||
[
|
||||
("ddl", {"plan", "apply", "verify"}),
|
||||
("report", {"index", "validate"}),
|
||||
("strategy", {"validate", "register", "list", "diff"}),
|
||||
(
|
||||
"sync",
|
||||
{"dividend", "financial", "index", "price", "trading", "backfill"},
|
||||
),
|
||||
("site", {"normalize", "archive", "build", "status"}),
|
||||
],
|
||||
)
|
||||
def test_documented_actions_exist(parser, cmd: str, expected: set[str]) -> None:
|
||||
sub = _subparsers(parser)[cmd]
|
||||
actual = _action_choices(sub)
|
||||
if not actual: # 用 add_subparsers 声明的命令
|
||||
actual = set(_subparsers(sub))
|
||||
missing = expected - actual
|
||||
assert not missing, f"hdiv {cmd} 缺少手册记载的动作:{sorted(missing)}"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"cmd,flags",
|
||||
[
|
||||
("sync", {"--only-missing", "--limit", "--symbols", "--apis",
|
||||
"--interleaved", "--start", "--end", "--no-resume", "--no-weight"}),
|
||||
("universe", {"-c", "--config", "--asof", "--no-persist", "--no-html"}),
|
||||
("profile", {"--universe-run", "--symbols", "--asof", "--html-limit"}),
|
||||
("backtest", {"-s", "--strategy", "--mode", "--start", "--end", "--universe-run"}),
|
||||
("sensitivity", {"-s", "--strategy", "--sweep"}),
|
||||
("strategy", {"-f", "--file", "--other"}),
|
||||
("audit", {"--no-persist", "--no-html"}),
|
||||
("report", {"--dir"}),
|
||||
("web", {"--host", "--port", "--api-only", "--out-dir"}),
|
||||
],
|
||||
)
|
||||
def test_documented_flags_exist(parser, cmd: str, flags: set[str]) -> None:
|
||||
sub = _subparsers(parser)[cmd]
|
||||
have = _opts(sub)
|
||||
missing = flags - have
|
||||
assert not missing, f"hdiv {cmd} 缺少手册记载的参数:{sorted(missing)}"
|
||||
|
||||
|
||||
def test_sync_interleaved_is_a_real_flag(parser) -> None:
|
||||
"""回归:手册要求 ``hdiv sync financial --interleaved``,CLI 必须真的支持。"""
|
||||
sub = _subparsers(parser)["sync"]
|
||||
assert "--interleaved" in _opts(sub)
|
||||
# 且必须真的能解析(不只是声明)
|
||||
args = parser.parse_args(["sync", "financial", "--interleaved"])
|
||||
assert args.interleaved is True
|
||||
args2 = parser.parse_args(["sync", "financial"])
|
||||
assert args2.interleaved is False
|
||||
|
||||
|
||||
def test_backtest_universe_run_flag(parser) -> None:
|
||||
"""股票池 ↔ 回测 的关联入口:``--universe-run`` 必须存在且可解析。"""
|
||||
args = parser.parse_args(["backtest", "--universe-run", "abc123"])
|
||||
assert args.universe_run == "abc123"
|
||||
args2 = parser.parse_args(["backtest"])
|
||||
assert args2.universe_run is None
|
||||
|
||||
|
||||
@pytest.mark.db
|
||||
def test_frozen_universe_is_used_when_run_id_given() -> None:
|
||||
"""指定 --universe-run 时引擎必须使用该股票池,而不是重新筛选。"""
|
||||
from datetime import date
|
||||
|
||||
from hdiv.backtest.engine import BacktestEngine
|
||||
from hdiv.core.config import load_config
|
||||
from hdiv.data import db
|
||||
|
||||
db.load_dotenv_once()
|
||||
try:
|
||||
cfg = load_config("datasource")
|
||||
df = db.read_sql(
|
||||
"SELECT run_id FROM hd_universe_run WHERE deleted_at IS NULL "
|
||||
"ORDER BY created_at DESC LIMIT 1", cfg=cfg,
|
||||
)
|
||||
if df.empty:
|
||||
pytest.skip("没有筛选记录")
|
||||
rid = df["run_id"].iloc[0]
|
||||
except Exception as exc:
|
||||
pytest.skip(f"数据库不可用:{exc}")
|
||||
|
||||
eng = BacktestEngine.from_strategy(
|
||||
"config/strategy/high_dividend_v1.yml", universe_run_id=rid
|
||||
)
|
||||
assert eng.universe_run_id == rid
|
||||
ctx = eng._prepare(eng.repo.trading_days(date(2024, 1, 1), date(2024, 6, 28)),
|
||||
verbose=False)
|
||||
sets = list(ctx["universe_by_refresh"].values())
|
||||
assert sets, "应产生股票池"
|
||||
assert all(x == sets[0] for x in sets), "冻结股票池在各调仓日必须完全相同"
|
||||
|
||||
|
||||
def test_backtest_mode_choices(parser) -> None:
|
||||
"""手册只承诺 single / walkforward 两种模式。"""
|
||||
sub = _subparsers(parser)["backtest"]
|
||||
for a in sub._actions:
|
||||
if "--mode" in a.option_strings:
|
||||
assert set(a.choices) == {"single", "walkforward"}
|
||||
return
|
||||
raise AssertionError("backtest 缺少 --mode 参数")
|
||||
|
||||
|
||||
def test_price_which_choices(parser) -> None:
|
||||
sub = _subparsers(parser)["sync"]
|
||||
for a in sub._actions:
|
||||
if "--which" in a.option_strings:
|
||||
assert set(a.choices) == {"daily", "adj_factor", "daily_basic"}
|
||||
return
|
||||
raise AssertionError("sync 缺少 --which 参数")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 手册记载的输出格式
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_ddl_verify_output_mentions_table_count() -> None:
|
||||
"""手册写的是「OK:30 张 hd_* 表结构全部符合 schema 定义」,输出须含数量。"""
|
||||
import inspect
|
||||
|
||||
from hdiv import cli
|
||||
from hdiv.data.schema import ALL_TABLES
|
||||
|
||||
src = inspect.getsource(cli.cmd_ddl)
|
||||
assert "len(ddl.ALL_TABLES)" in src, "ddl verify 输出应包含表的数量"
|
||||
assert len(ALL_TABLES) == 30
|
||||
|
||||
|
||||
def test_help_text_is_chinese(parser) -> None:
|
||||
"""手册面向中文用户,所有命令都应有中文说明。"""
|
||||
subs = _subparsers(parser)
|
||||
for name in DOCUMENTED_COMMANDS:
|
||||
assert subs[name].description, f"hdiv {name} 缺少说明文字"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 手册文件本身
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_user_guide_exists_and_covers_all_commands() -> None:
|
||||
from hdiv.core.paths import project_root
|
||||
|
||||
guide = project_root() / "docs" / "user-guide.md"
|
||||
assert guide.is_file(), "缺少使用手册 docs/user-guide.md"
|
||||
text = guide.read_text(encoding="utf-8")
|
||||
for cmd in DOCUMENTED_COMMANDS:
|
||||
assert f"hdiv {cmd}" in text, f"手册未记载命令:hdiv {cmd}"
|
||||
# 手册必须如实记录已知限制与部署要求(这两点最容易在文档里被淡化)
|
||||
assert "已知限制" in text
|
||||
assert "echarts.min.js" in text, "手册必须说明图表库的部署依赖"
|
||||
assert "asset_mode" in text, "手册必须说明自包含模式的开关"
|
||||
# 前后端分离的部署要点必须在手册里
|
||||
assert "/api" in text, "手册必须说明 API 反代"
|
||||
assert "nginx" in text, "手册必须给出 nginx 配置指引"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 参数组合与错误呈现
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
# 只有这三个命令同时具备 --no-persist 与 --html(audit 传内存结果,两者可共存)
|
||||
@pytest.mark.parametrize("cmd", ["universe", "backtest"])
|
||||
def test_no_persist_with_html_is_rejected_clearly(cmd: str) -> None:
|
||||
"""回归:--no-persist 与 --html 互相矛盾,必须给出可操作提示。
|
||||
|
||||
静态报告只从数据库读取(保证数字可追溯到 SQL),未落库的运行渲染不出来。
|
||||
早期实现会崩在报告生成器的 ValueError 里,用户看不出是参数冲突。
|
||||
"""
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
r = subprocess.run(
|
||||
[sys.executable, "-m", "hdiv", cmd, "--no-persist", "--html"],
|
||||
capture_output=True, text=True, env={**__import__("os").environ, "PYTHONPATH": "src"},
|
||||
)
|
||||
assert r.returncode == 1
|
||||
out = r.stdout + r.stderr
|
||||
assert "--no-persist 与 --html 不能同时使用" in out, out[:400]
|
||||
# 必须是友好提示,而不是 traceback
|
||||
assert "Traceback" not in out, "用户可理解的错误不应打印调用栈"
|
||||
|
||||
|
||||
def test_audit_can_combine_no_persist_and_html() -> None:
|
||||
"""audit 传的是内存结果而非按 run_id 读库,因此两者可以共存。"""
|
||||
import inspect
|
||||
|
||||
from hdiv import cli
|
||||
|
||||
src = inspect.getsource(cli.cmd_audit)
|
||||
assert "_reject_no_persist_with_html" not in src, \
|
||||
"audit 不需拦截该组合(报告用内存 summary 渲染)"
|
||||
|
||||
|
||||
def test_hdiv_error_is_caught_without_traceback() -> None:
|
||||
"""main() 应优雅处理 HdivError(数据缺失/配置冲突等),不打印调用栈。"""
|
||||
import inspect
|
||||
|
||||
from hdiv import cli
|
||||
|
||||
src = inspect.getsource(cli.main)
|
||||
assert "except HdivError" in src, "main() 未捕获 HdivError"
|
||||
branch = src.split("except HdivError")[1]
|
||||
assert "print_exc" not in branch, "HdivError 分支不应打印调用栈"
|
||||
assert "return 1" in branch, "应返回非零退出码"
|
||||
|
||||
|
||||
def test_html_flag_exists_on_all_report_producing_commands() -> None:
|
||||
"""所有会产出报告的命令都应有 --html 开关且默认关闭。"""
|
||||
from hdiv.cli import build_parser
|
||||
|
||||
p = build_parser()
|
||||
for cmd in ("universe", "profile", "backtest", "audit", "sensitivity"):
|
||||
args = p.parse_args([cmd])
|
||||
assert getattr(args, "html", None) is False, f"hdiv {cmd} --html 默认应为 False"
|
||||
@@ -0,0 +1,261 @@
|
||||
"""配置加载与严格校验测试(对应 development-plan.md P3 验收:非法配置 100% 被拒)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
|
||||
import pytest
|
||||
import yaml
|
||||
|
||||
from hdiv.core.config import (
|
||||
config_hash,
|
||||
load_all,
|
||||
load_config,
|
||||
resolve_strategy_path,
|
||||
)
|
||||
from hdiv.core.errors import ConfigError, ConfigNotFound
|
||||
from hdiv.core.paths import config_dir
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 正向:全部配置文件可加载
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"name", ["datasource", "universe", "profile", "cost", "backtest", "report"]
|
||||
)
|
||||
def test_all_configs_load(name: str) -> None:
|
||||
cfg = load_config(name)
|
||||
assert cfg is not None
|
||||
assert config_hash(cfg)
|
||||
|
||||
|
||||
def test_load_all() -> None:
|
||||
cfgs = load_all()
|
||||
assert set(cfgs) == {"datasource", "universe", "profile", "cost", "backtest", "report"}
|
||||
|
||||
|
||||
def test_strategy_config_loads() -> None:
|
||||
s = load_config("strategy:high_dividend_v1.yml")
|
||||
assert s.strategy.id == "HD_MR_V1"
|
||||
assert s.entry.yield_percentile == 75
|
||||
assert s.exit.yield_percentile == 25
|
||||
# 参数扁平化用于 hd_strategy_param
|
||||
pm = s.param_map()
|
||||
assert pm["strategy.id"] == "HD_MR_V1"
|
||||
assert pm["entry.yield_percentile"] == 75
|
||||
assert len(pm) > 20
|
||||
|
||||
|
||||
def test_strategy_path_resolution_variants() -> None:
|
||||
names = [
|
||||
"high_dividend_v1.yml",
|
||||
"strategy/high_dividend_v1.yml",
|
||||
"config/strategy/high_dividend_v1.yml",
|
||||
]
|
||||
paths = {resolve_strategy_path(n) for n in names}
|
||||
assert len(paths) == 1
|
||||
assert paths.pop().is_file()
|
||||
|
||||
|
||||
def test_resolve_strategy_path_missing() -> None:
|
||||
with pytest.raises(ConfigNotFound):
|
||||
resolve_strategy_path("不存在.yml")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 反向:拼写错误必须报错(不静默取默认值)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _load_modified(name: str, mutate) -> None:
|
||||
"""把配置改坏后写入临时文件并加载,断言抛错。"""
|
||||
raw = yaml.safe_load((config_dir() / f"{name}.yml").read_text(encoding="utf-8"))
|
||||
mutate(raw)
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
with tempfile.NamedTemporaryFile("w", suffix=".yml", delete=False, encoding="utf-8") as fh:
|
||||
yaml.safe_dump(raw, fh, allow_unicode=True)
|
||||
tmp = Path(fh.name)
|
||||
try:
|
||||
with pytest.raises(ConfigError):
|
||||
load_config(name, path=tmp)
|
||||
finally:
|
||||
tmp.unlink(missing_ok=True)
|
||||
|
||||
|
||||
def test_unknown_field_rejected() -> None:
|
||||
_load_modified("universe", lambda r: r["market"].update({"min_market_capp": 1}))
|
||||
|
||||
|
||||
def test_wrong_type_rejected() -> None:
|
||||
_load_modified("universe", lambda r: r["market"].update({"min_listing_years": "十年"}))
|
||||
|
||||
|
||||
def test_unknown_top_level_section_rejected() -> None:
|
||||
_load_modified("universe", lambda r: r.update({"bogus_section": {"a": 1}}))
|
||||
|
||||
|
||||
def test_bad_enum_rejected() -> None:
|
||||
_load_modified("cost", lambda r: r["slippage"].update({"mode": "percentage"}))
|
||||
|
||||
|
||||
def test_dividend_years_consistency_rejected() -> None:
|
||||
def mutate(r: dict) -> None:
|
||||
r["dividend"]["min_continuous_years"] = 10
|
||||
r["dividend"]["window_years"] = 6
|
||||
|
||||
_load_modified("universe", mutate)
|
||||
|
||||
|
||||
def test_negative_capital_rejected() -> None:
|
||||
_load_modified("backtest", lambda r: r["capital"].update({"initial": -1}))
|
||||
|
||||
|
||||
def test_period_order_rejected() -> None:
|
||||
_load_modified("backtest", lambda r: r["period"].update({"end": "2010-01-01"}))
|
||||
|
||||
|
||||
def test_freeze_params_disabled_rejected() -> None:
|
||||
"""plan.md §25:测试阶段禁止重新调参 —— 关闭该开关必须被拒绝。"""
|
||||
_load_modified(
|
||||
"backtest", lambda r: r["walk_forward"].update({"freeze_params_in_test": False})
|
||||
)
|
||||
|
||||
|
||||
def test_duplicate_benchmark_rejected() -> None:
|
||||
def mutate(r: dict) -> None:
|
||||
r["benchmark"] = [
|
||||
{"code": "000300.SH", "name": "沪深300"},
|
||||
{"code": "000300.SH", "name": "重复"},
|
||||
]
|
||||
|
||||
_load_modified("backtest", mutate)
|
||||
|
||||
|
||||
def test_profile_percentiles_must_be_sorted() -> None:
|
||||
_load_modified("profile", lambda r: r.update({"percentiles": [90, 10, 50]}))
|
||||
|
||||
|
||||
def test_profile_duplicate_metric_rejected() -> None:
|
||||
_load_modified(
|
||||
"profile", lambda r: r["metrics"].update({"valuation": ["pb", "pb"]})
|
||||
)
|
||||
|
||||
|
||||
def test_composite_weights_must_sum_to_one() -> None:
|
||||
def mutate(r: dict) -> None:
|
||||
r["safety_margin"]["mode"] = "composite"
|
||||
r["safety_margin"]["weights"] = {"dividend_yield": 0.5, "valuation": 0.2}
|
||||
|
||||
_load_modified("profile", mutate)
|
||||
|
||||
|
||||
def test_datasource_readonly_writeoverlap_rejected() -> None:
|
||||
def mutate(r: dict) -> None:
|
||||
r["database"]["allow_write_tables"] = ["stock"]
|
||||
|
||||
_load_modified("datasource", mutate)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 策略语义校验
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _load_strategy_mutated(mutate) -> None:
|
||||
raw = yaml.safe_load(
|
||||
(config_dir() / "strategy" / "high_dividend_v1.yml").read_text(encoding="utf-8")
|
||||
)
|
||||
mutate(raw)
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
with tempfile.NamedTemporaryFile("w", suffix=".yml", delete=False, encoding="utf-8") as fh:
|
||||
yaml.safe_dump(raw, fh, allow_unicode=True)
|
||||
tmp = Path(fh.name)
|
||||
try:
|
||||
with pytest.raises(ConfigError):
|
||||
load_config(f"strategy:{tmp}")
|
||||
finally:
|
||||
tmp.unlink(missing_ok=True)
|
||||
|
||||
|
||||
def test_strategy_entry_must_exceed_exit() -> None:
|
||||
_load_strategy_mutated(lambda r: r["entry"].update({"yield_percentile": 20}))
|
||||
|
||||
|
||||
def test_strategy_entry_percentile_range() -> None:
|
||||
_load_strategy_mutated(lambda r: r["entry"].update({"yield_percentile": 150}))
|
||||
|
||||
|
||||
def test_strategy_scale_in_must_be_ascending() -> None:
|
||||
def mutate(r: dict) -> None:
|
||||
r["entry"]["scale_in"] = [
|
||||
{"percentile": 80, "weight": 0.25},
|
||||
{"percentile": 75, "weight": 0.50},
|
||||
]
|
||||
|
||||
_load_strategy_mutated(mutate)
|
||||
|
||||
|
||||
def test_strategy_scale_out_must_descend() -> None:
|
||||
def mutate(r: dict) -> None:
|
||||
r["exit"]["scale_out"] = [
|
||||
{"percentile": 25, "weight": 0.50},
|
||||
{"percentile": 50, "weight": 0.00},
|
||||
]
|
||||
|
||||
_load_strategy_mutated(mutate)
|
||||
|
||||
|
||||
def test_strategy_scale_out_last_must_be_zero() -> None:
|
||||
def mutate(r: dict) -> None:
|
||||
r["exit"]["scale_out"] = [
|
||||
{"percentile": 50, "weight": 0.50},
|
||||
{"percentile": 25, "weight": 0.30},
|
||||
]
|
||||
|
||||
_load_strategy_mutated(mutate)
|
||||
|
||||
|
||||
def test_strategy_position_bounds() -> None:
|
||||
_load_strategy_mutated(lambda r: r["position"].update({"max_position": 1.5}))
|
||||
|
||||
|
||||
def test_strategy_bad_status() -> None:
|
||||
_load_strategy_mutated(lambda r: r["strategy"].update({"status": "RUNNING"}))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 可复现性:config_hash 稳定性
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_config_hash_is_stable_and_sensitive() -> None:
|
||||
a = load_config("universe")
|
||||
b = load_config("universe")
|
||||
assert config_hash(a) == config_hash(b)
|
||||
|
||||
raw = yaml.safe_load((config_dir() / "universe.yml").read_text(encoding="utf-8"))
|
||||
changed = copy.deepcopy(raw)
|
||||
changed["dividend"]["min_dividend_yield"] = 0.035
|
||||
assert config_hash(raw) != config_hash(changed)
|
||||
|
||||
|
||||
def test_config_hash_ignores_key_order() -> None:
|
||||
raw = yaml.safe_load((config_dir() / "universe.yml").read_text(encoding="utf-8"))
|
||||
reordered = {k: raw[k] for k in reversed(list(raw))}
|
||||
assert config_hash(raw) == config_hash(reordered)
|
||||
|
||||
|
||||
def test_missing_config_raises() -> None:
|
||||
with pytest.raises(ConfigNotFound):
|
||||
load_config("universe", path="/nonexistent/nope.yml")
|
||||
|
||||
|
||||
def test_unknown_config_name_raises() -> None:
|
||||
with pytest.raises(ConfigError):
|
||||
load_config("no_such_config")
|
||||
@@ -0,0 +1,263 @@
|
||||
"""数据安全硬约束测试(对应 development-plan.md §7.2 / §10 纪律 1)。
|
||||
|
||||
这些测试是「禁止删除数据」「只写自有前缀」的**执行保证**,
|
||||
不是靠自觉,而是靠断言。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from hdiv.core.config import load_config
|
||||
from hdiv.core.errors import SafetyViolation
|
||||
from hdiv.core.paths import project_root
|
||||
from hdiv.data.db import StatementGuard, classify_statement, split_statements
|
||||
|
||||
CFG = load_config("datasource")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def guard() -> StatementGuard:
|
||||
return StatementGuard(CFG, allow_backfill=False)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 语句分类
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"sql,verb",
|
||||
[
|
||||
("SELECT 1", "select"),
|
||||
("select * from stock", "select"),
|
||||
("WITH x AS (SELECT 1) SELECT * FROM x", "with"),
|
||||
("INSERT INTO hd_x VALUES (1)", "insert"),
|
||||
("insert ignore into hd_x values (1)", "insert"),
|
||||
("REPLACE INTO hd_x VALUES (1)", "replace"),
|
||||
("UPDATE hd_x SET a=1", "update"),
|
||||
("DELETE FROM hd_x", "delete"),
|
||||
("CREATE TABLE hd_x (a INT)", "create"),
|
||||
("ALTER TABLE hd_x ADD COLUMN b INT", "alter"),
|
||||
("DROP TABLE hd_x", "drop"),
|
||||
("TRUNCATE TABLE hd_x", "truncate"),
|
||||
("-- 注释\nSELECT 1", "select"),
|
||||
("/* 块注释 */ SELECT 1", "select"),
|
||||
],
|
||||
)
|
||||
def test_classify(sql: str, verb: str) -> None:
|
||||
assert classify_statement(sql)[0] == verb
|
||||
|
||||
|
||||
def test_split_statements() -> None:
|
||||
assert split_statements("SELECT 1; SELECT 2;") == ["SELECT 1", "SELECT 2"]
|
||||
assert split_statements("SELECT 1;;") == ["SELECT 1"]
|
||||
|
||||
|
||||
def test_classify_extracts_table() -> None:
|
||||
assert classify_statement("INSERT INTO `hd_dividend` (a) VALUES (1)")[1] == "hd_dividend"
|
||||
assert classify_statement("UPDATE stock SET a=1")[1] == "stock"
|
||||
assert classify_statement("DELETE FROM hd_x WHERE 1")[1] == "hd_x"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 禁止删除:任何形式的 DELETE / DROP / TRUNCATE 都必须被拒
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"sql",
|
||||
[
|
||||
"DELETE FROM stock",
|
||||
"DELETE FROM stock WHERE symbol='000001.SZ'",
|
||||
"delete from hd_dividend where id=1",
|
||||
"TRUNCATE TABLE stock",
|
||||
"TRUNCATE stock",
|
||||
"DROP TABLE hd_x",
|
||||
"DROP TABLE IF EXISTS hd_x",
|
||||
"RENAME TABLE hd_x TO hd_y",
|
||||
"DELETE FROM stock_daily",
|
||||
"SELECT 1; DELETE FROM stock",
|
||||
"UPDATE hd_x SET a=1; DROP TABLE hd_y",
|
||||
],
|
||||
)
|
||||
def test_delete_like_always_blocked(guard: StatementGuard, sql: str) -> None:
|
||||
with pytest.raises(SafetyViolation):
|
||||
guard.check_script(sql)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 只读白名单:qlib 既有表不可被写
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"sql",
|
||||
[
|
||||
"UPDATE stock SET name='x' WHERE symbol='000001.SZ'",
|
||||
"INSERT INTO stock_daily (symbol) VALUES ('000001.SZ')",
|
||||
"UPDATE daily_basic SET pe=1",
|
||||
"INSERT INTO financial_indicator (symbol) VALUES ('x')",
|
||||
"UPDATE trading_calendar SET is_open=0",
|
||||
"UPDATE stock_name_history SET name='x'",
|
||||
],
|
||||
)
|
||||
def test_readonly_tables_blocked(guard: StatementGuard, sql: str) -> None:
|
||||
with pytest.raises(SafetyViolation):
|
||||
guard.check_script(sql)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"sql",
|
||||
[
|
||||
"SELECT * FROM stock",
|
||||
"SELECT * FROM stock_daily WHERE trade_date > '2020-01-01'",
|
||||
"WITH d AS (SELECT * FROM daily_basic) SELECT COUNT(*) FROM d",
|
||||
],
|
||||
)
|
||||
def test_readonly_tables_readable(guard: StatementGuard, sql: str) -> None:
|
||||
guard.check_script(sql) # 不应抛错
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 自有前缀:只能写 hd_* 表
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"sql",
|
||||
[
|
||||
"INSERT INTO hd_dividend (symbol) VALUES ('x')",
|
||||
"UPDATE hd_strategy SET version='2'",
|
||||
"REPLACE INTO hd_sync_log (job) VALUES ('x')",
|
||||
],
|
||||
)
|
||||
def test_own_prefix_writes_allowed(guard: StatementGuard, sql: str) -> None:
|
||||
guard.check_script(sql)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"sql",
|
||||
[
|
||||
"CREATE TABLE foo (a INT)",
|
||||
"CREATE TABLE tmp_data (a INT)",
|
||||
"ALTER TABLE stock_extra ADD COLUMN a INT",
|
||||
"INSERT INTO random_table (a) VALUES (1)",
|
||||
],
|
||||
)
|
||||
def test_foreign_table_writes_blocked(guard: StatementGuard, sql: str) -> None:
|
||||
with pytest.raises(SafetyViolation):
|
||||
guard.check_script(sql)
|
||||
|
||||
|
||||
def test_own_prefix_create_allowed(guard: StatementGuard) -> None:
|
||||
guard.check_script("CREATE TABLE IF NOT EXISTS hd_new_table (id INT)")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 禁触表
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_alembic_version_protected(guard: StatementGuard) -> None:
|
||||
with pytest.raises(SafetyViolation):
|
||||
guard.check_script("UPDATE alembic_version SET version_num='x'")
|
||||
with pytest.raises(SafetyViolation):
|
||||
guard.check_script("ALTER TABLE alembic_version ADD COLUMN a INT")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 回补通道:默认关闭,显式开启后只允许白名单表
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_backfill_requires_explicit_optin() -> None:
|
||||
g = StatementGuard(CFG, allow_backfill=False)
|
||||
with pytest.raises(SafetyViolation):
|
||||
g.check_script("INSERT INTO stock_daily (symbol) VALUES ('x')")
|
||||
|
||||
|
||||
def test_backfill_allows_only_declared_tables() -> None:
|
||||
g = StatementGuard(CFG, allow_backfill=True)
|
||||
# 已声明的表放行
|
||||
g.check_script("INSERT IGNORE INTO stock_daily (symbol) VALUES ('x')")
|
||||
g.check_script("INSERT IGNORE INTO daily_basic (symbol) VALUES ('x')")
|
||||
# 未声明的表仍被拒
|
||||
with pytest.raises(SafetyViolation):
|
||||
g.check_script("INSERT INTO financial_indicator (symbol) VALUES ('x')")
|
||||
|
||||
|
||||
def test_backfill_still_forbids_delete() -> None:
|
||||
g = StatementGuard(CFG, allow_backfill=True)
|
||||
with pytest.raises(SafetyViolation):
|
||||
g.check_script("DELETE FROM stock_daily")
|
||||
|
||||
|
||||
def test_index_weight_is_writable() -> None:
|
||||
"""index_weight 是 qlib 的空表,经 allow_write_tables 允许纯新增填充。"""
|
||||
g = StatementGuard(CFG, allow_backfill=False)
|
||||
g.check_script("INSERT IGNORE INTO index_weight (index_code, trade_date, symbol) VALUES ('a','2020-01-01','b')")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 源码级扫描:代码库里不得出现删除语句
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_FORBIDDEN_PATTERNS = [
|
||||
(re.compile(r"\bDELETE\s+FROM\b", re.IGNORECASE), "DELETE FROM"),
|
||||
(re.compile(r"\bTRUNCATE\s+TABLE\b", re.IGNORECASE), "TRUNCATE TABLE"),
|
||||
(re.compile(r"\bDROP\s+TABLE\b", re.IGNORECASE), "DROP TABLE"),
|
||||
(re.compile(r"\.delete\s*\(", re.IGNORECASE), "session.delete("),
|
||||
(re.compile(r"\bDROP\s+INDEX\b(?!.*migrations)", re.IGNORECASE), "DROP INDEX"),
|
||||
]
|
||||
|
||||
# 允许出现这些词的文件(安全层自身、迁移定义、测试、文档)
|
||||
_ALLOWED_FILES = {
|
||||
"src/hdiv/data/db.py",
|
||||
"src/hdiv/data/schema.py",
|
||||
"src/hdiv/data/ddl.py",
|
||||
"src/hdiv/data/audit.py",
|
||||
"src/hdiv/data/sync/price.py",
|
||||
}
|
||||
|
||||
|
||||
def test_no_delete_statements_in_source() -> None:
|
||||
root = project_root() / "src"
|
||||
offenders: list[str] = []
|
||||
for py in root.rglob("*.py"):
|
||||
rel = py.relative_to(project_root()).as_posix()
|
||||
if rel in _ALLOWED_FILES:
|
||||
continue
|
||||
text = py.read_text(encoding="utf-8")
|
||||
for pat, label in _FORBIDDEN_PATTERNS:
|
||||
for m in pat.finditer(text):
|
||||
line = text[: m.start()].count("\n") + 1
|
||||
offenders.append(f"{rel}:{line} 含 {label}")
|
||||
assert not offenders, "源码中出现删除类语句:\n" + "\n".join(offenders)
|
||||
|
||||
|
||||
def test_no_qlib_code_import() -> None:
|
||||
"""决策 D5:不 import qlib 项目代码,只共享数据库。"""
|
||||
root = project_root() / "src"
|
||||
bad: list[str] = []
|
||||
pattern = re.compile(r"^\s*(?:from|import)\s+(app\.|qlib\.)", re.MULTILINE)
|
||||
for py in root.rglob("*.py"):
|
||||
text = py.read_text(encoding="utf-8")
|
||||
for m in pattern.finditer(text):
|
||||
line = text[: m.start()].count("\n") + 1
|
||||
bad.append(f"{py.relative_to(project_root())}:{line}")
|
||||
assert not bad, f"禁止 import qlib 项目代码:{bad}"
|
||||
|
||||
|
||||
def test_guard_readonly_lists_are_configured() -> None:
|
||||
"""白名单/前缀必须在配置里显式声明,不能靠默认值。"""
|
||||
db = CFG.database
|
||||
assert db.own_prefix == "hd_"
|
||||
assert "stock" in db.read_only_tables
|
||||
assert "stock_daily" in db.read_only_tables
|
||||
assert "alembic_version" in db.forbidden_tables
|
||||
assert db.forbid_delete is True
|
||||
assert Path(db.own_prefix) # 非空
|
||||
@@ -0,0 +1,203 @@
|
||||
"""表结构与幂等性测试(对应 development-plan.md §7.2 硬约束)。
|
||||
|
||||
需要真实数据库的用例标记为 ``db``;纯结构校验的用例无需数据库。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
|
||||
import pytest
|
||||
|
||||
from hdiv.core.config import load_config
|
||||
from hdiv.core.errors import SafetyViolation
|
||||
from hdiv.data.schema import ALL_TABLES, TABLE_NAMES
|
||||
|
||||
CFG = load_config("datasource")
|
||||
|
||||
# 需要库的表清单(无库时跳过)
|
||||
_PLAIN_TESTS_NEED_DB = pytest.mark.db
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 纯结构校验(无需数据库)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_table_count() -> None:
|
||||
assert len(ALL_TABLES) == 30
|
||||
assert len(set(TABLE_NAMES)) == 30
|
||||
|
||||
|
||||
def test_all_tables_use_own_prefix() -> None:
|
||||
bad = [t.name for t in ALL_TABLES if not t.name.startswith(CFG.database.own_prefix)]
|
||||
assert not bad, f"以下表未使用自有前缀:{bad}"
|
||||
|
||||
|
||||
def test_ddl_never_contains_destructive_statements() -> None:
|
||||
destructive = re.compile(r"\b(DROP|TRUNCATE|DELETE)\b", re.IGNORECASE)
|
||||
for t in ALL_TABLES:
|
||||
assert not destructive.search(t.ddl), f"{t.name} 的 DDL 含破坏性语句"
|
||||
for m in t.migrations:
|
||||
assert not re.search(r"\b(DROP\s+TABLE|TRUNCATE|DELETE)\b", m.sql, re.IGNORECASE), (
|
||||
f"{t.name} 迁移 {m.name} 含破坏性语句:{m.sql}"
|
||||
)
|
||||
|
||||
|
||||
def test_create_statements_are_idempotent() -> None:
|
||||
for t in ALL_TABLES:
|
||||
assert "CREATE TABLE IF NOT EXISTS" in t.ddl, f"{t.name} 未使用 IF NOT EXISTS"
|
||||
|
||||
|
||||
def test_every_table_has_primary_key() -> None:
|
||||
for t in ALL_TABLES:
|
||||
assert "PRIMARY KEY" in t.ddl, f"{t.name} 缺少主键"
|
||||
|
||||
|
||||
def test_tables_have_comment() -> None:
|
||||
for t in ALL_TABLES:
|
||||
assert t.comment, f"{t.name} 缺少注释"
|
||||
assert "COMMENT='" in t.ddl or 'COMMENT="' in t.ddl, f"{t.name} 的 DDL 缺少表注释"
|
||||
|
||||
|
||||
def test_no_nullable_column_in_unique_key() -> None:
|
||||
"""MySQL 唯一约束不约束 NULL —— 参与唯一键的列必须 NOT NULL。
|
||||
|
||||
唯一例外是显式构造的 dedup_key(本身 NOT NULL)。
|
||||
"""
|
||||
problems: list[str] = []
|
||||
for t in ALL_TABLES:
|
||||
for m in re.finditer(r"UNIQUE KEY `\w+` \(([^)]+)\)", t.ddl):
|
||||
cols = [c.strip().strip("`") for c in m.group(1).split(",")]
|
||||
for col in cols:
|
||||
# 找到该列的定义行
|
||||
pat = re.compile(rf"^\s*`{re.escape(col)}`\s+([A-Za-z]+(?:\([^)]*\))?)(.*)$", re.M)
|
||||
mm = pat.search(t.ddl)
|
||||
if not mm:
|
||||
problems.append(f"{t.name}.{col} 未找到列定义")
|
||||
continue
|
||||
rest = mm.group(2)
|
||||
if "NOT NULL" not in rest:
|
||||
problems.append(
|
||||
f"{t.name}.{col} 可空却参与唯一键 —— NULL 会绕过唯一约束造成重复"
|
||||
)
|
||||
assert not problems, "\n".join(problems)
|
||||
|
||||
|
||||
def test_character_set_is_utf8mb4() -> None:
|
||||
for t in ALL_TABLES:
|
||||
assert "utf8mb4" in t.ddl, f"{t.name} 未使用 utf8mb4"
|
||||
|
||||
|
||||
def test_engine_is_innodb() -> None:
|
||||
for t in ALL_TABLES:
|
||||
assert "InnoDB" in t.ddl, f"{t.name} 未使用 InnoDB"
|
||||
|
||||
|
||||
def test_migration_names_are_unique_within_table() -> None:
|
||||
for t in ALL_TABLES:
|
||||
names = [m.name for m in t.migrations]
|
||||
assert len(names) == len(set(names)), f"{t.name} 迁移名重复:{names}"
|
||||
|
||||
|
||||
def test_migrations_do_not_add_dropped_tables() -> None:
|
||||
"""迁移只允许 ADD/MODIFY,不允许改表名或删列。"""
|
||||
for t in ALL_TABLES:
|
||||
for m in t.migrations:
|
||||
assert not re.search(r"\bDROP\s+COLUMN\b", m.sql, re.IGNORECASE), m.sql
|
||||
assert not re.search(r"\bRENAME\b", m.sql, re.IGNORECASE), m.sql
|
||||
|
||||
|
||||
def test_own_prefix_guard_rejects_foreign_table() -> None:
|
||||
from hdiv.data.ddl import _assert_own_prefix
|
||||
|
||||
with pytest.raises(SafetyViolation):
|
||||
_assert_own_prefix(CFG, "stock")
|
||||
with pytest.raises(SafetyViolation):
|
||||
_assert_own_prefix(CFG, "hd_ok") if False else _assert_own_prefix(CFG, "alembic_version")
|
||||
|
||||
|
||||
def test_own_prefix_guard_accepts_own_table() -> None:
|
||||
from hdiv.data.ddl import _assert_own_prefix
|
||||
|
||||
_assert_own_prefix(CFG, "hd_dividend") # 不应抛错
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 需要数据库
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@_PLAIN_TESTS_NEED_DB
|
||||
def test_ddl_plan_is_idempotent_on_live_db() -> None:
|
||||
"""连续两次 plan 都应返回 0 个动作(结构已就绪且迁移已应用)。"""
|
||||
pytest.importorskip("pymysql")
|
||||
from hdiv.data import ddl, db
|
||||
|
||||
db.load_dotenv_once()
|
||||
try:
|
||||
db.list_tables(CFG)
|
||||
except Exception as exc: # pragma: no cover - 环境相关
|
||||
pytest.skip(f"数据库不可用:{exc}")
|
||||
|
||||
actions = ddl.plan(CFG)
|
||||
assert actions == [], f"结构未收敛,仍待执行:{[(a.kind, a.table, a.detail) for a in actions]}"
|
||||
|
||||
|
||||
@_PLAIN_TESTS_NEED_DB
|
||||
def test_all_declared_tables_exist() -> None:
|
||||
pytest.importorskip("pymysql")
|
||||
from hdiv.data import db
|
||||
|
||||
db.load_dotenv_once()
|
||||
try:
|
||||
existing = set(db.list_tables(CFG))
|
||||
except Exception as exc: # pragma: no cover
|
||||
pytest.skip(f"数据库不可用:{exc}")
|
||||
missing = [t for t in TABLE_NAMES if t not in existing]
|
||||
assert not missing, f"库中缺表:{missing}"
|
||||
|
||||
|
||||
@_PLAIN_TESTS_NEED_DB
|
||||
def test_verify_reports_no_problems() -> None:
|
||||
pytest.importorskip("pymysql")
|
||||
from hdiv.data import db, ddl
|
||||
|
||||
db.load_dotenv_once()
|
||||
try:
|
||||
problems = ddl.verify(CFG)
|
||||
except Exception as exc: # pragma: no cover
|
||||
pytest.skip(f"数据库不可用:{exc}")
|
||||
assert not problems, "\n".join(problems)
|
||||
|
||||
|
||||
@_PLAIN_TESTS_NEED_DB
|
||||
def test_alembic_version_untouched() -> None:
|
||||
"""决策 D1:本项目不得修改 qlib 的 Alembic 版本链。"""
|
||||
pytest.importorskip("pymysql")
|
||||
from hdiv.data import db
|
||||
|
||||
db.load_dotenv_once()
|
||||
try:
|
||||
df = db.read_sql("SELECT version_num FROM alembic_version", cfg=CFG)
|
||||
except Exception as exc: # pragma: no cover
|
||||
pytest.skip(f"数据库不可用:{exc}")
|
||||
# 只要存在且可读即可;本项目从未写入该表(由 SafetyViolation 保证)
|
||||
assert len(df) >= 1
|
||||
|
||||
|
||||
@_PLAIN_TESTS_NEED_DB
|
||||
def test_project_created_tables_all_have_own_prefix() -> None:
|
||||
"""本项目建的表必须 100% 匹配 hd_ 前缀。"""
|
||||
pytest.importorskip("pymysql")
|
||||
from hdiv.data import db
|
||||
from hdiv.data.schema import TABLE_NAMES as declared
|
||||
|
||||
db.load_dotenv_once()
|
||||
try:
|
||||
existing = set(db.list_tables(CFG))
|
||||
except Exception as exc: # pragma: no cover
|
||||
pytest.skip(f"数据库不可用:{exc}")
|
||||
hd_tables = {t for t in existing if t.startswith(CFG.database.own_prefix)}
|
||||
assert hd_tables <= set(declared), f"库中存在未声明的 hd_ 表:{hd_tables - set(declared)}"
|
||||
assert set(declared) <= existing, f"声明的表未全部创建:{set(declared) - existing}"
|
||||
@@ -0,0 +1,264 @@
|
||||
"""同步层测试:值转换、去重键、数据帧构造、限频器。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, datetime
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
from hdiv.data.db import _frame_to_records
|
||||
from hdiv.data.sync.base import nullify_zero, stable_id, to_date, to_float
|
||||
from hdiv.data.tushare_client import RateLimiter
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 值转换
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"raw,expected",
|
||||
[
|
||||
("20240102", date(2024, 1, 2)),
|
||||
("19991231", date(1999, 12, 31)),
|
||||
(date(2020, 5, 6), date(2020, 5, 6)),
|
||||
(datetime(2020, 5, 6, 1, 2, 3), date(2020, 5, 6)),
|
||||
("2020-05-06", date(2020, 5, 6)),
|
||||
(None, None),
|
||||
("", None),
|
||||
("nan", None),
|
||||
(float("nan"), None),
|
||||
],
|
||||
)
|
||||
def test_to_date(raw, expected) -> None:
|
||||
assert to_date(raw) == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"raw,expected",
|
||||
[(1, 1.0), ("2.5", 2.5), (None, None), ("", None), ("abc", None), (float("nan"), None)],
|
||||
)
|
||||
def test_to_float(raw, expected) -> None:
|
||||
assert to_float(raw) == expected
|
||||
|
||||
|
||||
def test_nullify_zero_treats_zero_as_missing() -> None:
|
||||
"""Tushare 用 0 表示「未实施」,需转 None 以免污染统计。"""
|
||||
assert nullify_zero(0) is None
|
||||
assert nullify_zero("0") is None
|
||||
assert nullify_zero(1.5) == 1.5
|
||||
assert nullify_zero(None) is None
|
||||
|
||||
|
||||
def test_stable_id_is_deterministic() -> None:
|
||||
a = stable_id("dividend", "600036.SH", "2020-01-01")
|
||||
b = stable_id("dividend", "600036.SH", "2020-01-01")
|
||||
c = stable_id("dividend", "600036.SH", "2020-01-02")
|
||||
assert a == b != c
|
||||
assert len(a) == 32 # 默认 length=32
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# NaN → SQL NULL(曾导致 "nan can not be used with MySQL" 的真实 bug)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_frame_to_records_converts_nan_to_none() -> None:
|
||||
df = pd.DataFrame(
|
||||
{
|
||||
"a": [1.0, float("nan"), 3.0],
|
||||
"b": ["x", None, "z"],
|
||||
"c": [pd.Timestamp("2024-01-01"), pd.NaT, pd.Timestamp("2024-01-03")],
|
||||
"d": pd.Series([1, 2, 3], dtype="Int64"),
|
||||
}
|
||||
)
|
||||
recs = _frame_to_records(df)
|
||||
assert recs[0]["a"] == 1.0
|
||||
assert recs[1]["a"] is None, "NaN 必须转成 None,否则 pymysql 报错"
|
||||
assert recs[1]["b"] is None
|
||||
assert recs[1]["c"] is None, "NaT 必须转成 None"
|
||||
# object dtype 保证 None 不被强制回 NaN
|
||||
assert all(not (isinstance(v, float) and v != v) for r in recs for v in r.values())
|
||||
|
||||
|
||||
def test_frame_to_records_preserves_dates_and_strings() -> None:
|
||||
df = pd.DataFrame({"d": [date(2024, 1, 2)], "s": ["银行"], "n": [1.5]})
|
||||
rec = _frame_to_records(df)[0]
|
||||
assert rec["d"] == date(2024, 1, 2)
|
||||
assert rec["s"] == "银行"
|
||||
assert rec["n"] == 1.5
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 分红去重键
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_dividend_dedup_key_handles_null_ann_date() -> None:
|
||||
"""ann_date 为 NULL 时必须有稳定的显式键 —— MySQL 唯一约束不约束 NULL。"""
|
||||
from hdiv.data.sync.dividend import build_dedup_key
|
||||
|
||||
k1 = build_dedup_key("600036.SH", "20260630", "股东提议", None)
|
||||
k2 = build_dedup_key("600036.SH", "20260630", "股东提议", None)
|
||||
assert k1 == k2, "同样的 NULL 输入必须产生同样的键,否则会重复入库"
|
||||
assert "NONE" in k1
|
||||
assert k1 != build_dedup_key("600036.SH", "20260630", "预案", None)
|
||||
assert k1 != build_dedup_key("600036.SH", "20250630", "股东提议", None)
|
||||
|
||||
|
||||
def test_dividend_dedup_key_distinguishes_ann_date() -> None:
|
||||
from hdiv.data.sync.dividend import build_dedup_key
|
||||
|
||||
a = build_dedup_key("600036.SH", "20251231", "实施", "20260328")
|
||||
b = build_dedup_key("600036.SH", "20251231", "实施", "20260626")
|
||||
assert a != b, "同一报告期的多次公告必须区分(预案 vs 实施)"
|
||||
|
||||
|
||||
def test_dividend_rows_to_frame_filters_incomplete() -> None:
|
||||
from hdiv.data.sync.dividend import rows_to_frame
|
||||
|
||||
rows = [
|
||||
{"ts_code": "600036.SH", "end_date": "20231231", "ann_date": "20240301",
|
||||
"div_proc": "实施", "cash_div_tax": 2.0, "ex_date": "20240710"},
|
||||
{"ts_code": "600036.SH", "end_date": None, "div_proc": "实施"}, # 缺报告期 → 丢
|
||||
{"ts_code": "", "end_date": "20231231", "div_proc": "实施"}, # 缺代码 → 丢
|
||||
{"ts_code": "600036.SH", "end_date": "20231231", "div_proc": ""}, # 缺状态 → 丢
|
||||
{"ts_code": "600036.SH", "end_date": "20230630", "ann_date": None,
|
||||
"div_proc": "股东提议", "cash_div_tax": None}, # 保留(状态完整)
|
||||
]
|
||||
df = rows_to_frame(rows)
|
||||
assert len(df) == 2
|
||||
assert set(df["div_proc"]) == {"实施", "股东提议"}
|
||||
assert df["symbol"].eq("600036.SH").all()
|
||||
# DataFrame 层缺失值表现为 NaN(pandas 的 float64 语义),
|
||||
# 但**写入数据库前**必须变成 NULL —— 这才是真正的约束(见 _frame_to_records)
|
||||
recs = _frame_to_records(df)
|
||||
proposer = next(r for r in recs if r["div_proc"] == "股东提议")
|
||||
assert proposer["cash_div_tax"] is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 财报:ann_date 为空必须丢弃(PIT 纪律)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_financial_rows_drop_missing_ann_date() -> None:
|
||||
from hdiv.data.sync.financial import SPECS, rows_to_frame
|
||||
|
||||
spec = SPECS["fina_indicator"]
|
||||
rows = [
|
||||
{"ts_code": "600036.SH", "end_date": "20231231", "ann_date": "20240301", "roe": 15.0},
|
||||
{"ts_code": "600036.SH", "end_date": "20230630", "ann_date": None, "roe": 8.0},
|
||||
{"ts_code": "600036.SH", "end_date": None, "ann_date": "20240101", "roe": 1.0},
|
||||
]
|
||||
df, dropped = rows_to_frame(rows, spec)
|
||||
assert len(df) == 1, "缺公告日的记录无法用于 PIT,必须丢弃"
|
||||
assert dropped == 2
|
||||
assert df.iloc[0]["roe"] == 15.0
|
||||
|
||||
|
||||
def test_financial_report_type_preserved() -> None:
|
||||
from hdiv.data.sync.financial import SPECS, rows_to_frame
|
||||
|
||||
spec = SPECS["income"]
|
||||
rows = [
|
||||
{"ts_code": "600036.SH", "end_date": "20231231", "ann_date": "20240301",
|
||||
"report_type": "1", "total_revenue": 100.0},
|
||||
{"ts_code": "600036.SH", "end_date": "20231231", "ann_date": "20240301",
|
||||
"report_type": "2", "total_revenue": 30.0},
|
||||
{"ts_code": "600036.SH", "end_date": "20230331", "ann_date": "20240401",
|
||||
"report_type": None, "total_revenue": 5.0},
|
||||
]
|
||||
df, _ = rows_to_frame(rows, spec)
|
||||
assert len(df) == 3, "report_type 不同不得被折叠(否则合并报表与单季混为一谈)"
|
||||
assert set(df["report_type"]) == {"1", "2", "1"}, "缺省应为 1(合并报表)"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 行情帧构造
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_price_frames_map_tushare_columns() -> None:
|
||||
from hdiv.data.sync.price import adj_frame, basic_frame, daily_frame
|
||||
|
||||
d = daily_frame([{"ts_code": "000001.SZ", "trade_date": "20150105", "open": 10,
|
||||
"high": 11, "low": 9, "close": 10.5, "vol": 100, "amount": 1000}])
|
||||
assert list(d.columns) == ["symbol", "trade_date", "open", "high", "low",
|
||||
"close", "volume", "amount", "source", "adjust"]
|
||||
assert d.iloc[0]["symbol"] == "000001.SZ"
|
||||
assert d.iloc[0]["volume"] == 100, "Tushare 的 vol 列映射为 volume"
|
||||
|
||||
a = adj_frame([{"ts_code": "000001.SZ", "trade_date": "20150105", "adj_factor": 1.2}])
|
||||
assert a.iloc[0]["factor"] == 1.2
|
||||
|
||||
b = basic_frame([{"ts_code": "000001.SZ", "trade_date": "20150105", "pe": 8.0,
|
||||
"dv_ttm": 4.5, "total_mv": 1e11}])
|
||||
assert b.iloc[0]["pe"] == 8.0
|
||||
assert b.iloc[0]["dv_ttm"] == 4.5
|
||||
assert b.iloc[0]["total_mv"] == 1e11
|
||||
|
||||
|
||||
def test_price_frames_drop_invalid_rows() -> None:
|
||||
from hdiv.data.sync.price import daily_frame
|
||||
|
||||
df = daily_frame([
|
||||
{"ts_code": "000001.SZ", "trade_date": "20150105", "close": 10},
|
||||
{"ts_code": "", "trade_date": "20150105", "close": 10},
|
||||
{"ts_code": "000002.SZ", "trade_date": None, "close": 10},
|
||||
])
|
||||
assert len(df) == 1
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 限频器
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_rate_limiter_enforces_budget() -> None:
|
||||
rl = RateLimiter(per_minute=3)
|
||||
for _ in range(3):
|
||||
rl.acquire()
|
||||
assert len(rl._hits) == 3
|
||||
rl.reset()
|
||||
assert len(rl._hits) == 0
|
||||
|
||||
|
||||
def test_rate_limiter_cooldown_clears_window() -> None:
|
||||
rl = RateLimiter(per_minute=2)
|
||||
rl.acquire()
|
||||
rl.acquire()
|
||||
rl.cooldown(0.01)
|
||||
assert len(rl._hits) == 0, "冷却必须清空窗口,否则会持续撞限频"
|
||||
|
||||
|
||||
def test_per_api_limits_are_independent() -> None:
|
||||
"""限频按接口计算:dividend 与 daily 应有不同额度。"""
|
||||
from hdiv.data.tushare_client import TushareClient
|
||||
|
||||
from hdiv.data import db
|
||||
|
||||
db.load_dotenv_once()
|
||||
cfg = __import__("hdiv.core.config", fromlist=["load_config"]).load_config("datasource")
|
||||
ts = cfg.tushare
|
||||
assert ts.limit_for("dividend") == 180
|
||||
assert ts.limit_for("daily") == 480
|
||||
assert ts.limit_for("未知接口") == ts.rate_limit_default
|
||||
# 构造客户端需要 token;此处仅验证配置层
|
||||
assert ts.rate_limit_cooldown_sec >= 60, "冷却必须覆盖 Tushare 的 60 秒滑动窗口"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# index_weight 的月度切分
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_month_starts_split() -> None:
|
||||
from hdiv.data.sync.index import _month_starts
|
||||
|
||||
segs = _month_starts(date(2024, 1, 15), date(2024, 4, 10))
|
||||
assert segs[0] == (date(2024, 1, 15), date(2024, 1, 31))
|
||||
assert segs[1] == (date(2024, 2, 1), date(2024, 2, 29)) # 闰年
|
||||
assert segs[-1] == (date(2024, 4, 1), date(2024, 4, 10))
|
||||
assert len(segs) == 4
|
||||
@@ -0,0 +1,562 @@
|
||||
"""单位换算与筛选滤网测试。
|
||||
|
||||
这里覆盖的都是**开发过程中真实踩到过**的坑,每个测试对应一次静默错误:
|
||||
|
||||
1. ``total_mv`` 是万元而阈值配的是元 → 筛选结果为空(不报错);
|
||||
2. 拿季报 ROE(年初至今累计)去比「5 年年均 ROE」→ 好公司全被误杀;
|
||||
3. 银行负债率天然 90%+ → 整个金融板块被误杀;
|
||||
4. 分红除权晚于 asof → 稳定分红公司被误判为「连续分红 0 年」。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
from hdiv.data.units import (
|
||||
normalize_financial_panel,
|
||||
normalize_market_panel,
|
||||
pct_to_ratio,
|
||||
verify_market_units,
|
||||
vol_shou_to_shares,
|
||||
wan_to_shares,
|
||||
wan_to_yuan,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 单位换算
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_wan_to_yuan() -> None:
|
||||
assert wan_to_yuan(pd.Series([1.0, 100.0])).tolist() == [1e4, 1e6]
|
||||
|
||||
|
||||
def test_wan_to_shares() -> None:
|
||||
assert wan_to_shares(pd.Series([125619.78])).iloc[0] == pytest.approx(1.2561978e9)
|
||||
|
||||
|
||||
def test_pct_to_ratio() -> None:
|
||||
assert pct_to_ratio(pd.Series([5.08])).iloc[0] == pytest.approx(0.0508)
|
||||
|
||||
|
||||
def test_vol_shou_to_shares() -> None:
|
||||
assert vol_shou_to_shares(pd.Series([100])).iloc[0] == 10000
|
||||
|
||||
|
||||
def test_normalize_market_panel_units() -> None:
|
||||
"""茅台 2024-06-28 实测值:总市值 1.84 万亿元,总股本 12.56 亿股。"""
|
||||
df = pd.DataFrame(
|
||||
{
|
||||
"symbol": ["600519.SH"],
|
||||
"close": [1467.39],
|
||||
"total_share": [125619.78], # 万股
|
||||
"total_mv": [184333208.97], # 万元
|
||||
"circ_mv": [184333208.97],
|
||||
"dv_ttm": [5.1720], # 百分数
|
||||
"turnover_rate": [0.25],
|
||||
}
|
||||
)
|
||||
out = normalize_market_panel(df)
|
||||
assert out["total_mv"].iloc[0] == pytest.approx(1.8433320897e12, rel=1e-6)
|
||||
assert out["total_share"].iloc[0] == pytest.approx(1.2561978e9, rel=1e-9)
|
||||
assert out["dv_ttm"].iloc[0] == pytest.approx(0.051720)
|
||||
assert out["_units"].iloc[0] == "yuan/shares/ratio"
|
||||
|
||||
|
||||
def test_normalize_market_panel_does_not_mutate_input() -> None:
|
||||
df = pd.DataFrame({"total_mv": [100.0], "total_share": [10.0]})
|
||||
before = df.copy()
|
||||
normalize_market_panel(df)
|
||||
pd.testing.assert_frame_equal(df, before)
|
||||
|
||||
|
||||
def test_normalize_financial_panel_converts_pct() -> None:
|
||||
df = pd.DataFrame({"roe": [15.13], "debt_to_assets": [90.23], "total_revenue": [1.78e11]})
|
||||
out = normalize_financial_panel(df)
|
||||
assert out["roe"].iloc[0] == pytest.approx(0.1513)
|
||||
assert out["debt_to_assets"].iloc[0] == pytest.approx(0.9023)
|
||||
assert out["total_revenue"].iloc[0] == 1.78e11, "金额列本就以元计,不得被换算"
|
||||
|
||||
|
||||
def test_verify_market_units_detects_correct_data() -> None:
|
||||
df = pd.DataFrame(
|
||||
{
|
||||
"close": [10.0, 20.0],
|
||||
"total_share": [1e9, 5e8],
|
||||
"total_mv": [1e10, 1e10],
|
||||
}
|
||||
)
|
||||
v = verify_market_units(df)
|
||||
assert v["checked"] == 2 and v["bad"] == 0
|
||||
assert v["median_ratio"] == pytest.approx(1.0)
|
||||
|
||||
|
||||
def test_identity_alone_cannot_detect_wan_vs_yuan() -> None:
|
||||
"""恒等式对「整体万元/元混淆」无效 —— 因为 元/股 × 万股 = 万元。
|
||||
|
||||
这是一个容易误以为「有了恒等式检查就安全」的陷阱,必须显式记录:
|
||||
原始单位下比值仍然恰好是 1。
|
||||
"""
|
||||
raw = pd.DataFrame(
|
||||
{"close": [1467.39], "total_share": [125619.78], "total_mv": [184333208.97]}
|
||||
)
|
||||
v = verify_market_units(raw)
|
||||
assert v["identity_ok"] is True, "恒等式在原始单位下同样成立"
|
||||
assert v["unit_ok"] is False, "但绝对量级检查必须发现单位错误"
|
||||
assert v["detected_unit"] == "wan"
|
||||
|
||||
|
||||
def test_verify_market_units_accepts_normalized_yuan() -> None:
|
||||
norm = normalize_market_panel(
|
||||
pd.DataFrame(
|
||||
{"close": [1467.39], "total_share": [125619.78], "total_mv": [184333208.97]}
|
||||
)
|
||||
)
|
||||
v = verify_market_units(norm)
|
||||
assert v["unit_ok"] is True and v["identity_ok"] is True
|
||||
assert v["detected_unit"] == "yuan"
|
||||
|
||||
|
||||
def test_verify_market_units_detects_partial_conversion() -> None:
|
||||
"""只换算一个字段(常见疏漏)也必须被抓出来。"""
|
||||
df = pd.DataFrame(
|
||||
{"close": [10.0], "total_share": [1e9], "total_mv": [1e10 / 1e4]} # 市值漏换算
|
||||
)
|
||||
v = verify_market_units(df)
|
||||
assert v["bad"] == 1
|
||||
|
||||
|
||||
def test_verify_market_units_empty() -> None:
|
||||
assert verify_market_units(pd.DataFrame())["checked"] == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 滤网:行业豁免
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _FakeRepo:
|
||||
def __init__(self, st: set[str] | None = None, susp: set[str] | None = None) -> None:
|
||||
self._st = st or set()
|
||||
self._susp = susp or set()
|
||||
|
||||
def st_symbols(self, asof, include_delisting=True): # noqa: ANN001, ARG002
|
||||
return self._st
|
||||
|
||||
def suspended_on(self, asof): # noqa: ANN001, ARG002
|
||||
return self._susp
|
||||
|
||||
|
||||
def _frame(**kwargs) -> pd.DataFrame:
|
||||
base = {
|
||||
"symbol": ["600036.SH", "601088.SH"],
|
||||
"name": ["招商银行", "中国神华"],
|
||||
"industry": ["银行", "煤炭开采"],
|
||||
"market": ["主板", "主板"],
|
||||
"exchange": ["SSE", "SSE"],
|
||||
"listed_years": [30.0, 17.0],
|
||||
"is_fresh": [True, True],
|
||||
"total_mv": [8.6e11, 8.8e11],
|
||||
"circ_mv": [7.0e11, 7.3e11],
|
||||
"avg_amount": [1e9, 8e8],
|
||||
"close": [34.19, 44.37],
|
||||
"debt_ratio": [0.902, 0.236],
|
||||
"total_hldr_eqy_exc_min_int": [3.5e11, 4.0e11],
|
||||
}
|
||||
base.update(kwargs)
|
||||
return pd.DataFrame(base)
|
||||
|
||||
|
||||
def test_risk_filter_exempts_banks_from_leverage() -> None:
|
||||
from hdiv.core.config import RiskFilterConfig
|
||||
from hdiv.universe.filters.risk import RiskFilter
|
||||
|
||||
f = RiskFilter(RiskFilterConfig(max_debt_to_assets=0.80), exempt_leverage=["银行"])
|
||||
out = f.compute(_frame(), _FakeRepo(), date(2024, 6, 28))
|
||||
assert bool(out.passed.iloc[0]) is True, "银行必须豁免负债率上限"
|
||||
assert bool(out.passed.iloc[1]) is True, "低负债公司自然通过"
|
||||
assert out.values["600036.SH"]["debt_ratio_exempt"] is True
|
||||
|
||||
|
||||
def test_risk_filter_still_rejects_high_leverage_non_exempt() -> None:
|
||||
from hdiv.core.config import RiskFilterConfig
|
||||
from hdiv.universe.filters.risk import RiskFilter
|
||||
|
||||
f = RiskFilter(RiskFilterConfig(max_debt_to_assets=0.80), exempt_leverage=["银行"])
|
||||
df = _frame(industry=["房地产", "煤炭开采"])
|
||||
out = f.compute(df, _FakeRepo(), date(2024, 6, 28))
|
||||
assert bool(out.passed.iloc[0]) is False
|
||||
assert "资产负债率" in out.reasons["600036.SH"]
|
||||
|
||||
|
||||
def test_quality_filter_exempts_banks_from_fcf_and_leverage() -> None:
|
||||
from hdiv.core.config import QualityFilterConfig
|
||||
from hdiv.universe.filters.quality import FinancialQualityFilter
|
||||
|
||||
cfg = QualityFilterConfig(min_ocf_to_profit=0.60, max_debt_to_assets=0.80)
|
||||
f = FinancialQualityFilter(cfg, exempt_leverage=["银行"], exempt_fcf=["银行"])
|
||||
df = _frame(roe_avg=[0.1513, 0.1406], ocf_to_profit=[-0.03, 1.80])
|
||||
out = f.compute(df, _FakeRepo(), date(2024, 6, 28))
|
||||
assert bool(out.passed.iloc[0]) is True, "银行豁免 FCF 与负债率"
|
||||
|
||||
|
||||
def test_quality_filter_uses_annual_average_not_quarterly() -> None:
|
||||
"""季报 ROE 3.47% 不该被拿去比年均 8% 的阈值。"""
|
||||
from hdiv.core.config import QualityFilterConfig
|
||||
from hdiv.universe.filters.quality import FinancialQualityFilter
|
||||
|
||||
cfg = QualityFilterConfig(min_roe_5y_avg=0.08, min_ocf_to_profit=None)
|
||||
f = FinancialQualityFilter(cfg)
|
||||
df = _frame(
|
||||
industry=["白酒", "煤炭开采"],
|
||||
roe=[0.0347, 0.0380], # 季报累计值(会被 avg 覆盖)
|
||||
roe_avg=[0.1513, 0.1406], # 年报 5 年平均
|
||||
)
|
||||
out = f.compute(df, _FakeRepo(), date(2024, 6, 28))
|
||||
assert bool(out.passed.iloc[0]) is True, "应使用 roe_avg 而非季报 roe"
|
||||
assert out.values["600036.SH"]["roe"] == pytest.approx(0.1513)
|
||||
|
||||
|
||||
def test_quality_filter_falls_back_to_latest_when_no_average() -> None:
|
||||
from hdiv.core.config import QualityFilterConfig
|
||||
from hdiv.universe.filters.quality import FinancialQualityFilter
|
||||
|
||||
cfg = QualityFilterConfig(min_roe_5y_avg=0.08, min_ocf_to_profit=None)
|
||||
f = FinancialQualityFilter(cfg)
|
||||
df = _frame(industry=["白酒", "煤炭开采"], roe=[0.20, 0.03]) # 无 roe_avg 列
|
||||
out = f.compute(df, _FakeRepo(), date(2024, 6, 28))
|
||||
assert bool(out.passed.iloc[0]) is True
|
||||
assert bool(out.passed.iloc[1]) is False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 分红连续性:一年宽限期
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _div(symbol: str, end_year: int, ex_year: int, dps: float = 1.0, month: int = 7) -> dict:
|
||||
return {
|
||||
"symbol": symbol,
|
||||
"end_date": date(end_year, 12, 31),
|
||||
"imp_ann_date": date(ex_year, month - 1 if month > 1 else 12, 1),
|
||||
"div_proc": "实施",
|
||||
"cash_div_tax": dps,
|
||||
"cash_div": dps,
|
||||
"stk_div": None,
|
||||
"base_share": 10000.0,
|
||||
"ex_date": date(ex_year, month, 15),
|
||||
}
|
||||
|
||||
|
||||
def test_continuity_within_target_year() -> None:
|
||||
"""FY2023 分红已在 2024-05 除权 → target=2023,连续 5 年。"""
|
||||
from hdiv.core.config import DividendFilterConfig
|
||||
from hdiv.universe.filters.dividend import DividendFilter
|
||||
|
||||
recs = {2023: 2024, 2022: 2023, 2021: 2022, 2020: 2021, 2019: 2020}
|
||||
rows = [_div("600036.SH", fy, ex, month=5) for fy, ex in recs.items()]
|
||||
s = DividendFilter._stats(rows, target_year=2023, asof=date(2024, 6, 28),
|
||||
cfg=DividendFilterConfig())
|
||||
assert s["dividend_continuity_years"] == 5
|
||||
assert s["continuity_grace_used"] is False
|
||||
|
||||
|
||||
def test_continuity_grace_when_ex_date_lags() -> None:
|
||||
"""神华实测情形:FY2023 分红要 2024-07 才除权,asof=2024-06-28 时不可见。
|
||||
|
||||
此时必须用一年宽限期从 FY2022 起算,而不是判定为「中断分红」。
|
||||
"""
|
||||
from hdiv.core.config import DividendFilterConfig
|
||||
from hdiv.universe.filters.dividend import DividendFilter
|
||||
|
||||
rows = [_div("601088.SH", fy, fy + 1, month=7) for fy in (2022, 2021, 2020, 2019, 2018)]
|
||||
s = DividendFilter._stats(rows, target_year=2023, asof=date(2024, 6, 28),
|
||||
cfg=DividendFilterConfig())
|
||||
assert s["latest_dividend_year"] == 2022
|
||||
assert s["continuity_start_year"] == 2022
|
||||
assert s["continuity_grace_used"] is True
|
||||
assert s["dividend_continuity_years"] == 5, "不应因为除权晚而误判中断"
|
||||
|
||||
|
||||
def test_continuity_zero_when_genuinely_stopped() -> None:
|
||||
"""最近可见分红比应考核财年早两年以上 → 视为真的中断。"""
|
||||
from hdiv.core.config import DividendFilterConfig
|
||||
from hdiv.universe.filters.dividend import DividendFilter
|
||||
|
||||
rows = [_div("000002.SZ", fy, fy + 1) for fy in (2019, 2018, 2017, 2016, 2015)]
|
||||
s = DividendFilter._stats(rows, target_year=2023, asof=date(2024, 6, 28),
|
||||
cfg=DividendFilterConfig())
|
||||
assert s["dividend_continuity_years"] == 0
|
||||
|
||||
|
||||
def test_continuity_breaks_on_gap() -> None:
|
||||
"""有断档:2022 有、2021 无 → 连续 1 年。"""
|
||||
from hdiv.core.config import DividendFilterConfig
|
||||
from hdiv.universe.filters.dividend import DividendFilter
|
||||
|
||||
rows = [_div("X.SZ", fy, fy + 1) for fy in (2022, 2020, 2019)]
|
||||
s = DividendFilter._stats(rows, target_year=2023, asof=date(2024, 6, 28),
|
||||
cfg=DividendFilterConfig())
|
||||
assert s["dividend_continuity_years"] == 1
|
||||
|
||||
|
||||
def test_ttm_dps_only_counts_ex_date_in_window() -> None:
|
||||
from hdiv.core.config import DividendFilterConfig
|
||||
from hdiv.universe.filters.dividend import DividendFilter
|
||||
|
||||
rows = [
|
||||
_div("X.SZ", 2023, 2024, dps=1.5, month=4), # 窗口内
|
||||
_div("X.SZ", 2022, 2023, dps=1.2, month=7), # 窗口内(>2023-06-28)
|
||||
_div("X.SZ", 2021, 2022, dps=1.0, month=7), # 窗口外
|
||||
]
|
||||
s = DividendFilter._stats(rows, target_year=2023, asof=date(2024, 6, 28),
|
||||
cfg=DividendFilterConfig())
|
||||
assert s["ttm_dps"] == pytest.approx(2.7)
|
||||
|
||||
|
||||
def test_shift_year_handles_leap_day() -> None:
|
||||
from hdiv.universe.filters.dividend import _shift_year
|
||||
|
||||
assert _shift_year(date(2024, 2, 29), -1) == date(2023, 2, 28)
|
||||
assert _shift_year(date(2024, 6, 28), -1) == date(2023, 6, 28)
|
||||
|
||||
|
||||
def test_target_year_uses_latest_annual_report() -> None:
|
||||
"""最新年报为 FY2023 时,考核目标年应为 2023。"""
|
||||
from hdiv.core.config import DividendFilterConfig
|
||||
from hdiv.universe.filters.dividend import DividendFilter
|
||||
|
||||
df = pd.DataFrame(
|
||||
{
|
||||
"symbol": ["A.SH", "B.SH"],
|
||||
"fin_end_date": [date(2024, 3, 31), date(2023, 12, 31)],
|
||||
}
|
||||
)
|
||||
t = DividendFilter._target_years(df, date(2024, 6, 28))
|
||||
# 能看到 2024Q1 报,说明 FY2023 年报必然已披露 → 目标年 2023
|
||||
assert t["A.SH"] == 2023, "有 2024Q1 报 → FY2023 年报已出"
|
||||
assert t["B.SH"] == 2023, "有 FY2023 年报 → 目标是 2023"
|
||||
|
||||
df2 = pd.DataFrame({"symbol": ["C.SH"], "fin_end_date": [date(2023, 9, 30)]})
|
||||
assert DividendFilter._target_years(df2, date(2024, 6, 28))["C.SH"] == 2022, (
|
||||
"只看到 2023Q3 → FY2023 年报未出,退回到 2022"
|
||||
)
|
||||
|
||||
|
||||
def test_payout_ratio_uses_base_share() -> None:
|
||||
"""总现金分红 = 每股分红 × 基准股本(万股→股)。"""
|
||||
from hdiv.core.config import DividendFilterConfig
|
||||
from hdiv.universe.filters.dividend import DividendFilter
|
||||
|
||||
rows = [
|
||||
{
|
||||
"symbol": "X.SH", "end_date": date(2023, 12, 31), "ex_date": date(2024, 6, 1),
|
||||
"div_proc": "实施", "cash_div_tax": 2.0, "base_share": 10000.0, # 1 亿股
|
||||
}
|
||||
]
|
||||
row = pd.Series({"n_income_attr_p": 4e8, "free_cashflow": 8e8})
|
||||
out = DividendFilter._payout_and_cover(rows, row, date(2024, 6, 28),
|
||||
DividendFilterConfig())
|
||||
assert out["total_cash_dividend"] == pytest.approx(2.0 * 10000.0 * 1e4)
|
||||
assert out["payout_ratio"] == pytest.approx(0.5)
|
||||
# 总现金分红 = 2.0 元 × 10000 万股 × 1e4 = 2e8 元;FCF 8e8 → 覆盖 4 倍
|
||||
assert out["fcf_dividend_cover"] == pytest.approx(4.0)
|
||||
|
||||
|
||||
def test_fcf_fallback_when_tushare_missing() -> None:
|
||||
from hdiv.core.config import DividendFilterConfig
|
||||
from hdiv.universe.filters.dividend import DividendFilter
|
||||
|
||||
rows = [{
|
||||
"symbol": "X.SH", "end_date": date(2023, 12, 31), "ex_date": date(2024, 6, 1),
|
||||
"div_proc": "实施", "cash_div_tax": 1.0, "base_share": 1000.0,
|
||||
}]
|
||||
row = pd.Series({
|
||||
"n_income_attr_p": 5e6, "free_cashflow": None,
|
||||
"n_cashflow_act": 1e7, "c_pay_dist_dpcp_int_exp": 2e6,
|
||||
})
|
||||
out = DividendFilter._payout_and_cover(rows, row, date(2024, 6, 28),
|
||||
DividendFilterConfig())
|
||||
assert out["fcf_source"] == "ocf_minus_dist"
|
||||
assert out["free_cashflow"] == pytest.approx(8e6)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 滤网结果索引契约(曾经把 DataFrame 索引当成 symbol 用的真实 bug)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_filter_outcome_index_is_dataframe_index_not_symbol() -> None:
|
||||
"""滤网的 passed 索引必须与传入 DataFrame 的索引一致。
|
||||
|
||||
选择器据此用 ``live.at[i, "symbol"]`` 映射;
|
||||
若误把索引当 symbol,会导致「全部淘汰」(本项目开发中确实发生过)。
|
||||
"""
|
||||
from hdiv.core.config import MarketFilterConfig
|
||||
from hdiv.universe.filters.market import MarketFilter
|
||||
|
||||
df = _frame()
|
||||
df.index = [10, 20] # 非默认索引
|
||||
f = MarketFilter(MarketFilterConfig(min_market_cap=1e10))
|
||||
out = f.compute(df, _FakeRepo(), date(2024, 6, 28))
|
||||
assert list(out.passed.index) == [10, 20]
|
||||
assert set(out.passed.values) <= {True, False}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 数据缺失策略(安全边际策略的关键取舍)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _div_cfg(**kw):
|
||||
from hdiv.core.config import DividendFilterConfig
|
||||
|
||||
base = {
|
||||
"min_dividend_yield": None,
|
||||
"min_continuous_years": 0,
|
||||
"min_dividend_years_in_window": 0,
|
||||
"max_payout_ratio": None,
|
||||
"require_positive_fcf": True,
|
||||
"min_fcf_dividend_cover": None,
|
||||
}
|
||||
base.update(kw)
|
||||
return DividendFilterConfig(**base)
|
||||
|
||||
|
||||
def test_missing_fcf_passes_by_default() -> None:
|
||||
"""默认宽松:数据缺失放行,避免因未同步而误杀。"""
|
||||
from hdiv.universe.filters.dividend import DividendFilter
|
||||
|
||||
cfg = _div_cfg(on_missing_data="pass")
|
||||
stats = {"free_cashflow": None, "dividend_continuity_years": 0,
|
||||
"dividend_years_in_window": 0, "dividend_yield": 0.05}
|
||||
assert DividendFilter._reject_reason(stats, cfg) is None
|
||||
|
||||
|
||||
def test_missing_fcf_rejected_when_strict() -> None:
|
||||
"""严格模式:无法验证现金流即淘汰 —— 忠于「安全边际」的策略逻辑。"""
|
||||
from hdiv.universe.filters.dividend import DividendFilter
|
||||
|
||||
cfg = _div_cfg(on_missing_data="reject")
|
||||
stats = {"free_cashflow": None, "dividend_continuity_years": 0,
|
||||
"dividend_years_in_window": 0, "dividend_yield": 0.05}
|
||||
why = DividendFilter._reject_reason(stats, cfg)
|
||||
assert why is not None and "缺失" in why
|
||||
|
||||
|
||||
def test_negative_fcf_rejected_in_both_modes() -> None:
|
||||
"""FCF 为负时两种模式都必须淘汰 —— 宽松不等于放行已知风险。"""
|
||||
from hdiv.universe.filters.dividend import DividendFilter
|
||||
|
||||
stats = {"free_cashflow": -1e8, "dividend_continuity_years": 0,
|
||||
"dividend_years_in_window": 0, "dividend_yield": 0.05}
|
||||
for mode in ("pass", "reject"):
|
||||
why = DividendFilter._reject_reason(stats, _div_cfg(on_missing_data=mode))
|
||||
assert why is not None and "负" in why, f"{mode} 模式必须拒绝负 FCF"
|
||||
|
||||
|
||||
def test_missing_coverage_strict_only() -> None:
|
||||
from hdiv.universe.filters.dividend import DividendFilter
|
||||
|
||||
stats = {"free_cashflow": 1e8, "fcf_dividend_cover": None,
|
||||
"dividend_continuity_years": 0, "dividend_years_in_window": 0,
|
||||
"dividend_yield": 0.05}
|
||||
assert DividendFilter._reject_reason(
|
||||
stats, _div_cfg(min_fcf_dividend_cover=1.0, on_missing_data="pass")
|
||||
) is None
|
||||
assert DividendFilter._reject_reason(
|
||||
stats, _div_cfg(min_fcf_dividend_cover=1.0, on_missing_data="reject")
|
||||
) is not None
|
||||
|
||||
|
||||
def test_on_missing_data_invalid_value_rejected() -> None:
|
||||
from hdiv.core.errors import SchemaValidationError
|
||||
|
||||
with pytest.raises(Exception) as ei:
|
||||
_div_cfg(on_missing_data="maybe")
|
||||
assert "on_missing_data" in str(ei.value) or "Input should be" in str(ei.value)
|
||||
del SchemaValidationError
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 分红支付率必须与分红**同财年**(真实踩到的错误)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_payout_uses_same_fiscal_year_not_latest_quarter() -> None:
|
||||
"""回归:曾用「FY2023 分红 ÷ 2024Q1 净利润」算出 230.9% 的荒谬支付率。
|
||||
|
||||
正确口径下美的集团 FY2023 为:分红 207.8 亿 ÷ 净利 337.2 亿 = 61.63%。
|
||||
"""
|
||||
from hdiv.core.config import DividendFilterConfig
|
||||
from hdiv.universe.filters.dividend import DividendFilter
|
||||
|
||||
recs = [{
|
||||
"symbol": "000333.SZ", "end_date": date(2023, 12, 31),
|
||||
"ex_date": date(2024, 5, 15), "div_proc": "实施",
|
||||
"cash_div_tax": 3.0, "base_share": 692675.9241,
|
||||
}]
|
||||
# 同财年(FY2023)财务
|
||||
fy_row = pd.Series({
|
||||
"n_income_attr_p": 3.372e10, "free_cashflow": 7.16e10,
|
||||
})
|
||||
out = DividendFilter._payout_and_cover(
|
||||
recs, fy_row, date(2024, 6, 28), DividendFilterConfig()
|
||||
)
|
||||
assert out["financial_year"] == 2023
|
||||
assert out["payout_ratio"] == pytest.approx(207.8 / 337.2, abs=0.01)
|
||||
assert out["payout_basis"] == "same_fiscal_year"
|
||||
|
||||
|
||||
def test_payout_not_computed_without_same_year_row() -> None:
|
||||
"""缺少同财年财务时必须返回「不可得」,而不是退回到最新季报。"""
|
||||
from hdiv.core.config import DividendFilterConfig
|
||||
from hdiv.universe.filters.dividend import DividendFilter
|
||||
|
||||
recs = [{
|
||||
"symbol": "X.SZ", "end_date": date(2023, 12, 31),
|
||||
"ex_date": date(2024, 5, 15), "div_proc": "实施",
|
||||
"cash_div_tax": 1.0, "base_share": 10000.0,
|
||||
}]
|
||||
out = DividendFilter._payout_and_cover(
|
||||
recs, None, date(2024, 6, 28), DividendFilterConfig()
|
||||
)
|
||||
assert out["payout_ratio"] is None
|
||||
assert out["fcf_dividend_cover"] is None
|
||||
assert out["payout_basis"] == "missing_same_year_financials"
|
||||
|
||||
|
||||
def test_total_cash_dividend_uses_base_share() -> None:
|
||||
"""base_share 是**万股**,漏乘 1e4 会把支付率缩小一万倍。"""
|
||||
from hdiv.core.config import DividendFilterConfig
|
||||
from hdiv.universe.filters.dividend import DividendFilter
|
||||
|
||||
recs = [{
|
||||
"symbol": "X.SZ", "end_date": date(2023, 12, 31),
|
||||
"ex_date": date(2024, 5, 15), "div_proc": "实施",
|
||||
"cash_div_tax": 2.0, "base_share": 10000.0, # 1 亿股
|
||||
}]
|
||||
row = pd.Series({"n_income_attr_p": 4e8, "free_cashflow": 8e8})
|
||||
out = DividendFilter._payout_and_cover(
|
||||
recs, row, date(2024, 6, 28), DividendFilterConfig()
|
||||
)
|
||||
assert out["total_cash_dividend"] == pytest.approx(2.0 * 10000.0 * 1e4)
|
||||
assert out["payout_ratio"] == pytest.approx(0.5)
|
||||
|
||||
|
||||
def test_dividend_records_include_base_share() -> None:
|
||||
"""回归:repo.dividend_records 曾漏选 base_share,
|
||||
|
||||
导致 payout_ratio 与 fcf_dividend_cover 在全库范围内静默为 NULL,
|
||||
进而使 max_payout_ratio / min_fcf_dividend_cover 两个筛选条件从未生效。
|
||||
"""
|
||||
import inspect
|
||||
|
||||
from hdiv.data.repo import Repo
|
||||
|
||||
src = inspect.getsource(Repo.dividend_records)
|
||||
assert "base_share" in src, "dividend_records 必须选出 base_share"
|
||||
@@ -0,0 +1,575 @@
|
||||
"""Web 层测试:API 契约、软删除语义、站点归一化、路径重写。
|
||||
|
||||
**为什么要有契约测试**:前端(``web/app.js``)与后端(``web/server.py``)是
|
||||
两个独立文件,改动其中一个很容易忘记另一个。这里把「前端会调用的接口」
|
||||
写成断言 —— 只要后端删掉某个路由,测试立刻失败,而不是等到页面上报错。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from hdiv.core.paths import project_root
|
||||
from hdiv.web import site
|
||||
from hdiv.web.server import ROUTES
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 前端 ↔ 后端 接口契约
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
#: 前端实际会调用的接口(method, 路径样例)
|
||||
#: 路径样例用于匹配路由正则;改前端时需同步此表。
|
||||
FRONTEND_CALLS: list[tuple[str, str]] = [
|
||||
("GET", "/api/health"),
|
||||
("GET", "/api/summary"),
|
||||
("GET", "/api/universes"),
|
||||
("GET", "/api/universes/abc123"),
|
||||
("GET", "/api/universes/abc123/members"),
|
||||
("GET", "/api/universes/abc123/backtests"),
|
||||
("PATCH", "/api/universes/abc123"),
|
||||
("GET", "/api/stocks/600519.SH"),
|
||||
("GET", "/api/backtests"),
|
||||
("GET", "/api/backtests/abc123"),
|
||||
("GET", "/api/backtests/abc123/metrics"),
|
||||
("GET", "/api/backtests/abc123/equity"),
|
||||
("GET", "/api/backtests/abc123/trades"),
|
||||
("GET", "/api/backtests/abc123/signals"),
|
||||
("PATCH", "/api/backtests/abc123"),
|
||||
]
|
||||
|
||||
|
||||
def _match(method: str, path: str) -> bool:
|
||||
return any(m == method and pat.match(path) for m, pat, _ in ROUTES)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("method,path", FRONTEND_CALLS)
|
||||
def test_frontend_api_calls_have_routes(method: str, path: str) -> None:
|
||||
assert _match(method, path), f"前端会调用 {method} {path},但后端没有对应路由"
|
||||
|
||||
|
||||
def test_app_js_only_calls_known_api_prefixes() -> None:
|
||||
"""app.js 里出现的接口路径必须在契约表中有覆盖。"""
|
||||
js = (project_root() / "web" / "app.js").read_text(encoding="utf-8")
|
||||
# 抓取 api(`xxx`) / api('xxx') / patch(`xxx`, ...) 的首段
|
||||
calls = re.findall(r"\b(?:api|patch)\(\s*[`'\"]([a-z][\w/\-${}.\[\]]*)", js)
|
||||
assert calls, "未从 app.js 中解析出任何接口调用(解析逻辑需更新)"
|
||||
# FRONTEND_CALLS 里是绝对路径 /api/xxx;app.js 里的参数是相对 api 基址的
|
||||
# (api('summary') 实际请求 /api/summary),因此要去掉 api/ 前缀再比较。
|
||||
known = {
|
||||
p.removeprefix("/api/").strip("/").split("/")[0]
|
||||
for _m, p in FRONTEND_CALLS
|
||||
}
|
||||
unknown = {
|
||||
c.split("/")[0].split("?")[0]
|
||||
for c in calls
|
||||
if c.split("/")[0].split("?")[0] not in known
|
||||
}
|
||||
# 允许 ${...} 模板片段(这些是动态拼出的 run_id / symbol,前缀已在表中)
|
||||
unknown = {u for u in unknown if not u.startswith("$")}
|
||||
assert not unknown, f"app.js 调用了契约表未覆盖的接口前缀:{sorted(unknown)}"
|
||||
|
||||
|
||||
def test_members_endpoint_defaults_to_selected() -> None:
|
||||
"""接口层默认值:不传 passed 时应等价于 passed=1,而非全部候选。"""
|
||||
import inspect
|
||||
|
||||
from hdiv.web import server
|
||||
|
||||
src = inspect.getsource(server._members)
|
||||
assert 'passed_raw == "all"' in src or "passed_raw == 'all'" in src, \
|
||||
"默认值应只把显式的 all 当作「全部候选」"
|
||||
assert 'passed_raw != "0"' in src or "passed_raw != '0'" in src, \
|
||||
"未指定 passed 时应视为 1(仅入选)"
|
||||
|
||||
|
||||
def test_every_route_has_a_frontend_or_cli_consumer() -> None:
|
||||
"""反向检查:后端不应暴露无人使用的接口(便于发现遗留死接口)。"""
|
||||
known_paths = {p for _m, p in FRONTEND_CALLS}
|
||||
orphans = []
|
||||
for _m, pat, fn in ROUTES:
|
||||
# 用契约表中的样例路径试探该路由是否有消费者
|
||||
sample = pat.pattern.replace("^", "").replace("$", "")
|
||||
sample = re.sub(r"\(\?P<\w+>\[[^\]]+\]\+?\)", "abc123", sample)
|
||||
if not any(pat.match(p) for p in known_paths) and "/stocks/" not in sample:
|
||||
orphans.append((_m, pat.pattern))
|
||||
# /stocks/ 由画像页使用;/universes/{id}/members/{sym} 为可选下钻
|
||||
assert len(orphans) <= 2, f"疑似无人使用的接口:{orphans}"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 资源路径重写(归档后图表不能失效)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_rewrite_asset_paths_depth0() -> None:
|
||||
html = '<script src="assets/echarts.min.js"></script>'
|
||||
assert site.rewrite_asset_paths(html, 0) == html
|
||||
|
||||
|
||||
def test_rewrite_asset_paths_depth2() -> None:
|
||||
html = '<script src="assets/echarts.min.js"></script>'
|
||||
out = site.rewrite_asset_paths(html, 2)
|
||||
assert 'src="../../assets/echarts.min.js"' in out
|
||||
|
||||
|
||||
def test_rewrite_does_not_touch_absolute_or_external() -> None:
|
||||
for html in (
|
||||
'<script src="/assets/echarts.min.js"></script>',
|
||||
'<script src="https://cdn.example.com/echarts.min.js"></script>',
|
||||
'<script src="//cdn.example.com/echarts.min.js"></script>',
|
||||
'<script src="../assets/echarts.min.js"></script>',
|
||||
):
|
||||
assert site.rewrite_asset_paths(html, 2) == html, f"不该改写:{html}"
|
||||
|
||||
|
||||
def test_rewrite_handles_href_too() -> None:
|
||||
html = '<link href="assets/app.css" rel="stylesheet">'
|
||||
assert '../../assets/app.css' in site.rewrite_asset_paths(html, 2)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 前端文件
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_frontend_source_files_exist() -> None:
|
||||
src = project_root() / "web"
|
||||
for name in ("index.html", "app.css", "app.js"):
|
||||
assert (src / name).is_file(), f"缺少前端文件 web/{name}"
|
||||
|
||||
|
||||
def test_frontend_index_references_local_assets_only() -> None:
|
||||
html = (project_root() / "web" / "index.html").read_text(encoding="utf-8")
|
||||
refs = re.findall(r'(?:src|href)="([^"]+)"', html)
|
||||
external = [r for r in refs if r.startswith(("http://", "https://", "//"))]
|
||||
assert not external, f"前端引用了外部资源(违反离线约束):{external}"
|
||||
assert "app/app.js" in refs and "app/app.css" in refs
|
||||
assert "assets/echarts.min.js" in refs
|
||||
|
||||
|
||||
def test_frontend_uses_hash_routing_only() -> None:
|
||||
"""hash 路由是刻意的:nginx 只需托管静态文件,不必配置 rewrite。"""
|
||||
js = (project_root() / "web" / "app.js").read_text(encoding="utf-8")
|
||||
assert "location.hash" in js
|
||||
# 不应出现 history.pushState 这类需要服务端配合的写法
|
||||
assert "pushState" not in js
|
||||
|
||||
|
||||
def test_frontend_escapes_html() -> None:
|
||||
"""用户可输入记录名称/备注,必须转义以避免 XSS。"""
|
||||
js = (project_root() / "web" / "app.js").read_text(encoding="utf-8")
|
||||
assert "function esc(" in js or "const esc =" in js
|
||||
assert """ in js and "<" in js
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# API 行为(需要数据库)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _db_ready() -> bool:
|
||||
try:
|
||||
from hdiv.core.config import load_config
|
||||
from hdiv.data import db
|
||||
|
||||
db.load_dotenv_once()
|
||||
db.list_tables(load_config("datasource"))
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
requires_db = pytest.mark.skipif(not _db_ready(), reason="数据库不可用")
|
||||
|
||||
|
||||
@requires_db
|
||||
def test_service_summary_shape() -> None:
|
||||
from hdiv.web import service
|
||||
|
||||
s = service.summary()
|
||||
for k in ("universes", "backtests", "profiles", "db"):
|
||||
assert k in s, f"summary 缺少字段 {k}"
|
||||
|
||||
|
||||
@requires_db
|
||||
def test_soft_delete_is_reversible_and_never_physical() -> None:
|
||||
"""软删除语义:记录仍在库中,可通过 include_deleted 找回。"""
|
||||
from hdiv.web import service
|
||||
|
||||
runs = service.list_universes()
|
||||
if not runs:
|
||||
pytest.skip("没有筛选记录可供测试")
|
||||
rid = runs[0]["run_id"]
|
||||
try:
|
||||
service.update_universe(rid, {"deleted": True})
|
||||
visible = {r["run_id"] for r in service.list_universes()}
|
||||
assert rid not in visible, "软删除后不应出现在默认列表"
|
||||
allruns = {r["run_id"] for r in service.list_universes(include_deleted=True)}
|
||||
assert rid in allruns, "软删除的记录必须仍可通过 include_deleted 找回(不做物理删除)"
|
||||
finally:
|
||||
service.update_universe(rid, {"deleted": False})
|
||||
assert rid in {r["run_id"] for r in service.list_universes()}
|
||||
|
||||
|
||||
@requires_db
|
||||
def test_archive_toggles_visibility() -> None:
|
||||
from hdiv.web import service
|
||||
|
||||
runs = service.list_universes()
|
||||
if not runs:
|
||||
pytest.skip("没有筛选记录可供测试")
|
||||
rid = runs[0]["run_id"]
|
||||
try:
|
||||
service.update_universe(rid, {"archived": True})
|
||||
assert rid not in {r["run_id"] for r in service.list_universes()}
|
||||
assert rid in {r["run_id"] for r in service.list_universes(include_archived=True)}
|
||||
finally:
|
||||
service.update_universe(rid, {"archived": False})
|
||||
|
||||
|
||||
@requires_db
|
||||
def test_rename_persists_and_nullable() -> None:
|
||||
from hdiv.web import service
|
||||
|
||||
runs = service.list_universes()
|
||||
if not runs:
|
||||
pytest.skip("没有筛选记录可供测试")
|
||||
rid = runs[0]["run_id"]
|
||||
original = runs[0]["display_name"]
|
||||
try:
|
||||
got = service.update_universe(rid, {"display_name": "单元测试名称"})
|
||||
assert got["display_name"] == "单元测试名称"
|
||||
assert got["title"] == "单元测试名称"
|
||||
# 清空后应回落到自动标题,而不是留空
|
||||
got = service.update_universe(rid, {"display_name": ""})
|
||||
assert got["display_name"] is None
|
||||
assert got["title"], "清空命名后应回落到自动标题"
|
||||
finally:
|
||||
service.update_universe(rid, {"display_name": original or ""})
|
||||
|
||||
|
||||
@requires_db
|
||||
def test_universe_backtest_link_is_bidirectional() -> None:
|
||||
from hdiv.web import service
|
||||
|
||||
runs = service.list_universes()
|
||||
bts = service.list_backtests()
|
||||
if not runs or not bts:
|
||||
pytest.skip("缺少筛选记录或回测记录")
|
||||
rid, bid = runs[0]["run_id"], bts[0]["run_id"]
|
||||
original = bts[0]["universe_run_id"]
|
||||
try:
|
||||
service.update_backtest(bid, {"universe_run_id": rid})
|
||||
linked = service.list_backtests(universe_run_id=rid)
|
||||
assert bid in {x["run_id"] for x in linked}, "回测侧应能按股票池过滤到"
|
||||
u = service.get_universe(rid)
|
||||
assert u["linked_backtests"] >= 1, "股票池侧应统计到关联回测数"
|
||||
finally:
|
||||
service.update_backtest(bid, {"universe_run_id": original})
|
||||
|
||||
|
||||
@requires_db
|
||||
def test_link_rejects_nonexistent_universe() -> None:
|
||||
from hdiv.web import service
|
||||
|
||||
bts = service.list_backtests()
|
||||
if not bts:
|
||||
pytest.skip("没有回测记录")
|
||||
with pytest.raises(ValueError):
|
||||
service.update_backtest(bts[0]["run_id"], {"universe_run_id": "不存在的runid"})
|
||||
|
||||
|
||||
@requires_db
|
||||
def test_update_rejects_empty_patch() -> None:
|
||||
from hdiv.web import service
|
||||
|
||||
runs = service.list_universes()
|
||||
if not runs:
|
||||
pytest.skip("没有筛选记录")
|
||||
with pytest.raises(ValueError):
|
||||
service.update_universe(runs[0]["run_id"], {})
|
||||
|
||||
|
||||
@requires_db
|
||||
def test_all_api_payloads_are_json_serializable() -> None:
|
||||
"""回归:pandas 的 numpy 标量曾导致接口 500。"""
|
||||
from hdiv.web import service
|
||||
from hdiv.web.server import _Encoder
|
||||
|
||||
runs = service.list_universes()
|
||||
payloads: list[object] = [service.summary(), runs]
|
||||
if runs:
|
||||
rid = runs[0]["run_id"]
|
||||
payloads += [
|
||||
service.get_universe(rid),
|
||||
service.list_members(rid, size=2),
|
||||
service.list_backtests(universe_run_id=rid),
|
||||
]
|
||||
bts = service.list_backtests()
|
||||
if bts:
|
||||
bid = bts[0]["run_id"]
|
||||
payloads += [
|
||||
service.get_backtest(bid),
|
||||
service.get_backtest_metrics(bid),
|
||||
service.get_backtest_equity(bid),
|
||||
service.list_backtest_trades(bid, size=2),
|
||||
service.list_backtest_signals(bid),
|
||||
]
|
||||
for p in payloads:
|
||||
if p is None:
|
||||
continue
|
||||
json.dumps(p, ensure_ascii=False, cls=_Encoder) # 不应抛异常
|
||||
|
||||
|
||||
@requires_db
|
||||
def test_reason_text_is_human_readable() -> None:
|
||||
"""成交理由必须渲染成人话,而不是丢一坨 JSON 给前端。"""
|
||||
from hdiv.web.service import _reason_text
|
||||
|
||||
txt = _reason_text({
|
||||
"dividend_yield": 0.0575, "yield_percentile": 83.4,
|
||||
"rule": "股息率历史分位 83.4% >= P75", "observation_count": 1200,
|
||||
})
|
||||
assert "股息率 5.75%" in txt
|
||||
assert "历史分位 83.4%" in txt
|
||||
assert "{" not in txt and "}" not in txt
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 部署资产(缺了它们就会出现「API 不可用」)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_nginx_example_has_api_proxy_before_static() -> None:
|
||||
"""回归:nginx 示例必须包含 /api 反代,且排在静态规则之前。
|
||||
|
||||
漏掉反代块时,/ggx/api/health 会被当静态文件去 output/api/health 找 → 404,
|
||||
页面能打开但显示「API 不可用」。
|
||||
"""
|
||||
cfg = (project_root() / "deploy" / "nginx.conf.example").read_text(encoding="utf-8")
|
||||
proxy = cfg.find("proxy_pass")
|
||||
assert proxy != -1, "nginx 示例缺少 proxy_pass(API 无法访问)"
|
||||
# 反代块必须出现在静态 location 之前(nginx 前缀匹配取最长者,但顺位更易读且不易误删)
|
||||
static = cfg.find("alias /srv/hddiv/site/")
|
||||
assert static == -1 or proxy < static, "API 反代块应排在静态规则之前"
|
||||
# location 与 proxy_pass 的末尾斜杠必须成对,否则路径会被改写错
|
||||
assert re.search(r"location\s+/ggx/api/\s*\{", cfg)
|
||||
assert "proxy_pass http://127.0.0.1:8099/api/;" in cfg
|
||||
|
||||
|
||||
def test_launchd_template_is_valid_and_placeholder_based() -> None:
|
||||
"""launchd 模板必须可渲染:占位符齐全、plist 结构完整。"""
|
||||
import plistlib
|
||||
|
||||
tpl = (project_root() / "deploy" / "com.hddiv.web.plist.example")
|
||||
assert tpl.is_file(), "缺少 launchd 模板"
|
||||
raw = tpl.read_text(encoding="utf-8")
|
||||
for ph in ("__PYTHON__", "__PROJECT_ROOT__"):
|
||||
assert ph in raw, f"模板缺少占位符 {ph}"
|
||||
# 注释里也有 __,解析时先剥掉注释
|
||||
body = re.sub(r"<!--.*?-->", "", raw, flags=re.DOTALL)
|
||||
data = plistlib.loads(body.encode("utf-8"))
|
||||
assert data["Label"] == "com.hddiv.web"
|
||||
assert data["RunAtLoad"] is True
|
||||
assert data["KeepAlive"] == {"SuccessfulExit": False}, "应支持崩溃自愈"
|
||||
assert "--api-only" in data["ProgramArguments"]
|
||||
|
||||
|
||||
def test_install_service_script_placeholder_guard() -> None:
|
||||
"""安装脚本必须检测未替换的占位符 —— 否则 launchd 会静默失败。"""
|
||||
sh = (project_root() / "deploy" / "install-service.sh").read_text(encoding="utf-8")
|
||||
assert "plutil -lint" in sh, "应校验 plist 格式"
|
||||
assert "未替换的占位符" in sh or "grep -q \"__\"" in sh
|
||||
assert "launchctl load" in sh and "launchctl unload" in sh
|
||||
|
||||
|
||||
def test_serve_script_is_syntax_valid() -> None:
|
||||
import subprocess
|
||||
|
||||
sh = project_root() / "deploy" / "serve.sh"
|
||||
r = subprocess.run(["bash", "-n", str(sh)], capture_output=True, text=True)
|
||||
assert r.returncode == 0, f"serve.sh 语法错误:{r.stderr}"
|
||||
|
||||
|
||||
@requires_db
|
||||
def test_members_default_returns_selected_only() -> None:
|
||||
"""回归:不传 passed 时应返回「入选」股票,而不是全部候选。
|
||||
|
||||
曾默认返回全部候选(5000+ 只),与「股票清单」的语义不符。
|
||||
"""
|
||||
from hdiv.web import service
|
||||
|
||||
runs = service.list_universes()
|
||||
if not runs:
|
||||
pytest.skip("没有筛选记录")
|
||||
rid = runs[0]["run_id"]
|
||||
default = service.list_members(rid, passed=True, size=1)
|
||||
allc = service.list_members(rid, passed=None, size=1)
|
||||
rejected = service.list_members(rid, passed=False, size=1)
|
||||
assert default["total"] <= allc["total"]
|
||||
assert default["total"] + rejected["total"] == allc["total"], \
|
||||
"入选数 + 淘汰数 应等于候选总数"
|
||||
assert default["total"] > 0, "示例记录应有入选股票"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# HTML 报告降级为「导出件」
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_html_flag_is_opt_in_in_cli() -> None:
|
||||
"""回归:静态报告已降级为导出件 —— 默认不生成,需显式 --html。
|
||||
|
||||
改造后 SPA 是主界面,默认再产出 HTML 会:
|
||||
1) 每次筛选多一个文件
|
||||
2) 同 asof 重跑时按 asof 命名互相覆盖,与「每次运行都留痕」矛盾
|
||||
"""
|
||||
from hdiv.cli import build_parser
|
||||
|
||||
p = build_parser()
|
||||
subs = {a.dest: a for a in p._actions if hasattr(a, "choices") and isinstance(a.choices, dict)}
|
||||
sub = subs["command"].choices
|
||||
for name in ("universe", "profile", "backtest", "audit", "sensitivity"):
|
||||
args = p.parse_args([name] if name != "universe" else [name])
|
||||
assert getattr(args, "html") is False, f"hdiv {name} 的 --html 应为 opt-in"
|
||||
assert hasattr(args, "no_html"), f"hdiv {name} 应保留 --no-html 以免旧命令报错"
|
||||
_ = sub # 断言子命令存在
|
||||
|
||||
|
||||
def test_default_universe_run_writes_no_html(monkeypatch) -> None:
|
||||
"""不传 --html 时不应调用 HTML 生成器。"""
|
||||
from hdiv import cli
|
||||
|
||||
called = {"n": 0}
|
||||
import hdiv.report.build as build
|
||||
|
||||
def fake(*a, **k):
|
||||
called["n"] += 1
|
||||
raise AssertionError("默认不应生成 HTML")
|
||||
|
||||
monkeypatch.setattr(build, "build_universe_report", fake)
|
||||
# 仅验证解析结果:默认 html=False
|
||||
p = cli.build_parser()
|
||||
assert p.parse_args(["universe"]).html is False
|
||||
assert p.parse_args(["universe", "--html"]).html is True
|
||||
assert called["n"] == 0
|
||||
|
||||
|
||||
def test_report_names_include_run_id() -> None:
|
||||
"""报告文件名必须带执行 id,否则同 asof/同 symbol 重跑会互相覆盖。"""
|
||||
from hdiv.core.config import load_config
|
||||
|
||||
n = load_config("report").naming
|
||||
for key in ("universe", "profile", "backtest", "walkforward", "sensitivity"):
|
||||
pattern = getattr(n, key)
|
||||
assert "{run_id}" in pattern or "{wf_id}" in pattern or "{sens_id}" in pattern, \
|
||||
f"naming.{key} 缺少执行 id 占位符,重跑会覆盖:{pattern}"
|
||||
|
||||
|
||||
@requires_db
|
||||
def test_universe_report_filenames_do_not_collide() -> None:
|
||||
"""同 asof 的两条记录必须产出两个不同文件(实测回归)。"""
|
||||
from hdiv.core.config import load_config
|
||||
from hdiv.report.renderer import Renderer
|
||||
|
||||
r = Renderer()
|
||||
a = r.name_from("universe", asof="2025-01-21", run_id="a" * 32)
|
||||
b = r.name_from("universe", asof="2025-01-21", run_id="b" * 32)
|
||||
assert a != b, "同 asof 不同 run 必须产生不同文件名"
|
||||
assert load_config("report").naming.universe.startswith("reports/")
|
||||
|
||||
|
||||
@requires_db
|
||||
def test_universe_detail_exposes_chart_data() -> None:
|
||||
"""前端图表所需数据由后端算好(前端不做业务计算)。"""
|
||||
from hdiv.web import service
|
||||
|
||||
runs = service.list_universes()
|
||||
if not runs:
|
||||
pytest.skip("没有筛选记录")
|
||||
u = service.get_universe(runs[0]["run_id"])
|
||||
f = u["funnel"]
|
||||
assert len(f["labels"]) == len(f["values"]) == 5
|
||||
# 漏斗最后一段必须等于入选数,否则 stats 不自洽
|
||||
assert f["values"][-1] == u["member_count"]
|
||||
assert f["values"][0] == u["candidate_count"]
|
||||
# 存活数必须单调不增
|
||||
assert all(f["values"][i] >= f["values"][i + 1] for i in range(len(f["values"]) - 1)), \
|
||||
f"漏斗存活数应单调不增:{f['values']}"
|
||||
dist = u["industry_distribution"]
|
||||
assert isinstance(dist, list)
|
||||
if dist:
|
||||
assert {"industry", "count"} <= set(dist[0])
|
||||
assert dist[0]["count"] >= dist[-1]["count"], "行业分布应按数量降序"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 站点根 index.html 的所有权(曾发生真实事故)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_naming_index_is_under_reports() -> None:
|
||||
"""回归:静态报告索引曾与前端首页抢 output/index.html。
|
||||
|
||||
命名配置里 index 漏配 reports/ 前缀 + 调用点硬编码 "index.html",
|
||||
导致跑一次 `hdiv audit --html`(内部会 build_index)就把前端首页
|
||||
覆盖成静态报告索引。
|
||||
"""
|
||||
from hdiv.core.config import load_config
|
||||
|
||||
naming = load_config("report").naming
|
||||
assert naming.index == "reports/index.html", \
|
||||
f"naming.index 必须是 reports/index.html(当前 {naming.index})"
|
||||
# 所有静态报告都应在 reports/ 下,不占用站点根
|
||||
for key in ("index", "audit", "universe", "profile", "backtest",
|
||||
"walkforward", "sensitivity"):
|
||||
pattern = getattr(naming, key)
|
||||
assert pattern.startswith("reports/"), \
|
||||
f"naming.{key} 应输出到 reports/ 下,不占用站点根:{pattern}"
|
||||
|
||||
|
||||
def test_renderer_refuses_to_write_site_root_index() -> None:
|
||||
"""渲染器必须拒绝把报告写到站点根 index.html(前置拦截)。"""
|
||||
from hdiv.core.errors import HdivError
|
||||
from hdiv.report.renderer import Renderer
|
||||
|
||||
r = Renderer()
|
||||
for bad in ("index.html", "./index.html"):
|
||||
with pytest.raises(HdivError, match="统一前端首页冲突"):
|
||||
r.render("reports/index.html", {}, bad, report_type="index")
|
||||
|
||||
|
||||
def test_site_root_index_is_spa() -> None:
|
||||
"""站点根 index.html 必须是前端外壳(引用 app/app.js),不是静态报告索引。"""
|
||||
idx = project_root() / "output" / "index.html"
|
||||
if not idx.is_file():
|
||||
pytest.skip("尚未生成站点")
|
||||
html = idx.read_text(encoding="utf-8")
|
||||
assert "app/app.js" in html, "output/index.html 不是统一前端外壳(可能被报告覆盖)"
|
||||
assert "报告索引" not in html or "app/app.js" in html
|
||||
|
||||
|
||||
def test_build_index_uses_naming_config() -> None:
|
||||
"""build_index 的输出路径必须来自 naming 配置,不能硬编码。"""
|
||||
import inspect
|
||||
|
||||
from hdiv.report import build
|
||||
|
||||
src = inspect.getsource(build.build_index)
|
||||
assert 'r.name_from("index")' in src, "build_index 应使用 naming 配置生成文件名"
|
||||
# 不应再把 "index.html" 当输出名直接传进去
|
||||
assert ' "index.html",\n' not in src, "build_index 仍在硬编码输出名"
|
||||
|
||||
|
||||
def test_site_build_does_not_clobber_spa() -> None:
|
||||
"""site.sync_frontend 之后,站点根 index.html 必须仍是 SPA。"""
|
||||
from hdiv.web import site
|
||||
|
||||
site.sync_frontend(verbose=False)
|
||||
html = (project_root() / "output" / "index.html").read_text(encoding="utf-8")
|
||||
assert "app/app.js" in html
|
||||
Reference in New Issue
Block a user