test_event_agent.py 14 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475
  1. import json
  2. from collections.abc import AsyncIterator
  3. import pytest
  4. from agent_lab.application.contracts import AgentParams
  5. from agent_lab.application.event_agent import EventAgent
  6. from agent_lab.application.tools import ToolDefinition, ToolExecutionContext, ToolRegistry
  7. from agent_lab.domain.events import ToolCallEvent
  8. from agent_lab.domain.messages import ChatMessage, StreamItem
  9. class ToolCallingChatClient:
  10. def __init__(self, arguments: dict) -> None:
  11. self.arguments = arguments
  12. self.calls: list[dict] = []
  13. async def stream_chat(
  14. self,
  15. messages: list[ChatMessage],
  16. tools: list[dict],
  17. params: AgentParams,
  18. tool_choice: dict | None = None,
  19. ) -> AsyncIterator[StreamItem]:
  20. self.calls.append(
  21. {
  22. "messages": list(messages),
  23. "tools": list(tools),
  24. "params": params,
  25. "tool_choice": tool_choice,
  26. }
  27. )
  28. tool_name = tools[0]["function"]["name"]
  29. yield StreamItem.raw_response_chunk(
  30. {
  31. "choices": [
  32. {
  33. "delta": {
  34. "tool_calls": [
  35. {
  36. "index": 0,
  37. "function": {"name": tool_name},
  38. }
  39. ]
  40. },
  41. "finish_reason": None,
  42. }
  43. ]
  44. }
  45. )
  46. yield StreamItem.provider_tool_call(
  47. ToolCallEvent(
  48. id="llm_call_1",
  49. name=tool_name,
  50. arguments=self.arguments,
  51. raw_arguments=json.dumps(self.arguments),
  52. )
  53. )
  54. class NoToolCallChatClient:
  55. def __init__(self) -> None:
  56. self.calls: list[dict] = []
  57. async def stream_chat(
  58. self,
  59. messages: list[ChatMessage],
  60. tools: list[dict],
  61. params: AgentParams,
  62. tool_choice: dict | None = None,
  63. ) -> AsyncIterator[StreamItem]:
  64. self.calls.append(
  65. {
  66. "messages": list(messages),
  67. "tools": list(tools),
  68. "params": params,
  69. "tool_choice": tool_choice,
  70. }
  71. )
  72. yield StreamItem.message_delta("I should have called the tool.")
  73. class WrongSourceToolCallChatClient:
  74. def __init__(self, item: StreamItem) -> None:
  75. self.item = item
  76. async def stream_chat(
  77. self,
  78. messages: list[ChatMessage],
  79. tools: list[dict],
  80. params: AgentParams,
  81. tool_choice: dict | None = None,
  82. ) -> AsyncIterator[StreamItem]:
  83. yield self.item
  84. @pytest.mark.asyncio
  85. async def test_event_agent_uses_deterministic_resolver_before_llm_fallback():
  86. chat_client = ToolCallingChatClient({"message": "LLM generated handoff"})
  87. agent = EventAgent(
  88. enabled_tools=["handoff_note"],
  89. chat_client=chat_client,
  90. params=AgentParams(model="event-model", temperature=0, max_tokens=80),
  91. )
  92. event = ToolCallEvent(
  93. id="call_1",
  94. name="handoff_note",
  95. arguments={"message": "chat agent argument should be ignored"},
  96. raw_arguments='{"message":"chat agent argument should be ignored"}',
  97. )
  98. history = [
  99. ChatMessage(role="user", content="debug this event flow"),
  100. ChatMessage(role="assistant", content="I need the event agent."),
  101. ]
  102. reply = await agent.handle(event, history=history)
  103. assert reply.role == "tool"
  104. assert reply.tool_call_id == "call_1"
  105. assert reply.name == "handoff_note"
  106. assert json.loads(reply.content) == {
  107. "tool": "handoff_note",
  108. "message": "I need the event agent.",
  109. }
  110. assert chat_client.calls == []
  111. assert agent.raw_model_chunks([event]) == [
  112. {
  113. "event_id": "call_1",
  114. "event_name": "handoff_note",
  115. "chunks": [],
  116. }
  117. ]
  118. @pytest.mark.asyncio
  119. async def test_event_agent_deterministic_mock_resolver_skips_llm():
  120. chat_client = NoToolCallChatClient()
  121. agent = EventAgent(
  122. enabled_tools=["mock_search"],
  123. chat_client=chat_client,
  124. params=AgentParams(model="event-model", temperature=0, max_tokens=80),
  125. )
  126. event = ToolCallEvent(
  127. id="call_1",
  128. name="mock_search",
  129. arguments={},
  130. raw_arguments="{}",
  131. )
  132. history = [
  133. ChatMessage(role="user", content="Find latency docs"),
  134. ChatMessage(role="assistant", content="Need a search for latency docs"),
  135. ]
  136. reply = await agent.handle(event, history=history)
  137. payload = json.loads(reply.content)
  138. assert payload["tool"] == "mock_search"
  139. assert payload["query"] == "Need a search for latency docs"
  140. assert "event agent did not return arguments" not in reply.content
  141. assert chat_client.calls == []
  142. @pytest.mark.asyncio
  143. async def test_event_agent_calls_llm_once_for_incomplete_deterministic_arguments():
  144. chat_client = ToolCallingChatClient({"message": "LLM generated handoff"})
  145. registry = ToolRegistry(
  146. [
  147. ToolDefinition(
  148. name="ambiguous_handoff",
  149. description="Resolve an ambiguous handoff.",
  150. parameters={
  151. "type": "object",
  152. "properties": {"message": {"type": "string"}},
  153. "required": ["message"],
  154. },
  155. handler=lambda event: {
  156. "tool": event.name,
  157. "message": event.arguments["message"],
  158. },
  159. argument_resolver=lambda event, context: {},
  160. )
  161. ]
  162. )
  163. event = ToolCallEvent(
  164. id="call_1",
  165. name="ambiguous_handoff",
  166. arguments={},
  167. raw_arguments="{}",
  168. )
  169. reply = await EventAgent(
  170. enabled_tools=["ambiguous_handoff"],
  171. registry=registry,
  172. chat_client=chat_client,
  173. ).handle(event, history=[ChatMessage(role="user", content="ambiguous")])
  174. assert json.loads(reply.content) == {
  175. "tool": "ambiguous_handoff",
  176. "message": "LLM generated handoff",
  177. }
  178. assert len(chat_client.calls) == 1
  179. assert chat_client.calls[0]["tool_choice"] == {
  180. "type": "function",
  181. "function": {"name": "ambiguous_handoff"},
  182. }
  183. @pytest.mark.asyncio
  184. @pytest.mark.parametrize(
  185. "item",
  186. [
  187. StreamItem.text_event(
  188. ToolCallEvent(
  189. id="text_event_1",
  190. name="mock_search",
  191. arguments={"query": "wrong source"},
  192. raw_arguments='{"query":"wrong source"}',
  193. )
  194. ),
  195. StreamItem.event(
  196. ToolCallEvent(
  197. id="legacy_event_1",
  198. name="mock_search",
  199. arguments={"query": "legacy arguments"},
  200. raw_arguments='{"query":"legacy arguments"}',
  201. )
  202. ),
  203. ],
  204. ids=["text_event", "legacy_event"],
  205. )
  206. async def test_event_agent_ignores_non_provider_tool_call_sources(item: StreamItem):
  207. event = ToolCallEvent(
  208. id="call_1",
  209. name="mock_search",
  210. arguments={},
  211. raw_arguments="{}",
  212. )
  213. registry = ToolRegistry(
  214. [
  215. ToolDefinition(
  216. name="mock_search",
  217. description="Search mock external knowledge for the current turn.",
  218. parameters={
  219. "type": "object",
  220. "properties": {"query": {"type": "string"}},
  221. "required": ["query"],
  222. },
  223. handler=lambda resolved: {
  224. "tool": resolved.name,
  225. "query": resolved.arguments["query"],
  226. },
  227. argument_resolver=lambda event, context: {},
  228. )
  229. ]
  230. )
  231. agent = EventAgent(
  232. enabled_tools=["mock_search"],
  233. registry=registry,
  234. chat_client=WrongSourceToolCallChatClient(item),
  235. )
  236. reply = await agent.handle(
  237. event,
  238. history=[ChatMessage(role="assistant", content="fallback query")],
  239. )
  240. payload = json.loads(reply.content)
  241. assert payload["tool"] == "mock_search"
  242. assert payload == {
  243. "tool": "mock_search",
  244. "error": "missing required arguments: query",
  245. }
  246. @pytest.mark.asyncio
  247. async def test_event_agent_returns_registry_errors_for_disabled_and_unknown_tools():
  248. registry = ToolRegistry(
  249. [
  250. ToolDefinition(
  251. name="handoff_note",
  252. description="Send a note to the event agent.",
  253. parameters={"type": "object"},
  254. handler=lambda event: {"tool": event.name, "message": "handled"},
  255. )
  256. ]
  257. )
  258. disabled_reply = await EventAgent(
  259. enabled_tools=[],
  260. registry=registry,
  261. ).handle(
  262. ToolCallEvent(
  263. id="call_1",
  264. name="handoff_note",
  265. arguments={"message": "inspect this event"},
  266. raw_arguments='{"message":"inspect this event"}',
  267. )
  268. )
  269. unknown_reply = await EventAgent(
  270. enabled_tools=["missing_tool"],
  271. registry=registry,
  272. ).handle(
  273. ToolCallEvent(
  274. id="call_2",
  275. name="missing_tool",
  276. arguments={},
  277. raw_arguments="{}",
  278. )
  279. )
  280. assert json.loads(disabled_reply.content) == {
  281. "tool": "handoff_note",
  282. "error": "tool disabled",
  283. }
  284. assert json.loads(unknown_reply.content) == {
  285. "tool": "missing_tool",
  286. "error": "unknown tool",
  287. }
  288. @pytest.mark.asyncio
  289. async def test_event_agent_returns_structured_error_when_tool_handler_raises():
  290. def fail_tool(event: ToolCallEvent) -> dict:
  291. raise RuntimeError("boom")
  292. registry = ToolRegistry(
  293. [
  294. ToolDefinition(
  295. name="handoff_note",
  296. description="Send a note to the event agent.",
  297. parameters={"type": "object"},
  298. handler=fail_tool,
  299. )
  300. ]
  301. )
  302. event = ToolCallEvent(
  303. id="call_1",
  304. name="handoff_note",
  305. arguments={"message": "inspect this event"},
  306. raw_arguments='{"message":"inspect this event"}',
  307. )
  308. reply = await EventAgent(
  309. enabled_tools=["handoff_note"],
  310. registry=registry,
  311. ).handle(event)
  312. assert reply.role == "tool"
  313. assert reply.tool_call_id == "call_1"
  314. assert json.loads(reply.content) == {
  315. "tool": "handoff_note",
  316. "error": "tool handler failed: boom",
  317. }
  318. @pytest.mark.asyncio
  319. async def test_event_agent_llm_receives_history_and_agent_config_context():
  320. chat_client = ToolCallingChatClient(
  321. {"message": "Use strict tool parameters.", "thinking": "disabled"}
  322. )
  323. registry = ToolRegistry(
  324. [
  325. ToolDefinition(
  326. name="handoff_note",
  327. description="Send a note to the event agent.",
  328. parameters={
  329. "type": "object",
  330. "properties": {
  331. "message": {"type": "string"},
  332. "thinking": {"type": "string"},
  333. },
  334. "required": ["message", "thinking"],
  335. },
  336. handler=lambda event: {
  337. "tool": event.name,
  338. "message": event.arguments["message"],
  339. "thinking": event.arguments["thinking"],
  340. },
  341. )
  342. ]
  343. )
  344. event = ToolCallEvent(
  345. id="call_1",
  346. name="handoff_note",
  347. arguments={},
  348. raw_arguments="{}",
  349. )
  350. reply = await EventAgent(
  351. enabled_tools=["handoff_note"],
  352. registry=registry,
  353. chat_client=chat_client,
  354. params=AgentParams(model="event-model", temperature=0.4, max_tokens=120),
  355. ).handle(
  356. event,
  357. history=[ChatMessage(role="user", content="debug this")],
  358. system_prompt="Use strict tool parameters.",
  359. extra_body={"thinking": {"type": "disabled"}},
  360. )
  361. messages = chat_client.calls[0]["messages"]
  362. assert any(message.content == "debug this" for message in messages)
  363. assert any(message.content == "Use strict tool parameters." for message in messages)
  364. assert chat_client.calls[0]["params"].extra_body == {"thinking": {"type": "disabled"}}
  365. assert json.loads(reply.content) == {
  366. "tool": "handoff_note",
  367. "message": "Use strict tool parameters.",
  368. "thinking": "disabled",
  369. }
  370. @pytest.mark.asyncio
  371. async def test_event_agent_projects_complete_tool_round_to_visible_history():
  372. chat_client = ToolCallingChatClient({"message": "projected history"})
  373. registry = ToolRegistry(
  374. [
  375. ToolDefinition(
  376. name="handoff_note",
  377. description="Send a note to the event agent.",
  378. parameters={
  379. "type": "object",
  380. "properties": {"message": {"type": "string"}},
  381. "required": ["message"],
  382. },
  383. handler=lambda resolved: {
  384. "tool": resolved.name,
  385. "message": resolved.arguments["message"],
  386. },
  387. argument_resolver=lambda event, context: {},
  388. )
  389. ]
  390. )
  391. event = ToolCallEvent(
  392. id="call_current",
  393. name="handoff_note",
  394. arguments={},
  395. raw_arguments="{}",
  396. )
  397. prior_call = ToolCallEvent(
  398. id="call_prior",
  399. name="mock_search",
  400. arguments={"query": "latency"},
  401. raw_arguments='{"query":"latency"}',
  402. )
  403. await EventAgent(
  404. enabled_tools=["handoff_note"],
  405. registry=registry,
  406. chat_client=chat_client,
  407. ).handle(
  408. event,
  409. history=[
  410. ChatMessage(role="user", content="Find latency docs"),
  411. ChatMessage(
  412. role="assistant",
  413. content="I checked the latency sources.",
  414. tool_calls=[prior_call],
  415. ),
  416. ChatMessage(
  417. role="tool",
  418. content='{"results":["doc"]}',
  419. tool_call_id="call_prior",
  420. ),
  421. ChatMessage(role="user", content="Prepare a handoff"),
  422. ],
  423. )
  424. projected_messages = chat_client.calls[0]["messages"]
  425. visible_assistant = next(
  426. message
  427. for message in projected_messages
  428. if message.content == "I checked the latency sources."
  429. )
  430. assert visible_assistant.tool_calls == []
  431. assert not any(message.role == "tool" for message in projected_messages)