AgentEvalTool/tests/unit/test_engine.py
sinohqb 867d4e3ff1 fix(engine): dynamic 生成失败原因持久化到 run.summary.case_errors
## 背景
用户反馈动态问诊评测「执行不下去」。诊断发现:dynamic 用例的 LLM 消息
生成 API 调用失败(0.2s 瞬间 failed,凭证/参数问题),引擎正确地标记
case 失败——但失败的具体原因(如 401 详情)只 emit 到 WebSocket,从不
写入 run.summary。导致 run 记录只有 total_rules:0 failed,DB/报告查不到
任何原因,用户和排查者都无从下手。

## 修复
- EvalEngine 新增 self._case_errors 收集致命的 case 级错误
- _generate_messages 的 8 个失败点统一走 _fail() helper:既 emit 到
  WebSocket,也记录到 _case_errors(含 case_id + stage + 具体 error)
- run() 汇总时把 _case_errors 写入 summary["case_errors"]
- 新增测试:dynamic 生成失败时 summary.case_errors 必须含原因(补上
  之前 KNOWN-2 记录的 _generate_messages 测试盲区)

## 注
这不是导致失败的 bug(失败源于外部 API 凭证/参数),而是让失败「可诊断」
的可用性修复。用户需自查 llm_config 的 api_key 是否有效/model 是否被
该端点接受。

Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-17 16:29:51 +08:00

335 lines
11 KiB
Python
Raw 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.

"""Unit tests for the async EvalEngine.
Focus on the new async behavior: cooperative cancellation, configurable
timeouts, and concurrent case execution via the semaphore.
"""
import asyncio
from typing import Any
import pytest
from agenteval.channels.base import SendResult
from agenteval.evaluation.engine import CancelledError, EvalEngine, TimeoutConfig
from agenteval.models import (
Case, CaseType, ChannelType, EvalTarget, Expectation, PlatformType,
RunStatus, Scenario, TargetStatus,
)
from agenteval.storage.repository import RunRepository
from tests.unit.mock_channel import MockChannel
def _make_target() -> EvalTarget:
return EvalTarget(
id="t-1",
name="mock-target",
platform=PlatformType.AI_DIGITAL_EMPLOYEE,
channel_type=ChannelType.TUTU_API,
channel_config={
"base_url": "http://mock",
"token": "x",
"tenant": "t",
"chat_channel_id": "c",
"chat_contact_id": "u",
},
status=TargetStatus.ACTIVE,
)
def _build_engine(
scenario: Scenario,
channel: MockChannel,
*,
cancel_token: asyncio.Event | None = None,
timeout_config: TimeoutConfig | None = None,
max_concurrent_cases: int = 1,
session=None,
) -> EvalEngine:
"""Construct an EvalEngine with a MockChannel injected.
The factory still runs (needs a valid channel_config on the target) but
we immediately replace the channel with the mock.
"""
engine = EvalEngine(
target=_make_target(),
scenario=scenario,
session=session,
cancel_token=cancel_token or asyncio.Event(),
timeout_config=timeout_config,
max_concurrent_cases=max_concurrent_cases,
)
engine.channel = channel
return engine
# ── basic happy path ─────────────────────────────────────────────────────
async def test_run_single_case_completes(db_session):
scenario = Scenario(
id="s-1", name="single",
cases=[Case(id="c1", type=CaseType.SINGLE, messages=["hello"])],
)
channel = MockChannel()
engine = _build_engine(scenario, channel, session=db_session)
run = await engine.run()
assert run.status == RunStatus.COMPLETED
assert run.summary["total_cases"] == 1
assert run.summary["passed_cases"] == 1
assert channel.send_calls == 1
assert channel.poll_calls == 1
async def test_run_multi_turn_collects_all_turns(db_session):
scenario = Scenario(
id="s-1", name="multi",
cases=[Case(id="c1", type=CaseType.MULTI_TURN, messages=["a", "b", "c"])],
)
channel = MockChannel()
engine = _build_engine(scenario, channel, session=db_session)
run = await engine.run()
assert run.status == RunStatus.COMPLETED
assert channel.sent == ["a", "b", "c"]
assert channel.send_calls == 3
# Each turn polls once.
assert channel.poll_calls == 3
async def test_run_multiple_cases(db_session):
scenario = Scenario(
id="s-1", name="multi-case",
cases=[
Case(id="c1", type=CaseType.SINGLE, messages=["one"]),
Case(id="c2", type=CaseType.SINGLE, messages=["two"]),
Case(id="c3", type=CaseType.SINGLE, messages=["three"]),
],
)
channel = MockChannel()
engine = _build_engine(scenario, channel, session=db_session)
run = await engine.run()
assert run.status == RunStatus.COMPLETED
assert run.summary["total_cases"] == 3
assert channel.sent == ["one", "two", "three"]
# ── progress callback ────────────────────────────────────────────────────
async def test_progress_callback_receives_events(db_session):
scenario = Scenario(
id="s-1", name="single",
cases=[Case(id="c1", type=CaseType.SINGLE, messages=["hi"])],
)
events: list[tuple[str, dict]] = []
async def cb(event: str, data: dict[str, Any]) -> None:
events.append((event, data))
channel = MockChannel()
engine = _build_engine(scenario, channel, session=db_session)
await engine.run(progress_callback=cb)
event_names = [e[0] for e in events]
assert "case_start" in event_names
assert "turn_start" in event_names
assert "turn_end" in event_names
assert "case_end" in event_names
assert "run_completed" in event_names
async def test_sync_progress_callback_also_works(db_session):
"""Engine must accept sync callbacks (CLI uses them)."""
scenario = Scenario(
id="s-1", name="single",
cases=[Case(id="c1", type=CaseType.SINGLE, messages=["hi"])],
)
seen: list[str] = []
def sync_cb(event: str, data: dict) -> None:
seen.append(event)
channel = MockChannel()
engine = _build_engine(scenario, channel, session=db_session)
await engine.run(progress_callback=sync_cb)
assert "run_completed" in seen
# ── cancellation ─────────────────────────────────────────────────────────
async def test_cancel_before_run_marks_failed(db_session):
scenario = Scenario(
id="s-1", name="single",
cases=[Case(id="c1", type=CaseType.SINGLE, messages=["hi"])],
)
cancel_token = asyncio.Event()
cancel_token.set() # already cancelled
channel = MockChannel()
engine = _build_engine(scenario, channel, cancel_token=cancel_token, session=db_session)
run = await engine.run()
assert run.status == RunStatus.FAILED
assert run.summary["error"]["code"] == "cancelled_by_user"
# Engine must NOT have called the channel.
assert channel.send_calls == 0
async def test_cancel_mid_run_stops_after_current_case(db_session):
"""Cancelling between cases should stop further cases from running."""
scenario = Scenario(
id="s-1", name="multi",
cases=[
Case(id="c1", type=CaseType.SINGLE, messages=["a"]),
Case(id="c2", type=CaseType.SINGLE, messages=["b"]),
Case(id="c3", type=CaseType.SINGLE, messages=["c"]),
],
)
cancel_token = asyncio.Event()
# Slow channel so we have time to set the cancel token from another task.
channel = MockChannel(reply_delay=0.05)
async def cancel_soon() -> None:
await asyncio.sleep(0.08)
cancel_token.set()
engine = _build_engine(scenario, channel, cancel_token=cancel_token, session=db_session)
run, _ = await asyncio.gather(
engine.run(),
cancel_soon(),
)
assert run.status == RunStatus.FAILED
assert run.summary["error"]["code"] == "cancelled_by_user"
# At least one case ran, but not all three.
assert 1 <= channel.send_calls < 3
# ── timeouts ─────────────────────────────────────────────────────────────
async def test_poll_timeout_records_missing_reply(db_session):
"""When the channel returns no reply within the timeout, the turn is saved
with reply=None but the engine keeps going (doesn't crash)."""
scenario = Scenario(
id="s-1", name="single",
cases=[Case(id="c1", type=CaseType.SINGLE, messages=["hi"])],
)
channel = MockChannel(missing_reply=True)
engine = _build_engine(
scenario, channel,
timeout_config=TimeoutConfig(poll_reply=0.1),
session=db_session,
)
run = await engine.run()
# Engine completes (doesn't hang), turn has no reply.
assert run.status == RunStatus.COMPLETED
turns = RunRepository(db_session).get_turns(run.id)
assert len(turns) == 1
assert turns[0].reply is None
# ── concurrency ──────────────────────────────────────────────────────────
async def test_concurrent_cases_respect_semaphore(db_session):
"""With max_concurrent_cases=2, at most 2 cases should be in-flight."""
scenario = Scenario(
id="s-1", name="concurrent",
cases=[
Case(id=f"c{i}", type=CaseType.SINGLE, messages=[f"m{i}"])
for i in range(5)
],
)
in_flight = {"n": 0, "max": 0}
lock = asyncio.Lock()
class TrackedChannel(MockChannel):
async def send(self, content: str, **kwargs):
async with lock:
in_flight["n"] += 1
in_flight["max"] = max(in_flight["max"], in_flight["n"])
await asyncio.sleep(0.05)
result = await super().send(content, **kwargs)
async with lock:
in_flight["n"] -= 1
return result
channel = TrackedChannel()
engine = _build_engine(
scenario, channel,
max_concurrent_cases=2,
session=db_session,
)
run = await engine.run()
assert run.status == RunStatus.COMPLETED
assert channel.send_calls == 5
assert in_flight["max"] <= 2
# ── send failure ─────────────────────────────────────────────────────────
async def test_send_failure_aborts_case(db_session):
"""When channel.send fails, the case is marked failed but the engine
continues with the next case."""
scenario = Scenario(
id="s-1", name="mixed",
cases=[
Case(id="c1", type=CaseType.SINGLE, messages=["a"]),
Case(id="c2", type=CaseType.SINGLE, messages=["b"]),
],
)
call_count = {"n": 0}
class FlakeyChannel(MockChannel):
async def send(self, content, **kwargs):
call_count["n"] += 1
if call_count["n"] == 1:
return SendResult(ok=False, error="boom")
return await super().send(content, **kwargs)
channel = FlakeyChannel()
engine = _build_engine(scenario, channel, session=db_session)
run = await engine.run()
assert run.status == RunStatus.COMPLETED
assert run.summary["failed_cases"] == 1
assert run.summary["passed_cases"] == 1
# ── dynamic case generation failure ──────────────────────────────────────
async def test_dynamic_generation_failure_records_case_error(db_session):
"""A dynamic case whose message generation fails must persist the reason
into run.summary.case_errors — not just emit it transiently. Otherwise a
run shows 0 rules / failed with no discoverable cause."""
scenario = Scenario(
id="s-1", name="dynamic",
cases=[Case(id="dyn-1", type=CaseType.DYNAMIC, prompt="生成问题", turns=3)],
llm_config=None, # 缺 llm_config → 生成消息立即失败
)
channel = MockChannel()
engine = _build_engine(scenario, channel, session=db_session)
run = await engine.run()
assert run.status == RunStatus.COMPLETED
assert run.summary["failed_cases"] == 1
assert run.summary["total_rules"] == 0
# 关键:失败原因被持久化到 summary可在报告 / DB 查看
assert "case_errors" in run.summary
assert run.summary["case_errors"][0]["case_id"] == "dyn-1"
assert "llm_config" in run.summary["case_errors"][0]["error"]
# 被测通道不应被调用(生成阶段就失败了)
assert channel.send_calls == 0