test_debug_runtime.py 3.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125
  1. from collections.abc import AsyncIterator
  2. import pytest
  3. from agent_lab.application.contracts import AgentParams, DebugRunRequest, EventAgentParams
  4. from agent_lab.application.runtime import DebugRuntime
  5. from agent_lab.domain.events import ToolCallEvent
  6. from agent_lab.domain.messages import ChatMessage, StreamItem
  7. class FakeChatClient:
  8. def __init__(self) -> None:
  9. self.calls = 0
  10. async def stream_chat(
  11. self,
  12. messages: list[ChatMessage],
  13. tools: list[dict],
  14. params: AgentParams,
  15. ) -> AsyncIterator[StreamItem]:
  16. self.calls += 1
  17. if self.calls == 1:
  18. yield StreamItem.event(
  19. ToolCallEvent(
  20. id="call_1",
  21. name="handoff_note",
  22. arguments={"message": "need event agent"},
  23. raw_arguments='{"message":"need event agent"}',
  24. )
  25. )
  26. return
  27. assert any(message.role == "tool" for message in messages)
  28. yield StreamItem.message_delta("final answer")
  29. class StrictHistoryChatClient:
  30. def __init__(self) -> None:
  31. self.calls = 0
  32. self.second_call_messages: list[ChatMessage] = []
  33. async def stream_chat(
  34. self,
  35. messages: list[ChatMessage],
  36. tools: list[dict],
  37. params: AgentParams,
  38. ) -> AsyncIterator[StreamItem]:
  39. self.calls += 1
  40. if self.calls == 1:
  41. yield StreamItem.event(
  42. ToolCallEvent(
  43. id="call_1",
  44. name="handoff_note",
  45. arguments={"message": "need event agent"},
  46. raw_arguments='{"message":"need event agent"}',
  47. )
  48. )
  49. return
  50. self.second_call_messages = list(messages)
  51. yield StreamItem.message_delta("final answer")
  52. @pytest.mark.asyncio
  53. async def test_runtime_routes_chat_events_through_event_agent_then_continues_chat():
  54. request = DebugRunRequest(
  55. user_message="debug this",
  56. system_prompts=["You are a debugger."],
  57. pre_messages=[],
  58. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  59. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
  60. )
  61. client = FakeChatClient()
  62. runtime = DebugRuntime(client)
  63. outputs = [message async for message in runtime.run(request)]
  64. assert client.calls == 2
  65. assert [message["type"] for message in outputs] == [
  66. "session_started",
  67. "event",
  68. "tool_result",
  69. "message_delta",
  70. "done",
  71. ]
  72. assert outputs[1]["event"]["name"] == "handoff_note"
  73. assert outputs[3]["content"] == "final answer"
  74. @pytest.mark.asyncio
  75. async def test_runtime_preserves_assistant_tool_calls_before_tool_reply():
  76. request = DebugRunRequest(
  77. user_message="debug this",
  78. system_prompts=["You are a debugger."],
  79. pre_messages=[],
  80. chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
  81. event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
  82. )
  83. client = StrictHistoryChatClient()
  84. runtime = DebugRuntime(client)
  85. outputs = [message async for message in runtime.run(request)]
  86. assert client.calls == 2
  87. assert [message.role for message in client.second_call_messages] == [
  88. "system",
  89. "user",
  90. "assistant",
  91. "tool",
  92. ]
  93. assistant_message = client.second_call_messages[2]
  94. tool_message = client.second_call_messages[3]
  95. assert assistant_message.content == ""
  96. assert assistant_message.tool_calls == [
  97. {
  98. "id": "call_1",
  99. "type": "function",
  100. "function": {
  101. "name": "handoff_note",
  102. "arguments": '{"message":"need event agent"}',
  103. },
  104. }
  105. ]
  106. assert tool_message.tool_call_id == "call_1"
  107. assert outputs[-1] == {"type": "done"}