286 lines
12 KiB
Python
286 lines
12 KiB
Python
"""压测入口(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}
|
|
|
|
|
|
_SIGNED_MAX_CONCURRENCY = 100 # 网关保护签字值(设计 §8.1,2026-07-20 人类)
|
|
_SIGNED_MAX_RPM = 600
|
|
|
|
|
|
def _guard(env: dict[str, str], workers: int, scope: str) -> 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("拒跑: 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")
|
|
# 网关保护(签字值 100/600): 压测 scope 必须配全局闸且不超签字上限
|
|
conc = env.get(f"{scope}__GLOBAL__MAX_CONCURRENCY")
|
|
rpm = env.get(f"{scope}__GLOBAL__RPM")
|
|
if not conc or not rpm:
|
|
raise SystemExit(
|
|
f"拒跑: 压测须配网关保护 {scope}__GLOBAL__MAX_CONCURRENCY(≤{_SIGNED_MAX_CONCURRENCY})"
|
|
f" 与 {scope}__GLOBAL__RPM(≤{_SIGNED_MAX_RPM})——设计 §8.1 签字值"
|
|
)
|
|
if int(conc) > _SIGNED_MAX_CONCURRENCY or int(rpm) > _SIGNED_MAX_RPM:
|
|
raise SystemExit(
|
|
f"拒跑: 全局闸 {conc}/{rpm} 超签字上限 {_SIGNED_MAX_CONCURRENCY}/{_SIGNED_MAX_RPM}"
|
|
)
|
|
|
|
|
|
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()]
|
|
from tools.soak.scoreboard import capped_budget
|
|
|
|
deadline = time.monotonic() + args.max_hours * 3600
|
|
budget_calls = capped_budget(args.scenario, 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)
|
|
# 不变量 1(记账归零 + gate 可再准入): 活后端检查,仅 redis 后端可跨进程复查
|
|
if env.get("PGW_LIMITER_BACKEND") == "redis":
|
|
asyncio.run(_live_checks(args, env, verdicts))
|
|
else:
|
|
verdicts.append(("记账归零/gate 可再准入", "SKIP — memory 后端跨进程不可查"))
|
|
# 不变量 2c: P5/P6 且配置了故障源 → 故障必须真实发生
|
|
if args.scenario in ("P5", "P6"):
|
|
fault_names = [
|
|
name for name in rpm_conf if rpm_conf[name] <= 10
|
|
] # 紧闸源;坏 key/黑洞源由错误分布人工核对(findings §4 条 2 校准留待 P5 实跑)
|
|
_check("故障混编生效(P5/P6)", sb.inv_fault_errors_present, rows, fault_source_names=fault_names)
|
|
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("硬不变量被击穿,见报告")
|
|
|
|
|
|
async def _live_checks(
|
|
args: argparse.Namespace, env: dict[str, str], verdicts: list[tuple[str, str]]
|
|
) -> None:
|
|
"""不变量 1 活后端复查(跑后): 限流 inflight 归零 + 熔断门可再准入。"""
|
|
from polygateway.backends.redis.breaker import RedisGate
|
|
from polygateway.backends.redis.limiter import RedisLimiter
|
|
from polygateway.config import GatewaySettings
|
|
|
|
from tools.soak import scoreboard as sb
|
|
|
|
settings = GatewaySettings.from_env(args.scope, env=env)
|
|
names = [s.name for s in settings.sources]
|
|
limiter = RedisLimiter.from_url(
|
|
env["REDIS_URL"],
|
|
scope=settings.scope,
|
|
sources={s.name: s for s in settings.sources},
|
|
global_limits=settings.global_limits,
|
|
lease_ttl_s=settings.lease_ttl_s,
|
|
)
|
|
gate = RedisGate.from_url(env["REDIS_URL"], config=settings.breaker, scope=settings.scope)
|
|
try:
|
|
for name, invariant in (
|
|
("记账归零(inflight)", sb.inv_accounting_zeroed(limiter, names)),
|
|
("gate 可再准入(探针不悬挂)", sb.inv_gate_reenterable(gate, names)),
|
|
):
|
|
try:
|
|
await invariant
|
|
verdicts.append((name, "PASS"))
|
|
except AssertionError as exc:
|
|
verdicts.append((name, f"FAIL — {exc}"))
|
|
finally:
|
|
await limiter.aclose()
|
|
await gate.aclose()
|
|
|
|
|
|
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, args.scope)
|
|
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()
|