Files

241 lines
9.1 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""测试引擎:编排测试用例、统计指标、保存结果并生成报告"""
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)