diff --git a/sanguo_portfolio/runner_backtest.py b/sanguo_portfolio/runner_backtest.py index 54b8af9..8bfb328 100644 --- a/sanguo_portfolio/runner_backtest.py +++ b/sanguo_portfolio/runner_backtest.py @@ -75,6 +75,10 @@ def parse_args() -> argparse.Namespace: "--json", action="store_true", help="JSON 模式:print(json.dumps(result)) 到 stdout,供 SSH 捕获", ) + p.add_argument( + "--initial-positions", default="", + help='初始持仓 JSON(影子柜台 checkpoint 续跑用): [{"security":"600519.SH","amount":100,"avg_cost":1700.5}]', + ) return p.parse_args() @@ -283,6 +287,10 @@ def run_backtest(args: argparse.Namespace) -> Dict[str, Any]: ) print(f"[runner] ENGINE_BUILD_PRE strategy={args.strategy}", flush=True) + initial_positions = None + ip_raw = getattr(args, "initial_positions", "") or "" + if ip_raw: + initial_positions = json.loads(ip_raw) # malformed 直接抛,续跑账目不能静默丢 engine = BacktestEngine( initialize=initialize, start_date=args.start, @@ -290,11 +298,29 @@ def run_backtest(args: argparse.Namespace) -> Dict[str, Any]: frequency=args.frequency, initial_cash=args.cash, benchmark=args.benchmark, + initial_positions=initial_positions, ) print("[runner] RUN_START", flush=True) result = engine.run() print("[runner] RUN_DONE type=%s" % type(result).__name__, flush=True) + # 期末组合状态(checkpoint 续跑对账用): 现金/持仓/总值从引擎 context 直取 + try: + pf = engine.context.portfolio + if isinstance(result, dict): + result["final_portfolio"] = { + "cash": float(pf.available_cash), + "positions_value": float(pf.positions_value), + "total_value": float(pf.total_value), + "positions": [ + {"security": pos.security, "amount": int(pos.total_amount), + "avg_cost": float(pos.avg_cost or 0.0), "price": float(pos.price or 0.0)} + for pos in pf.positions.values() if pos.total_amount > 0 + ], + } + except Exception as exc: + logger.warning("提取期末组合状态失败: %s", exc) + # 引擎不把基准序列放进 results——这里带出(引擎已按区间加载 benchmark_data) try: bd = getattr(engine, "benchmark_data", None) @@ -386,6 +412,8 @@ def run_backtest_json(params: Dict[str, Any]) -> Dict[str, Any]: stamp_duty_rate=float(params.get("stamp_duty_rate", 0.001)), min_commission=float(params.get("min_commission", 5.0)), slippage=float(params.get("slippage", 0.0)), + initial_positions=json.dumps(params["initial_positions"]) + if params.get("initial_positions") else "", ) raw = run_backtest(args) @@ -422,6 +450,7 @@ def run_backtest_json(params: Dict[str, Any]) -> Dict[str, Any]: "drawdown_curve": drawdown_curve, "holdings_curve": holdings_curve, "metrics": metrics, + "final_portfolio": raw.get("final_portfolio"), "raw_summary": summary, } diff --git a/scripts/shadow_desk/spike_p0_checkpoint_replay.py b/scripts/shadow_desk/spike_p0_checkpoint_replay.py new file mode 100644 index 0000000..5d7d1d0 --- /dev/null +++ b/scripts/shadow_desk/spike_p0_checkpoint_replay.py @@ -0,0 +1,182 @@ +# -*- coding: utf-8 -*- +"""影子柜台 P0 对账 spike(docs/design/paper-shadow-desk-design.md §6 P0)。 + +验证目标:checkpoint 续跑(分段跑,段间用 initial_positions + 段末现金恢复) +与全量重放(一次跑完整个区间)账目一致。 + +判定标准: + - 期末持仓逐只证券数量完全一致;均价差 < 1e-6(相对) + - 期末现金 / 总值差 < 1e-6(相对) + - 两边净值曲线在共同交易日上的偏差 < 1e-6(相对) +跨除权:区间拉长到年,段边界自然落在除权日前后的概率极高;可通过 +--split 显式把边界压在已知除权日上(前复权因子是否随段变化由对账暴露)。 + +用法(VPS,真实数据): + cd C:\\sanguo_vnpy_v2 + C:\\Python310\\python.exe -X utf8 scripts\\shadow_desk\\spike_p0_checkpoint_replay.py ^ + --strategy all_weather --start 2024-01-01 --end 2024-12-31 --cash 1000000 + +输出末行 SPIKE_P0_PASS / SPIKE_P0_FAIL。 +""" +from __future__ import annotations + +import argparse +import json +import sys + +TOL = 1e-6 # 相对容差 + + +def _run_segment(params: dict) -> dict: + """跑一段回测,返回含 final_portfolio 的结果 dict。""" + from sanguo_portfolio.runner_backtest import run_backtest_json + return run_backtest_json(params) + + +def _positions_map(final: dict | None) -> dict[str, dict]: + fp = (final or {}).get("final_portfolio") or {} + return {p["security"]: p for p in fp.get("positions", [])} + + +def _rel_diff(a: float, b: float) -> float: + denom = max(abs(a), abs(b), 1e-12) + return abs(a - b) / denom + + +def compare(full: dict, segmented: dict) -> list[str]: + """返回问题清单(空 = 通过)。""" + problems: list[str] = [] + fpf = (full.get("final_portfolio") or {}) + fps = (segmented.get("final_portfolio") or {}) + if not fpf or not fps: + return [f"缺 final_portfolio: full={bool(fpf)} segmented={bool(fps)}"] + + # 1. 期末持仓 + pf, ps = _positions_map(full), _positions_map(segmented) + for code in sorted(set(pf) | set(ps)): + a, b = pf.get(code), ps.get(code) + if a is None or b is None: + problems.append(f"持仓 {code}: 单边缺失 full={a is not None} seg={b is not None}") + continue + if int(a["amount"]) != int(b["amount"]): + problems.append(f"持仓 {code}: 数量 full={a['amount']} seg={b['amount']}") + if _rel_diff(a["avg_cost"], b["avg_cost"]) > TOL: + problems.append( + f"持仓 {code}: 均价 full={a['avg_cost']:.6f} seg={b['avg_cost']:.6f}") + # 2. 现金 / 总值 + for key in ("cash", "total_value"): + d = _rel_diff(fpf[key], fps[key]) + if d > TOL: + problems.append(f"{key}: full={fpf[key]:.4f} seg={fps[key]:.4f} rel_diff={d:.2e}") + + # 3. 净值曲线共同交易日 + ef = {p["date"]: float(p["equity"]) for p in full.get("equity_curve") or []} + es = {p["date"]: float(p["equity"]) for p in segmented.get("equity_curve") or []} + common = sorted(set(ef) & set(es)) + if not common: + problems.append("净值曲线无共同交易日") + worst, worst_d = "", 0.0 + for d in common: + rd = _rel_diff(ef[d], es[d]) + if rd > worst_d: + worst, worst_d = d, rd + if worst_d > TOL: + problems.append(f"净值曲线最大偏差 {worst_d:.2e} @ {worst}") + return problems + + +def main() -> int: + ap = argparse.ArgumentParser(description="影子柜台 P0 对账 spike") + ap.add_argument("--strategy", default="all_weather") + ap.add_argument("--start", default="2024-01-01") + ap.add_argument("--end", default="2024-12-31") + ap.add_argument("--cash", type=float, default=1_000_000.0) + ap.add_argument("--benchmark", default="000300.XSHG") + ap.add_argument("--max-pool", type=int, default=30) + ap.add_argument( + "--split", nargs="*", default=[], + help="段边界日期(升序, 2 个 = 3 段)。缺省自动取 1/3, 2/3 处的月初。") + ap.add_argument("--provider", default="unified") + ap.add_argument("--provider-config", default="{}") + ap.add_argument("--slippage", type=float, default=0.0) + args = ap.parse_args() + + base = { + "strategy": args.strategy, + "start_date": args.start, + "end_date": args.end, + "initial_cash": args.cash, + "benchmark": args.benchmark, + "max_pool": args.max_pool, + "commission_rate": 0.0003, + "stamp_duty_rate": 0.001, + "min_commission": 5.0, + "slippage": args.slippage, + "provider": args.provider, + "provider_config": args.provider_config, + } + + splits = args.split or _auto_splits(args.start, args.end) + print(f"[spike] 全量重放 {args.start} ~ {args.end} ...", flush=True) + full = _run_segment(base) + + # 分段续跑: 每段用上一段期末现金 + 期末持仓恢复 + bounds = [args.start] + list(splits) + [args.end] + carry_cash, carry_positions = args.cash, None + seg_results = [] + for i in range(len(bounds) - 1): + seg = dict(base) + seg["start_date"], seg["end_date"] = bounds[i], bounds[i + 1] + seg["initial_cash"] = carry_cash + if carry_positions: + seg["initial_positions"] = carry_positions + print(f"[spike] 段{i + 1}/{len(bounds) - 1}: {seg['start_date']} ~ {seg['end_date']} " + f"cash={carry_cash:.2f} positions={len(carry_positions or [])}", flush=True) + r = _run_segment(seg) + seg_results.append(r) + fp = r.get("final_portfolio") or {} + carry_cash = fp.get("cash", 0.0) + carry_positions = [ + {"security": p["security"], "amount": p["amount"], "avg_cost": p["avg_cost"]} + for p in fp.get("positions", []) + ] + + last_seg = seg_results[-1] + # 分段侧期末 end_date 是最后一段区间,与全量一致 + problems = compare(full, last_seg) + + fpf = (full.get("final_portfolio") or {}) + print("\n[spike] ===== 对账结果 =====", flush=True) + print(f" 全量期末: 现金={fpf.get('cash', 0):.2f} 总值={fpf.get('total_value', 0):.2f} " + f"持仓数={len(fpf.get('positions', []))}", flush=True) + if problems: + for p in problems: + print(f" ❌ {p}", flush=True) + print("SPIKE_P0_FAIL", flush=True) + return 1 + print(" ✅ 期末持仓逐只一致; 现金/总值/净值曲线相对偏差 < 1e-6", flush=True) + print("SPIKE_P0_PASS", flush=True) + return 0 + + +def _auto_splits(start: str, end: str) -> list[str]: + """缺省段边界: 区间 1/3 与 2/3 处最近月初(月度调仓策略段边界取月初更干净)。""" + from datetime import date + + def to_date(s: str) -> date: + y, m, d = map(int, s.split("-")) + return date(y, m, d) + + s, e = to_date(start), to_date(end) + total = (e - s).days + marks = [] + for frac in (1 / 3, 2 / 3): + target = s + __import__("datetime").timedelta(days=int(total * frac)) + # 回退到月初 + first = target.replace(day=1) + marks.append(first.isoformat()) + return marks + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/tests/shadow_desk/__init__.py b/tests/shadow_desk/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/shadow_desk/test_spike_p0_compare.py b/tests/shadow_desk/test_spike_p0_compare.py new file mode 100644 index 0000000..a568159 --- /dev/null +++ b/tests/shadow_desk/test_spike_p0_compare.py @@ -0,0 +1,84 @@ +"""P0 对账 spike 的纯逻辑测试(compare / _auto_splits)。 + +跑真实数据的对账在 VPS 上执行(见 scripts/shadow_desk/spike_p0_checkpoint_replay.py)。 +""" +from __future__ import annotations + +import importlib.util +import sys +from pathlib import Path + +import pytest + +_SPEC = importlib.util.spec_from_file_location( + "spike_p0", + Path(__file__).resolve().parents[2] / "scripts" / "shadow_desk" / "spike_p0_checkpoint_replay.py", +) +spike = importlib.util.module_from_spec(_SPEC) +sys.modules.setdefault("spike_p0", spike) +_SPEC.loader.exec_module(spike) + + +def _full(consistency: str = "same") -> dict: + if consistency == "same": + positions = [{"security": "600519.SH", "amount": 100, "avg_cost": 1700.5, "price": 1750.0}] + cash, total = 820_000.0, 995_000.0 + equity = [("2024-01-05", 1_000_000.0), ("2024-02-01", 1_010_000.0)] + elif consistency == "amount_diff": + positions = [{"security": "600519.SH", "amount": 200, "avg_cost": 1700.5, "price": 1750.0}] + cash, total = 820_000.0, 995_000.0 + equity = [("2024-01-05", 1_000_000.0), ("2024-02-01", 1_010_000.0)] + else: # equity_drift + positions = [{"security": "600519.SH", "amount": 100, "avg_cost": 1700.5, "price": 1750.0}] + cash, total = 820_000.0, 995_000.0 + equity = [("2024-01-05", 1_000_000.0), ("2024-02-01", 1_010_500.0)] + return { + "final_portfolio": {"cash": cash, "total_value": total, "positions": positions}, + "equity_curve": [{"date": d, "equity": v} for d, v in equity], + } + + +def _seg() -> dict: + return { + "final_portfolio": { + "cash": 820_000.000_000_4, # 浮点级噪声 + "total_value": 995_000.0, + "positions": [{"security": "600519.SH", "amount": 100, + "avg_cost": 1700.500_000_1, "price": 1750.0}], + }, + "equity_curve": [{"date": "2024-01-05", "equity": 1_000_000.0}, + {"date": "2024-02-01", "equity": 1_010_000.0}], + } + + +def test_compare_pass_on_float_noise(): + assert spike.compare(_full(), _seg()) == [] + + +def test_compare_catches_amount_diff(): + problems = spike.compare(_full("amount_diff"), _seg()) + assert any("数量" in p for p in problems) + + +def test_compare_catches_equity_drift(): + problems = spike.compare(_full("equity_drift"), _seg()) + assert any("净值曲线最大偏差" in p for p in problems) + + +def test_compare_catches_missing_position(): + seg = _seg() + seg["final_portfolio"]["positions"] = [] + problems = spike.compare(_full(), seg) + assert any("单边缺失" in p for p in problems) + + +def test_compare_catches_missing_final_portfolio(): + problems = spike.compare({"equity_curve": []}, _seg()) + assert problems and "final_portfolio" in problems[0] + + +def test_auto_splits_month_first_and_ascending(): + splits = spike._auto_splits("2024-01-01", "2024-12-31") + assert len(splits) == 2 + assert splits[0] < splits[1] + assert all(s.endswith("-01") for s in splits) # 月初