| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461 |
- import asyncio
- import importlib
- from collections.abc import AsyncIterator
- from typing import Any
- import pytest
- from agent_lab.application.contracts import AgentParams, DebugRunRequest, EventAgentParams
- from agent_lab.application.runtime import DebugRuntime
- from agent_lab.application.tools import ToolDefinition, ToolRegistry
- from agent_lab.domain.events import ToolCallEvent
- from agent_lab.domain.messages import ChatMessage, StreamItem, TokenUsage
- def _runtime_queues_class():
- module = importlib.import_module("agent_lab.application.queues")
- return module.RuntimeQueues
- class RecordingQueue(asyncio.Queue):
- def __init__(self, name: str, log: list[tuple[str, str, str]]) -> None:
- super().__init__()
- self.name = name
- self.log = log
- async def put(self, item: Any) -> None:
- self.log.append((self.name, "put", self._describe(item)))
- await super().put(item)
- async def get(self) -> Any:
- item = await super().get()
- self.log.append((self.name, "get", self._describe(item)))
- return item
- def _describe(self, item: Any) -> str:
- if isinstance(item, ChatMessage):
- if item.role == "tool":
- return f"tool:{item.tool_call_id}"
- return item.role
- if isinstance(item, ToolCallEvent):
- return f"event:{item.name}:{item.id}"
- if isinstance(item, dict):
- return f"output:{item.get('type')}"
- return type(item).__name__
- 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")
- class ToolCapturingChatClient:
- def __init__(self) -> None:
- self.tools: list[dict[str, Any]] = []
- async def stream_chat(
- self,
- messages: list[ChatMessage],
- tools: list[dict],
- params: AgentParams,
- ) -> AsyncIterator[StreamItem]:
- self.tools = list(tools)
- yield StreamItem.message_delta("final answer")
- class RoundStatsChatClient:
- async def stream_chat(
- self,
- messages: list[ChatMessage],
- tools: list[dict],
- params: AgentParams,
- ) -> AsyncIterator[StreamItem]:
- yield StreamItem.message_delta("hello")
- yield StreamItem.usage_item(
- TokenUsage(
- prompt_tokens=10,
- completion_tokens=20,
- total_tokens=30,
- cached_tokens=5,
- )
- )
- class EventRoundStatsChatClient:
- 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"}',
- )
- )
- yield StreamItem.usage_item(
- TokenUsage(prompt_tokens=3, completion_tokens=0, total_tokens=3)
- )
- return
- yield StreamItem.message_delta("final answer")
- yield StreamItem.usage_item(
- TokenUsage(prompt_tokens=4, completion_tokens=6, total_tokens=10)
- )
- def test_runtime_queues_exposes_input_output_and_events_queues():
- RuntimeQueues = _runtime_queues_class()
- queues = RuntimeQueues()
- assert isinstance(queues.input, asyncio.Queue)
- assert isinstance(queues.output, asyncio.Queue)
- assert isinstance(queues.events, asyncio.Queue)
- assert queues.input is not queues.output
- assert queues.input is not queues.events
- assert queues.output is not queues.events
- @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",
- "round_stats",
- "message_delta",
- "round_stats",
- "done",
- ]
- assert outputs[1]["event"]["name"] == "handoff_note"
- assert outputs[4]["content"] == "final answer"
- @pytest.mark.asyncio
- async def test_runtime_uses_event_and_input_queues_for_event_agent_handoff():
- RuntimeQueues = _runtime_queues_class()
- queue_log: list[tuple[str, str, str]] = []
- queues = RuntimeQueues(
- input=RecordingQueue("input", queue_log),
- output=RecordingQueue("output", queue_log),
- events=RecordingQueue("events", queue_log),
- )
- 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),
- )
- runtime = DebugRuntime(FakeChatClient(), queues=queues)
- outputs = [message async for message in runtime.run(request)]
- assert [message["type"] for message in outputs] == [
- "session_started",
- "event",
- "tool_result",
- "round_stats",
- "message_delta",
- "round_stats",
- "done",
- ]
- assert queue_log.index(("input", "put", "user")) < queue_log.index(
- ("input", "get", "user")
- )
- assert queue_log.index(("events", "put", "event:handoff_note:call_1")) < queue_log.index(
- ("events", "get", "event:handoff_note:call_1")
- )
- assert queue_log.index(("events", "get", "event:handoff_note:call_1")) < queue_log.index(
- ("input", "put", "tool:call_1")
- )
- assert queue_log.index(("input", "put", "tool:call_1")) < queue_log.index(
- ("input", "get", "tool:call_1")
- )
- @pytest.mark.asyncio
- async def test_runtime_yields_existing_output_order_from_output_queue():
- RuntimeQueues = _runtime_queues_class()
- queue_log: list[tuple[str, str, str]] = []
- queues = RuntimeQueues(
- input=RecordingQueue("input", queue_log),
- output=RecordingQueue("output", queue_log),
- events=RecordingQueue("events", queue_log),
- )
- 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),
- )
- runtime = DebugRuntime(FakeChatClient(), queues=queues)
- outputs = [message async for message in runtime.run(request)]
- assert [message["type"] for message in outputs] == [
- "session_started",
- "event",
- "tool_result",
- "round_stats",
- "message_delta",
- "round_stats",
- "done",
- ]
- assert [
- entry
- for entry in queue_log
- if entry[0] == "output" and entry[1] in {"put", "get"}
- ] == [
- ("output", "put", "output:session_started"),
- ("output", "get", "output:session_started"),
- ("output", "put", "output:event"),
- ("output", "get", "output:event"),
- ("output", "put", "output:tool_result"),
- ("output", "get", "output:tool_result"),
- ("output", "put", "output:round_stats"),
- ("output", "get", "output:round_stats"),
- ("output", "put", "output:message_delta"),
- ("output", "get", "output:message_delta"),
- ("output", "put", "output:round_stats"),
- ("output", "get", "output:round_stats"),
- ("output", "put", "output:done"),
- ("output", "get", "output:done"),
- ]
- @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"}
- @pytest.mark.asyncio
- async def test_runtime_passes_selected_tool_schema_from_registry_to_chat_agent():
- registry = ToolRegistry(
- [
- ToolDefinition(
- name="handoff_note",
- description="Registry-owned handoff tool.",
- parameters={
- "type": "object",
- "properties": {
- "message": {"type": "string"},
- "priority": {"type": "number"},
- },
- "required": ["message"],
- },
- handler=lambda event: {"tool": event.name, "message": "handled"},
- )
- ]
- )
- request = DebugRunRequest(
- user_message="debug this",
- system_prompts=[],
- 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 = ToolCapturingChatClient()
- runtime = DebugRuntime(client, registry=registry)
- outputs = [message async for message in runtime.run(request)]
- assert outputs[-1] == {"type": "done"}
- assert client.tools == [
- {
- "type": "function",
- "function": {
- "name": "handoff_note",
- "description": "Registry-owned handoff tool.",
- "parameters": {
- "type": "object",
- "properties": {
- "message": {"type": "string"},
- "priority": {"type": "number"},
- },
- "required": ["message"],
- },
- },
- }
- ]
- @pytest.mark.asyncio
- async def test_runtime_emits_round_stats_with_clock_and_usage_after_model_turn():
- request = DebugRunRequest(
- user_message="debug this",
- system_prompts=[],
- pre_messages=[],
- chat_agent=AgentParams(model="fake-model", temperature=0.1, max_tokens=200),
- event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
- )
- ticks = iter([1.0, 1.123, 1.456])
- runtime = DebugRuntime(RoundStatsChatClient(), clock=lambda: next(ticks))
- outputs = [message async for message in runtime.run(request)]
- assert [message["type"] for message in outputs] == [
- "session_started",
- "message_delta",
- "usage",
- "round_stats",
- "done",
- ]
- assert outputs[3] == {
- "type": "round_stats",
- "round_index": 1,
- "ttft_ms": 123,
- "elapsed_ms": 456,
- "prompt_tokens": 10,
- "completion_tokens": 20,
- "total_tokens": 30,
- "cached_tokens": 5,
- "had_event": False,
- }
- @pytest.mark.asyncio
- async def test_runtime_emits_round_stats_for_each_chat_call_in_event_handoff():
- request = DebugRunRequest(
- user_message="debug this",
- system_prompts=[],
- 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),
- )
- ticks = iter([2.0, 2.25, 3.0, 3.05, 3.2])
- client = EventRoundStatsChatClient()
- runtime = DebugRuntime(client, clock=lambda: next(ticks))
- outputs = [message async for message in runtime.run(request)]
- stats = [message for message in outputs if message["type"] == "round_stats"]
- assert client.calls == 2
- assert stats == [
- {
- "type": "round_stats",
- "round_index": 1,
- "ttft_ms": None,
- "elapsed_ms": 250,
- "prompt_tokens": 3,
- "completion_tokens": 0,
- "total_tokens": 3,
- "cached_tokens": 0,
- "had_event": True,
- },
- {
- "type": "round_stats",
- "round_index": 2,
- "ttft_ms": 50,
- "elapsed_ms": 200,
- "prompt_tokens": 4,
- "completion_tokens": 6,
- "total_tokens": 10,
- "cached_tokens": 0,
- "had_event": False,
- },
- ]
|