test_tool_invocation_comparison.py 15 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482
  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.contracts import AgentParams, DebugRunRequest, EventAgentParams
  8. from agent_lab.application.runtime import DebugRuntime
  9. from agent_lab.application.tools import build_default_tool_registry
  10. from agent_lab.domain.events import ToolCallEvent
  11. from agent_lab.domain.messages import ChatMessage, StreamItem
  12. MODES = ("dual_agent", "chat_agent_tools")
  13. SCENARIOS = (
  14. "ordinary_chat",
  15. "device_volume_silent",
  16. "calendar_schedule_template",
  17. "web_search_two_answers",
  18. "session_terminate",
  19. "parallel_volume_schedule",
  20. )
  21. BUILTIN_EVENT_NAMES = (
  22. "session.terminate",
  23. "device.volume.adjust",
  24. "calendar.schedule.create",
  25. "knowledge.web.search",
  26. )
  27. SCHEDULE_ARGUMENTS = {
  28. "title": "Design review",
  29. "start_at": "2026-07-14T09:30:00+08:00",
  30. "timezone": "Asia/Shanghai",
  31. }
  32. SCHEDULE_CONFIRMATION = (
  33. "Scheduled Design review for 2026-07-14T09:30:00+08:00 "
  34. "(Asia/Shanghai)."
  35. )
  36. @dataclass(frozen=True)
  37. class ExpectedSemantics:
  38. visible_messages: tuple[str, ...]
  39. tool_names: tuple[str, ...]
  40. model_calls: int
  41. terminal_count: int = 0
  42. EXPECTED = {
  43. "ordinary_chat": ExpectedSemantics(("Ordinary answer.",), (), 1),
  44. "device_volume_silent": ExpectedSemantics(
  45. (),
  46. ("device.volume.adjust",),
  47. 1,
  48. ),
  49. "calendar_schedule_template": ExpectedSemantics(
  50. (SCHEDULE_CONFIRMATION,),
  51. ("calendar.schedule.create",),
  52. 1,
  53. ),
  54. "web_search_two_answers": ExpectedSemantics(
  55. ("I will check.", "Grounded search update."),
  56. ("knowledge.web.search",),
  57. 2,
  58. ),
  59. "session_terminate": ExpectedSemantics(
  60. ("Goodbye.",),
  61. ("session.terminate",),
  62. 1,
  63. terminal_count=1,
  64. ),
  65. "parallel_volume_schedule": ExpectedSemantics(
  66. (SCHEDULE_CONFIRMATION,),
  67. ("device.volume.adjust", "calendar.schedule.create"),
  68. 1,
  69. ),
  70. }
  71. class ScriptedChatClient:
  72. def __init__(self, rounds: list[list[StreamItem]]) -> None:
  73. self.rounds = rounds
  74. self.calls = 0
  75. async def stream_chat(
  76. self,
  77. messages: list[ChatMessage],
  78. tools: list[dict[str, Any]],
  79. params: AgentParams,
  80. tool_choice: dict[str, Any] | None = None,
  81. ) -> AsyncIterator[StreamItem]:
  82. del messages, tools, params, tool_choice
  83. items = self.rounds[self.calls]
  84. self.calls += 1
  85. for item in items:
  86. yield item
  87. class RecordingPorts:
  88. def __init__(self, *, parallel_barrier: bool = False) -> None:
  89. self.parallel_barrier = parallel_barrier
  90. self.session_calls: list[tuple[str, str | None]] = []
  91. self.volume_calls: list[tuple[str, str, int | None, int | None]] = []
  92. self.calendar_calls: list[
  93. tuple[str, str, str, str, str | None, int | None]
  94. ] = []
  95. self.search_calls: list[tuple[str, str, int]] = []
  96. self.started: set[str] = set()
  97. self.both_started = asyncio.Event()
  98. self.volume_release = asyncio.Event()
  99. self.schedule_release = asyncio.Event()
  100. self.schedule_finished = asyncio.Event()
  101. if not parallel_barrier:
  102. self.volume_release.set()
  103. self.schedule_release.set()
  104. async def terminate(
  105. self,
  106. event_id: str,
  107. *,
  108. reason: str | None = None,
  109. ) -> dict[str, Any]:
  110. self.session_calls.append((event_id, reason))
  111. return {
  112. "tool": "session.terminate",
  113. "status": "terminated",
  114. "event_id": event_id,
  115. "reason": reason,
  116. }
  117. async def adjust(
  118. self,
  119. event_id: str,
  120. *,
  121. mode: str,
  122. value: int | None = None,
  123. delta: int | None = None,
  124. ) -> dict[str, Any]:
  125. self.volume_calls.append((event_id, mode, value, delta))
  126. await self._wait_at_barrier("volume", self.volume_release)
  127. payload: dict[str, Any] = {
  128. "tool": "device.volume.adjust",
  129. "status": "applied",
  130. "event_id": event_id,
  131. "mode": mode,
  132. }
  133. if value is not None:
  134. payload["value"] = value
  135. if delta is not None:
  136. payload["delta"] = delta
  137. return payload
  138. async def create(
  139. self,
  140. event_id: str,
  141. *,
  142. title: str,
  143. start_at: str,
  144. timezone: str,
  145. recurrence: str | None = None,
  146. reminder_minutes: int | None = None,
  147. ) -> dict[str, Any]:
  148. self.calendar_calls.append(
  149. (
  150. event_id,
  151. title,
  152. start_at,
  153. timezone,
  154. recurrence,
  155. reminder_minutes,
  156. )
  157. )
  158. await self._wait_at_barrier("schedule", self.schedule_release)
  159. self.schedule_finished.set()
  160. return {
  161. "tool": "calendar.schedule.create",
  162. "status": "created",
  163. "event_id": event_id,
  164. "schedule": {
  165. "title": title,
  166. "start_at": start_at,
  167. "timezone": timezone,
  168. },
  169. }
  170. async def search(
  171. self,
  172. event_id: str,
  173. *,
  174. query: str,
  175. max_results: int = 3,
  176. ) -> dict[str, Any]:
  177. self.search_calls.append((event_id, query, max_results))
  178. return {
  179. "tool": "knowledge.web.search",
  180. "query": query,
  181. "sources": [
  182. {
  183. "title": "Matrix source",
  184. "url": "https://example.invalid/matrix",
  185. "snippet": f"Deterministic context for: {query}",
  186. }
  187. ],
  188. "retrieved_at": "2026-07-13T00:00:00Z",
  189. }
  190. async def _wait_at_barrier(
  191. self,
  192. name: str,
  193. release: asyncio.Event,
  194. ) -> None:
  195. if not self.parallel_barrier:
  196. return
  197. self.started.add(name)
  198. if len(self.started) == 2:
  199. self.both_started.set()
  200. await release.wait()
  201. def _event(call_id: str, name: str, arguments: dict[str, Any]) -> ToolCallEvent:
  202. return ToolCallEvent(
  203. id=call_id,
  204. name=name,
  205. arguments=arguments,
  206. raw_arguments=json.dumps(arguments, sort_keys=True, separators=(",", ":")),
  207. )
  208. def _event_item(mode: str, event: ToolCallEvent) -> StreamItem:
  209. if mode == "dual_agent":
  210. return StreamItem.text_event(event)
  211. return StreamItem.provider_tool_call(event)
  212. def _scenario_setup(
  213. mode: str,
  214. scenario: str,
  215. ) -> tuple[DebugRunRequest, list[list[StreamItem]]]:
  216. if scenario == "ordinary_chat":
  217. user_message = "Hello."
  218. enabled_tools = list(BUILTIN_EVENT_NAMES)
  219. rounds = [[StreamItem.message_delta("Ordinary answer.")]]
  220. elif scenario == "device_volume_silent":
  221. user_message = "set volume to 30"
  222. enabled_tools = ["device.volume.adjust"]
  223. rounds = [
  224. [
  225. _event_item(
  226. mode,
  227. _event(
  228. "volume-1",
  229. "device.volume.adjust",
  230. {"mode": "absolute", "value": 30},
  231. ),
  232. )
  233. ]
  234. ]
  235. elif scenario == "calendar_schedule_template":
  236. user_message = (
  237. 'schedule "Design review" at 2026-07-14T09:30:00+08:00 '
  238. "timezone Asia/Shanghai"
  239. )
  240. enabled_tools = ["calendar.schedule.create"]
  241. rounds = [
  242. [
  243. _event_item(
  244. mode,
  245. _event(
  246. "schedule-1",
  247. "calendar.schedule.create",
  248. SCHEDULE_ARGUMENTS,
  249. ),
  250. )
  251. ]
  252. ]
  253. elif scenario == "web_search_two_answers":
  254. user_message = "event batch safety"
  255. enabled_tools = ["knowledge.web.search"]
  256. rounds = [
  257. [
  258. StreamItem.message_delta("I will check."),
  259. _event_item(
  260. mode,
  261. _event(
  262. "search-1",
  263. "knowledge.web.search",
  264. {"query": user_message},
  265. ),
  266. ),
  267. ],
  268. [StreamItem.message_delta("Grounded search update.")],
  269. ]
  270. elif scenario == "session_terminate":
  271. user_message = "end this session"
  272. enabled_tools = ["session.terminate"]
  273. rounds = [
  274. [
  275. _event_item(
  276. mode,
  277. _event("terminate-1", "session.terminate", {}),
  278. )
  279. ]
  280. ]
  281. elif scenario == "parallel_volume_schedule":
  282. user_message = (
  283. 'set volume to 30 and schedule "Design review" at '
  284. "2026-07-14T09:30:00+08:00 timezone Asia/Shanghai"
  285. )
  286. enabled_tools = ["device.volume.adjust", "calendar.schedule.create"]
  287. rounds = [
  288. [
  289. _event_item(
  290. mode,
  291. _event(
  292. "parallel-volume",
  293. "device.volume.adjust",
  294. {"mode": "absolute", "value": 30},
  295. ),
  296. ),
  297. _event_item(
  298. mode,
  299. _event(
  300. "parallel-schedule",
  301. "calendar.schedule.create",
  302. SCHEDULE_ARGUMENTS,
  303. ),
  304. ),
  305. ]
  306. ]
  307. else:
  308. raise AssertionError(f"unknown comparison scenario: {scenario}")
  309. request = DebugRunRequest(
  310. user_message=user_message,
  311. system_prompts=["Exercise the requested business scenario."],
  312. pre_messages=[],
  313. chat_agent=AgentParams(
  314. model="comparison-model",
  315. temperature=0.0,
  316. max_tokens=128,
  317. ),
  318. event_agent=EventAgentParams(
  319. model="comparison-model",
  320. temperature=0.0,
  321. max_tokens=128,
  322. enabled_tools=enabled_tools,
  323. max_event_loops=1,
  324. max_parallel_events=2,
  325. batch_timeout_seconds=1.0,
  326. ),
  327. tool_invocation_mode=mode,
  328. )
  329. return request, rounds
  330. async def _collect_outputs(
  331. stream: AsyncIterator[dict[str, Any]],
  332. ) -> list[dict[str, Any]]:
  333. return [message async for message in stream]
  334. def _tool_payloads(outputs: list[dict[str, Any]]) -> list[dict[str, Any]]:
  335. return [
  336. json.loads(message["message"]["content"])
  337. for message in outputs
  338. if message["type"] == "tool_result"
  339. ]
  340. def _assert_port_calls(scenario: str, ports: RecordingPorts) -> None:
  341. expected_session: list[tuple[str, str | None]] = []
  342. expected_volume: list[tuple[str, str, int | None, int | None]] = []
  343. expected_calendar: list[
  344. tuple[str, str, str, str, str | None, int | None]
  345. ] = []
  346. expected_search: list[tuple[str, str, int]] = []
  347. if scenario == "device_volume_silent":
  348. expected_volume = [("volume-1", "absolute", 30, None)]
  349. elif scenario == "calendar_schedule_template":
  350. expected_calendar = [
  351. (
  352. "schedule-1",
  353. "Design review",
  354. "2026-07-14T09:30:00+08:00",
  355. "Asia/Shanghai",
  356. None,
  357. None,
  358. )
  359. ]
  360. elif scenario == "web_search_two_answers":
  361. expected_search = [("search-1", "event batch safety", 3)]
  362. elif scenario == "session_terminate":
  363. expected_session = [("terminate-1", None)]
  364. elif scenario == "parallel_volume_schedule":
  365. expected_volume = [("parallel-volume", "absolute", 30, None)]
  366. expected_calendar = [
  367. (
  368. "parallel-schedule",
  369. "Design review",
  370. "2026-07-14T09:30:00+08:00",
  371. "Asia/Shanghai",
  372. None,
  373. None,
  374. )
  375. ]
  376. assert ports.session_calls == expected_session
  377. assert ports.volume_calls == expected_volume
  378. assert ports.calendar_calls == expected_calendar
  379. assert ports.search_calls == expected_search
  380. @pytest.mark.asyncio
  381. @pytest.mark.parametrize("mode", MODES)
  382. @pytest.mark.parametrize("scenario", SCENARIOS)
  383. async def test_invocation_modes_share_business_semantics(
  384. mode: str,
  385. scenario: str,
  386. ) -> None:
  387. request, rounds = _scenario_setup(mode, scenario)
  388. ports = RecordingPorts(parallel_barrier=scenario == "parallel_volume_schedule")
  389. registry = build_default_tool_registry(
  390. session_termination_port=ports,
  391. device_volume_port=ports,
  392. calendar_schedule_port=ports,
  393. web_search_port=ports,
  394. )
  395. client = ScriptedChatClient(rounds)
  396. run_task = asyncio.create_task(
  397. _collect_outputs(DebugRuntime(client, registry=registry).run(request))
  398. )
  399. try:
  400. if scenario == "parallel_volume_schedule":
  401. await asyncio.wait_for(ports.both_started.wait(), timeout=1)
  402. assert ports.started == {"volume", "schedule"}
  403. ports.schedule_release.set()
  404. await asyncio.wait_for(ports.schedule_finished.wait(), timeout=1)
  405. assert not run_task.done()
  406. ports.volume_release.set()
  407. outputs = await asyncio.wait_for(run_task, timeout=2)
  408. finally:
  409. ports.schedule_release.set()
  410. ports.volume_release.set()
  411. if not run_task.done():
  412. run_task.cancel()
  413. await asyncio.gather(run_task, return_exceptions=True)
  414. expected = EXPECTED[scenario]
  415. visible_messages = tuple(
  416. message["content"]
  417. for message in outputs
  418. if message["type"] == "message_delta"
  419. )
  420. tool_payloads = _tool_payloads(outputs)
  421. terminal_count = sum(
  422. message.get("event") == "terminal_completed" for message in outputs
  423. )
  424. assert visible_messages == expected.visible_messages
  425. assert tuple(payload["tool"] for payload in tool_payloads) == expected.tool_names
  426. assert client.calls == expected.model_calls
  427. assert terminal_count == expected.terminal_count
  428. assert outputs[-1] == {"type": "done"}
  429. _assert_port_calls(scenario, ports)
  430. if scenario == "web_search_two_answers":
  431. assert tool_payloads[0]["sources"][0]["title"] == "Matrix source"
  432. business = [message for message in outputs if message["type"] != "audit"]
  433. visible_indexes = [
  434. index
  435. for index, message in enumerate(business)
  436. if message["type"] == "message_delta"
  437. ]
  438. tool_index = next(
  439. index
  440. for index, message in enumerate(business)
  441. if message["type"] == "tool_result"
  442. )
  443. assert visible_indexes[0] < tool_index < visible_indexes[1]