test_event_agent.py 9.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304
  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. ) -> AsyncIterator[StreamItem]:
  19. self.calls.append(
  20. {
  21. "messages": list(messages),
  22. "tools": list(tools),
  23. "params": params,
  24. }
  25. )
  26. tool_name = tools[0]["function"]["name"]
  27. yield StreamItem.raw_response_chunk(
  28. {
  29. "choices": [
  30. {
  31. "delta": {
  32. "tool_calls": [
  33. {
  34. "index": 0,
  35. "function": {"name": tool_name},
  36. }
  37. ]
  38. },
  39. "finish_reason": None,
  40. }
  41. ]
  42. }
  43. )
  44. yield StreamItem.event(
  45. ToolCallEvent(
  46. id="llm_call_1",
  47. name=tool_name,
  48. arguments=self.arguments,
  49. raw_arguments=json.dumps(self.arguments),
  50. )
  51. )
  52. class NoToolCallChatClient:
  53. def __init__(self) -> None:
  54. self.calls: list[dict] = []
  55. async def stream_chat(
  56. self,
  57. messages: list[ChatMessage],
  58. tools: list[dict],
  59. params: AgentParams,
  60. ) -> AsyncIterator[StreamItem]:
  61. self.calls.append(
  62. {
  63. "messages": list(messages),
  64. "tools": list(tools),
  65. "params": params,
  66. }
  67. )
  68. yield StreamItem.message_delta("I should have called the tool.")
  69. @pytest.mark.asyncio
  70. async def test_event_agent_resolves_tool_arguments_with_llm_tool_call():
  71. chat_client = ToolCallingChatClient({"message": "LLM generated handoff"})
  72. agent = EventAgent(
  73. enabled_tools=["handoff_note"],
  74. chat_client=chat_client,
  75. params=AgentParams(model="event-model", temperature=0, max_tokens=80),
  76. )
  77. event = ToolCallEvent(
  78. id="call_1",
  79. name="handoff_note",
  80. arguments={"message": "chat agent argument should be ignored"},
  81. raw_arguments='{"message":"chat agent argument should be ignored"}',
  82. )
  83. history = [
  84. ChatMessage(role="user", content="debug this event flow"),
  85. ChatMessage(role="assistant", content="I need the event agent."),
  86. ]
  87. reply = await agent.handle(event, history=history)
  88. assert reply.role == "tool"
  89. assert reply.tool_call_id == "call_1"
  90. assert reply.name == "handoff_note"
  91. assert json.loads(reply.content) == {
  92. "tool": "handoff_note",
  93. "message": "LLM generated handoff",
  94. }
  95. assert chat_client.calls[0]["tools"] == [
  96. {
  97. "type": "function",
  98. "function": {
  99. "name": "handoff_note",
  100. "description": "Send a note to the event agent.",
  101. "parameters": {
  102. "type": "object",
  103. "properties": {
  104. "message": {"type": "string"},
  105. },
  106. "required": ["message"],
  107. },
  108. },
  109. }
  110. ]
  111. assert chat_client.calls[0]["params"].model == "event-model"
  112. assert agent.raw_model_chunks([event]) == [
  113. {
  114. "event_id": "call_1",
  115. "event_name": "handoff_note",
  116. "chunks": [
  117. {
  118. "choices": [
  119. {
  120. "delta": {
  121. "tool_calls": [
  122. {
  123. "index": 0,
  124. "function": {"name": "handoff_note"},
  125. }
  126. ]
  127. },
  128. "finish_reason": None,
  129. }
  130. ]
  131. }
  132. ],
  133. }
  134. ]
  135. @pytest.mark.asyncio
  136. async def test_event_agent_falls_back_to_context_arguments_when_llm_returns_no_tool_call():
  137. chat_client = NoToolCallChatClient()
  138. agent = EventAgent(
  139. enabled_tools=["mock_search"],
  140. chat_client=chat_client,
  141. params=AgentParams(model="event-model", temperature=0, max_tokens=80),
  142. )
  143. event = ToolCallEvent(
  144. id="call_1",
  145. name="mock_search",
  146. arguments={},
  147. raw_arguments="{}",
  148. )
  149. history = [
  150. ChatMessage(role="user", content="Find latency docs"),
  151. ChatMessage(role="assistant", content="Need a search for latency docs"),
  152. ]
  153. reply = await agent.handle(event, history=history)
  154. payload = json.loads(reply.content)
  155. assert payload["tool"] == "mock_search"
  156. assert payload["query"] == "Need a search for latency docs"
  157. assert "event agent did not return arguments" not in reply.content
  158. assert chat_client.calls[0]["tools"][0]["function"]["name"] == "mock_search"
  159. @pytest.mark.asyncio
  160. async def test_event_agent_returns_registry_errors_for_disabled_and_unknown_tools():
  161. registry = ToolRegistry(
  162. [
  163. ToolDefinition(
  164. name="handoff_note",
  165. description="Send a note to the event agent.",
  166. parameters={"type": "object"},
  167. handler=lambda event: {"tool": event.name, "message": "handled"},
  168. )
  169. ]
  170. )
  171. disabled_reply = await EventAgent(
  172. enabled_tools=[],
  173. registry=registry,
  174. ).handle(
  175. ToolCallEvent(
  176. id="call_1",
  177. name="handoff_note",
  178. arguments={"message": "inspect this event"},
  179. raw_arguments='{"message":"inspect this event"}',
  180. )
  181. )
  182. unknown_reply = await EventAgent(
  183. enabled_tools=["missing_tool"],
  184. registry=registry,
  185. ).handle(
  186. ToolCallEvent(
  187. id="call_2",
  188. name="missing_tool",
  189. arguments={},
  190. raw_arguments="{}",
  191. )
  192. )
  193. assert json.loads(disabled_reply.content) == {
  194. "tool": "handoff_note",
  195. "error": "tool disabled",
  196. }
  197. assert json.loads(unknown_reply.content) == {
  198. "tool": "missing_tool",
  199. "error": "unknown tool",
  200. }
  201. @pytest.mark.asyncio
  202. async def test_event_agent_returns_structured_error_when_tool_handler_raises():
  203. def fail_tool(event: ToolCallEvent) -> dict:
  204. raise RuntimeError("boom")
  205. registry = ToolRegistry(
  206. [
  207. ToolDefinition(
  208. name="handoff_note",
  209. description="Send a note to the event agent.",
  210. parameters={"type": "object"},
  211. handler=fail_tool,
  212. )
  213. ]
  214. )
  215. event = ToolCallEvent(
  216. id="call_1",
  217. name="handoff_note",
  218. arguments={"message": "inspect this event"},
  219. raw_arguments='{"message":"inspect this event"}',
  220. )
  221. reply = await EventAgent(
  222. enabled_tools=["handoff_note"],
  223. registry=registry,
  224. ).handle(event)
  225. assert reply.role == "tool"
  226. assert reply.tool_call_id == "call_1"
  227. assert json.loads(reply.content) == {
  228. "tool": "handoff_note",
  229. "error": "tool handler failed: boom",
  230. }
  231. @pytest.mark.asyncio
  232. async def test_event_agent_llm_receives_history_and_agent_config_context():
  233. chat_client = ToolCallingChatClient(
  234. {"message": "Use strict tool parameters.", "thinking": "disabled"}
  235. )
  236. registry = ToolRegistry(
  237. [
  238. ToolDefinition(
  239. name="handoff_note",
  240. description="Send a note to the event agent.",
  241. parameters={"type": "object"},
  242. handler=lambda event: {
  243. "tool": event.name,
  244. "message": event.arguments["message"],
  245. "thinking": event.arguments["thinking"],
  246. },
  247. )
  248. ]
  249. )
  250. event = ToolCallEvent(
  251. id="call_1",
  252. name="handoff_note",
  253. arguments={},
  254. raw_arguments="{}",
  255. )
  256. reply = await EventAgent(
  257. enabled_tools=["handoff_note"],
  258. registry=registry,
  259. chat_client=chat_client,
  260. params=AgentParams(model="event-model", temperature=0.4, max_tokens=120),
  261. ).handle(
  262. event,
  263. history=[ChatMessage(role="user", content="debug this")],
  264. system_prompt="Use strict tool parameters.",
  265. extra_body={"thinking": {"type": "disabled"}},
  266. )
  267. messages = chat_client.calls[0]["messages"]
  268. assert any(message.content == "debug this" for message in messages)
  269. assert any(message.content == "Use strict tool parameters." for message in messages)
  270. assert chat_client.calls[0]["params"].extra_body == {"thinking": {"type": "disabled"}}
  271. assert json.loads(reply.content) == {
  272. "tool": "handoff_note",
  273. "message": "Use strict tool parameters.",
  274. "thinking": "disabled",
  275. }