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:
2026-09-24 20:53:10 +08:00
parent 3cedcf7d53
commit d73bbe6841
2 changed files with 280 additions and 0 deletions
+155
View File
@@ -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())
+125
View File
@@ -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