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 "" 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]