"""用例8:并发压力测试 —— 测量高并发下的 TPS 保持率、P95/P99 延迟与错误率 并发数取 config.DEFAULT_CONCURRENCY(默认 100),repeats 控制并发批次数量。 """ import concurrent.futures import logging from .. import config from ..stats import percentile from .base import case_meta, make_result logger = logging.getLogger(__name__) CONCURRENT_PROMPT = "用一句话回答:1+1等于几?" BASELINE_SAMPLES = 5 # 基准串行采样次数 def _run_one(client) -> dict: """单次并发请求,失败时返回 None 标记""" try: infer = client.generate( CONCURRENT_PROMPT, options={"temperature": 0.0, "num_predict": 32}, ) return { "ok": True, "infer": infer, "elapsed_ms": infer.e2e_ms, } except Exception as e: logger.debug("并发请求失败: %s", e) return {"ok": False, "error": str(e)} def _collect_serial(client, n) -> list: """串行基准采样""" entries = [] for i in range(n): r = _run_one(client) if r["ok"]: entries.append(r) return entries def _collect_concurrent(client, n, stop_event=None) -> list: """并发采集(线程池)""" entries = [] completed = 0 def _wrapped(): return _run_one(client) with concurrent.futures.ThreadPoolExecutor(max_workers=n) as executor: futures = [executor.submit(_wrapped) for _ in range(n)] for fut in concurrent.futures.as_completed(futures): try: r = fut.result(timeout=300) except Exception as e: r = {"ok": False, "error": str(e)} entries.append(r) completed += 1 if stop_event and stop_event.is_set(): for f in futures: f.cancel() break return entries def run_test(client, repeats: int = 1, stop_event=None) -> dict: """执行并发压力测试""" if repeats > 3: repeats = 3 # 并发批次过多会非常耗时,限制上限 n_concurrent = config.DEFAULT_CONCURRENCY # ---- Step 1: 串行基准 ---- logger.info("[并发] 计算基准 TPS (串行 %d 次)...", BASELINE_SAMPLES) baseline = _collect_serial(client, BASELINE_SAMPLES) valid_base = [e for e in baseline if e["ok"]] if not valid_base: return case_meta( "case_08_concurrency", "并发压力测试", "基准测试全部失败", "压力", 30, ) | {"status": "failed", "error": "基准测试全部失败", "results": []} base_tps = sum(e["infer"].decode_speed_tok_s for e in valid_base) / len(valid_base) base_e2e = sorted(e["elapsed_ms"] for e in valid_base) base_p50 = percentile(base_e2e, 50) logger.info("[并发] 基准 TPS: %.2f tok/s, E2E P50: %.0fms", base_tps, base_p50) # ---- Step 2: 并发测试 ---- logger.info("[并发] 启动 %d 并发请求...", n_concurrent) all_concurrent = [] for batch in range(repeats): if stop_event and stop_event.is_set(): break batch_results = _collect_concurrent(client, n_concurrent, stop_event) all_concurrent.extend(batch_results) ok_count = sum(1 for r in batch_results if r["ok"]) logger.info("[并发] 批次 %d 完成: %d/%d 成功", batch + 1, ok_count, len(batch_results)) # ---- Step 3: 汇总 ---- valid_con = [e for e in all_concurrent if e["ok"]] failed_count = len(all_concurrent) - len(valid_con) con_tps = ( sum(e["infer"].decode_speed_tok_s for e in valid_con) / len(valid_con) if valid_con else 0.0 ) e2e_vals = sorted(e["elapsed_ms"] for e in valid_con) e2e_p50 = percentile(e2e_vals, 50) e2e_p95 = percentile(e2e_vals, 95) e2e_p99 = percentile(e2e_vals, 99) tps_keep = (con_tps / base_tps * 100) if base_tps > 0 else 0.0 total = len(all_concurrent) or 1 error_rate = failed_count / total * 100 # 构造标准结果条目(供引擎统计/前端展示),并附加并发摘要 results = [ make_result( e["infer"], i + 1, batch="concurrent", error=("" if e["ok"] else e.get("error", "unknown")), ) for i, e in enumerate(all_concurrent) if e["ok"] ] meta = case_meta( "case_08_concurrency", "并发压力测试", "高并发下的 TPS 保持率、P95/P99 延迟、错误率", "压力", 30, ) meta["results"] = results meta["concurrency"] = n_concurrent meta["baseline_tps"] = round(base_tps, 2) meta["concurrent_tps"] = round(con_tps, 2) meta["tps_keep_rate_pct"] = round(tps_keep, 2) meta["e2e_p50_ms"] = round(e2e_p50, 2) meta["e2e_p95_ms"] = round(e2e_p95, 2) meta["e2e_p99_ms"] = round(e2e_p99, 2) meta["success_count"] = len(valid_con) meta["failed_count"] = failed_count meta["error_rate_pct"] = round(error_rate, 2) return meta