test_tool_invocation_comparison.py 12 KB

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