241 lines
9.1 KiB
Python
241 lines
9.1 KiB
Python
"""测试引擎:编排测试用例、统计指标、保存结果并生成报告"""
|
||
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)
|