| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125 |
- from collections.abc import AsyncIterator
- import pytest
- from agent_lab.application.contracts import AgentParams, DebugRunRequest, EventAgentParams
- from agent_lab.application.runtime import DebugRuntime
- from agent_lab.domain.events import ToolCallEvent
- from agent_lab.domain.messages import ChatMessage, StreamItem
- class FakeChatClient:
- def __init__(self) -> None:
- self.calls = 0
- async def stream_chat(
- self,
- messages: list[ChatMessage],
- tools: list[dict],
- params: AgentParams,
- ) -> AsyncIterator[StreamItem]:
- self.calls += 1
- if self.calls == 1:
- yield StreamItem.event(
- ToolCallEvent(
- id="call_1",
- name="handoff_note",
- arguments={"message": "need event agent"},
- raw_arguments='{"message":"need event agent"}',
- )
- )
- return
- assert any(message.role == "tool" for message in messages)
- yield StreamItem.message_delta("final answer")
- class StrictHistoryChatClient:
- def __init__(self) -> None:
- self.calls = 0
- self.second_call_messages: list[ChatMessage] = []
- async def stream_chat(
- self,
- messages: list[ChatMessage],
- tools: list[dict],
- params: AgentParams,
- ) -> AsyncIterator[StreamItem]:
- self.calls += 1
- if self.calls == 1:
- yield StreamItem.event(
- ToolCallEvent(
- id="call_1",
- name="handoff_note",
- arguments={"message": "need event agent"},
- raw_arguments='{"message":"need event agent"}',
- )
- )
- return
- self.second_call_messages = list(messages)
- yield StreamItem.message_delta("final answer")
- @pytest.mark.asyncio
- async def test_runtime_routes_chat_events_through_event_agent_then_continues_chat():
- request = DebugRunRequest(
- user_message="debug this",
- system_prompts=["You are a debugger."],
- pre_messages=[],
- chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
- event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=3),
- )
- client = FakeChatClient()
- runtime = DebugRuntime(client)
- outputs = [message async for message in runtime.run(request)]
- assert client.calls == 2
- assert [message["type"] for message in outputs] == [
- "session_started",
- "event",
- "tool_result",
- "message_delta",
- "done",
- ]
- assert outputs[1]["event"]["name"] == "handoff_note"
- assert outputs[3]["content"] == "final answer"
- @pytest.mark.asyncio
- async def test_runtime_preserves_assistant_tool_calls_before_tool_reply():
- request = DebugRunRequest(
- user_message="debug this",
- system_prompts=["You are a debugger."],
- pre_messages=[],
- chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
- event_agent=EventAgentParams(enabled_tools=["handoff_note"], max_event_loops=2),
- )
- client = StrictHistoryChatClient()
- runtime = DebugRuntime(client)
- outputs = [message async for message in runtime.run(request)]
- assert client.calls == 2
- assert [message.role for message in client.second_call_messages] == [
- "system",
- "user",
- "assistant",
- "tool",
- ]
- assistant_message = client.second_call_messages[2]
- tool_message = client.second_call_messages[3]
- assert assistant_message.content == ""
- assert assistant_message.tool_calls == [
- {
- "id": "call_1",
- "type": "function",
- "function": {
- "name": "handoff_note",
- "arguments": '{"message":"need event agent"}',
- },
- }
- ]
- assert tool_message.tool_call_id == "call_1"
- assert outputs[-1] == {"type": "done"}
|