147 lines
4.9 KiB
Python
147 lines
4.9 KiB
Python
"""用例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
|