| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884885886887888889890891892893894895896897898899900901902903904905906907908909910911912913914915916917918919920921922923924925926927928929930931932933934935936937938939940941942943944945946947948949950951952953954955956957958959960961962963964965966967968969970971972973974975976977978979980981 |
- import asyncio
- import importlib
- import json
- import logging
- from collections.abc import AsyncIterator
- from typing import Any
- import pytest
- from agent_lab.application.contracts import AgentParams, DebugRunRequest, EventAgentParams
- from agent_lab.application.event_agent import EventAgentRequest
- from agent_lab.application.runtime import DebugRuntime
- from agent_lab.application.tools import ToolDefinition, ToolExecutionContext, 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
- async def _collect_outputs(stream: AsyncIterator[dict[str, Any]]) -> list[dict[str, Any]]:
- return [message async for message in stream]
- def _without_audit(outputs: list[dict[str, Any]]) -> list[dict[str, Any]]:
- return [message for message in outputs if message["type"] != "audit"]
- def _message_types(outputs: list[dict[str, Any]]) -> list[str]:
- return [message["type"] for message in _without_audit(outputs)]
- def _event_tool_call_from_tools(
- tools: list[dict],
- messages: list[ChatMessage],
- ) -> StreamItem:
- tool_name = tools[0]["function"]["name"]
- content = ""
- for role in ("assistant", "user"):
- content = next(
- (
- message.content
- for message in reversed(messages)
- if message.role == role and message.content.strip()
- ),
- "",
- )
- if content:
- break
- arguments = {
- "message": content,
- "query": content,
- "title": content,
- }
- return StreamItem.event(
- ToolCallEvent(
- id="event_agent_call_1",
- name=tool_name,
- arguments=arguments,
- raw_arguments=json.dumps(arguments),
- )
- )
- async def _next_non_audit(
- stream: AsyncIterator[dict[str, Any]],
- ) -> dict[str, Any]:
- while True:
- message = await anext(stream)
- if message["type"] != "audit":
- return message
- async def _next_non_audit_from_queue(queues: Any) -> dict[str, Any]:
- while True:
- message = await queues.output.get()
- if message["type"] != "audit":
- return message
- 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, EventAgentRequest):
- events = ",".join(f"{event.name}:{event.id}" for event in item.events)
- return f"event_request:{events}"
- 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]:
- if tools:
- yield _event_tool_call_from_tools(tools, messages)
- return
- self.calls += 1
- if self.calls == 1:
- yield StreamItem.event(
- ToolCallEvent(
- id="call_1",
- name="handoff_note",
- arguments={},
- raw_arguments="{}",
- )
- )
- return
- assert not any(message.role == "tool" for message in messages)
- assert any(message.name == "event_agent" 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]:
- if tools:
- yield _event_tool_call_from_tools(tools, messages)
- return
- self.calls += 1
- if self.calls == 1:
- yield StreamItem.event(
- ToolCallEvent(
- id="call_1",
- name="handoff_note",
- arguments={},
- raw_arguments="{}",
- )
- )
- return
- self.second_call_messages = list(messages)
- yield StreamItem.message_delta("final answer")
- class MultiEventChatClient:
- 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]:
- if tools:
- yield _event_tool_call_from_tools(tools, messages)
- return
- self.calls += 1
- if self.calls == 1:
- yield StreamItem.message_delta("Checking events.")
- yield StreamItem.event(
- ToolCallEvent(
- id="call_1",
- name="handoff_note",
- arguments={"message": "ignored chat argument"},
- raw_arguments='{"message":"ignored chat argument"}',
- )
- )
- yield StreamItem.event(
- ToolCallEvent(
- id="call_2",
- name="audit_note",
- arguments={"message": "ignored chat argument"},
- raw_arguments='{"message":"ignored chat argument"}',
- )
- )
- 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]] = []
- self.messages: list[ChatMessage] = []
- async def stream_chat(
- self,
- messages: list[ChatMessage],
- tools: list[dict],
- params: AgentParams,
- ) -> AsyncIterator[StreamItem]:
- self.messages = list(messages)
- self.tools = list(tools)
- yield StreamItem.message_delta("final answer")
- class EventLoopLimitChatClient:
- def __init__(self) -> None:
- self.calls = 0
- self.tools_by_call: list[list[dict[str, Any]]] = []
- self.messages_by_call: list[list[ChatMessage]] = []
- async def stream_chat(
- self,
- messages: list[ChatMessage],
- tools: list[dict],
- params: AgentParams,
- ) -> AsyncIterator[StreamItem]:
- if tools:
- yield _event_tool_call_from_tools(tools, messages)
- return
- self.calls += 1
- self.tools_by_call.append(list(tools))
- self.messages_by_call.append(list(messages))
- if self.calls == 1:
- yield StreamItem.event(
- ToolCallEvent(
- id="call_1",
- name="handoff_note",
- arguments={},
- raw_arguments="{}",
- )
- )
- return
- yield StreamItem.message_delta("final after event limit")
- 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]:
- if tools:
- yield _event_tool_call_from_tools(tools, messages)
- return
- self.calls += 1
- if self.calls == 1:
- yield StreamItem.event(
- ToolCallEvent(
- id="call_1",
- name="handoff_note",
- arguments={},
- raw_arguments="{}",
- )
- )
- 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)
- )
- class TwoTurnSessionChatClient:
- def __init__(self) -> None:
- self.calls = 0
- self.messages_by_call: list[list[ChatMessage]] = []
- async def stream_chat(
- self,
- messages: list[ChatMessage],
- tools: list[dict],
- params: AgentParams,
- ) -> AsyncIterator[StreamItem]:
- if tools:
- yield _event_tool_call_from_tools(tools, messages)
- return
- self.calls += 1
- self.messages_by_call.append(list(messages))
- if self.calls in {1, 3}:
- yield StreamItem.event(
- ToolCallEvent(
- id=f"call_{self.calls}",
- name="handoff_note",
- arguments={},
- raw_arguments="{}",
- )
- )
- return
- yield StreamItem.message_delta(f"final answer {self.calls}")
- class SlowAfterEventChatClient:
- def __init__(self) -> None:
- self.calls = 0
- self.event_seen = asyncio.Event()
- self.release_stream = asyncio.Event()
- async def stream_chat(
- self,
- messages: list[ChatMessage],
- tools: list[dict],
- params: AgentParams,
- ) -> AsyncIterator[StreamItem]:
- if tools:
- yield _event_tool_call_from_tools(tools, messages)
- return
- self.calls += 1
- if self.calls == 1:
- yield StreamItem.message_delta("Need event.")
- yield StreamItem.event(
- ToolCallEvent(
- id="call_1",
- name="handoff_note",
- arguments={},
- raw_arguments="{}",
- )
- )
- self.event_seen.set()
- await self.release_stream.wait()
- return
- yield StreamItem.message_delta("final answer")
- 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)]
- business_outputs = _without_audit(outputs)
- assert client.calls == 2
- assert _message_types(outputs) == [
- "session_started",
- "event",
- "tool_result",
- "round_stats",
- "message_delta",
- "round_stats",
- "done",
- ]
- assert business_outputs[1]["event"]["name"] == "handoff_note"
- assert business_outputs[4]["content"] == "final answer"
- @pytest.mark.asyncio
- async def test_runtime_emits_audit_events_and_backend_logs(caplog):
- 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=1),
- )
- runtime = DebugRuntime(FakeChatClient())
- with caplog.at_level(logging.INFO, logger="agent_lab.application.runtime"):
- outputs = [message async for message in runtime.run(request)]
- audit_events = [
- message["event"]
- for message in outputs
- if message["type"] == "audit"
- ]
- assert audit_events == [
- "session_started",
- "chat_round_started",
- "chat_event_detected",
- "event_agent_completed",
- "chat_round_finished",
- "chat_round_started",
- "chat_round_finished",
- "session_finished",
- ]
- assert "chat_event_detected" in caplog.text
- assert "event_agent_completed" in caplog.text
- @pytest.mark.asyncio
- async def test_runtime_outputs_event_as_soon_as_chat_stream_detects_it():
- RuntimeQueues = _runtime_queues_class()
- queues = RuntimeQueues()
- client = SlowAfterEventChatClient()
- runtime = DebugRuntime(client, queues=queues)
- 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=1),
- )
- runtime.start(request)
- try:
- assert await _next_non_audit_from_queue(queues) == {"type": "session_started"}
- assert await _next_non_audit_from_queue(queues) == {
- "type": "message_delta",
- "content": "Need event.",
- }
- await asyncio.wait_for(client.event_seen.wait(), timeout=1)
- event_message = await asyncio.wait_for(
- _next_non_audit_from_queue(queues),
- timeout=0.2,
- )
- assert event_message["type"] == "event"
- assert event_message["event"]["name"] == "handoff_note"
- finally:
- client.release_stream.set()
- await runtime.aclose()
- @pytest.mark.asyncio
- async def test_runtime_batches_round_events_before_continuing_chat_agent():
- def resolve_from_history(
- event: ToolCallEvent,
- context: ToolExecutionContext,
- ) -> dict[str, Any]:
- return {"message": context.history[-1].content, "event": event.name}
- registry = ToolRegistry(
- [
- ToolDefinition(
- name="handoff_note",
- description="Send a handoff note.",
- parameters={
- "type": "object",
- "properties": {"message": {"type": "string"}},
- "required": ["message"],
- },
- handler=lambda event: {
- "tool": event.name,
- "message": event.arguments["message"],
- },
- argument_resolver=resolve_from_history,
- ),
- ToolDefinition(
- name="audit_note",
- description="Send an audit note.",
- parameters={
- "type": "object",
- "properties": {"message": {"type": "string"}},
- "required": ["message"],
- },
- handler=lambda event: {
- "tool": event.name,
- "message": event.arguments["message"],
- },
- argument_resolver=resolve_from_history,
- ),
- ]
- )
- 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", "audit_note"],
- max_event_loops=2,
- ),
- )
- client = MultiEventChatClient()
- runtime = DebugRuntime(client, registry=registry)
- outputs = [message async for message in runtime.run(request)]
- assert _message_types(outputs) == [
- "session_started",
- "message_delta",
- "event",
- "event",
- "tool_result",
- "tool_result",
- "round_stats",
- "message_delta",
- "round_stats",
- "done",
- ]
- assert client.calls == 2
- assert [message.role for message in client.second_call_messages] == [
- "system",
- "user",
- "assistant",
- "user",
- ]
- assistant_message = client.second_call_messages[2]
- assert assistant_message.content == "Checking events."
- assert not any(message.role == "tool" for message in client.second_call_messages)
- assert client.second_call_messages[-1].content == (
- "EventAgent results:\n"
- '{"tool": "handoff_note", "message": "Checking events."}\n'
- '{"tool": "audit_note", "message": "Checking events."}'
- )
- assert client.second_call_messages[-1].name == "event_agent"
- @pytest.mark.asyncio
- async def test_runtime_start_returns_queues_for_downstream_output_consumer():
- 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),
- )
- runtime = DebugRuntime(RoundStatsChatClient())
- queues = runtime.start(request)
- outputs: list[dict[str, Any]] = []
- while True:
- message = await asyncio.wait_for(queues.output.get(), timeout=1)
- outputs.append(message)
- if message["type"] == "done":
- break
- assert _message_types(outputs) == [
- "session_started",
- "message_delta",
- "usage",
- "round_stats",
- "done",
- ]
- @pytest.mark.asyncio
- async def test_runtime_finalizes_chat_after_reaching_event_loop_limit():
- 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=1),
- )
- client = EventLoopLimitChatClient()
- runtime = DebugRuntime(client)
- outputs = [message async for message in runtime.run(request)]
- business_outputs = _without_audit(outputs)
- assert client.calls == 2
- assert client.tools_by_call[0] == []
- assert client.tools_by_call[1] == []
- assert "Available events:" in client.messages_by_call[0][0].content
- assert not any(
- "Available events:" in message.content
- for message in client.messages_by_call[1]
- if message.role == "system"
- )
- assert _message_types(outputs) == [
- "session_started",
- "event",
- "tool_result",
- "round_stats",
- "message_delta",
- "round_stats",
- "done",
- ]
- assert business_outputs[4]["content"] == "final after event limit"
- @pytest.mark.asyncio
- async def test_runtime_buffers_upstream_user_input_until_after_matching_tool_reply():
- RuntimeQueues = _runtime_queues_class()
- queues = RuntimeQueues()
- 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 = StrictHistoryChatClient()
- runtime = DebugRuntime(client, queues=queues)
- stream = runtime.run(request)
- assert await _next_non_audit(stream) == {"type": "session_started"}
- event_message = await _next_non_audit(stream)
- assert event_message["type"] == "event"
- await queues.input.put(ChatMessage(role="user", content="follow-up while tool runs"))
- remaining = [message async for message in stream]
- assert remaining[-1] == {"type": "done"}
- assert [message.role for message in client.second_call_messages] == [
- "system",
- "user",
- "assistant",
- "user",
- "user",
- ]
- assert not any(message.role == "tool" for message in client.second_call_messages)
- assert client.second_call_messages[3].content.startswith("EventAgent results:\n")
- assert client.second_call_messages[3].name == "event_agent"
- assert client.second_call_messages[4].content == "follow-up while tool runs"
- @pytest.mark.asyncio
- async def test_runtime_continues_when_event_agent_tool_handler_raises():
- def fail_tool(event: ToolCallEvent) -> dict[str, Any]:
- raise RuntimeError("boom")
- registry = ToolRegistry(
- [
- ToolDefinition(
- name="handoff_note",
- description="Broken handoff tool.",
- parameters={"type": "object"},
- handler=fail_tool,
- )
- ]
- )
- 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),
- )
- runtime = DebugRuntime(FakeChatClient(), registry=registry)
- outputs = await asyncio.wait_for(
- _collect_outputs(runtime.run(request)),
- timeout=1,
- )
- business_outputs = _without_audit(outputs)
- assert _message_types(outputs) == [
- "session_started",
- "event",
- "tool_result",
- "round_stats",
- "message_delta",
- "round_stats",
- "done",
- ]
- assert json.loads(business_outputs[2]["message"]["content"]) == {
- "tool": "handoff_note",
- "error": "tool handler failed: boom",
- }
- @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_types(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_request:handoff_note:call_1")
- ) < queue_log.index(("events", "get", "event_request:handoff_note:call_1"))
- assert queue_log.index(
- ("events", "get", "event_request: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_run_consumes_output_queue_in_stream_order():
- 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_types(outputs) == [
- "session_started",
- "event",
- "tool_result",
- "round_stats",
- "message_delta",
- "round_stats",
- "done",
- ]
- output_puts = [
- entry[2] for entry in queue_log if entry[0] == "output" and entry[1] == "put"
- ]
- output_gets = [
- entry[2] for entry in queue_log if entry[0] == "output" and entry[1] == "get"
- ]
- assert [message for message in output_puts if message != "output:audit"] == [
- "output:session_started",
- "output:event",
- "output:tool_result",
- "output:round_stats",
- "output:message_delta",
- "output:round_stats",
- "output:done",
- ]
- assert output_gets == output_puts
- @pytest.mark.asyncio
- async def test_runtime_continues_with_event_summary_without_tool_call_history():
- 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",
- "system",
- "user",
- "assistant",
- "user",
- ]
- assistant_message = client.second_call_messages[3]
- assert assistant_message.content == ""
- assert "handoff_note" in client.second_call_messages[1].content
- assert not any(message.role == "tool" for message in client.second_call_messages)
- assert client.second_call_messages[4].content.startswith("EventAgent results:\n")
- assert client.second_call_messages[4].name == "event_agent"
- assert outputs[-1] == {"type": "done"}
- @pytest.mark.asyncio
- async def test_runtime_passes_event_catalog_system_message_without_chat_tools():
- 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 == []
- assert client.messages[0].role == "system"
- assert "Available events:" in client.messages[0].content
- assert "- handoff_note: Registry-owned handoff tool." in client.messages[0].content
- assert "message" not in client.messages[0].content
- assert "priority" not in client.messages[0].content
- @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)]
- business_outputs = _without_audit(outputs)
- assert _message_types(outputs) == [
- "session_started",
- "message_delta",
- "usage",
- "round_stats",
- "done",
- ]
- assert business_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,
- },
- ]
- @pytest.mark.asyncio
- async def test_runtime_session_resets_event_budget_for_each_user_turn():
- request = DebugRunRequest(
- user_message="first turn",
- 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=1),
- )
- client = TwoTurnSessionChatClient()
- runtime = DebugRuntime(client)
- queues = runtime.start_session(request)
- outputs: list[dict[str, Any]] = []
- while len([message for message in outputs if message["type"] == "turn_completed"]) < 1:
- outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
- await queues.input.put(ChatMessage(role="user", content="second turn"))
- while len([message for message in outputs if message["type"] == "turn_completed"]) < 2:
- outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
- await runtime.aclose()
- business_types = [message["type"] for message in _without_audit(outputs)]
- assert business_types.count("turn_started") == 2
- assert business_types.count("turn_completed") == 2
- assert client.calls == 4
- assert "Available events:" in client.messages_by_call[0][0].content
- assert "Available events:" in client.messages_by_call[2][0].content
|