"""压测入口(M2 设计 §8): 预算硬顶、隔离守卫、多进程 worker、跑后记分。 用法(conda 环境内,项目根目录): python tools/soak/run_soak.py --scenario P3 --budget-calls 6 \\ --budget-tokens 50000 --workers 2 [--scope LLM] [--concurrency 8] 守卫(违反即拒跑): REDIS_URL 必须指向 db3(实验室 db0 有在用键); PGW_TELEMETRY_PG_DSN 若配必须指向 polygateway 专用库;--workers>1 时 限流/熔断后端必须是 redis(内存后端不跨进程,多 worker 无共享闸即超发)。 预算双上限(请求数/token)任一命中优雅停;签字硬顶(设计 §8.1): 全程 token ≤ 2 亿。soak 与 pytest 不并跑(FLUSHDB 清 db3 测试状态)。 """ from __future__ import annotations import argparse import asyncio import json import multiprocessing as mp import os import resource import sys import time from pathlib import Path _ROOT = Path(__file__).resolve().parents[2] sys.path.insert(0, str(_ROOT)) sys.path.insert(0, str(_ROOT / "src")) TOKEN_HARD_CAP = 200_000_000 # 2 亿,2026-07-20 人类签字(设计 §8.1) def _merged_env() -> dict[str, str]: from dotenv import dotenv_values return {k: v for k, v in {**dotenv_values(_ROOT / ".env"), **os.environ}.items() if v} def _guard(env: dict[str, str], workers: int) -> None: redis_url = env.get("REDIS_URL", "") if not redis_url.rstrip("/").endswith("/3"): raise SystemExit(f"拒跑: REDIS_URL 必须指向专用 db3,当前 {redis_url!r}") pg = env.get("PGW_TELEMETRY_PG_DSN", "") if pg and not pg.rstrip("/").endswith("/polygateway"): raise SystemExit(f"拒跑: PG DSN 必须指向 polygateway 专用库,当前库名不符") if workers > 1 and ( env.get("PGW_LIMITER_BACKEND") != "redis" or env.get("PGW_BREAKER_BACKEND") != "redis" ): raise SystemExit("拒跑: --workers>1 需要 PGW_LIMITER_BACKEND/PGW_BREAKER_BACKEND=redis") async def _flush_db3(redis_url: str) -> None: import redis.asyncio as aioredis client = aioredis.from_url(redis_url) try: await client.flushdb() finally: await client.aclose() def _rss_mb() -> float: peak = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss return peak / 1e6 if sys.platform == "darwin" else peak / 1024 async def _worker_async(args: argparse.Namespace, worker_idx: int) -> None: """单 worker: 消费场景生成器,信号量限并发,双预算任一命中即停。""" from polygateway import GatewayClient from polygateway.telemetry.sqlite import SQLiteRecorder from tools.soak.scenarios import SCENARIOS, SoakCorpus env = _merged_env() run_id = args.run_id telemetry_path = _ROOT / f"data/soak/telemetry_{run_id}_{worker_idx}.db" recorder = SQLiteRecorder(telemetry_path) client = GatewayClient.from_env(args.scope, telemetry=recorder, env=env) corpus = SoakCorpus( harness_db=_ROOT / "data/soak/harness.db", telemetry_db=_ROOT / "data/soak/generate_questions_telemetry.db", frames_root=_ROOT / "data/soak/vt_frames", images_root=_ROOT / "data/soak/chs_images", ) generator = SCENARIOS[args.scenario](corpus, f"{run_id}-w{worker_idx}") sem = asyncio.Semaphore(args.concurrency) stats = {"calls": 0, "ok": 0, "failed": 0, "cancelled": 0, "tokens": 0} rss_samples = [_rss_mb()] deadline = time.monotonic() + args.max_hours * 3600 budget_calls = args.budget_calls // args.workers budget_tokens = min(args.budget_tokens, TOKEN_HARD_CAP) // args.workers inflight: set[asyncio.Task] = set() async def _one(kind: str, kwargs: dict) -> None: async with sem: try: # run 级双层隔离: FLUSHDB(入口)+ 每 run 独立 namespace(此处) kwargs.setdefault("cache_namespace", run_id) resp = await client.chat(**kwargs) stats["ok"] += 1 stats["tokens"] += resp.prompt_tokens + resp.completion_tokens except asyncio.CancelledError: stats["cancelled"] += 1 raise except Exception as exc: # 记分板从遥测读错误分布;这里只计数 stats["failed"] += 1 print(f"[w{worker_idx}] 调用失败: {type(exc).__name__}: {exc}", flush=True) async for kind, kwargs in generator: if ( stats["calls"] >= budget_calls or stats["tokens"] >= budget_tokens or time.monotonic() > deadline ): break stats["calls"] += 1 task = asyncio.create_task(_one(kind, kwargs)) inflight.add(task) task.add_done_callback(inflight.discard) if stats["calls"] % 20 == 0: rss_samples.append(_rss_mb()) print(f"[w{worker_idx}] {stats}", flush=True) if inflight: await asyncio.gather(*inflight, return_exceptions=True) # 优雅收尾 in-flight rss_samples.append(_rss_mb()) await client.aclose() result = {"stats": stats, "rss_mb": rss_samples, "telemetry": str(telemetry_path)} (_ROOT / f"data/soak/result_{run_id}_{worker_idx}.json").write_text( json.dumps(result), encoding="utf-8" ) print(f"[w{worker_idx}] 完成: {stats}", flush=True) def _worker_main(args: argparse.Namespace, worker_idx: int) -> None: asyncio.run(_worker_async(args, worker_idx)) def _scoreboard(args: argparse.Namespace, env: dict[str, str]) -> None: """跑后记分: 合并 worker 遥测,断言硬不变量,产出报告。""" from tools.soak import scoreboard as sb results = [] for w in range(args.workers): path = _ROOT / f"data/soak/result_{args.run_id}_{w}.json" results.append(json.loads(path.read_text(encoding="utf-8"))) rows = sb.load_rows(*(r["telemetry"] for r in results)) calls = sum(r["stats"]["calls"] for r in results) max_attempts = int( env.get(f"{args.scope}__RETRY__MAX_ATTEMPTS", env.get("LLM_MAX_RETRIES", "3")) ) verdicts: list[tuple[str, str]] = [] def _check(name: str, fn, *fargs, **fkwargs) -> None: try: fn(*fargs, **fkwargs) verdicts.append((name, "PASS")) except AssertionError as exc: verdicts.append((name, f"FAIL — {exc}")) # 行数下界=请求数(每请求至少 1 行),上界容重试放大(每请求 ≤ max_attempts 行) _check( "遥测完备(行数≥请求数,重试容差内)", sb.inv_rows_match_calls, rows, expected_calls=calls, tolerance=calls * max(max_attempts - 1, 0) + sum(r["stats"]["cancelled"] for r in results), ) _check("call_id 唯一", sb.inv_call_ids_unique, rows) rpm_conf = {} for key, val in env.items(): parts = key.split("__") # 仅四段源键 {SCOPE}__{PROVIDER}__{N}__RPM(排除 {SCOPE}__GLOBAL__RPM) if len(parts) == 4 and parts[0] == args.scope and parts[3] == "RPM": rpm_conf[f"{parts[1].lower()}_{parts[2]}"] = int(val) _check("RPM 从未击穿(分钟桶)", sb.inv_rpm_never_exceeded, rows, rpm_conf) for r in results: _check(f"RSS 平稳(w)", sb.inv_rss_stable, r["rss_mb"], max_growth_mb=args.max_rss_growth_mb) rate = sb.structured_success_rate(rows, session_suffix="-p3") verdicts.append(("P3 结构化成功率", f"{rate:.3f}(基线首跑建立)")) report = sb.render_report(args.run_id, rows, verdicts) path = sb.write_report(args.run_id, report) print(report) print(f"报告: {path}") if any(v.startswith("FAIL") for _, v in verdicts): raise SystemExit("硬不变量被击穿,见报告") def main() -> None: parser = argparse.ArgumentParser(description="PolyGateway 真实数据压测") parser.add_argument("--scenario", required=True, choices=["P1", "P2", "P3", "P4", "P5", "P6"]) parser.add_argument("--budget-calls", type=int, required=True) parser.add_argument("--budget-tokens", type=int, required=True) parser.add_argument("--workers", type=int, default=1) parser.add_argument("--concurrency", type=int, default=8, help="单 worker 并发上限") parser.add_argument("--scope", default="SOAK") parser.add_argument("--run-id", default=None) parser.add_argument("--max-hours", type=float, default=3.0) parser.add_argument("--max-rss-growth-mb", type=float, default=500.0) args = parser.parse_args() if args.run_id is None: args.run_id = time.strftime("soak_%Y%m%d_%H%M%S") env = _merged_env() _guard(env, args.workers) print(f"run_id={args.run_id}: FLUSHDB db3 + namespace 隔离") asyncio.run(_flush_db3(env["REDIS_URL"])) if args.workers == 1: _worker_main(args, 0) else: ctx = mp.get_context("spawn") procs = [ctx.Process(target=_worker_main, args=(args, w)) for w in range(args.workers)] for p in procs: p.start() for p in procs: p.join() if any(p.exitcode != 0 for p in procs): raise SystemExit(f"worker 退出码异常: {[p.exitcode for p in procs]}") _scoreboard(args, env) if __name__ == "__main__": main()