test_event_agent.py 17 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577
  1. import asyncio
  2. import json
  3. from collections.abc import AsyncIterator
  4. import pytest
  5. from agent_lab.application.contracts import AgentParams
  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. replies = await asyncio.wait_for(
  261. EventAgent(
  262. enabled_tools=[event.name for event in events],
  263. registry=registry,
  264. ).handle_many(events, history=[]),
  265. timeout=0.5,
  266. )
  267. assert started == ["event.first", "event.second"]
  268. assert [reply.name for reply in replies] == ["event.first", "event.second"]
  269. assert [json.loads(reply.content) for reply in replies] == [
  270. {"tool": "event.first"},
  271. {"tool": "event.second"},
  272. ]
  273. @pytest.mark.asyncio
  274. @pytest.mark.parametrize(
  275. "item",
  276. [
  277. StreamItem.text_event(
  278. ToolCallEvent(
  279. id="text_event_1",
  280. name="mock_search",
  281. arguments={"query": "wrong source"},
  282. raw_arguments='{"query":"wrong source"}',
  283. )
  284. ),
  285. StreamItem.event(
  286. ToolCallEvent(
  287. id="legacy_event_1",
  288. name="mock_search",
  289. arguments={"query": "legacy arguments"},
  290. raw_arguments='{"query":"legacy arguments"}',
  291. )
  292. ),
  293. ],
  294. ids=["text_event", "legacy_event"],
  295. )
  296. async def test_event_agent_ignores_non_provider_tool_call_sources(item: StreamItem):
  297. event = ToolCallEvent(
  298. id="call_1",
  299. name="mock_search",
  300. arguments={},
  301. raw_arguments="{}",
  302. )
  303. registry = ToolRegistry(
  304. [
  305. ToolDefinition(
  306. name="mock_search",
  307. description="Search mock external knowledge for the current turn.",
  308. parameters={
  309. "type": "object",
  310. "properties": {"query": {"type": "string"}},
  311. "required": ["query"],
  312. },
  313. handler=lambda resolved: {
  314. "tool": resolved.name,
  315. "query": resolved.arguments["query"],
  316. },
  317. argument_resolver=lambda event, context: {},
  318. )
  319. ]
  320. )
  321. agent = EventAgent(
  322. enabled_tools=["mock_search"],
  323. registry=registry,
  324. chat_client=WrongSourceToolCallChatClient(item),
  325. )
  326. reply = await agent.handle(
  327. event,
  328. history=[ChatMessage(role="assistant", content="fallback query")],
  329. )
  330. payload = json.loads(reply.content)
  331. assert payload["tool"] == "mock_search"
  332. assert payload == {
  333. "tool": "mock_search",
  334. "error": "missing required arguments: query",
  335. }
  336. @pytest.mark.asyncio
  337. async def test_event_agent_returns_registry_errors_for_disabled_and_unknown_tools():
  338. registry = ToolRegistry(
  339. [
  340. ToolDefinition(
  341. name="handoff_note",
  342. description="Send a note to the event agent.",
  343. parameters={"type": "object"},
  344. handler=lambda event: {"tool": event.name, "message": "handled"},
  345. )
  346. ]
  347. )
  348. disabled_reply = await EventAgent(
  349. enabled_tools=[],
  350. registry=registry,
  351. ).handle(
  352. ToolCallEvent(
  353. id="call_1",
  354. name="handoff_note",
  355. arguments={"message": "inspect this event"},
  356. raw_arguments='{"message":"inspect this event"}',
  357. )
  358. )
  359. unknown_reply = await EventAgent(
  360. enabled_tools=["missing_tool"],
  361. registry=registry,
  362. ).handle(
  363. ToolCallEvent(
  364. id="call_2",
  365. name="missing_tool",
  366. arguments={},
  367. raw_arguments="{}",
  368. )
  369. )
  370. assert json.loads(disabled_reply.content) == {
  371. "tool": "handoff_note",
  372. "error": "tool disabled",
  373. }
  374. assert json.loads(unknown_reply.content) == {
  375. "tool": "missing_tool",
  376. "error": "unknown tool",
  377. }
  378. @pytest.mark.asyncio
  379. async def test_event_agent_returns_structured_error_when_tool_handler_raises():
  380. def fail_tool(event: ToolCallEvent) -> dict:
  381. raise RuntimeError("boom")
  382. registry = ToolRegistry(
  383. [
  384. ToolDefinition(
  385. name="handoff_note",
  386. description="Send a note to the event agent.",
  387. parameters={"type": "object"},
  388. handler=fail_tool,
  389. )
  390. ]
  391. )
  392. event = ToolCallEvent(
  393. id="call_1",
  394. name="handoff_note",
  395. arguments={"message": "inspect this event"},
  396. raw_arguments='{"message":"inspect this event"}',
  397. )
  398. reply = await EventAgent(
  399. enabled_tools=["handoff_note"],
  400. registry=registry,
  401. ).handle(event)
  402. assert reply.role == "tool"
  403. assert reply.tool_call_id == "call_1"
  404. assert json.loads(reply.content) == {
  405. "tool": "handoff_note",
  406. "error": "tool handler failed: boom",
  407. }
  408. @pytest.mark.asyncio
  409. async def test_event_agent_llm_receives_history_and_agent_config_context():
  410. chat_client = ToolCallingChatClient(
  411. {"message": "Use strict tool parameters.", "thinking": "disabled"}
  412. )
  413. registry = ToolRegistry(
  414. [
  415. ToolDefinition(
  416. name="handoff_note",
  417. description="Send a note to the event agent.",
  418. parameters={
  419. "type": "object",
  420. "properties": {
  421. "message": {"type": "string"},
  422. "thinking": {"type": "string"},
  423. },
  424. "required": ["message", "thinking"],
  425. },
  426. handler=lambda event: {
  427. "tool": event.name,
  428. "message": event.arguments["message"],
  429. "thinking": event.arguments["thinking"],
  430. },
  431. )
  432. ]
  433. )
  434. event = ToolCallEvent(
  435. id="call_1",
  436. name="handoff_note",
  437. arguments={},
  438. raw_arguments="{}",
  439. )
  440. reply = await EventAgent(
  441. enabled_tools=["handoff_note"],
  442. registry=registry,
  443. chat_client=chat_client,
  444. params=AgentParams(model="event-model", temperature=0.4, max_tokens=120),
  445. ).handle(
  446. event,
  447. history=[ChatMessage(role="user", content="debug this")],
  448. system_prompt="Use strict tool parameters.",
  449. extra_body={"thinking": {"type": "disabled"}},
  450. )
  451. messages = chat_client.calls[0]["messages"]
  452. assert any(message.content == "debug this" for message in messages)
  453. assert any(message.content == "Use strict tool parameters." for message in messages)
  454. assert chat_client.calls[0]["params"].extra_body == {"thinking": {"type": "disabled"}}
  455. assert json.loads(reply.content) == {
  456. "tool": "handoff_note",
  457. "message": "Use strict tool parameters.",
  458. "thinking": "disabled",
  459. }
  460. @pytest.mark.asyncio
  461. async def test_event_agent_projects_complete_tool_round_to_visible_history():
  462. chat_client = ToolCallingChatClient({"message": "projected history"})
  463. registry = ToolRegistry(
  464. [
  465. ToolDefinition(
  466. name="handoff_note",
  467. description="Send a note to the event agent.",
  468. parameters={
  469. "type": "object",
  470. "properties": {"message": {"type": "string"}},
  471. "required": ["message"],
  472. },
  473. handler=lambda resolved: {
  474. "tool": resolved.name,
  475. "message": resolved.arguments["message"],
  476. },
  477. argument_resolver=lambda event, context: {},
  478. )
  479. ]
  480. )
  481. event = ToolCallEvent(
  482. id="call_current",
  483. name="handoff_note",
  484. arguments={},
  485. raw_arguments="{}",
  486. )
  487. prior_call = ToolCallEvent(
  488. id="call_prior",
  489. name="mock_search",
  490. arguments={"query": "latency"},
  491. raw_arguments='{"query":"latency"}',
  492. )
  493. await EventAgent(
  494. enabled_tools=["handoff_note"],
  495. registry=registry,
  496. chat_client=chat_client,
  497. ).handle(
  498. event,
  499. history=[
  500. ChatMessage(role="user", content="Find latency docs"),
  501. ChatMessage(
  502. role="assistant",
  503. content="I checked the latency sources.",
  504. tool_calls=[prior_call],
  505. ),
  506. ChatMessage(
  507. role="tool",
  508. content='{"results":["doc"]}',
  509. tool_call_id="call_prior",
  510. ),
  511. ChatMessage(role="user", content="Prepare a handoff"),
  512. ],
  513. )
  514. projected_messages = chat_client.calls[0]["messages"]
  515. visible_assistant = next(
  516. message
  517. for message in projected_messages
  518. if message.content == "I checked the latency sources."
  519. )
  520. assert visible_assistant.tool_calls == []
  521. assert not any(message.role == "tool" for message in projected_messages)