test_event_agent.py 22 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722
  1. import asyncio
  2. import json
  3. from collections.abc import AsyncIterator
  4. import pytest
  5. from agent_lab.application.contracts import AgentParams, EventAgentParams
  6. from agent_lab.application.event_agent import EventAgent
  7. from agent_lab.application.tools import ToolDefinition, ToolExecutionContext, ToolRegistry
  8. from agent_lab.domain.events import ToolCallEvent
  9. from agent_lab.domain.messages import ChatMessage, StreamItem
  10. class ToolCallingChatClient:
  11. def __init__(self, arguments: dict) -> None:
  12. self.arguments = arguments
  13. self.calls: list[dict] = []
  14. async def stream_chat(
  15. self,
  16. messages: list[ChatMessage],
  17. tools: list[dict],
  18. params: AgentParams,
  19. tool_choice: dict | None = None,
  20. ) -> AsyncIterator[StreamItem]:
  21. self.calls.append(
  22. {
  23. "messages": list(messages),
  24. "tools": list(tools),
  25. "params": params,
  26. "tool_choice": tool_choice,
  27. }
  28. )
  29. tool_name = tools[0]["function"]["name"]
  30. yield StreamItem.raw_response_chunk(
  31. {
  32. "choices": [
  33. {
  34. "delta": {
  35. "tool_calls": [
  36. {
  37. "index": 0,
  38. "function": {"name": tool_name},
  39. }
  40. ]
  41. },
  42. "finish_reason": None,
  43. }
  44. ]
  45. }
  46. )
  47. yield StreamItem.provider_tool_call(
  48. ToolCallEvent(
  49. id="llm_call_1",
  50. name=tool_name,
  51. arguments=self.arguments,
  52. raw_arguments=json.dumps(self.arguments),
  53. )
  54. )
  55. class NoToolCallChatClient:
  56. def __init__(self) -> None:
  57. self.calls: list[dict] = []
  58. async def stream_chat(
  59. self,
  60. messages: list[ChatMessage],
  61. tools: list[dict],
  62. params: AgentParams,
  63. tool_choice: dict | None = None,
  64. ) -> AsyncIterator[StreamItem]:
  65. self.calls.append(
  66. {
  67. "messages": list(messages),
  68. "tools": list(tools),
  69. "params": params,
  70. "tool_choice": tool_choice,
  71. }
  72. )
  73. yield StreamItem.message_delta("I should have called the tool.")
  74. class WrongSourceToolCallChatClient:
  75. def __init__(self, item: StreamItem) -> None:
  76. self.item = item
  77. async def stream_chat(
  78. self,
  79. messages: list[ChatMessage],
  80. tools: list[dict],
  81. params: AgentParams,
  82. tool_choice: dict | None = None,
  83. ) -> AsyncIterator[StreamItem]:
  84. yield self.item
  85. class WrongNameToolCallingChatClient(ToolCallingChatClient):
  86. async def stream_chat(
  87. self,
  88. messages: list[ChatMessage],
  89. tools: list[dict],
  90. params: AgentParams,
  91. tool_choice: dict | None = None,
  92. ) -> AsyncIterator[StreamItem]:
  93. self.calls.append({"messages": list(messages), "tools": list(tools)})
  94. yield StreamItem.provider_tool_call(
  95. ToolCallEvent(
  96. id="wrong-call",
  97. name="wrong.tool",
  98. arguments=self.arguments,
  99. raw_arguments=json.dumps(self.arguments),
  100. )
  101. )
  102. @pytest.mark.asyncio
  103. async def test_event_agent_uses_deterministic_resolver_before_llm_fallback():
  104. chat_client = ToolCallingChatClient({"message": "LLM generated handoff"})
  105. agent = EventAgent(
  106. enabled_tools=["handoff_note"],
  107. chat_client=chat_client,
  108. params=AgentParams(model="event-model", temperature=0, max_tokens=80),
  109. )
  110. event = ToolCallEvent(
  111. id="call_1",
  112. name="handoff_note",
  113. arguments={"message": "chat agent argument should be ignored"},
  114. raw_arguments='{"message":"chat agent argument should be ignored"}',
  115. )
  116. history = [
  117. ChatMessage(role="user", content="debug this event flow"),
  118. ChatMessage(role="assistant", content="I need the event agent."),
  119. ]
  120. reply = await agent.handle(event, history=history)
  121. assert reply.role == "tool"
  122. assert reply.tool_call_id == "call_1"
  123. assert reply.name == "handoff_note"
  124. assert json.loads(reply.content) == {
  125. "tool": "handoff_note",
  126. "message": "I need the event agent.",
  127. }
  128. assert chat_client.calls == []
  129. assert agent.raw_model_chunks([event]) == [
  130. {
  131. "event_id": "call_1",
  132. "event_name": "handoff_note",
  133. "chunks": [],
  134. }
  135. ]
  136. @pytest.mark.asyncio
  137. async def test_event_agent_deterministic_mock_resolver_skips_llm():
  138. chat_client = NoToolCallChatClient()
  139. agent = EventAgent(
  140. enabled_tools=["mock_search"],
  141. chat_client=chat_client,
  142. params=AgentParams(model="event-model", temperature=0, max_tokens=80),
  143. )
  144. event = ToolCallEvent(
  145. id="call_1",
  146. name="mock_search",
  147. arguments={},
  148. raw_arguments="{}",
  149. )
  150. history = [
  151. ChatMessage(role="user", content="Find latency docs"),
  152. ChatMessage(role="assistant", content="Need a search for latency docs"),
  153. ]
  154. reply = await agent.handle(event, history=history)
  155. payload = json.loads(reply.content)
  156. assert payload["tool"] == "mock_search"
  157. assert payload["query"] == "Need a search for latency docs"
  158. assert "event agent did not return arguments" not in reply.content
  159. assert chat_client.calls == []
  160. @pytest.mark.asyncio
  161. async def test_event_agent_calls_llm_once_for_incomplete_deterministic_arguments():
  162. chat_client = ToolCallingChatClient({"message": "LLM generated handoff"})
  163. registry = ToolRegistry(
  164. [
  165. ToolDefinition(
  166. name="ambiguous_handoff",
  167. description="Resolve an ambiguous handoff.",
  168. parameters={
  169. "type": "object",
  170. "properties": {"message": {"type": "string"}},
  171. "required": ["message"],
  172. },
  173. handler=lambda event: {
  174. "tool": event.name,
  175. "message": event.arguments["message"],
  176. },
  177. argument_resolver=lambda event, context: {},
  178. )
  179. ]
  180. )
  181. event = ToolCallEvent(
  182. id="call_1",
  183. name="ambiguous_handoff",
  184. arguments={},
  185. raw_arguments="{}",
  186. )
  187. reply = await EventAgent(
  188. enabled_tools=["ambiguous_handoff"],
  189. registry=registry,
  190. chat_client=chat_client,
  191. ).handle(event, history=[ChatMessage(role="user", content="ambiguous")])
  192. assert json.loads(reply.content) == {
  193. "tool": "ambiguous_handoff",
  194. "message": "LLM generated handoff",
  195. }
  196. assert len(chat_client.calls) == 1
  197. assert chat_client.calls[0]["tool_choice"] == {
  198. "type": "function",
  199. "function": {"name": "ambiguous_handoff"},
  200. }
  201. @pytest.mark.asyncio
  202. async def test_event_agent_rejects_fallback_provider_call_for_wrong_tool_name():
  203. chat_client = WrongNameToolCallingChatClient({"message": "wrong"})
  204. registry = ToolRegistry(
  205. [
  206. ToolDefinition(
  207. name="expected.tool",
  208. description="Expected tool.",
  209. parameters={
  210. "type": "object",
  211. "properties": {"message": {"type": "string"}},
  212. "required": ["message"],
  213. },
  214. handler=lambda event: {"tool": event.name},
  215. argument_resolver=lambda event, context: {},
  216. )
  217. ]
  218. )
  219. reply = await EventAgent(
  220. enabled_tools=["expected.tool"],
  221. registry=registry,
  222. chat_client=chat_client,
  223. ).handle(
  224. ToolCallEvent(
  225. id="call-1",
  226. name="expected.tool",
  227. arguments={},
  228. raw_arguments="{}",
  229. )
  230. )
  231. assert json.loads(reply.content) == {
  232. "tool": "expected.tool",
  233. "error": "fallback returned tool wrong.tool for expected.tool",
  234. }
  235. @pytest.mark.asyncio
  236. async def test_event_agent_handle_many_starts_independent_handlers_concurrently():
  237. started: list[str] = []
  238. both_started = asyncio.Event()
  239. async def handler(event: ToolCallEvent) -> dict:
  240. started.append(event.name)
  241. if len(started) == 2:
  242. both_started.set()
  243. await asyncio.wait_for(both_started.wait(), timeout=0.2)
  244. return {"tool": event.name}
  245. registry = ToolRegistry(
  246. [
  247. ToolDefinition(
  248. name=name,
  249. description=f"Handle {name}.",
  250. parameters={"type": "object"},
  251. handler=handler,
  252. )
  253. for name in ("event.first", "event.second")
  254. ]
  255. )
  256. events = [
  257. ToolCallEvent(id=f"call-{index}", name=name, arguments={}, raw_arguments="{}")
  258. for index, name in enumerate(("event.first", "event.second"), start=1)
  259. ]
  260. batch = await asyncio.wait_for(
  261. EventAgent(
  262. enabled_tools=[event.name for event in events],
  263. registry=registry,
  264. params=EventAgentParams(max_parallel_events=2),
  265. ).handle_many(events, history=[]),
  266. timeout=0.5,
  267. )
  268. assert started == ["event.first", "event.second"]
  269. assert [result.event_name for result in batch.results] == [
  270. "event.first",
  271. "event.second",
  272. ]
  273. assert [result.payload for result in batch.results] == [
  274. {"tool": "event.first"},
  275. {"tool": "event.second"},
  276. ]
  277. @pytest.mark.asyncio
  278. @pytest.mark.parametrize(
  279. "item",
  280. [
  281. StreamItem.text_event(
  282. ToolCallEvent(
  283. id="text_event_1",
  284. name="mock_search",
  285. arguments={"query": "wrong source"},
  286. raw_arguments='{"query":"wrong source"}',
  287. )
  288. ),
  289. StreamItem.event(
  290. ToolCallEvent(
  291. id="legacy_event_1",
  292. name="mock_search",
  293. arguments={"query": "legacy arguments"},
  294. raw_arguments='{"query":"legacy arguments"}',
  295. )
  296. ),
  297. ],
  298. ids=["text_event", "legacy_event"],
  299. )
  300. async def test_event_agent_ignores_non_provider_tool_call_sources(item: StreamItem):
  301. event = ToolCallEvent(
  302. id="call_1",
  303. name="mock_search",
  304. arguments={},
  305. raw_arguments="{}",
  306. )
  307. registry = ToolRegistry(
  308. [
  309. ToolDefinition(
  310. name="mock_search",
  311. description="Search mock external knowledge for the current turn.",
  312. parameters={
  313. "type": "object",
  314. "properties": {"query": {"type": "string"}},
  315. "required": ["query"],
  316. },
  317. handler=lambda resolved: {
  318. "tool": resolved.name,
  319. "query": resolved.arguments["query"],
  320. },
  321. argument_resolver=lambda event, context: {},
  322. )
  323. ]
  324. )
  325. agent = EventAgent(
  326. enabled_tools=["mock_search"],
  327. registry=registry,
  328. chat_client=WrongSourceToolCallChatClient(item),
  329. )
  330. reply = await agent.handle(
  331. event,
  332. history=[ChatMessage(role="assistant", content="fallback query")],
  333. )
  334. payload = json.loads(reply.content)
  335. assert payload["tool"] == "mock_search"
  336. assert payload == {
  337. "tool": "mock_search",
  338. "error": "missing required arguments: query",
  339. }
  340. @pytest.mark.asyncio
  341. async def test_event_agent_returns_registry_errors_for_disabled_and_unknown_tools():
  342. registry = ToolRegistry(
  343. [
  344. ToolDefinition(
  345. name="handoff_note",
  346. description="Send a note to the event agent.",
  347. parameters={"type": "object"},
  348. handler=lambda event: {"tool": event.name, "message": "handled"},
  349. )
  350. ]
  351. )
  352. disabled_reply = await EventAgent(
  353. enabled_tools=[],
  354. registry=registry,
  355. ).handle(
  356. ToolCallEvent(
  357. id="call_1",
  358. name="handoff_note",
  359. arguments={"message": "inspect this event"},
  360. raw_arguments='{"message":"inspect this event"}',
  361. )
  362. )
  363. unknown_reply = await EventAgent(
  364. enabled_tools=["missing_tool"],
  365. registry=registry,
  366. ).handle(
  367. ToolCallEvent(
  368. id="call_2",
  369. name="missing_tool",
  370. arguments={},
  371. raw_arguments="{}",
  372. )
  373. )
  374. assert json.loads(disabled_reply.content) == {
  375. "tool": "handoff_note",
  376. "error": "tool disabled",
  377. }
  378. assert json.loads(unknown_reply.content) == {
  379. "tool": "missing_tool",
  380. "error": "unknown tool",
  381. }
  382. @pytest.mark.asyncio
  383. async def test_event_agent_returns_structured_error_when_tool_handler_raises():
  384. def fail_tool(event: ToolCallEvent) -> dict:
  385. raise RuntimeError("boom")
  386. registry = ToolRegistry(
  387. [
  388. ToolDefinition(
  389. name="handoff_note",
  390. description="Send a note to the event agent.",
  391. parameters={"type": "object"},
  392. handler=fail_tool,
  393. )
  394. ]
  395. )
  396. event = ToolCallEvent(
  397. id="call_1",
  398. name="handoff_note",
  399. arguments={"message": "inspect this event"},
  400. raw_arguments='{"message":"inspect this event"}',
  401. )
  402. reply = await EventAgent(
  403. enabled_tools=["handoff_note"],
  404. registry=registry,
  405. ).handle(event)
  406. assert reply.role == "tool"
  407. assert reply.tool_call_id == "call_1"
  408. assert json.loads(reply.content) == {
  409. "tool": "handoff_note",
  410. "error": "tool handler failed: boom",
  411. }
  412. @pytest.mark.asyncio
  413. async def test_event_agent_serializes_non_json_handler_payload_as_tool_error():
  414. registry = ToolRegistry(
  415. [
  416. ToolDefinition(
  417. name="bad_payload",
  418. description="Return an invalid payload.",
  419. parameters={"type": "object"},
  420. handler=lambda event: {"invalid": object()},
  421. )
  422. ]
  423. )
  424. reply = await EventAgent(
  425. enabled_tools=["bad_payload"],
  426. registry=registry,
  427. ).handle(
  428. ToolCallEvent(
  429. id="call-1",
  430. name="bad_payload",
  431. arguments={},
  432. raw_arguments="{}",
  433. )
  434. )
  435. assert json.loads(reply.content) == {
  436. "tool": "bad_payload",
  437. "error": "event handler returned non-JSON payload",
  438. }
  439. @pytest.mark.asyncio
  440. async def test_event_agent_normalizes_handler_payload_snapshot_exceptions():
  441. class ExplodingItemsDict(dict):
  442. def items(self):
  443. raise RuntimeError("payload items failed")
  444. registry = ToolRegistry(
  445. [
  446. ToolDefinition(
  447. name="bad_payload",
  448. description="Return a payload that fails during snapshot.",
  449. parameters={"type": "object"},
  450. handler=lambda event: ExplodingItemsDict(ok=True),
  451. )
  452. ]
  453. )
  454. reply = await EventAgent(
  455. enabled_tools=["bad_payload"],
  456. registry=registry,
  457. ).handle(
  458. ToolCallEvent(
  459. id="call-1",
  460. name="bad_payload",
  461. arguments={},
  462. raw_arguments="{}",
  463. )
  464. )
  465. assert reply.role == "tool"
  466. assert reply.tool_call_id == "call-1"
  467. assert json.loads(reply.content) == {
  468. "tool": "bad_payload",
  469. "error": "event handler returned non-JSON payload",
  470. }
  471. @pytest.mark.asyncio
  472. async def test_event_agent_llm_receives_history_and_agent_config_context():
  473. chat_client = ToolCallingChatClient(
  474. {"message": "Use strict tool parameters.", "thinking": "disabled"}
  475. )
  476. registry = ToolRegistry(
  477. [
  478. ToolDefinition(
  479. name="handoff_note",
  480. description="Send a note to the event agent.",
  481. parameters={
  482. "type": "object",
  483. "properties": {
  484. "message": {"type": "string"},
  485. "thinking": {"type": "string"},
  486. },
  487. "required": ["message", "thinking"],
  488. },
  489. handler=lambda event: {
  490. "tool": event.name,
  491. "message": event.arguments["message"],
  492. "thinking": event.arguments["thinking"],
  493. },
  494. )
  495. ]
  496. )
  497. event = ToolCallEvent(
  498. id="call_1",
  499. name="handoff_note",
  500. arguments={},
  501. raw_arguments="{}",
  502. )
  503. reply = await EventAgent(
  504. enabled_tools=["handoff_note"],
  505. registry=registry,
  506. chat_client=chat_client,
  507. params=AgentParams(model="event-model", temperature=0.4, max_tokens=120),
  508. ).handle(
  509. event,
  510. history=[ChatMessage(role="user", content="debug this")],
  511. system_prompt="Use strict tool parameters.",
  512. extra_body={"thinking": {"type": "disabled"}},
  513. )
  514. messages = chat_client.calls[0]["messages"]
  515. assert any(message.content == "debug this" for message in messages)
  516. assert any(message.content == "Use strict tool parameters." for message in messages)
  517. assert chat_client.calls[0]["params"].extra_body == {"thinking": {"type": "disabled"}}
  518. assert json.loads(reply.content) == {
  519. "tool": "handoff_note",
  520. "message": "Use strict tool parameters.",
  521. "thinking": "disabled",
  522. }
  523. @pytest.mark.asyncio
  524. async def test_event_agent_uses_configured_extra_body_when_override_is_omitted():
  525. chat_client = ToolCallingChatClient({"message": "resolved"})
  526. registry = ToolRegistry(
  527. [
  528. ToolDefinition(
  529. name="handoff_note",
  530. description="Send a note to the event agent.",
  531. parameters={
  532. "type": "object",
  533. "properties": {"message": {"type": "string"}},
  534. "required": ["message"],
  535. },
  536. handler=lambda event: {"tool": event.name},
  537. argument_resolver=lambda event, context: {},
  538. )
  539. ]
  540. )
  541. await EventAgent(
  542. enabled_tools=["handoff_note"],
  543. registry=registry,
  544. chat_client=chat_client,
  545. params=AgentParams(extra_body={"configured": True}),
  546. ).handle(
  547. ToolCallEvent(
  548. id="call-1",
  549. name="handoff_note",
  550. arguments={},
  551. raw_arguments="{}",
  552. )
  553. )
  554. assert chat_client.calls[0]["params"].extra_body == {"configured": True}
  555. @pytest.mark.asyncio
  556. async def test_event_agent_preserves_explicit_empty_extra_body_override():
  557. chat_client = ToolCallingChatClient({"message": "resolved"})
  558. registry = ToolRegistry(
  559. [
  560. ToolDefinition(
  561. name="handoff_note",
  562. description="Send a note to the event agent.",
  563. parameters={
  564. "type": "object",
  565. "properties": {"message": {"type": "string"}},
  566. "required": ["message"],
  567. },
  568. handler=lambda event: {"tool": event.name},
  569. argument_resolver=lambda event, context: {},
  570. )
  571. ]
  572. )
  573. await EventAgent(
  574. enabled_tools=["handoff_note"],
  575. registry=registry,
  576. chat_client=chat_client,
  577. params=AgentParams(extra_body={"configured": True}),
  578. ).handle(
  579. ToolCallEvent(
  580. id="call-1",
  581. name="handoff_note",
  582. arguments={},
  583. raw_arguments="{}",
  584. ),
  585. extra_body={},
  586. )
  587. assert chat_client.calls[0]["params"].extra_body == {}
  588. @pytest.mark.asyncio
  589. async def test_event_agent_projects_complete_tool_round_to_visible_history():
  590. chat_client = ToolCallingChatClient({"message": "projected history"})
  591. registry = ToolRegistry(
  592. [
  593. ToolDefinition(
  594. name="handoff_note",
  595. description="Send a note to the event agent.",
  596. parameters={
  597. "type": "object",
  598. "properties": {"message": {"type": "string"}},
  599. "required": ["message"],
  600. },
  601. handler=lambda resolved: {
  602. "tool": resolved.name,
  603. "message": resolved.arguments["message"],
  604. },
  605. argument_resolver=lambda event, context: {},
  606. )
  607. ]
  608. )
  609. event = ToolCallEvent(
  610. id="call_current",
  611. name="handoff_note",
  612. arguments={},
  613. raw_arguments="{}",
  614. )
  615. prior_call = ToolCallEvent(
  616. id="call_prior",
  617. name="mock_search",
  618. arguments={"query": "latency"},
  619. raw_arguments='{"query":"latency"}',
  620. )
  621. await EventAgent(
  622. enabled_tools=["handoff_note"],
  623. registry=registry,
  624. chat_client=chat_client,
  625. ).handle(
  626. event,
  627. history=[
  628. ChatMessage(role="user", content="Find latency docs"),
  629. ChatMessage(
  630. role="assistant",
  631. content="I checked the latency sources.",
  632. tool_calls=[prior_call],
  633. ),
  634. ChatMessage(
  635. role="tool",
  636. content='{"results":["doc"]}',
  637. tool_call_id="call_prior",
  638. ),
  639. ChatMessage(role="user", content="Prepare a handoff"),
  640. ],
  641. )
  642. projected_messages = chat_client.calls[0]["messages"]
  643. visible_assistant = next(
  644. message
  645. for message in projected_messages
  646. if message.content == "I checked the latency sources."
  647. )
  648. assert visible_assistant.tool_calls == []
  649. assert not any(message.role == "tool" for message in projected_messages)