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