feat(portfolio): 晋级闸门纯函数——默认阈值+回测/影子两套判定+DB覆盖汇合点 [vps] [no-doc]
spec §4.5 决议 L 六件之 2:一处定义三处调用;frozen dataclass;压线=过; thresholds_from_overrides 支撑「代码默认+每机运行值」两级阈值。 Co-Authored-By: Claude Code <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,155 @@
|
||||
# sanguo_portfolio/promotion_gate.py
|
||||
"""晋级闸门纯函数——流水线 spec §4.5 决议 L 六件之 2.
|
||||
|
||||
一处定义三处调用: strategy_registry 转移硬闸 / routes_pipeline 梯子页展示 /
|
||||
(批量)回测报告。阈值两级=代码默认常量(本档)+每机 pipeline 库运行值
|
||||
(pipeline_store 覆盖,配置页可编辑)——thresholds_from_overrides 是两级汇合点.
|
||||
默认值=spec 原文:回测夏普>=1.0+最大回撤<=30%(2022-2024 全窗);影子 4 周+TE 年化
|
||||
<=5%+fillRate>=95%+周均收益偏差<=3%(起步默认,首年校准).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
from dataclasses import dataclass, fields, replace
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class GateThresholds:
|
||||
backtest_sharpe_min: float = 1.0
|
||||
backtest_max_dd: float = 0.30
|
||||
shadow_weeks_min: int = 4
|
||||
shadow_te_annual_max: float = 0.05
|
||||
shadow_fill_rate_min: float = 0.95
|
||||
shadow_weekly_dev_max: float = 0.03
|
||||
|
||||
|
||||
DEFAULT_THRESHOLDS = GateThresholds()
|
||||
|
||||
THRESHOLD_KEYS: tuple[str, ...] = tuple(f.name for f in fields(GateThresholds))
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class GateCheck:
|
||||
key: str
|
||||
label: str
|
||||
value: float
|
||||
threshold: float
|
||||
passed: bool
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class GateResult:
|
||||
passed: bool
|
||||
checks: tuple[GateCheck, ...]
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {"passed": self.passed,
|
||||
"checks": [{"key": c.key, "label": c.label, "value": c.value,
|
||||
"threshold": c.threshold, "passed": c.passed}
|
||||
for c in self.checks]}
|
||||
|
||||
|
||||
def evaluate_backtest_gate(sharpe: float, max_dd: float,
|
||||
thresholds: GateThresholds = DEFAULT_THRESHOLDS
|
||||
) -> GateResult:
|
||||
"""回测 gate:夏普>=阈值 且 最大回撤<=阈值(压线=过)."""
|
||||
checks = (
|
||||
GateCheck("backtest_sharpe_min", "回测夏普", sharpe,
|
||||
thresholds.backtest_sharpe_min,
|
||||
sharpe >= thresholds.backtest_sharpe_min),
|
||||
GateCheck("backtest_max_dd", "最大回撤", max_dd,
|
||||
thresholds.backtest_max_dd,
|
||||
max_dd <= thresholds.backtest_max_dd),
|
||||
)
|
||||
return GateResult(all(c.passed for c in checks), checks)
|
||||
|
||||
|
||||
def evaluate_shadow_gate(weeks: int, te_annual: float, fill_rate: float,
|
||||
weekly_dev: float,
|
||||
thresholds: GateThresholds = DEFAULT_THRESHOLDS
|
||||
) -> GateResult:
|
||||
"""影子毕业 gate:周数够+TE 达标+fillRate 达标+周偏差达标(四查全过才毕业)."""
|
||||
checks = (
|
||||
GateCheck("shadow_weeks_min", "影子周数", weeks,
|
||||
float(thresholds.shadow_weeks_min),
|
||||
weeks >= thresholds.shadow_weeks_min),
|
||||
GateCheck("shadow_te_annual_max", "跟踪误差年化", te_annual,
|
||||
thresholds.shadow_te_annual_max,
|
||||
te_annual <= thresholds.shadow_te_annual_max),
|
||||
GateCheck("shadow_fill_rate_min", "fillRate", fill_rate,
|
||||
thresholds.shadow_fill_rate_min,
|
||||
fill_rate >= thresholds.shadow_fill_rate_min),
|
||||
GateCheck("shadow_weekly_dev_max", "周均收益偏差", weekly_dev,
|
||||
thresholds.shadow_weekly_dev_max,
|
||||
weekly_dev <= thresholds.shadow_weekly_dev_max),
|
||||
)
|
||||
return GateResult(all(c.passed for c in checks), checks)
|
||||
|
||||
|
||||
def thresholds_from_overrides(overrides: dict[str, float]) -> GateThresholds:
|
||||
"""DB 运行值覆盖代码默认(配置页写入的键才覆盖,其余保持默认)."""
|
||||
unknown = set(overrides) - set(THRESHOLD_KEYS)
|
||||
if unknown:
|
||||
raise ValueError(f"未知阈值键: {sorted(unknown)}")
|
||||
return replace(DEFAULT_THRESHOLDS, **overrides)
|
||||
|
||||
|
||||
def _read_task_metrics(db_path: str, task_id: str) -> tuple[float, float]:
|
||||
"""backtest→paper 转移的硬闸数据源:backtest_stats.statistics 键缺则报错.
|
||||
|
||||
只读 statistics 的 sharpe_ratio/max_drawdown 键(组合与 CTA 回测统计同键名);
|
||||
缺键不做兜底猜测——回测报告页有全套指标,人读数后用显式旗标,防编数.
|
||||
"""
|
||||
import json
|
||||
import sqlite3
|
||||
with sqlite3.connect(db_path) as conn:
|
||||
row = conn.execute(
|
||||
"SELECT statistics FROM backtest_stats WHERE task_id=? "
|
||||
"ORDER BY id DESC LIMIT 1", (task_id,)).fetchone()
|
||||
if row is None or not row[0]:
|
||||
raise ValueError(f"backtest_stats 无 task_id={task_id} 的统计行")
|
||||
stats = json.loads(row[0])
|
||||
sharpe, max_dd = stats.get("sharpe_ratio"), stats.get("max_drawdown")
|
||||
if sharpe is None or max_dd is None:
|
||||
raise ValueError(
|
||||
f"task_id={task_id} statistics 缺 sharpe_ratio/max_drawdown 键——"
|
||||
"回测报告页读数后用 --sharpe/--max-dd 显式旗标")
|
||||
return float(sharpe), float(max_dd)
|
||||
|
||||
|
||||
def main(argv: list[str] | None = None) -> int:
|
||||
"""CLI: 回测 gate 判定→打印 gate JSON(喂 strategy_registry transition --gate-json)."""
|
||||
import argparse
|
||||
import json as _json
|
||||
ap = argparse.ArgumentParser(description="晋级闸门判定(决议 L 六件之 2)")
|
||||
ap.add_argument("kind", choices=["backtest"])
|
||||
ap.add_argument("--sharpe", type=float, default=None)
|
||||
ap.add_argument("--max-dd", type=float, default=None)
|
||||
ap.add_argument("--task-id", default=None,
|
||||
help="从 backtest_stats.statistics 读 sharpe_ratio/max_drawdown")
|
||||
ap.add_argument("--db", default=None,
|
||||
help="主库路径(--task-id 时必给;缺省 env SANGUO_DB_PATH)")
|
||||
args = ap.parse_args(argv)
|
||||
if args.task_id:
|
||||
db = args.db or os.environ.get("SANGUO_DB_PATH")
|
||||
if not db:
|
||||
print("[gate] --task-id 需要 --db 或 env SANGUO_DB_PATH", file=sys.stderr)
|
||||
return 1
|
||||
try:
|
||||
sharpe, max_dd = _read_task_metrics(db, args.task_id)
|
||||
except ValueError as e:
|
||||
print(f"[gate] {e}", file=sys.stderr)
|
||||
return 1
|
||||
elif args.sharpe is not None and args.max_dd is not None:
|
||||
sharpe, max_dd = args.sharpe, args.max_dd
|
||||
else:
|
||||
print("[gate] 需要 --task-id 或 --sharpe+--max-dd", file=sys.stderr)
|
||||
return 1
|
||||
result = evaluate_backtest_gate(sharpe, max_dd)
|
||||
print(_json.dumps(result.to_dict(), ensure_ascii=False))
|
||||
return 0 if result.passed else 2
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -0,0 +1,125 @@
|
||||
# tests/portfolio/test_promotion_gate.py
|
||||
"""gate 纯函数 TDD——spec §4.5 决议 L 六件之 2(默认值逐字抄 spec).
|
||||
|
||||
回测 gate=夏普>=1.0+最大回撤<=30%;影子毕业 gate=4 周+TE 年化<=5%+fillRate>=95%
|
||||
+周均收益偏差<=3%。一处定义三处调用(转移闸门/梯子页/回测报告)。
|
||||
"""
|
||||
import pytest
|
||||
|
||||
from sanguo_portfolio.promotion_gate import (
|
||||
DEFAULT_THRESHOLDS,
|
||||
THRESHOLD_KEYS,
|
||||
GateThresholds,
|
||||
evaluate_backtest_gate,
|
||||
evaluate_shadow_gate,
|
||||
thresholds_from_overrides,
|
||||
)
|
||||
|
||||
|
||||
def test_default_thresholds_match_spec():
|
||||
t = DEFAULT_THRESHOLDS
|
||||
assert t.backtest_sharpe_min == 1.0
|
||||
assert t.backtest_max_dd == 0.30
|
||||
assert t.shadow_weeks_min == 4
|
||||
assert t.shadow_te_annual_max == 0.05
|
||||
assert t.shadow_fill_rate_min == 0.95
|
||||
assert t.shadow_weekly_dev_max == 0.03
|
||||
assert THRESHOLD_KEYS == ("backtest_sharpe_min", "backtest_max_dd",
|
||||
"shadow_weeks_min", "shadow_te_annual_max",
|
||||
"shadow_fill_rate_min", "shadow_weekly_dev_max")
|
||||
|
||||
|
||||
def test_backtest_gate_pass_and_fail():
|
||||
ok = evaluate_backtest_gate(sharpe=1.4, max_dd=0.22)
|
||||
assert ok.passed and len(ok.checks) == 2
|
||||
bad_sharpe = evaluate_backtest_gate(sharpe=0.9, max_dd=0.22)
|
||||
assert not bad_sharpe.passed
|
||||
bad_dd = evaluate_backtest_gate(sharpe=1.4, max_dd=0.35)
|
||||
assert not bad_dd.passed
|
||||
|
||||
|
||||
def test_shadow_gate_all_four_checks():
|
||||
ok = evaluate_shadow_gate(weeks=4, te_annual=0.04, fill_rate=0.97,
|
||||
weekly_dev=0.01)
|
||||
assert ok.passed and len(ok.checks) == 4
|
||||
short = evaluate_shadow_gate(weeks=3, te_annual=0.04, fill_rate=0.97,
|
||||
weekly_dev=0.01)
|
||||
assert not short.passed
|
||||
te = evaluate_shadow_gate(weeks=4, te_annual=0.06, fill_rate=0.97,
|
||||
weekly_dev=0.01)
|
||||
assert not te.passed
|
||||
fill = evaluate_shadow_gate(weeks=4, te_annual=0.04, fill_rate=0.90,
|
||||
weekly_dev=0.01)
|
||||
assert not fill.passed
|
||||
dev = evaluate_shadow_gate(weeks=4, te_annual=0.04, fill_rate=0.97,
|
||||
weekly_dev=0.05)
|
||||
assert not dev.passed
|
||||
|
||||
|
||||
def test_boundary_is_inclusive():
|
||||
# 恰好压线=过(>=1.0 / <=0.30 / >=4 / <=0.05 / >=0.95 / <=0.03)
|
||||
assert evaluate_backtest_gate(sharpe=1.0, max_dd=0.30).passed
|
||||
assert evaluate_shadow_gate(weeks=4, te_annual=0.05, fill_rate=0.95,
|
||||
weekly_dev=0.03).passed
|
||||
|
||||
|
||||
def test_to_dict_roundtrip_for_snapshot():
|
||||
r = evaluate_shadow_gate(weeks=4, te_annual=0.04, fill_rate=0.97,
|
||||
weekly_dev=0.01)
|
||||
d = r.to_dict()
|
||||
assert d["passed"] is True and d["checks"][0]["key"] == "shadow_weeks_min"
|
||||
|
||||
|
||||
def test_thresholds_from_overrides_and_unknown_key():
|
||||
t = thresholds_from_overrides({"backtest_sharpe_min": 1.2})
|
||||
assert t.backtest_sharpe_min == 1.2 and t.shadow_weeks_min == 4
|
||||
with pytest.raises(ValueError, match="未知阈值键"):
|
||||
thresholds_from_overrides({"nope": 1.0})
|
||||
|
||||
|
||||
def test_frozen():
|
||||
with pytest.raises(Exception):
|
||||
DEFAULT_THRESHOLDS.backtest_sharpe_min = 2.0 # type: ignore[misc]
|
||||
|
||||
|
||||
def test_cli_backtest_explicit_flags(capsys):
|
||||
from sanguo_portfolio.promotion_gate import main
|
||||
rc = main(["backtest", "--sharpe", "1.4", "--max-dd", "0.22"])
|
||||
assert rc == 0
|
||||
out = capsys.readouterr().out
|
||||
assert '"passed": true' in out
|
||||
|
||||
|
||||
def test_cli_backtest_task_id_reads_stats(tmp_path, capsys):
|
||||
import json
|
||||
import sqlite3
|
||||
p = str(tmp_path / "bt.db")
|
||||
conn = sqlite3.connect(p)
|
||||
conn.execute("CREATE TABLE backtest_stats (id INTEGER PRIMARY KEY "
|
||||
"AUTOINCREMENT, task_id TEXT, statistics TEXT, equity_path TEXT)")
|
||||
conn.execute("INSERT INTO backtest_stats(task_id, statistics) VALUES(?,?)",
|
||||
("bt1", json.dumps({"sharpe_ratio": 1.05, "max_drawdown": 0.28})))
|
||||
conn.commit()
|
||||
conn.close()
|
||||
from sanguo_portfolio.promotion_gate import main
|
||||
rc = main(["backtest", "--task-id", "bt1", "--db", p])
|
||||
assert rc == 0
|
||||
out = capsys.readouterr().out
|
||||
assert '"passed": true' in out and '"value": 1.05' in out
|
||||
|
||||
|
||||
def test_cli_task_id_missing_keys_errors(tmp_path, capsys):
|
||||
import json
|
||||
import sqlite3
|
||||
p = str(tmp_path / "bt.db")
|
||||
conn = sqlite3.connect(p)
|
||||
conn.execute("CREATE TABLE backtest_stats (id INTEGER PRIMARY KEY "
|
||||
"AUTOINCREMENT, task_id TEXT, statistics TEXT, equity_path TEXT)")
|
||||
conn.execute("INSERT INTO backtest_stats(task_id, statistics) VALUES(?,?)",
|
||||
("bt2", json.dumps({"total_return": 0.2})))
|
||||
conn.commit()
|
||||
conn.close()
|
||||
from sanguo_portfolio.promotion_gate import main
|
||||
rc = main(["backtest", "--task-id", "bt2", "--db", p])
|
||||
assert rc == 1
|
||||
assert "显式" in capsys.readouterr().err
|
||||
Reference in New Issue
Block a user