## 背景 用户反馈动态问诊评测「执行不下去」。诊断发现: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>
335 lines
11 KiB
Python
335 lines
11 KiB
Python
"""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
|
||
|