| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384 |
- import asyncio
- import json
- from collections.abc import AsyncIterator
- from dataclasses import dataclass
- from enum import StrEnum
- from typing import Any
- import pytest
- from agent_lab.application.benchmark import (
- BENCHMARK_CASE_CATALOG,
- BENCHMARK_CASE_IDS,
- BENCHMARK_MODES,
- BenchmarkCaseId,
- BenchmarkMode,
- build_benchmark_request,
- build_mock_rounds,
- )
- from agent_lab.application.contracts import AgentParams
- from agent_lab.application.runtime import DebugRuntime
- from agent_lab.application.tools import build_default_tool_registry
- from agent_lab.domain.messages import ChatMessage, StreamItem
- @dataclass(frozen=True)
- class RecordedChatRequest:
- messages: tuple[ChatMessage, ...]
- tools: tuple[dict[str, Any], ...]
- tool_choice: dict[str, Any] | None
- class ScriptedChatClient:
- def __init__(self, rounds: list[list[StreamItem]]) -> None:
- self.rounds = rounds
- self.calls = 0
- self.requests: list[RecordedChatRequest] = []
- async def stream_chat(
- self,
- messages: list[ChatMessage],
- tools: list[dict[str, Any]],
- params: AgentParams,
- tool_choice: dict[str, Any] | None = None,
- ) -> AsyncIterator[StreamItem]:
- del params
- self.requests.append(
- RecordedChatRequest(
- messages=tuple(messages),
- tools=tuple(tools),
- tool_choice=dict(tool_choice) if tool_choice is not None else None,
- )
- )
- items = self.rounds[self.calls]
- self.calls += 1
- for item in items:
- yield item
- class RecordingPorts:
- def __init__(self, *, parallel_barrier: bool = False) -> None:
- self.parallel_barrier = parallel_barrier
- self.session_calls: list[tuple[str, str | None]] = []
- self.volume_calls: list[tuple[str, str, int | None, int | None]] = []
- self.calendar_calls: list[
- tuple[str, str, str, str, str | None, int | None]
- ] = []
- self.search_calls: list[tuple[str, str, int]] = []
- self.started: set[str] = set()
- self.both_started = asyncio.Event()
- self.volume_release = asyncio.Event()
- self.schedule_release = asyncio.Event()
- self.schedule_finished = asyncio.Event()
- if not parallel_barrier:
- self.volume_release.set()
- self.schedule_release.set()
- async def terminate(
- self,
- event_id: str,
- *,
- reason: str | None = None,
- ) -> dict[str, Any]:
- self.session_calls.append((event_id, reason))
- return {
- "tool": "session.terminate",
- "status": "terminated",
- "event_id": event_id,
- "reason": reason,
- }
- async def adjust(
- self,
- event_id: str,
- *,
- mode: str,
- value: int | None = None,
- delta: int | None = None,
- ) -> dict[str, Any]:
- self.volume_calls.append((event_id, mode, value, delta))
- await self._wait_at_barrier("volume", self.volume_release)
- payload: dict[str, Any] = {
- "tool": "device.volume.adjust",
- "status": "applied",
- "event_id": event_id,
- "mode": mode,
- }
- if value is not None:
- payload["value"] = value
- if delta is not None:
- payload["delta"] = delta
- return payload
- async def create(
- self,
- event_id: str,
- *,
- title: str,
- start_at: str,
- timezone: str,
- recurrence: str | None = None,
- reminder_minutes: int | None = None,
- ) -> dict[str, Any]:
- self.calendar_calls.append(
- (
- event_id,
- title,
- start_at,
- timezone,
- recurrence,
- reminder_minutes,
- )
- )
- await self._wait_at_barrier("schedule", self.schedule_release)
- self.schedule_finished.set()
- return {
- "tool": "calendar.schedule.create",
- "status": "created",
- "event_id": event_id,
- "schedule": {
- "title": title,
- "start_at": start_at,
- "timezone": timezone,
- },
- }
- async def search(
- self,
- event_id: str,
- *,
- query: str,
- max_results: int = 3,
- ) -> dict[str, Any]:
- self.search_calls.append((event_id, query, max_results))
- return {
- "tool": "knowledge.web.search",
- "query": query,
- "sources": [
- {
- "title": "Matrix source",
- "url": "https://example.invalid/matrix",
- "snippet": f"Deterministic context for: {query}",
- }
- ],
- "retrieved_at": "2026-07-13T00:00:00Z",
- }
- async def _wait_at_barrier(
- self,
- name: str,
- release: asyncio.Event,
- ) -> None:
- if not self.parallel_barrier:
- return
- self.started.add(name)
- if len(self.started) == 2:
- self.both_started.set()
- await release.wait()
- async def _collect_outputs(
- stream: AsyncIterator[dict[str, Any]],
- ) -> list[dict[str, Any]]:
- return [message async for message in stream]
- def _tool_payloads(outputs: list[dict[str, Any]]) -> list[dict[str, Any]]:
- return [
- json.loads(message["message"]["content"])
- for message in outputs
- if message["type"] == "tool_result"
- ]
- def _generated_event_catalog_names(
- messages: tuple[ChatMessage, ...],
- ) -> tuple[str, ...]:
- catalogs = [
- message.content
- for message in messages
- if message.role == "system"
- and "Available events:" in message.content
- and "<agent_events>" in message.content
- ]
- if not catalogs:
- return ()
- assert len(catalogs) == 1
- return tuple(
- line.removeprefix("- ").split(":", 1)[0]
- for line in catalogs[0].splitlines()
- if line.startswith("- ")
- )
- def _provider_tool_names(
- tools: tuple[dict[str, Any], ...],
- ) -> tuple[str, ...]:
- return tuple(tool["function"]["name"] for tool in tools)
- def _assert_port_calls(scenario: BenchmarkCaseId, ports: RecordingPorts) -> None:
- expected_session: list[tuple[str, str | None]] = []
- expected_volume: list[tuple[str, str, int | None, int | None]] = []
- expected_calendar: list[
- tuple[str, str, str, str, str | None, int | None]
- ] = []
- expected_search: list[tuple[str, str, int]] = []
- if scenario == "device_volume_silent":
- expected_volume = [("volume-1", "absolute", 30, None)]
- elif scenario == "calendar_schedule_template":
- expected_calendar = [
- (
- "schedule-1",
- "Design review",
- "2026-07-14T09:30:00+08:00",
- "Asia/Shanghai",
- None,
- None,
- )
- ]
- elif scenario == "web_search_two_answers":
- expected_search = [("search-1", "event batch safety", 3)]
- elif scenario == "session_terminate":
- expected_session = [("terminate-1", None)]
- elif scenario == "parallel_volume_schedule":
- expected_volume = [("parallel-volume", "absolute", 30, None)]
- expected_calendar = [
- (
- "parallel-schedule",
- "Design review",
- "2026-07-14T09:30:00+08:00",
- "Asia/Shanghai",
- None,
- None,
- )
- ]
- assert ports.session_calls == expected_session
- assert ports.volume_calls == expected_volume
- assert ports.calendar_calls == expected_calendar
- assert ports.search_calls == expected_search
- def test_invocation_matrix_remains_twelve_enum_runs():
- matrix = tuple(
- (mode, scenario)
- for scenario in BENCHMARK_CASE_IDS
- for mode in BENCHMARK_MODES
- )
- assert len(matrix) == 12
- assert len(set(matrix)) == 12
- assert all(isinstance(mode, StrEnum) for mode, _ in matrix)
- assert all(isinstance(scenario, StrEnum) for _, scenario in matrix)
- assert {type(mode) for mode, _ in matrix} == {BenchmarkMode}
- assert {type(scenario) for _, scenario in matrix} == {BenchmarkCaseId}
- @pytest.mark.asyncio
- @pytest.mark.parametrize("mode", BENCHMARK_MODES)
- @pytest.mark.parametrize("scenario", BENCHMARK_CASE_IDS)
- async def test_invocation_modes_share_business_semantics(
- mode: BenchmarkMode,
- scenario: BenchmarkCaseId,
- ) -> None:
- case = BENCHMARK_CASE_CATALOG[scenario]
- request = build_benchmark_request(case, mode, "comparison-model")
- rounds = build_mock_rounds(case, mode)
- ports = RecordingPorts(parallel_barrier=scenario == "parallel_volume_schedule")
- registry = build_default_tool_registry(
- session_termination_port=ports,
- device_volume_port=ports,
- calendar_schedule_port=ports,
- web_search_port=ports,
- )
- client = ScriptedChatClient(rounds)
- run_task = asyncio.create_task(
- _collect_outputs(DebugRuntime(client, registry=registry).run(request))
- )
- try:
- if scenario == "parallel_volume_schedule":
- await asyncio.wait_for(ports.both_started.wait(), timeout=1)
- assert ports.started == {"volume", "schedule"}
- ports.schedule_release.set()
- await asyncio.wait_for(ports.schedule_finished.wait(), timeout=1)
- assert not run_task.done()
- ports.volume_release.set()
- outputs = await asyncio.wait_for(run_task, timeout=2)
- finally:
- ports.schedule_release.set()
- ports.volume_release.set()
- if not run_task.done():
- run_task.cancel()
- await asyncio.gather(run_task, return_exceptions=True)
- visible_messages = tuple(
- message["content"]
- for message in outputs
- if message["type"] == "message_delta"
- )
- tool_payloads = _tool_payloads(outputs)
- terminal_count = sum(
- message.get("event") == "terminal_completed" for message in outputs
- )
- assert visible_messages == case.expectation.visible_messages
- assert tuple(payload["tool"] for payload in tool_payloads) == (
- case.expectation.event_names
- )
- assert client.calls == len(rounds)
- assert len(client.requests) == len(rounds)
- assert terminal_count == int(case.expectation.terminal)
- assert outputs[-1] == {"type": "done"}
- _assert_port_calls(scenario, ports)
- initial_request = client.requests[0]
- enabled_names = tuple(request.event_agent.enabled_tools)
- assert initial_request.tool_choice is None
- if mode == "dual_agent":
- assert initial_request.tools == ()
- assert _generated_event_catalog_names(initial_request.messages) == enabled_names
- else:
- assert _provider_tool_names(initial_request.tools) == enabled_names
- assert _generated_event_catalog_names(initial_request.messages) == ()
- if case.expectation.event_names:
- business = [message for message in outputs if message["type"] != "audit"]
- first_message_index = next(
- index
- for index, message in enumerate(business)
- if message["type"] == "message_delta"
- )
- first_tool_result_index = next(
- index
- for index, message in enumerate(business)
- if message["type"] == "tool_result"
- )
- assert first_message_index < first_tool_result_index
- detected_events = [
- message
- for message in outputs
- if message.get("event") == "chat_event_detected"
- ]
- expected_source = (
- "text_event" if mode == "dual_agent" else "provider_resolved"
- )
- assert tuple(
- message["details"]["event_source"] for message in detected_events
- ) == (expected_source,) * len(case.expectation.event_names)
- if scenario == "web_search_two_answers":
- assert tool_payloads[0]["sources"][0]["title"] == "Matrix source"
- visible_indexes = [
- index
- for index, message in enumerate(business)
- if message["type"] == "message_delta"
- ]
- tool_index = next(
- index
- for index, message in enumerate(business)
- if message["type"] == "tool_result"
- )
- assert visible_indexes[0] < tool_index < visible_indexes[1]
|