import asyncio import json from collections.abc import AsyncIterator from dataclasses import dataclass from typing import Any import pytest from agent_lab.application.contracts import AgentParams, DebugRunRequest, EventAgentParams from agent_lab.application.runtime import DebugRuntime from agent_lab.application.tools import build_default_tool_registry from agent_lab.domain.events import ToolCallEvent from agent_lab.domain.messages import ChatMessage, StreamItem MODES = ("dual_agent", "chat_agent_tools") SCENARIOS = ( "ordinary_chat", "device_volume_silent", "calendar_schedule_template", "web_search_two_answers", "session_terminate", "parallel_volume_schedule", ) BUILTIN_EVENT_NAMES = ( "session.terminate", "device.volume.adjust", "calendar.schedule.create", "knowledge.web.search", ) SCHEDULE_ARGUMENTS = { "title": "Design review", "start_at": "2026-07-14T09:30:00+08:00", "timezone": "Asia/Shanghai", } SCHEDULE_CONFIRMATION = ( "Scheduled Design review for 2026-07-14T09:30:00+08:00 " "(Asia/Shanghai)." ) @dataclass(frozen=True) class ExpectedSemantics: visible_messages: tuple[str, ...] tool_names: tuple[str, ...] model_calls: int terminal_count: int = 0 EXPECTED = { "ordinary_chat": ExpectedSemantics(("Ordinary answer.",), (), 1), "device_volume_silent": ExpectedSemantics( (), ("device.volume.adjust",), 1, ), "calendar_schedule_template": ExpectedSemantics( (SCHEDULE_CONFIRMATION,), ("calendar.schedule.create",), 1, ), "web_search_two_answers": ExpectedSemantics( ("I will check.", "Grounded search update."), ("knowledge.web.search",), 2, ), "session_terminate": ExpectedSemantics( ("Goodbye.",), ("session.terminate",), 1, terminal_count=1, ), "parallel_volume_schedule": ExpectedSemantics( (SCHEDULE_CONFIRMATION,), ("device.volume.adjust", "calendar.schedule.create"), 1, ), } class ScriptedChatClient: def __init__(self, rounds: list[list[StreamItem]]) -> None: self.rounds = rounds self.calls = 0 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 messages, tools, params, tool_choice 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() def _event(call_id: str, name: str, arguments: dict[str, Any]) -> ToolCallEvent: return ToolCallEvent( id=call_id, name=name, arguments=arguments, raw_arguments=json.dumps(arguments, sort_keys=True, separators=(",", ":")), ) def _event_item(mode: str, event: ToolCallEvent) -> StreamItem: if mode == "dual_agent": return StreamItem.text_event(event) return StreamItem.provider_tool_call(event) def _scenario_setup( mode: str, scenario: str, ) -> tuple[DebugRunRequest, list[list[StreamItem]]]: if scenario == "ordinary_chat": user_message = "Hello." enabled_tools = list(BUILTIN_EVENT_NAMES) rounds = [[StreamItem.message_delta("Ordinary answer.")]] elif scenario == "device_volume_silent": user_message = "set volume to 30" enabled_tools = ["device.volume.adjust"] rounds = [ [ _event_item( mode, _event( "volume-1", "device.volume.adjust", {"mode": "absolute", "value": 30}, ), ) ] ] elif scenario == "calendar_schedule_template": user_message = ( 'schedule "Design review" at 2026-07-14T09:30:00+08:00 ' "timezone Asia/Shanghai" ) enabled_tools = ["calendar.schedule.create"] rounds = [ [ _event_item( mode, _event( "schedule-1", "calendar.schedule.create", SCHEDULE_ARGUMENTS, ), ) ] ] elif scenario == "web_search_two_answers": user_message = "event batch safety" enabled_tools = ["knowledge.web.search"] rounds = [ [ StreamItem.message_delta("I will check."), _event_item( mode, _event( "search-1", "knowledge.web.search", {"query": user_message}, ), ), ], [StreamItem.message_delta("Grounded search update.")], ] elif scenario == "session_terminate": user_message = "end this session" enabled_tools = ["session.terminate"] rounds = [ [ _event_item( mode, _event("terminate-1", "session.terminate", {}), ) ] ] elif scenario == "parallel_volume_schedule": user_message = ( 'set volume to 30 and schedule "Design review" at ' "2026-07-14T09:30:00+08:00 timezone Asia/Shanghai" ) enabled_tools = ["device.volume.adjust", "calendar.schedule.create"] rounds = [ [ _event_item( mode, _event( "parallel-volume", "device.volume.adjust", {"mode": "absolute", "value": 30}, ), ), _event_item( mode, _event( "parallel-schedule", "calendar.schedule.create", SCHEDULE_ARGUMENTS, ), ), ] ] else: raise AssertionError(f"unknown comparison scenario: {scenario}") request = DebugRunRequest( user_message=user_message, system_prompts=["Exercise the requested business scenario."], pre_messages=[], chat_agent=AgentParams( model="comparison-model", temperature=0.0, max_tokens=128, ), event_agent=EventAgentParams( model="comparison-model", temperature=0.0, max_tokens=128, enabled_tools=enabled_tools, max_event_loops=1, max_parallel_events=2, batch_timeout_seconds=1.0, ), tool_invocation_mode=mode, ) return request, rounds 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 _assert_port_calls(scenario: str, 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 @pytest.mark.asyncio @pytest.mark.parametrize("mode", MODES) @pytest.mark.parametrize("scenario", SCENARIOS) async def test_invocation_modes_share_business_semantics( mode: str, scenario: str, ) -> None: request, rounds = _scenario_setup(mode, scenario) 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) expected = EXPECTED[scenario] 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 == expected.visible_messages assert tuple(payload["tool"] for payload in tool_payloads) == expected.tool_names assert client.calls == expected.model_calls assert terminal_count == expected.terminal_count assert outputs[-1] == {"type": "done"} _assert_port_calls(scenario, ports) if scenario == "web_search_two_answers": assert tool_payloads[0]["sources"][0]["title"] == "Matrix source" business = [message for message in outputs if message["type"] != "audit"] 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]