Initial commit of LLM speed test app

This commit is contained in:
2026-09-06 11:45:52 +08:00
commit 5544b5f9f9
38 changed files with 8587 additions and 0 deletions
+35
View File
@@ -0,0 +1,35 @@
"""客户端公共基类:推理指标与错误类型
Ollama 原生客户端与 OpenAI 兼容客户端共享同一数据模型,
测试用例通过统一接口(generate/chat)调用,无需关心后端差异。
"""
from dataclasses import dataclass
class ClientError(Exception):
"""推理服务调用相关错误(统一异常基类)"""
@dataclass
class InferenceResult:
"""单次推理的完整指标(两套后端共用)"""
ttft_ms: float = 0.0 # 首 Token 延迟(毫秒)
prefill_ms: float = 0.0 # 预填充/处理 prompt 时间(毫秒)
decode_speed_tok_s: float = 0.0 # 解码速度(tokens/s)
total_tokens: int = 0 # 总 token 数
prompt_tokens: int = 0 # prompt token 数
completion_tokens: int = 0 # 生成 token 数
e2e_ms: float = 0.0 # 端到端耗时(毫秒)
response: str = "" # 完整回复文本
def to_dict(self) -> dict:
return {
"ttft_ms": round(self.ttft_ms, 2),
"prefill_ms": round(self.prefill_ms, 2),
"decode_speed_tok_s": round(self.decode_speed_tok_s, 2),
"total_tokens": self.total_tokens,
"prompt_tokens": self.prompt_tokens,
"completion_tokens": self.completion_tokens,
"e2e_ms": round(self.e2e_ms, 2),
"response_length": len(self.response),
}
+31
View File
@@ -0,0 +1,31 @@
"""全局配置"""
import os
# ---- 推理后端选择 ----
# 可选值: "ollama"(本地 Ollama)| "openai"(任意 OpenAI 兼容 /v1 服务)
LLM_BACKEND = os.environ.get("LLM_BACKEND", "openai")
# ---- Ollama 服务(backend=ollama)----
# 本地 Ollama 默认监听端口 11434,可通过环境变量覆盖
OLLAMA_BASE_URL = os.environ.get("OLLAMA_BASE_URL", "http://localhost:11434")
# 默认模型(前端会从 /api/models 动态加载实际可用的模型列表)
DEFAULT_MODEL = os.environ.get("DEFAULT_MODEL", "qwen3.6:35b-a3b")
# ---- OpenAI 兼容服务(backend=openai)----
OPENAI_BASE_URL = os.environ.get("OPENAI_BASE_URL", "https://ai.lebiztrips.com/v1")
OPENAI_API_KEY = os.environ.get("OPENAI_API_KEY", "123")
OPENAI_DEFAULT_MODEL = os.environ.get("OPENAI_DEFAULT_MODEL", "Q3.6-35B-A3B-Orig-Thi")
# ---- 并发压力测试 ----
DEFAULT_CONCURRENCY = int(os.environ.get("DEFAULT_CONCURRENCY", "100"))
# ---- HTTP 服务 ----
SERVER_HOST = os.environ.get("SPEED_TEST_HOST", "0.0.0.0")
SERVER_PORT = int(os.environ.get("SPEED_TEST_PORT", "8000"))
# ---- 目录结构 ----
ROOT_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
BACKEND_DIR = os.path.dirname(os.path.abspath(__file__))
FRONTEND_DIR = os.path.join(ROOT_DIR, "frontend")
RESULTS_DIR = os.path.join(ROOT_DIR, "results")
REPORT_DIR = os.path.join(ROOT_DIR, "report")
+240
View File
@@ -0,0 +1,240 @@
"""测试引擎:编排测试用例、统计指标、保存结果并生成报告"""
import json
import logging
import os
import sys
import time
from datetime import datetime
from glob import glob
from . import config
from . import report as report_module
from .stats import mean, std, percentile, percentiles, min_val, max_val
from .ollama_client import OllamaClient, OllamaError
from .test_cases.case_01_generation import run_test as run_case_01
from .test_cases.case_02_simple_tool import run_test as run_case_02
from .test_cases.case_03_file_analysis import run_test as run_case_03
from .test_cases.case_04_parallel_tool import run_test as run_case_04
from .test_cases.case_05_long_context import run_test as run_case_05
from .test_cases.case_06_reasoning import run_test as run_case_06
from .test_cases.case_07_multiturn import run_test as run_case_07
from .test_cases.case_08_concurrency import run_test as run_case_08
logger = logging.getLogger(__name__)
# 用例注册表
TEST_CASES = {
"case_01_generation": {"func": run_case_01, "name": "纯文本生成基准", "difficulty": "简单", "estimated_seconds": 5},
"case_02_simple_tool": {"func": run_case_02, "name": "简单工具调用延迟", "difficulty": "简单", "estimated_seconds": 3},
"case_03_file_analysis": {"func": run_case_03, "name": "文件读取+分析", "difficulty": "中等", "estimated_seconds": 2},
"case_04_parallel_tool": {"func": run_case_04, "name": "并行工具调用", "difficulty": "中等", "estimated_seconds": 5},
"case_05_long_context": {"func": run_case_05, "name": "长上下文处理", "difficulty": "耗时", "estimated_seconds": 10},
"case_06_reasoning": {"func": run_case_06, "name": "复杂推理任务", "difficulty": "耗时", "estimated_seconds": 15},
"case_07_multiturn": {"func": run_case_07, "name": "多轮对话累积", "difficulty": "耗时", "estimated_seconds": 10},
"case_08_concurrency": {"func": run_case_08, "name": "并发压力测试", "difficulty": "压力", "estimated_seconds": 30},
}
def create_client(model: str):
"""根据配置创建对应的推理客户端"""
if config.LLM_BACKEND == "openai":
from .openai_client import OpenAIClient
return OpenAIClient(model=model)
from .ollama_client import OllamaClient
return OllamaClient(model=model)
# 保留最近 N 次运行结果,超出则删除旧文件
_MAX_RESULTS = 50
def calculate_statistics(results_list: list) -> dict:
"""从用例结果列表计算延迟统计(基于 elapsed_ms)"""
values = [float(r["elapsed_ms"]) for r in results_list if r.get("elapsed_ms") is not None]
if not values:
return {}
pcts = percentiles(values, [50, 95, 99])
return {
"count": len(values),
"mean_ms": round(mean(values), 2),
"std_ms": round(std(values), 2),
"min_ms": round(min_val(values), 2),
"max_ms": round(max_val(values), 2),
"p50_ms": round(pcts[50], 2),
"p95_ms": round(pcts[95], 2),
"p99_ms": round(pcts[99], 2),
}
def calculate_inference_metrics(results_list: list) -> dict:
"""单次遍历聚合推理指标(TTFT / Prefill / Decode / E2E)"""
accumulators = {}
for r in results_list:
for key in ("ttft_ms", "prefill_ms", "e2e_ms"):
val = r.get(key, 0)
if val and float(val) > 0:
accumulators.setdefault(key, []).append(float(val))
decode = r.get("decode_speed_tok_s", 0)
if decode and float(decode) > 0:
accumulators.setdefault("decode_speed_tok_s", []).append(float(decode))
total_tok = r.get("total_tokens", 0)
if total_tok and int(total_tok) > 0:
accumulators.setdefault("total_tokens", []).append(int(total_tok))
result = {}
for key, vals in accumulators.items():
result[f"{key}_mean"] = round(mean(vals), 2)
result[f"{key}_p95"] = round(percentile(vals, 95), 2)
return result
def _cleanup_old_results() -> None:
"""清理过期的运行结果,只保留最近 _MAX_RESULTS 个 run_*.json"""
pattern = os.path.join(config.RESULTS_DIR, "run_*.json")
files = sorted(glob(pattern))
while len(files) > _MAX_RESULTS:
old = files.pop(0)
try:
os.remove(old)
logger.info("清理旧结果: %s", old)
except OSError:
pass
def run_all_tests(cases: list, repeats: int, model: str, skip_heavy: bool,
status: dict = None, stop_event=None) -> dict:
"""运行选定的测试用例,返回完整报告
status: 运行状态字典(由 RunManager 共享,用于向前端汇报进度)
stop_event: threading.Event,置位后停止
"""
client = create_client(model)
# 预检:确认推理服务可用
if not client.ping():
backend = config.LLM_BACKEND
hint = "请确认已运行 'ollama serve' 并拉取模型" if backend == "ollama" else "请检查 OPENAI_BASE_URL / OPENAI_API_KEY 配置"
raise OllamaError(f"无法连接 {backend} 后端服务。{hint}")
# 过滤耗时用例
if skip_heavy:
cases = [c for c in cases if TEST_CASES[c]["difficulty"] != "耗时"]
total_cases = len(cases)
all_results = []
stopped = False
total_start = time.time()
def _update(**kw):
if status is not None:
status.update(kw)
for i, case_id in enumerate(cases, 1):
if stop_event and stop_event.is_set():
stopped = True
_update(message="用户停止测试")
break
case_info = TEST_CASES[case_id]
_update(
status="running",
current_case_index=i - 1,
total_cases=total_cases,
current_case_name=case_info["name"],
progress=round((i - 1) / total_cases * 100, 1),
message=f"正在执行: {case_info['name']} ({case_info['difficulty']})",
)
case_start = time.time()
try:
raw_result = case_info["func"](client=client, repeats=repeats, stop_event=stop_event)
stats = {}
inference_metrics = {}
if raw_result.get("results"):
stats = calculate_statistics(raw_result["results"])
inference_metrics = calculate_inference_metrics(raw_result["results"])
all_results.append({
"case_id": case_id,
"case_name": case_info["name"],
"difficulty": case_info["difficulty"],
"status": "passed",
"raw_data": raw_result,
"statistics": stats,
"inference_metrics": inference_metrics,
"elapsed_seconds": round(time.time() - case_start, 2),
})
_update(
current_case_index=i,
progress=round(i / total_cases * 100, 1),
message=f"完成: {case_info['name']}",
)
logger.info("用例 %s 完成: %s", case_id, case_info["name"])
except Exception as e:
all_results.append({
"case_id": case_id,
"case_name": case_info["name"],
"difficulty": case_info["difficulty"],
"status": "failed",
"error": str(e),
"elapsed_seconds": round(time.time() - case_start, 2),
})
_update(message=f"失败: {case_info['name']} - {e}")
logger.error("用例 %s 失败: %s", case_id, e)
total_elapsed = time.time() - total_start
report = {
"version": "1.0",
"timestamp": datetime.now().isoformat(),
"run_id": datetime.now().strftime("%Y%m%d-%H%M%S"),
"environment": {
"python_version": f"{sys.version_info.major}.{sys.version_info.minor}.{sys.version_info.micro}",
"platform": sys.platform,
"model": model,
"backend": config.LLM_BACKEND,
},
"config": {
"cases_requested": cases,
"repeats": repeats,
"skip_heavy": skip_heavy,
"total_cases": len(cases),
},
"summary": {
"total_cases": len(cases),
"passed": sum(1 for r in all_results if r["status"] == "passed"),
"failed": sum(1 for r in all_results if r["status"] == "failed"),
"stopped": stopped,
"total_elapsed_seconds": round(total_elapsed, 2),
},
"results": all_results,
}
# 保存结果 + 生成 HTML 报告
save_results(report)
return report
def save_results(report: dict) -> None:
"""保存 JSON 结果并生成 HTML 报告"""
os.makedirs(config.RESULTS_DIR, exist_ok=True)
latest_path = os.path.join(config.RESULTS_DIR, "latest.json")
with open(latest_path, "w", encoding="utf-8") as f:
json.dump(report, f, ensure_ascii=False, indent=2)
archive_path = os.path.join(config.RESULTS_DIR, f"run_{report['run_id']}.json")
with open(archive_path, "w", encoding="utf-8") as f:
json.dump(report, f, ensure_ascii=False, indent=2)
# 清理旧结果
_cleanup_old_results()
try:
report_module.generate_report(report)
except Exception as e:
logger.warning("生成 HTML 报告失败: %s", e)
+154
View File
@@ -0,0 +1,154 @@
"""Ollama API 客户端
通过 HTTP 流式调用本地 Ollama 服务,并测量真实推理指标:
- TTFT: 首个 token 到达时间(从发起请求开始计时)
- Prefill: prompt 预填充耗时(取自 Ollama 返回的 prompt_eval_duration)
- Decode Speed: 解码速度 tokens/s(eval_count / eval_duration)
- E2E: 端到端总耗时
依赖:requests(Ollama 服务需已启动,默认 http://localhost:11434)
"""
import json
import logging
import time
import requests
from . import config
from .client_base import InferenceResult, ClientError
logger = logging.getLogger(__name__)
class OllamaError(ClientError):
"""Ollama 调用相关错误"""
class OllamaClient:
"""封装对本地 Ollama 服务的调用"""
def __init__(self, base_url: str = None, model: str = None, timeout: int = 180):
self.base_url = (base_url or config.OLLAMA_BASE_URL).rstrip("/")
self.model = model or config.DEFAULT_MODEL
self.timeout = timeout
self._session = requests.Session()
# ---------- 服务探测 ----------
def ping(self) -> bool:
"""探测 Ollama 服务是否可用"""
try:
return self._session.get(f"{self.base_url}/api/tags", timeout=5).ok
except Exception:
return False
def list_models(self) -> list:
"""列出本地已拉取的模型"""
try:
resp = self._session.get(f"{self.base_url}/api/tags", timeout=5)
resp.raise_for_status()
return [m["name"] for m in resp.json().get("models", [])]
except requests.exceptions.ConnectionError as e:
raise OllamaError(
f"无法连接 Ollama 服务({self.base_url}),请确认已运行 'ollama serve'"
) from e
except Exception as e:
raise OllamaError(f"获取模型列表失败: {e}") from e
# ---------- 文本生成(/api/generate)----------
def generate(self, prompt: str, system: str = None, options: dict = None) -> InferenceResult:
"""单轮文本生成"""
payload = {
"model": self.model,
"prompt": prompt,
"stream": True,
"options": options or {},
}
if system:
payload["system"] = system
return self._request(payload)
# ---------- 多轮对话(/api/chat)----------
def chat(self, messages: list, options: dict = None) -> InferenceResult:
"""多轮对话,messages 为 [{"role": "user"/"assistant", "content": ...}]"""
payload = {
"model": self.model,
"messages": messages,
"stream": True,
"options": options or {},
}
return self._request(payload, chat=True)
# ---------- 核心请求逻辑 ----------
def _request(self, payload: dict, chat: bool = False) -> InferenceResult:
result = InferenceResult()
chunks = []
final = {}
ttft_start = None
start = time.perf_counter()
endpoint = "chat" if chat else "generate"
try:
resp = self._session.post(
f"{self.base_url}/api/{endpoint}",
json=payload,
stream=True,
timeout=self.timeout,
)
resp.raise_for_status()
for line in resp.iter_lines(decode_unicode=True):
if not line:
continue
data = json.loads(line)
if chat:
piece = (data.get("message") or {}).get("content", "")
else:
piece = data.get("response", "")
done = data.get("done", False)
if not done:
# 首个非空片段出现时刻记为 TTFT
if ttft_start is None and piece:
ttft_start = time.perf_counter() - start
if piece:
chunks.append(piece)
else:
final = data
except requests.exceptions.ConnectionError as e:
raise OllamaError(
f"无法连接 Ollama 服务({self.base_url}),请确认已运行 'ollama serve' 且模型 {self.model} 已拉取"
) from e
except requests.exceptions.Timeout as e:
raise OllamaError(f"Ollama 请求超时({self.timeout}s): {e}") from e
except (json.JSONDecodeError, KeyError) as e:
raise OllamaError(f"Ollama 返回数据格式错误: {e}") from e
result.response = "".join(chunks)
result.e2e_ms = (time.perf_counter() - start) * 1000
result.ttft_ms = (ttft_start * 1000) if ttft_start else 0.0
# Ollama 结束块中的指标(单位为纳秒)
prompt_eval_count = final.get("prompt_eval_count", 0) or 0
prompt_eval_dur_ns = final.get("prompt_eval_duration", 0) or 0
eval_count = final.get("eval_count", 0) or 0
eval_dur_ns = final.get("eval_duration", 0) or 0
result.prompt_tokens = int(prompt_eval_count)
result.completion_tokens = int(eval_count)
result.total_tokens = result.prompt_tokens + result.completion_tokens
result.prefill_ms = prompt_eval_dur_ns / 1_000_000
if eval_dur_ns:
result.decode_speed_tok_s = result.completion_tokens / (eval_dur_ns / 1e9)
# TTFT 未测到时(如空回复),回退到 prefill 时间
if result.ttft_ms <= 0 and result.prefill_ms > 0:
result.ttft_ms = result.prefill_ms
logger.debug(
"推理完成: e2e=%.1fms ttft=%.1fms decode=%.1f tok/s tokens=%d/%d",
result.e2e_ms, result.ttft_ms, result.decode_speed_tok_s,
result.prompt_tokens, result.completion_tokens,
)
return result
+213
View File
@@ -0,0 +1,213 @@
"""OpenAI 兼容 API 客户端
对接任意 OpenAI /v1 兼容服务(OpenAI、Together、硅基流动、llama.cpp server 等),
与 OllamaClient 保持相同的 generate/chat 接口,测试用例无需区分后端。
指标来源:
- TTFT: 流式请求中首个 content 片段到达时间(真实测量)
- Prefill: 优先取服务端 timings.prompt_ms(llama.cpp server 暴露),否则回退 TTFT
- Decode Speed: 优先取服务端 timings.predicted_per_second,否则按 completion_tokens/e2e 估算
- E2E: 客户端请求总耗时
"""
import json
import logging
import time
import requests
from . import config
from .client_base import InferenceResult, ClientError
logger = logging.getLogger(__name__)
class OpenAIError(ClientError):
"""OpenAI 兼容 API 调用错误"""
class OpenAIClient:
"""封装对 OpenAI 兼容 /v1 服务的调用"""
def __init__(self, base_url: str = None, model: str = None, api_key: str = None,
timeout: int = 300):
self.base_url = (base_url or config.OPENAI_BASE_URL).rstrip("/")
self.model = model or config.OPENAI_DEFAULT_MODEL
self.api_key = api_key if api_key not in (None, "") else config.OPENAI_API_KEY
self.timeout = timeout
self._session = requests.Session()
# ---------- 请求头 ----------
@property
def _headers(self) -> dict:
headers = {"Content-Type": "application/json"}
if self.api_key:
headers["Authorization"] = f"Bearer {self.api_key}"
return headers
# ---------- 服务探测 ----------
def ping(self) -> bool:
"""探测服务是否可用"""
try:
return self._session.get(
f"{self.base_url}/models", headers=self._headers, timeout=5
).ok
except Exception:
return False
def list_models(self) -> list:
"""列出可用模型"""
try:
resp = self._session.get(
f"{self.base_url}/models", headers=self._headers, timeout=10
)
resp.raise_for_status()
data = resp.json().get("data", [])
models = []
for m in data:
mid = m.get("id") or m.get("name")
if mid:
models.append(mid)
return models
except requests.exceptions.ConnectionError as e:
raise OpenAIError(
f"无法连接服务({self.base_url}),请检查地址与网络"
) from e
except requests.exceptions.HTTPError as e:
if resp.status_code in (401, 403):
raise OpenAIError("API Key 无效或被拒绝(401/403)") from e
raise OpenAIError(f"获取模型列表失败: {resp.status_code} {resp.text[:200]}") from e
except Exception as e:
raise OpenAIError(f"获取模型列表失败: {e}") from e
# ---------- 文本生成 ----------
def generate(self, prompt: str, system: str = None, options: dict = None) -> InferenceResult:
"""单轮文本生成(流式),等价于单条 user 消息的 chat"""
messages = []
if system:
messages.append({"role": "system", "content": system})
messages.append({"role": "user", "content": prompt})
return self.chat(messages, options)
# ---------- 多轮对话 ----------
def chat(self, messages: list, options: dict = None) -> InferenceResult:
"""多轮对话,messages 为 [{"role": ..., "content": ...}]"""
options = options or {}
payload = {
"model": self.model,
"messages": messages,
"stream": True,
"temperature": options.get("temperature", 0.3),
"max_tokens": options.get("num_predict", 256),
}
start = time.perf_counter()
ttft_ms = 0.0
chunks = []
final = {}
timings = {}
try:
resp = self._session.post(
f"{self.base_url}/chat/completions",
headers=self._headers,
json=payload,
stream=True,
timeout=self.timeout,
)
except requests.exceptions.ConnectionError as e:
raise OpenAIError(
f"无法连接服务({self.base_url}),请检查地址与网络"
) from e
except requests.exceptions.Timeout as e:
raise OpenAIError(f"请求超时({self.timeout}s): {e}") from e
if resp.status_code != 200:
err_text = resp.text[:300]
raise OpenAIError(f"API 错误 {resp.status_code}: {err_text}")
# 流式解析 SSE
try:
for line in resp.iter_lines(decode_unicode=True):
if not line:
continue
line = line.strip()
if line.startswith("data:"):
line = line[5:].strip()
if line == "[DONE]":
break
if not line:
continue
try:
data = json.loads(line)
except json.JSONDecodeError:
continue
choices = data.get("choices") or []
if not choices:
continue
delta = choices[0].get("delta") or {}
# 兼容推理模型:token 可能放在 reasoning_content(思考过程)或 content
piece = delta.get("content") or delta.get("reasoning_content") or ""
if piece:
# 首个内容片段到达时间 = TTFT
if ttft_ms == 0.0:
ttft_ms = (time.perf_counter() - start) * 1000
chunks.append(piece)
# 服务端最终块可能携带 usage / timings
if choices[0].get("finish_reason"):
final = data
if data.get("usage"):
final = data
except (requests.exceptions.ConnectionError, requests.exceptions.ChunkedEncodingError) as e:
raise OpenAIError(f"流式读取中断: {e}") from e
e2e_ms = (time.perf_counter() - start) * 1000
result = InferenceResult()
result.response = "".join(chunks)
result.e2e_ms = e2e_ms
result.ttft_ms = ttft_ms
# usage: token 计数(部分服务端流式响应不含 usage,回退用 timings)
usage = final.get("usage", {}) or {}
prompt_tok = int(usage.get("prompt_tokens", 0) or 0)
completion_tok = int(usage.get("completion_tokens", 0) or 0)
# timings: llama.cpp / llama-server 暴露的服务端计时(单位毫秒)
timings = final.get("timings", {}) or {}
prompt_n = int(timings.get("prompt_n", 0) or 0)
predicted_n = int(timings.get("predicted_n", 0) or 0)
prompt_ms = float(timings.get("prompt_ms", 0) or 0)
predicted_ms = float(timings.get("predicted_ms", 0) or 0)
predicted_pps = float(timings.get("predicted_per_second", 0) or 0)
result.prompt_tokens = prompt_tok or prompt_n
result.completion_tokens = completion_tok or predicted_n
result.total_tokens = result.prompt_tokens + result.completion_tokens
result.prefill_ms = prompt_ms
# 解码速度:优先服务端计时
if predicted_pps > 0:
result.decode_speed_tok_s = predicted_pps
elif result.completion_tokens > 0 and result.e2e_ms > 0:
result.decode_speed_tok_s = result.completion_tokens / (result.e2e_ms / 1000)
# TTFT 回退:流式被缓冲或未测到时,用服务端 prompt_ms
if result.ttft_ms <= 0 or result.ttft_ms > e2e_ms * 0.95:
if prompt_ms > 0:
result.ttft_ms = prompt_ms
else:
result.ttft_ms = e2e_ms if e2e_ms > 0 else 0.0
logger.debug(
"推理完成(openai): e2e=%.1fms ttft=%.1fms decode=%.1f tok/s tokens=%d/%d",
result.e2e_ms, result.ttft_ms, result.decode_speed_tok_s,
result.prompt_tokens, result.completion_tokens,
)
return result
+279
View File
@@ -0,0 +1,279 @@
"""生成 HTML 可视化报告"""
import json
import logging
import os
from . import config
logger = logging.getLogger(__name__)
def generate_report(report: dict) -> str:
"""从报告数据生成 HTML 报告文件,返回输出路径"""
html = build_html(report)
os.makedirs(config.REPORT_DIR, exist_ok=True)
output_path = os.path.join(config.REPORT_DIR, "report.html")
with open(output_path, "w", encoding="utf-8") as f:
f.write(html)
logger.info("报告已生成: %s", output_path)
return output_path
def _diff_color(difficulty: str) -> str:
"""难度对应的颜色"""
return {
"简单": "#22c55e",
"中等": "#f59e0b",
"耗时": "#ef4444",
"压力": "#8b5cf6",
}.get(difficulty, "#6b7280")
def _stats_rows(results: list) -> str:
"""构建详细数据表格行 HTML"""
rows = []
for r in results:
stats = r.get("statistics", {})
if not stats:
continue
rows.append(
f'<tr>'
f'<td><strong>{r["case_name"]}</strong></td>'
f'<td><span class="diff-badge" style="background:{_diff_color(r["difficulty"])}">{r["difficulty"]}</span></td>'
f'<td class="status-{r["status"]}" style="font-weight:bold">{r["status"]}</td>'
f'<td class="mono">{stats.get("mean_ms", "N/A")}</td>'
f'<td class="mono">{stats.get("std_ms", "N/A")}</td>'
f'<td class="mono">{stats.get("min_ms", "N/A")}</td>'
f'<td class="mono">{stats.get("max_ms", "N/A")}</td>'
f'<td class="mono">{stats.get("p95_ms", "N/A")}</td>'
f'</tr>'
)
return "\n".join(rows)
def _chart_data(results: list) -> tuple:
"""提取图表数据:(labels, means, p95s)"""
labels, means, p95s = [], [], []
for r in results:
stats = r.get("statistics", {})
if stats:
labels.append(r["case_name"])
means.append(stats.get("mean_ms", 0))
p95s.append(stats.get("p95_ms", 0))
return labels, means, p95s
def _performance_analysis(results: list) -> str:
"""性能分析摘要"""
_, means, _ = _chart_data(results)
if len(means) < 2:
return ""
best_idx = means.index(min(means))
worst_idx = means.index(max(means))
ratio = (means[worst_idx] / means[best_idx]) if means[best_idx] > 0 else "N/A"
avg = sum(means) / len(means)
return (
f'<div class="analysis-box">'
f'<h3>性能分析</h3>'
f'<ul>'
f'<li><strong>最快用例:</strong> {results[best_idx]["case_name"]} ({means[best_idx]:.2f}ms)</li>'
f'<li><strong>最慢用例:</strong> {results[worst_idx]["case_name"]} ({means[worst_idx]:.2f}ms)</li>'
f'<li><strong>最快/最慢比:</strong> {ratio}x</li>'
f'<li><strong>平均延迟:</strong> {avg:.2f}ms</li>'
f'</ul></div>'
)
def _concurrency_summary(results: list) -> str:
"""提取并发压力测试摘要(case_08),无则返回空串"""
for r in results:
raw = r.get("raw_data", {}) or {}
if r.get("case_id") == "case_08_concurrency" and raw:
def g(key, suffix=""):
val = raw.get(key)
return "N/A" if val is None else f"{val}{suffix}"
keep = raw.get("tps_keep_rate_pct", 0)
keep_color = "#22c55e" if keep >= 60 else ("#f59e0b" if keep > 0 else "#ef4444")
err = raw.get("error_rate_pct", 0)
err_color = "#22c55e" if err == 0 else "#ef4444"
return (
f'<div class="section">'
f'<h2>⚡ 并发压力测试摘要</h2>'
f'<div class="dashboard">'
f'<div class="stat-card"><div class="label">并发数</div><div class="value" style="color:#8b5cf6">{g("concurrency")}</div></div>'
f'<div class="stat-card"><div class="label">基准 TPS</div><div class="value" style="color:#6366f1">{g("baseline_tps")}</div></div>'
f'<div class="stat-card"><div class="label">并发 TPS</div><div class="value" style="color:#6366f1">{g("concurrent_tps")}</div></div>'
f'<div class="stat-card"><div class="label">TPS 保持率</div><div class="value" style="color:{keep_color}">{g("tps_keep_rate_pct", "%")}</div></div>'
f'<div class="stat-card"><div class="label">P95 延迟</div><div class="value" style="color:#ef4444">{g("e2e_p95_ms", "ms")}</div></div>'
f'<div class="stat-card"><div class="label">P99 延迟</div><div class="value" style="color:#ef4444">{g("e2e_p99_ms", "ms")}</div></div>'
f'<div class="stat-card"><div class="label">错误率</div><div class="value" style="color:{err_color}">{g("error_rate_pct", "%")}</div></div>'
f'</div></div>'
)
return ""
def build_html(report: dict) -> str:
"""构建 HTML 报告"""
summary = report.get("summary", {})
results = report.get("results", [])
env = report.get("environment", {})
config_ = report.get("config", {})
timestamp = report.get("timestamp", "")
run_id = report.get("run_id", "")
case_names, case_means, case_p95s = _chart_data(results)
return f"""<!DOCTYPE html>
<html lang="zh-CN">
<head>
<meta charset="UTF-8">
<meta name="viewport" content="width=device-width, initial-scale=1.0">
<title>大模型速度测试报告 - {run_id}</title>
<script src="https://cdn.jsdelivr.net/npm/chart.js@4.4.0/dist/chart.umd.min.js"></script>
<style>
* {{ margin: 0; padding: 0; box-sizing: border-box; }}
body {{
font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', 'PingFang SC', 'Microsoft YaHei', sans-serif;
background: #f8fafc; color: #1e293b; line-height: 1.6;
}}
.container {{ max-width: 1400px; margin: 0 auto; padding: 20px; }}
.header {{ background: linear-gradient(135deg, #667eea 0%, #764ba2 100%); color: white; padding: 30px; border-radius: 12px; margin-bottom: 20px; }}
.header h1 {{ font-size: 28px; margin-bottom: 10px; }}
.header .meta {{ font-size: 14px; opacity: 0.9; }}
.dashboard {{ display: grid; grid-template-columns: repeat(auto-fit, minmax(200px, 1fr)); gap: 15px; margin-bottom: 20px; }}
.stat-card {{ background: white; padding: 20px; border-radius: 10px; box-shadow: 0 2px 8px rgba(0,0,0,0.1); }}
.stat-card .label {{ font-size: 12px; color: #64748b; text-transform: uppercase; }}
.stat-card .value {{ font-size: 28px; font-weight: bold; margin-top: 5px; }}
.section {{ background: white; border-radius: 10px; padding: 25px; margin-bottom: 20px; box-shadow: 0 2px 8px rgba(0,0,0,0.1); }}
.section h2 {{ font-size: 20px; margin-bottom: 15px; padding-bottom: 10px; border-bottom: 2px solid #e2e8f0; }}
table {{ width: 100%; border-collapse: collapse; }}
th, td {{ padding: 12px; text-align: left; border-bottom: 1px solid #e2e8f0; }}
th {{ background: #f1f5f9; font-weight: 600; font-size: 13px; text-transform: uppercase; }}
tr:hover {{ background: #f8fafc; }}
.mono {{ font-family: 'SF Mono', 'Fira Code', monospace; }}
.diff-badge {{ padding: 4px 10px; border-radius: 12px; color: white; font-size: 12px; font-weight: 600; }}
.status-passed {{ color: #22c55e; }}
.status-failed {{ color: #ef4444; }}
.chart-container {{ position: relative; height: 400px; margin: 20px 0; }}
.analysis-box {{ background: #f0fdf4; border-left: 4px solid #22c55e; padding: 15px; border-radius: 6px; }}
.analysis-box ul {{ padding-left: 20px; margin-top: 10px; }}
.analysis-box li {{ margin: 5px 0; }}
footer {{ text-align: center; padding: 20px; color: #64748b; font-size: 13px; }}
</style>
</head>
<body>
<div class="container">
<div class="header">
<h1>大模型速度测试报告</h1>
<div class="meta">
Run ID: {run_id} | 时间: {timestamp} | 用例: {config_.get('total_cases', 0)} | 重复: {config_.get('repeats', 0)}次
<br>模型: {env.get('model', 'N/A')} | 平台: {env.get('platform', 'N/A')}
</div>
</div>
<div class="dashboard">
<div class="stat-card"><div class="label">总用例数</div><div class="value">{summary.get('total_cases', 0)}</div></div>
<div class="stat-card"><div class="label">通过</div><div class="value" style="color:#22c55e">{summary.get('passed', 0)}</div></div>
<div class="stat-card"><div class="label">失败</div><div class="value" style="color:#ef4444">{summary.get('failed', 0)}</div></div>
<div class="stat-card"><div class="label">总耗时</div><div class="value">{summary.get('total_elapsed_seconds', 0):.2f}s</div></div>
</div>
{_concurrency_summary(results)}
{_performance_analysis(results)}
<div class="section">
<h2>性能对比柱状图</h2>
<div class="chart-container"><canvas id="barChart"></canvas></div>
</div>
<div class="section">
<h2>P95延迟对比</h2>
<div class="chart-container"><canvas id="p95Chart"></canvas></div>
</div>
<div class="section">
<h2>详细数据</h2>
<table>
<thead>
<tr>
<th>用例名称</th><th>难度</th><th>状态</th><th>平均延迟 (ms)</th>
<th>标准差</th><th>最小</th><th>最大</th><th>P95</th>
</tr>
</thead>
<tbody>
{_stats_rows(results)}
</tbody>
</table>
</div>
<div class="section">
<h2>环境信息</h2>
<table>
<tr><td>模型</td><td class="mono">{env.get('model', 'N/A')}</td></tr>
<tr><td>运行平台</td><td class="mono">{env.get('platform', 'N/A')}</td></tr>
<tr><td>运行ID</td><td class="mono">{run_id}</td></tr>
<tr><td>测试时间</td><td class="mono">{timestamp}</td></tr>
</table>
</div>
<footer>大模型速度测试助手 | 报告由 report.py 自动生成</footer>
</div>
<script>
const barCtx = document.getElementById('barChart').getContext('2d');
new Chart(barCtx, {{
type: 'bar',
data: {{
labels: {json.dumps(case_names, ensure_ascii=False)},
datasets: [{{
label: '平均延迟 (ms)',
data: {json.dumps(case_means)},
backgroundColor: 'rgba(102, 126, 234, 0.7)',
borderColor: 'rgba(102, 126, 234, 1)',
borderWidth: 1
}}]
}},
options: {{
responsive: true,
maintainAspectRatio: false,
plugins: {{
legend: {{ display: false }},
title: {{ display: true, text: '各用例平均延迟对比 (ms)', font: {{ size: 16 }} }}
}},
scales: {{ y: {{ beginAtZero: true, title: {{ display: true, text: '延迟 (ms)' }} }} }}
}}
}});
const p95Ctx = document.getElementById('p95Chart').getContext('2d');
new Chart(p95Ctx, {{
type: 'bar',
data: {{
labels: {json.dumps(case_names, ensure_ascii=False)},
datasets: [{{
label: 'P95延迟 (ms)',
data: {json.dumps(case_p95s)},
backgroundColor: 'rgba(239, 68, 68, 0.7)',
borderColor: 'rgba(239, 68, 68, 1)',
borderWidth: 1
}}]
}},
options: {{
responsive: true,
maintainAspectRatio: false,
plugins: {{
legend: {{ display: false }},
title: {{ display: true, text: '各用例P95延迟 (ms)', font: {{ size: 16 }} }}
}},
scales: {{ y: {{ beginAtZero: true, title: {{ display: true, text: '延迟 (ms)' }} }} }}
}}
}});
</script>
</body>
</html>"""
+380
View File
@@ -0,0 +1,380 @@
"""HTTP 服务器:提供 API + 托管前端静态文件
启动后访问 http://localhost:8000 即可打开测试界面。
"""
import json
import logging
import mimetypes
import os
import threading
from dataclasses import asdict, dataclass, field
from datetime import datetime
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from urllib.parse import urlparse
from . import config
from .client_base import ClientError
from .engine import TEST_CASES, run_all_tests
from .ollama_client import OllamaClient, OllamaError
logger = logging.getLogger(__name__)
@dataclass
class RunState:
"""测试运行状态(线程安全:所有读写通过 RunManager 加锁)"""
status: str = "idle" # idle | running | done | stopped | failed
run_id: str = ""
model: str = ""
repeats: int = 0
skip_heavy: bool = False
cases: list = field(default_factory=list)
current_case_index: int = 0
total_cases: int = 0
current_case_name: str = ""
message: str = "准备中"
progress: float = 0.0
started_at: str = ""
finished_at: str = ""
summary: dict = field(default_factory=dict)
def update(self, values=None, **kwargs) -> None:
"""批量更新字段(引擎通过 status.update(kw) 或 status.update(**kw) 汇报进度)"""
updates = dict(values or {})
updates.update(kwargs)
for key, value in updates.items():
if hasattr(self, key):
setattr(self, key, value)
def to_dict(self) -> dict:
return asdict(self)
class RunManager:
"""管理测试运行状态(单实例、同一时间只允许一次运行)"""
def __init__(self):
self._lock = threading.Lock()
self._stop_event = threading.Event()
self._state = RunState()
self._thread = None
@property
def is_running(self) -> bool:
return self._state.status == "running"
def start(self, cases: list, repeats: int, model: str, skip_heavy: bool) -> str:
with self._lock:
if self.is_running:
raise RuntimeError("已有测试正在运行,请等待完成或先停止")
self._stop_event.clear()
run_id = datetime.now().strftime("%Y%m%d-%H%M%S")
self._state = RunState(
status="running",
run_id=run_id,
model=model,
repeats=repeats,
skip_heavy=skip_heavy,
cases=list(cases),
total_cases=len(cases),
started_at=datetime.now().isoformat(),
)
self._thread = threading.Thread(
target=self._run, args=(cases, repeats, model, skip_heavy), daemon=True
)
self._thread.start()
return run_id
def stop(self) -> None:
"""请求停止(置位事件,引擎在用例之间检查)"""
self._stop_event.set()
def status(self) -> dict:
"""返回当前运行状态(对外只读拷贝)"""
with self._lock:
return self._state.to_dict()
def _run(self, cases, repeats, model, skip_heavy):
try:
report = run_all_tests(
cases=cases,
repeats=repeats,
model=model,
skip_heavy=skip_heavy,
status=self._state,
stop_event=self._stop_event,
)
stopped = report["summary"].get("stopped", False)
with self._lock:
self._state.status = "stopped" if stopped else "done"
self._state.progress = 100
self._state.message = "测试已停止" if stopped else "测试完成"
self._state.finished_at = datetime.now().isoformat()
self._state.summary = report["summary"]
except Exception as e:
logger.exception("测试引擎异常")
with self._lock:
self._state.status = "failed"
self._state.message = f"测试失败: {e}"
self._state.finished_at = datetime.now().isoformat()
manager = RunManager()
class Handler(BaseHTTPRequestHandler):
"""API + 静态文件处理"""
# ---------- HTTP 基础 ----------
def _send_json(self, obj, code=200):
body = json.dumps(obj, ensure_ascii=False).encode("utf-8")
self.send_response(code)
self.send_header("Content-Type", "application/json; charset=utf-8")
self.send_header("Content-Length", str(len(body)))
self.end_headers()
self.wfile.write(body)
def _read_json(self) -> dict:
try:
length = int(self.headers.get("Content-Length", 0) or 0)
except ValueError:
length = 0
if length <= 0:
return {}
try:
return json.loads(self.rfile.read(length).decode("utf-8"))
except json.JSONDecodeError:
return {}
def log_message(self, format, *args):
print(f"[API] {args[0]}")
# ---------- GET ----------
def do_GET(self):
path = urlparse(self.path).path
if path == "/api/health":
self._send_json({"status": "ok"})
elif path == "/api/config":
self._handle_config()
elif path == "/api/models":
self._handle_models()
elif path == "/api/cases":
cases = {
cid: {
"name": info["name"],
"difficulty": info["difficulty"],
"estimated_seconds": info["estimated_seconds"],
}
for cid, info in TEST_CASES.items()
}
self._send_json(cases)
elif path == "/api/status":
self._send_json(manager.status())
elif path == "/api/results/latest.json":
self._handle_latest()
elif path == "/api/results/list":
self._handle_results_list()
elif path.startswith("/api/results/"):
self._handle_result_file(path)
elif path == "/report.html":
self._serve_file(os.path.join(config.REPORT_DIR, "report.html"))
else:
self._serve_static(path)
# ---------- POST ----------
def do_POST(self):
path = urlparse(self.path).path
body = self._read_json()
if path == "/api/run":
self._handle_run(body)
elif path == "/api/stop":
manager.stop()
self._send_json({"ok": True, "message": "已请求停止"})
else:
self._send_json({"error": "Not found"}, 404)
# ---------- API 处理器 ----------
def _handle_config(self):
self._send_json({
"backend": config.LLM_BACKEND,
"ollama_base_url": config.OLLAMA_BASE_URL,
"openai_base_url": config.OPENAI_BASE_URL,
"default_model": config.DEFAULT_MODEL,
"openai_default_model": config.OPENAI_DEFAULT_MODEL,
"concurrency": config.DEFAULT_CONCURRENCY,
})
def _handle_models(self):
try:
if config.LLM_BACKEND == "openai":
from .openai_client import OpenAIClient
models = OpenAIClient().list_models()
else:
models = OllamaClient().list_models()
self._send_json({"models": models, "error": None})
except ClientError as e:
self._send_json({"models": [], "error": str(e)})
def _handle_run(self, body):
cases = body.get("cases") or []
repeats = int(body.get("repeats", 3) or 3)
model = body.get("model") or config.DEFAULT_MODEL
skip_heavy = bool(body.get("skip_heavy", False))
# 校验用例
unknown = [c for c in cases if c not in TEST_CASES]
if not cases:
self._send_json({"error": "未选择测试用例"}, 400)
return
if unknown:
self._send_json({"error": f"未知用例: {', '.join(unknown)}"}, 400)
return
if repeats < 1 or repeats > 50:
self._send_json({"error": "重复次数需在 1-50 之间"}, 400)
return
try:
run_id = manager.start(cases, repeats, model, skip_heavy)
self._send_json({"run_id": run_id})
except RuntimeError as e:
self._send_json({"error": str(e)}, 409)
def _handle_latest(self):
"""返回最新测试结果,文件损坏时优雅降级"""
latest_path = os.path.join(config.RESULTS_DIR, "latest.json")
if not os.path.exists(latest_path):
self._send_json({"error": "暂无测试结果,请先运行测试"}, 404)
return
try:
with open(latest_path, "r", encoding="utf-8") as f:
data = json.load(f)
self._send_json(data)
except (json.JSONDecodeError, OSError) as e:
logger.error("读取 latest.json 失败: %s", e)
self._send_json({"error": "结果文件损坏,请先运行新测试"}, 500)
def _handle_results_list(self):
"""列出所有 run_*.json 结果文件"""
import glob as glob_mod
pattern = os.path.join(config.RESULTS_DIR, "run_*.json")
files = sorted(glob_mod.glob(pattern))
result = []
for f in files:
name = os.path.basename(f)
try:
with open(f, "r", encoding="utf-8") as fh:
data = json.load(fh)
summary = data.get("summary", {})
result.append({
"filename": name,
"timestamp": data.get("timestamp", ""),
"model": data.get("environment", {}).get("model", ""),
"total_cases": summary.get("total_cases", 0),
"passed": summary.get("passed", 0),
"failed": summary.get("failed", 0),
"elapsed_seconds": summary.get("total_elapsed_seconds", 0),
})
except (json.JSONDecodeError, OSError):
result.append({"filename": name, "error": "无法读取"})
self._send_json({"results": result})
def _handle_result_file(self, path):
"""返回单个历史结果文件 /api/results/<filename>"""
filename = path.rsplit("/", 1)[-1]
if not filename.endswith(".json") or "/" in filename or "\\" in filename:
self._send_json({"error": "Forbidden"}, 403)
return
full = os.path.join(config.RESULTS_DIR, filename)
if not os.path.isfile(full):
self._send_json({"error": "Not found"}, 404)
return
try:
with open(full, "r", encoding="utf-8") as f:
data = json.load(f)
self._send_json(data)
except (json.JSONDecodeError, OSError):
self._send_json({"error": "文件损坏"}, 500)
# ---------- 静态文件 ----------
def _serve_static(self, path):
# 只允许从 frontend/ 目录读取
if path == "/" or path == "/index.html":
rel = "index.html"
else:
rel = path.lstrip("/")
full = os.path.normpath(os.path.join(config.FRONTEND_DIR, rel))
frontend_root = os.path.normpath(config.FRONTEND_DIR)
# 防止路径遍历攻击
if not os.path.abspath(full).startswith(os.path.abspath(frontend_root + os.sep)):
self._send_json({"error": "Forbidden"}, 403)
return
# 防止符号链接指向外部目录
# 逐段检查路径中每个组件是否安全
parts = os.path.relpath(full, frontend_root).split(os.sep)
current = frontend_root
for part in parts:
current = os.path.join(current, part)
if os.path.islink(current):
link_target = os.path.realpath(current)
if not link_target.startswith(os.path.realpath(frontend_root) + os.sep):
self._send_json({"error": "Forbidden"}, 403)
return
if not os.path.exists(current):
break
self._serve_file(full)
def _serve_file(self, full_path):
if not os.path.isfile(full_path):
self._send_json({"error": "Not found"}, 404)
return
content_type = mimetypes.guess_type(full_path)[0] or "application/octet-stream"
with open(full_path, "rb") as f:
data = f.read()
self.send_response(200)
self.send_header("Content-Type", content_type)
self.send_header("Content-Length", str(len(data)))
self.end_headers()
self.wfile.write(data)
def main():
# 配置根日志
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
datefmt="%H:%M:%S",
)
server = ThreadingHTTPServer((config.SERVER_HOST, config.SERVER_PORT), Handler)
backend = config.LLM_BACKEND
print("=" * 60)
print(f"大模型速度测试助手 v2.1 ({'OpenAI 兼容' if backend == 'openai' else 'Ollama'})")
print("=" * 60)
if backend == "openai":
print(f"服务地址 : {config.OPENAI_BASE_URL}")
print(f"默认模型 : {config.OPENAI_DEFAULT_MODEL}")
print(f"并发测试 : {config.DEFAULT_CONCURRENCY} 并发")
else:
print(f"Ollama 地址 : {config.OLLAMA_BASE_URL}")
print(f"默认模型 : {config.DEFAULT_MODEL}")
print(f"前端界面 : http://localhost:{config.SERVER_PORT}")
print("按 Ctrl+C 停止服务器")
print("=" * 60)
try:
server.serve_forever()
except KeyboardInterrupt:
print("\n服务器已停止")
server.server_close()
if __name__ == "__main__":
main()
+61
View File
@@ -0,0 +1,61 @@
"""统计计算工具"""
import math
from typing import List
def _percentile_at(sorted_data: list[float], pct: float) -> float:
"""在线性插值下,返回排序后数据的指定百分位值。
注意:调用方必须确保 data 已排序。
"""
n = len(sorted_data)
if n == 0:
return 0.0
k = (pct / 100.0) * (n - 1)
f = math.floor(k)
c = math.ceil(k)
if f == c:
return sorted_data[int(k)]
d0 = sorted_data[int(f)] * (c - k)
d1 = sorted_data[int(c)] * (k - f)
return d0 + d1
def percentile(data: List[float], pct: float) -> float:
"""计算指定百分位数的值(线性插值)
空数据返回 0.0,调用方应通过检查 count 来判断是否有数据。
"""
return _percentile_at(sorted(data), pct)
def percentiles(data: List[float], pcts: List[float]) -> dict:
"""单次排序,批量计算多个百分位,返回 {pct: value}。"""
sorted_data = sorted(data)
return {p: _percentile_at(sorted_data, p) for p in pcts}
def mean(data: List[float]) -> float:
"""算术平均值"""
if not data:
return 0.0
return sum(data) / len(data)
def std(data: List[float]) -> float:
"""样本标准差(n-1)"""
if len(data) < 2:
return 0.0
m = mean(data)
variance = sum((x - m) ** 2 for x in data) / (len(data) - 1)
return math.sqrt(variance)
def min_val(data: List[float]) -> float:
"""最小值"""
return min(data) if data else 0.0
def max_val(data: List[float]) -> float:
"""最大值"""
return max(data) if data else 0.0
+1
View File
@@ -0,0 +1 @@
"""测试用例包:每个用例实现 run_test(client, repeats, stop_event)"""
+32
View File
@@ -0,0 +1,32 @@
"""测试用例公共工具"""
from ..ollama_client import InferenceResult
def make_result(infer: InferenceResult, iteration: int, **extra) -> dict:
"""将一次推理结果转换为用例数据条目"""
item = {
"iteration": iteration,
"ttft_ms": round(infer.ttft_ms, 2),
"prefill_ms": round(infer.prefill_ms, 2),
"decode_speed_tok_s": round(infer.decode_speed_tok_s, 2),
"total_tokens": infer.total_tokens,
"prompt_tokens": infer.prompt_tokens,
"completion_tokens": infer.completion_tokens,
"e2e_ms": round(infer.e2e_ms, 2),
"elapsed_ms": round(infer.e2e_ms, 2),
"response_length": len(infer.response),
}
item.update(extra)
return item
def case_meta(case_id: str, name: str, description: str,
difficulty: str, estimated_seconds: int) -> dict:
"""构造用例元数据"""
return {
"case_id": case_id,
"name": name,
"description": description,
"difficulty": difficulty,
"estimated_seconds": estimated_seconds,
}
+23
View File
@@ -0,0 +1,23 @@
"""用例1:纯文本生成基准 —— 基础生成速度基准测试"""
from ..ollama_client import OllamaClient
from .base import case_meta, make_result
PROMPT = "请用200字左右概述人工智能从诞生至今的关键发展阶段,每段不超过50字。"
def run_test(client: OllamaClient, repeats: int = 3, stop_event=None) -> dict:
"""执行纯文本生成测试"""
results = []
for i in range(repeats):
if stop_event and stop_event.is_set():
break
infer = client.generate(PROMPT, options={"temperature": 0.7, "num_predict": 300})
results.append(make_result(infer, i + 1))
meta = case_meta(
"case_01_generation", "纯文本生成基准",
"基础生成速度基准测试", "简单", 5,
)
meta["results"] = results
meta["prompt_used"] = PROMPT
return meta
+41
View File
@@ -0,0 +1,41 @@
"""用例2:简单工具调用延迟 —— 测量结构化工具输出的调用延迟"""
from ..ollama_client import OllamaClient
from .base import case_meta, make_result
# 模拟"工具调用"场景:要求模型只输出结构化结果(短输出、低复杂度)
SYSTEM = (
"你是一个工具调用助手。你只能输出一个JSON对象,格式为 "
'{"tool": "工具名", "params": {"参数名": "参数值"}},不要输出任何其它文字。'
)
TOOL_PROMPTS = [
("查询当前日期", {"tool": "get_date", "params": {}}),
("获取CPU使用率", {"tool": "get_cpu_usage", "params": {}}),
("列出当前目录文件", {"tool": "list_files", "params": {"path": "/tmp"}}),
]
def run_test(client: OllamaClient, repeats: int = 3, stop_event=None) -> dict:
"""执行简单工具调用延迟测试"""
results = []
for i in range(repeats):
if stop_event and stop_event.is_set():
break
desc, expected = TOOL_PROMPTS[i % len(TOOL_PROMPTS)]
infer = client.generate(
f"请调用工具:{desc}",
system=SYSTEM,
options={"temperature": 0.0, "num_predict": 60},
)
results.append(make_result(
infer, i + 1,
tool=expected["tool"],
matched=(expected["tool"] in infer.response),
))
meta = case_meta(
"case_02_simple_tool", "简单工具调用延迟",
"测量简单工具调用的额外延迟", "简单", 3,
)
meta["results"] = results
return meta
@@ -0,0 +1,72 @@
"""用例3:文件读取+分析 —— 测试文件写入、读取、分析的完整链路"""
import os
import tempfile
from ..ollama_client import OllamaClient
from .base import case_meta, make_result
SAMPLE_CODE = '''"""示例模块:一个简单的计算器"""
class Calculator:
"""四则运算计算器"""
def __init__(self):
self.history = []
def add(self, a, b):
self.history.append(("add", a, b))
return a + b
def divide(self, a, b):
if b == 0:
raise ValueError("除数不能为零")
self.history.append(("divide", a, b))
return a / b
def main():
calc = Calculator()
print(calc.add(1, 2))
print(calc.divide(10, 2))
if __name__ == "__main__":
main()
'''
ANALYSIS_PROMPT = "请分析以下Python代码的结构、功能,并指出其中可能存在的问题:\n\n```python\n{code}\n```"
def run_test(client: OllamaClient, repeats: int = 3, stop_event=None) -> dict:
"""执行文件读取+分析测试"""
results = []
for i in range(repeats):
if stop_event and stop_event.is_set():
break
# 写入临时文件,模拟"生成测试文件"环节
fd, test_path = tempfile.mkstemp(prefix="llm_speed_test_", suffix=".py")
try:
with os.fdopen(fd, "w", encoding="utf-8") as f:
f.write(SAMPLE_CODE)
# 读取文件内容作为模型输入
with open(test_path, "r", encoding="utf-8") as f:
content = f.read()
finally:
try:
os.remove(test_path)
except OSError:
pass
infer = client.generate(
ANALYSIS_PROMPT.format(code=content),
options={"temperature": 0.3, "num_predict": 300},
)
results.append(make_result(infer, i + 1, file_chars=len(content)))
meta = case_meta(
"case_03_file_analysis", "文件读取+分析",
"文件读取与分析完整链路计时", "中等", 2,
)
meta["results"] = results
return meta
@@ -0,0 +1,64 @@
"""用例4:并行工具调用 —— 对比串行 vs 并行的加速比"""
import time
from concurrent.futures import ThreadPoolExecutor
from ..ollama_client import OllamaClient
from .base import case_meta
# 三个独立、互不依赖的"工具调用"请求
PARALLEL_PROMPTS = [
"请用一句话说明今天的天气如何。",
"请用一句话说明Python的优缺点。",
"请用一句话说明如何提高代码质量。",
]
NUM_PARALLEL = len(PARALLEL_PROMPTS)
def run_test(client: OllamaClient, repeats: int = 3, stop_event=None) -> dict:
"""执行串行 vs 并行对比测试"""
results = []
def call_one(prompt: str):
return client.generate(prompt, options={"temperature": 0.3, "num_predict": 80})
# 线程池在循环外创建,复用 worker 线程
with ThreadPoolExecutor(max_workers=NUM_PARALLEL) as executor:
for i in range(repeats):
if stop_event and stop_event.is_set():
break
# 串行:依次执行
serial_start = time.perf_counter()
serial_results = [call_one(p) for p in PARALLEL_PROMPTS]
serial_ms = (time.perf_counter() - serial_start) * 1000
# 并行:线程池并发执行
parallel_start = time.perf_counter()
parallel_results = list(executor.map(call_one, PARALLEL_PROMPTS))
parallel_ms = (time.perf_counter() - parallel_start) * 1000
speedup = (serial_ms / parallel_ms) if parallel_ms > 0 else 0.0
avg_ttft = sum(r.ttft_ms for r in parallel_results) / NUM_PARALLEL
avg_decode = sum(r.decode_speed_tok_s for r in parallel_results) / NUM_PARALLEL
results.append({
"iteration": i + 1,
"parallel_tasks": NUM_PARALLEL,
"serial_ms": round(serial_ms, 2),
"parallel_ms": round(parallel_ms, 2),
"speedup_ratio": round(speedup, 2),
"elapsed_ms": round(parallel_ms, 2), # 用并行耗时作为该用例的代表延迟
"avg_ttft_ms": round(avg_ttft, 2),
"avg_decode_speed_tok_s": round(avg_decode, 2),
"ttft_ms": round(avg_ttft, 2),
"decode_speed_tok_s": round(avg_decode, 2),
})
meta = case_meta(
"case_04_parallel_tool", "并行工具调用",
"对比串行 vs 并行工具调用的延迟差异", "中等", 5,
)
meta["results"] = results
return meta
@@ -0,0 +1,45 @@
"""用例5:长上下文处理 —— 测试不同上下文大小下的处理延迟"""
from ..ollama_client import OllamaClient
from .base import case_meta, make_result
# 约 15 个 token/句的填充句,用于按目标 token 数拼接上下文
FILLER_SENTENCE = "这是一个用于测试长上下文处理能力的示例句子,包含一些常见的词汇和表述方式。\n"
TOKENS_PER_SENTENCE = 15
CONTEXT_SIZES = [500, 1000, 2000, 4000] # 目标 prompt token 数
QUESTION = "请用一句话总结:上面这段文本主要介绍了什么?"
def _build_context(target_tokens: int) -> str:
"""按目标 token 数近似生成一段填充文本"""
n_sentences = max(1, target_tokens // TOKENS_PER_SENTENCE)
return FILLER_SENTENCE * n_sentences
def run_test(client: OllamaClient, repeats: int = 2, stop_event=None) -> dict:
"""执行长上下文处理测试"""
results = []
for context_size in CONTEXT_SIZES:
for r in range(repeats):
if stop_event and stop_event.is_set():
break
context = _build_context(context_size)
prompt = f"{context}\n\n{QUESTION}"
infer = client.generate(
prompt,
options={"temperature": 0.3, "num_predict": 80},
)
results.append(make_result(
infer, r + 1,
context_tokens_target=context_size,
context_chars=len(context),
))
meta = case_meta(
"case_05_long_context", "长上下文处理",
"不同上下文长度对延迟的影响", "耗时", 10,
)
meta["results"] = results
meta["context_sizes_tested"] = CONTEXT_SIZES
return meta
+33
View File
@@ -0,0 +1,33 @@
"""用例6:复杂推理任务 —— 对比不同难度推理任务的延迟差异"""
from ..ollama_client import OllamaClient
from .base import case_meta, make_result
REASONING_LEVELS = [
("简单问答", "什么是人工智能?", 80),
("中等推理", "解释Transformer架构中的自注意力机制是如何工作的。", 200),
("复杂推理", "分析大语言模型在长链推理中的主要局限性,并提出三种改进方案。", 400),
]
def run_test(client: OllamaClient, repeats: int = 2, stop_event=None) -> dict:
"""执行复杂推理任务测试"""
results = []
for level_name, prompt, num_predict in REASONING_LEVELS:
for i in range(repeats):
if stop_event and stop_event.is_set():
break
infer = client.generate(
prompt,
options={"temperature": 0.3, "num_predict": num_predict},
)
results.append(make_result(
infer, i + 1,
reasoning_level=level_name,
))
meta = case_meta(
"case_06_reasoning", "复杂推理任务",
"不同推理复杂度下的延迟对比", "耗时", 15,
)
meta["results"] = results
return meta
+49
View File
@@ -0,0 +1,49 @@
"""用例7:多轮对话累积 —— 追踪连续多轮对话中延迟随上下文累积的变化"""
from ..ollama_client import OllamaClient
from .base import case_meta, make_result
DEBUG_SCENARIOS = [
"修复Python中的索引越界错误",
"修复JavaScript闭包变量捕获问题",
"修复SQL注入漏洞",
"修复多线程竞态条件",
"修复递归栈溢出",
"修复正则表达式灾难回溯",
]
def run_test(client: OllamaClient, repeats: int = 1, stop_event=None) -> dict:
"""执行多轮对话累积测试(使用 /api/chat 累积上下文)
repeats 控制总轮次数,每轮切换不同调试场景。
"""
results = []
messages = []
for i in range(repeats):
if stop_event and stop_event.is_set():
break
scenario = DEBUG_SCENARIOS[i % len(DEBUG_SCENARIOS)]
user_msg = f"(对话第{i + 1}轮)请协助:{scenario}。请简要给出你的分析。"
messages.append({"role": "user", "content": user_msg})
infer = client.chat(messages, options={"temperature": 0.3, "num_predict": 120})
assistant_reply = infer.response
messages.append({"role": "assistant", "content": assistant_reply})
results.append(make_result(
infer, i + 1,
round=i + 1,
scenario=scenario,
messages_in_context=len(messages),
))
meta = case_meta(
"case_07_multiturn", "多轮对话累积",
"追踪连续多轮对话中延迟随上下文累积的变化", "耗时", 10,
)
meta["results"] = results
meta["total_rounds"] = len(results)
return meta
+146
View File
@@ -0,0 +1,146 @@
"""用例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