test_tool_invocation_comparison.py 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366
  1. import asyncio
  2. import json
  3. from collections.abc import AsyncIterator
  4. from dataclasses import dataclass
  5. from typing import Any
  6. import pytest
  7. from agent_lab.application.benchmark import (
  8. BENCHMARK_CASE_CATALOG,
  9. BENCHMARK_CASE_IDS,
  10. BENCHMARK_MODES,
  11. build_benchmark_request,
  12. build_mock_rounds,
  13. )
  14. from agent_lab.application.contracts import AgentParams
  15. from agent_lab.application.runtime import DebugRuntime
  16. from agent_lab.application.tools import build_default_tool_registry
  17. from agent_lab.domain.messages import ChatMessage, StreamItem
  18. @dataclass(frozen=True)
  19. class RecordedChatRequest:
  20. messages: tuple[ChatMessage, ...]
  21. tools: tuple[dict[str, Any], ...]
  22. tool_choice: dict[str, Any] | None
  23. class ScriptedChatClient:
  24. def __init__(self, rounds: list[list[StreamItem]]) -> None:
  25. self.rounds = rounds
  26. self.calls = 0
  27. self.requests: list[RecordedChatRequest] = []
  28. async def stream_chat(
  29. self,
  30. messages: list[ChatMessage],
  31. tools: list[dict[str, Any]],
  32. params: AgentParams,
  33. tool_choice: dict[str, Any] | None = None,
  34. ) -> AsyncIterator[StreamItem]:
  35. del params
  36. self.requests.append(
  37. RecordedChatRequest(
  38. messages=tuple(messages),
  39. tools=tuple(tools),
  40. tool_choice=dict(tool_choice) if tool_choice is not None else None,
  41. )
  42. )
  43. items = self.rounds[self.calls]
  44. self.calls += 1
  45. for item in items:
  46. yield item
  47. class RecordingPorts:
  48. def __init__(self, *, parallel_barrier: bool = False) -> None:
  49. self.parallel_barrier = parallel_barrier
  50. self.session_calls: list[tuple[str, str | None]] = []
  51. self.volume_calls: list[tuple[str, str, int | None, int | None]] = []
  52. self.calendar_calls: list[
  53. tuple[str, str, str, str, str | None, int | None]
  54. ] = []
  55. self.search_calls: list[tuple[str, str, int]] = []
  56. self.started: set[str] = set()
  57. self.both_started = asyncio.Event()
  58. self.volume_release = asyncio.Event()
  59. self.schedule_release = asyncio.Event()
  60. self.schedule_finished = asyncio.Event()
  61. if not parallel_barrier:
  62. self.volume_release.set()
  63. self.schedule_release.set()
  64. async def terminate(
  65. self,
  66. event_id: str,
  67. *,
  68. reason: str | None = None,
  69. ) -> dict[str, Any]:
  70. self.session_calls.append((event_id, reason))
  71. return {
  72. "tool": "session.terminate",
  73. "status": "terminated",
  74. "event_id": event_id,
  75. "reason": reason,
  76. }
  77. async def adjust(
  78. self,
  79. event_id: str,
  80. *,
  81. mode: str,
  82. value: int | None = None,
  83. delta: int | None = None,
  84. ) -> dict[str, Any]:
  85. self.volume_calls.append((event_id, mode, value, delta))
  86. await self._wait_at_barrier("volume", self.volume_release)
  87. payload: dict[str, Any] = {
  88. "tool": "device.volume.adjust",
  89. "status": "applied",
  90. "event_id": event_id,
  91. "mode": mode,
  92. }
  93. if value is not None:
  94. payload["value"] = value
  95. if delta is not None:
  96. payload["delta"] = delta
  97. return payload
  98. async def create(
  99. self,
  100. event_id: str,
  101. *,
  102. title: str,
  103. start_at: str,
  104. timezone: str,
  105. recurrence: str | None = None,
  106. reminder_minutes: int | None = None,
  107. ) -> dict[str, Any]:
  108. self.calendar_calls.append(
  109. (
  110. event_id,
  111. title,
  112. start_at,
  113. timezone,
  114. recurrence,
  115. reminder_minutes,
  116. )
  117. )
  118. await self._wait_at_barrier("schedule", self.schedule_release)
  119. self.schedule_finished.set()
  120. return {
  121. "tool": "calendar.schedule.create",
  122. "status": "created",
  123. "event_id": event_id,
  124. "schedule": {
  125. "title": title,
  126. "start_at": start_at,
  127. "timezone": timezone,
  128. },
  129. }
  130. async def search(
  131. self,
  132. event_id: str,
  133. *,
  134. query: str,
  135. max_results: int = 3,
  136. ) -> dict[str, Any]:
  137. self.search_calls.append((event_id, query, max_results))
  138. return {
  139. "tool": "knowledge.web.search",
  140. "query": query,
  141. "sources": [
  142. {
  143. "title": "Matrix source",
  144. "url": "https://example.invalid/matrix",
  145. "snippet": f"Deterministic context for: {query}",
  146. }
  147. ],
  148. "retrieved_at": "2026-07-13T00:00:00Z",
  149. }
  150. async def _wait_at_barrier(
  151. self,
  152. name: str,
  153. release: asyncio.Event,
  154. ) -> None:
  155. if not self.parallel_barrier:
  156. return
  157. self.started.add(name)
  158. if len(self.started) == 2:
  159. self.both_started.set()
  160. await release.wait()
  161. async def _collect_outputs(
  162. stream: AsyncIterator[dict[str, Any]],
  163. ) -> list[dict[str, Any]]:
  164. return [message async for message in stream]
  165. def _tool_payloads(outputs: list[dict[str, Any]]) -> list[dict[str, Any]]:
  166. return [
  167. json.loads(message["message"]["content"])
  168. for message in outputs
  169. if message["type"] == "tool_result"
  170. ]
  171. def _generated_event_catalog_names(
  172. messages: tuple[ChatMessage, ...],
  173. ) -> tuple[str, ...]:
  174. catalogs = [
  175. message.content
  176. for message in messages
  177. if message.role == "system"
  178. and "Available events:" in message.content
  179. and "<agent_events>" in message.content
  180. ]
  181. if not catalogs:
  182. return ()
  183. assert len(catalogs) == 1
  184. return tuple(
  185. line.removeprefix("- ").split(":", 1)[0]
  186. for line in catalogs[0].splitlines()
  187. if line.startswith("- ")
  188. )
  189. def _provider_tool_names(
  190. tools: tuple[dict[str, Any], ...],
  191. ) -> tuple[str, ...]:
  192. return tuple(tool["function"]["name"] for tool in tools)
  193. def _assert_port_calls(scenario: str, ports: RecordingPorts) -> None:
  194. expected_session: list[tuple[str, str | None]] = []
  195. expected_volume: list[tuple[str, str, int | None, int | None]] = []
  196. expected_calendar: list[
  197. tuple[str, str, str, str, str | None, int | None]
  198. ] = []
  199. expected_search: list[tuple[str, str, int]] = []
  200. if scenario == "device_volume_silent":
  201. expected_volume = [("volume-1", "absolute", 30, None)]
  202. elif scenario == "calendar_schedule_template":
  203. expected_calendar = [
  204. (
  205. "schedule-1",
  206. "Design review",
  207. "2026-07-14T09:30:00+08:00",
  208. "Asia/Shanghai",
  209. None,
  210. None,
  211. )
  212. ]
  213. elif scenario == "web_search_two_answers":
  214. expected_search = [("search-1", "event batch safety", 3)]
  215. elif scenario == "session_terminate":
  216. expected_session = [("terminate-1", None)]
  217. elif scenario == "parallel_volume_schedule":
  218. expected_volume = [("parallel-volume", "absolute", 30, None)]
  219. expected_calendar = [
  220. (
  221. "parallel-schedule",
  222. "Design review",
  223. "2026-07-14T09:30:00+08:00",
  224. "Asia/Shanghai",
  225. None,
  226. None,
  227. )
  228. ]
  229. assert ports.session_calls == expected_session
  230. assert ports.volume_calls == expected_volume
  231. assert ports.calendar_calls == expected_calendar
  232. assert ports.search_calls == expected_search
  233. @pytest.mark.asyncio
  234. @pytest.mark.parametrize("mode", BENCHMARK_MODES)
  235. @pytest.mark.parametrize("scenario", BENCHMARK_CASE_IDS)
  236. async def test_invocation_modes_share_business_semantics(
  237. mode: str,
  238. scenario: str,
  239. ) -> None:
  240. case = BENCHMARK_CASE_CATALOG[scenario]
  241. request = build_benchmark_request(case, mode, "comparison-model")
  242. rounds = build_mock_rounds(case, mode)
  243. ports = RecordingPorts(parallel_barrier=scenario == "parallel_volume_schedule")
  244. registry = build_default_tool_registry(
  245. session_termination_port=ports,
  246. device_volume_port=ports,
  247. calendar_schedule_port=ports,
  248. web_search_port=ports,
  249. )
  250. client = ScriptedChatClient(rounds)
  251. run_task = asyncio.create_task(
  252. _collect_outputs(DebugRuntime(client, registry=registry).run(request))
  253. )
  254. try:
  255. if scenario == "parallel_volume_schedule":
  256. await asyncio.wait_for(ports.both_started.wait(), timeout=1)
  257. assert ports.started == {"volume", "schedule"}
  258. ports.schedule_release.set()
  259. await asyncio.wait_for(ports.schedule_finished.wait(), timeout=1)
  260. assert not run_task.done()
  261. ports.volume_release.set()
  262. outputs = await asyncio.wait_for(run_task, timeout=2)
  263. finally:
  264. ports.schedule_release.set()
  265. ports.volume_release.set()
  266. if not run_task.done():
  267. run_task.cancel()
  268. await asyncio.gather(run_task, return_exceptions=True)
  269. visible_messages = tuple(
  270. message["content"]
  271. for message in outputs
  272. if message["type"] == "message_delta"
  273. )
  274. tool_payloads = _tool_payloads(outputs)
  275. terminal_count = sum(
  276. message.get("event") == "terminal_completed" for message in outputs
  277. )
  278. assert visible_messages == case.expectation.visible_messages
  279. assert tuple(payload["tool"] for payload in tool_payloads) == (
  280. case.expectation.event_names
  281. )
  282. assert client.calls == len(rounds)
  283. assert len(client.requests) == len(rounds)
  284. assert terminal_count == int(case.expectation.terminal)
  285. assert outputs[-1] == {"type": "done"}
  286. _assert_port_calls(scenario, ports)
  287. initial_request = client.requests[0]
  288. enabled_names = tuple(request.event_agent.enabled_tools)
  289. assert initial_request.tool_choice is None
  290. if mode == "dual_agent":
  291. assert initial_request.tools == ()
  292. assert _generated_event_catalog_names(initial_request.messages) == enabled_names
  293. else:
  294. assert _provider_tool_names(initial_request.tools) == enabled_names
  295. assert _generated_event_catalog_names(initial_request.messages) == ()
  296. if case.expectation.event_names:
  297. business = [message for message in outputs if message["type"] != "audit"]
  298. first_message_index = next(
  299. index
  300. for index, message in enumerate(business)
  301. if message["type"] == "message_delta"
  302. )
  303. first_tool_result_index = next(
  304. index
  305. for index, message in enumerate(business)
  306. if message["type"] == "tool_result"
  307. )
  308. assert first_message_index < first_tool_result_index
  309. detected_events = [
  310. message
  311. for message in outputs
  312. if message.get("event") == "chat_event_detected"
  313. ]
  314. expected_source = (
  315. "text_event" if mode == "dual_agent" else "provider_resolved"
  316. )
  317. assert tuple(
  318. message["details"]["event_source"] for message in detected_events
  319. ) == (expected_source,) * len(case.expectation.event_names)
  320. if scenario == "web_search_two_answers":
  321. assert tool_payloads[0]["sources"][0]["title"] == "Matrix source"
  322. visible_indexes = [
  323. index
  324. for index, message in enumerate(business)
  325. if message["type"] == "message_delta"
  326. ]
  327. tool_index = next(
  328. index
  329. for index, message in enumerate(business)
  330. if message["type"] == "tool_result"
  331. )
  332. assert visible_indexes[0] < tool_index < visible_indexes[1]