test_tool_invocation_comparison.py 18 KB

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