| 12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879808182838485868788899091929394959697989910010110210310410510610710810911011111211311411511611711811912012112212312412512612712812913013113213313413513613713813914014114214314414514614714814915015115215315415515615715815916016116216316416516616716816917017117217317417517617717817918018118218318418518618718818919019119219319419519619719819920020120220320420520620720820921021121221321421521621721821922022122222322422522622722822923023123223323423523623723823924024124224324424524624724824925025125225325425525625725825926026126226326426526626726826927027127227327427527627727827928028128228328428528628728828929029129229329429529629729829930030130230330430530630730830931031131231331431531631731831932032132232332432532632732832933033133233333433533633733833934034134234334434534634734834935035135235335435535635735835936036136236336436536636736836937037137237337437537637737837938038138238338438538638738838939039139239339439539639739839940040140240340440540640740840941041141241341441541641741841942042142242342442542642742842943043143243343443543643743843944044144244344444544644744844945045145245345445545645745845946046146246346446546646746846947047147247347447547647747847948048148248348448548648748848949049149249349449549649749849950050150250350450550650750850951051151251351451551651751851952052152252352452552652752852953053153253353453553653753853954054154254354454554654754854955055155255355455555655755855956056156256356456556656756856957057157257357457557657757857958058158258358458558658758858959059159259359459559659759859960060160260360460560660760860961061161261361461561661761861962062162262362462562662762862963063163263363463563663763863964064164264364464564664764864965065165265365465565665765865966066166266366466566666766866967067167267367467567667767867968068168268368468568668768868969069169269369469569669769869970070170270370470570670770870971071171271371471571671771871972072172272372472572672772872973073173273373473573673773873974074174274374474574674774874975075175275375475575675775875976076176276376476576676776876977077177277377477577677777877978078178278378478578678778878979079179279379479579679779879980080180280380480580680780880981081181281381481581681781881982082182282382482582682782882983083183283383483583683783883984084184284384484584684784884985085185285385485585685785885986086186286386486586686786886987087187287387487587687787887988088188288388488588688788888989089189289389489589689789889990090190290390490590690790890991091191291391491591691791891992092192292392492592692792892993093193293393493593693793893994094194294394494594694794894995095195295395495595695795895996096196296396496596696796896997097197297397497597697797897998098198298398498598698798898999099199299399499599699799899910001001100210031004100510061007100810091010101110121013101410151016101710181019102010211022102310241025102610271028102910301031103210331034103510361037103810391040104110421043104410451046104710481049105010511052105310541055105610571058105910601061106210631064106510661067106810691070107110721073107410751076107710781079108010811082108310841085108610871088108910901091109210931094109510961097109810991100110111021103110411051106110711081109111011111112111311141115111611171118111911201121112211231124112511261127112811291130113111321133113411351136113711381139114011411142114311441145114611471148114911501151115211531154115511561157115811591160116111621163116411651166116711681169117011711172117311741175117611771178117911801181118211831184118511861187118811891190119111921193119411951196119711981199120012011202120312041205120612071208120912101211121212131214121512161217121812191220122112221223122412251226122712281229123012311232123312341235123612371238123912401241124212431244124512461247124812491250125112521253125412551256125712581259126012611262126312641265126612671268126912701271127212731274127512761277127812791280128112821283128412851286128712881289129012911292129312941295129612971298129913001301130213031304130513061307130813091310131113121313131413151316131713181319132013211322132313241325132613271328132913301331133213331334133513361337133813391340134113421343134413451346134713481349135013511352135313541355135613571358135913601361136213631364136513661367136813691370137113721373137413751376137713781379138013811382138313841385138613871388138913901391139213931394139513961397139813991400140114021403140414051406140714081409141014111412141314141415141614171418141914201421142214231424142514261427142814291430143114321433143414351436143714381439144014411442144314441445144614471448144914501451145214531454145514561457145814591460146114621463146414651466146714681469147014711472147314741475147614771478147914801481148214831484148514861487148814891490149114921493149414951496149714981499150015011502150315041505150615071508150915101511151215131514151515161517151815191520 |
- 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
- from agent_lab.infrastructure.sqlite_store import SQLiteSessionStore
- 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 StreamItem.raw_response_chunk(
- {
- "choices": [
- {
- "delta": {
- "tool_calls": [
- {
- "index": 0,
- "function": {"name": tools[0]["function"]["name"]},
- }
- ]
- },
- "finish_reason": None,
- }
- ]
- }
- )
- yield _event_tool_call_from_tools(tools, messages)
- return
- self.calls += 1
- if self.calls == 1:
- yield StreamItem.raw_response_chunk(
- {
- "choices": [
- {
- "delta": {"content": "<agent_events>handoff_note</agent_events>"},
- "finish_reason": None,
- }
- ]
- }
- )
- 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 IncrementingClock:
- def __init__(self, current: float = 100.0, step: float = 0.01) -> None:
- self.current = current
- self.step = step
- def __call__(self) -> float:
- self.current += self.step
- return self.current
- 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 StreamItem.raw_response_chunk(
- {
- "choices": [
- {
- "delta": {
- "tool_calls": [
- {
- "index": 0,
- "function": {"name": tools[0]["function"]["name"]},
- }
- ]
- },
- "finish_reason": None,
- }
- ]
- }
- )
- yield _event_tool_call_from_tools(tools, messages)
- return
- self.calls += 1
- if self.calls == 1:
- yield StreamItem.raw_response_chunk(
- {
- "choices": [
- {
- "delta": {"content": "<agent_events>handoff_note</agent_events>"},
- "finish_reason": None,
- }
- ]
- }
- )
- 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 StreamItem.raw_response_chunk(
- {
- "choices": [
- {
- "delta": {
- "tool_calls": [
- {
- "index": 0,
- "function": {
- "name": tools[0]["function"]["name"],
- },
- }
- ]
- },
- "finish_reason": None,
- }
- ]
- }
- )
- yield _event_tool_call_from_tools(tools, messages)
- return
- self.calls += 1
- if self.calls == 1:
- yield StreamItem.raw_response_chunk(
- {
- "choices": [
- {
- "delta": {
- "content": "<agent_events>handoff_note</agent_events>",
- },
- "finish_reason": None,
- }
- ]
- }
- )
- 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.raw_response_chunk(
- {
- "choices": [
- {
- "delta": {"content": "final answer"},
- "finish_reason": None,
- }
- ]
- }
- )
- yield StreamItem.message_delta("final answer")
- yield StreamItem.usage_item(
- TokenUsage(prompt_tokens=4, completion_tokens=6, total_tokens=10)
- )
- class ContextBoundaryChatClient:
- def __init__(self) -> None:
- self.calls = 0
- self.event_agent_messages: list[list[ChatMessage]] = []
- async def stream_chat(
- self,
- messages: list[ChatMessage],
- tools: list[dict],
- params: AgentParams,
- ) -> AsyncIterator[StreamItem]:
- if tools:
- self.event_agent_messages.append(list(messages))
- yield _event_tool_call_from_tools(tools, messages)
- return
- self.calls += 1
- if self.calls == 1:
- yield StreamItem.message_delta("I will check.")
- yield StreamItem.event(
- ToolCallEvent(
- id="call_1",
- name="handoff_note",
- arguments={},
- raw_arguments="{}",
- )
- )
- return
- yield StreamItem.message_delta("final answer")
- 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_agent_request",
- "chat_event_detected",
- "chat_agent_response",
- "event_agent_request",
- "event_agent_response",
- "event_agent_completed",
- "chat_round_finished",
- "chat_round_started",
- "chat_agent_request",
- "chat_message_stream_started",
- "chat_message_stream_finished",
- "chat_agent_response",
- "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_audits_chat_message_stream_boundaries_in_output_order():
- 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())
- outputs = [message async for message in runtime.run(request)]
- ordered_labels = [
- message["event"] if message["type"] == "audit" else message["type"]
- for message in outputs
- ]
- assert ordered_labels.index("chat_message_stream_started") < ordered_labels.index(
- "message_delta"
- )
- assert ordered_labels.index("message_delta") < ordered_labels.index(
- "chat_message_stream_finished"
- )
- assert ordered_labels.index("chat_message_stream_finished") < ordered_labels.index(
- "chat_agent_response"
- )
- stream_started = next(
- message
- for message in outputs
- if message.get("event") == "chat_message_stream_started"
- )
- stream_finished = next(
- message
- for message in outputs
- if message.get("event") == "chat_message_stream_finished"
- )
- assert stream_started["details"]["agent"] == "chat_agent"
- assert stream_started["details"]["round_index"] == 1
- assert stream_finished["details"]["delta_count"] == 1
- assert stream_finished["details"]["content_length"] == len("hello")
- @pytest.mark.asyncio
- async def test_runtime_audit_events_include_turn_relative_elapsed_time():
- 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(), clock=IncrementingClock())
- queues = runtime.start_session(request)
- outputs: list[dict[str, Any]] = []
- while not any(message["type"] == "turn_completed" for message in outputs):
- outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
- await runtime.aclose()
- turn_audits = [
- message
- for message in outputs
- if message["type"] == "audit"
- ]
- elapsed_values = [
- message["details"].get("turn_elapsed_ms") for message in turn_audits
- ]
- assert turn_audits
- assert all(isinstance(value, int) for value in elapsed_values)
- assert all(value >= 0 for value in elapsed_values)
- assert elapsed_values == sorted(elapsed_values)
- assert next(
- message
- for message in turn_audits
- if message["event"] == "event_agent_request"
- )["details"]["turn_elapsed_ms"] >= 0
- @pytest.mark.asyncio
- async def test_runtime_audit_includes_model_params_prompts_results_and_usage():
- request = DebugRunRequest(
- user_message="debug this",
- system_prompts=["You are a debugger."],
- pre_messages=[],
- chat_agent=AgentParams(model="chat-model", temperature=0.1, max_tokens=200),
- event_agent=EventAgentParams(
- model="event-model",
- temperature=0.4,
- max_tokens=80,
- enabled_tools=["handoff_note"],
- max_event_loops=1,
- ),
- )
- runtime = DebugRuntime(EventRoundStatsChatClient())
- outputs = [message async for message in runtime.run(request)]
- audits = [message for message in outputs if message["type"] == "audit"]
- chat_request = next(
- message for message in audits if message["event"] == "chat_agent_request"
- )
- assert chat_request["details"]["agent"] == "chat_agent"
- assert chat_request["details"]["params"]["model"] == "chat-model"
- assert chat_request["details"]["params"]["temperature"] == 0.1
- assert chat_request["details"]["params"]["max_tokens"] == 200
- assert chat_request["details"]["tools"] == []
- assert chat_request["details"]["messages"][0] == {
- "role": "system",
- "content": "You are a debugger.",
- "name": None,
- "tool_call_id": None,
- }
- chat_response = next(
- message for message in audits if message["event"] == "chat_agent_response"
- )
- assert chat_response["details"]["event_names"] == ["handoff_note"]
- assert chat_response["details"]["raw_chunks"] == [
- {
- "choices": [
- {
- "delta": {"content": "<agent_events>handoff_note</agent_events>"},
- "finish_reason": None,
- }
- ]
- }
- ]
- assert chat_response["details"]["usage"] == {
- "prompt_tokens": 3,
- "completion_tokens": 0,
- "total_tokens": 3,
- "cached_tokens": 0,
- }
- event_request = next(
- message for message in audits if message["event"] == "event_agent_request"
- )
- assert event_request["details"]["agent"] == "event_agent"
- assert event_request["details"]["params"]["model"] == "event-model"
- assert event_request["details"]["params"]["temperature"] == 0.4
- assert event_request["details"]["events"][0]["name"] == "handoff_note"
- assert event_request["details"]["tools"][0]["function"]["name"] == "handoff_note"
- assert "Authorization" not in str(event_request["details"])
- event_response = next(
- message for message in audits if message["event"] == "event_agent_response"
- )
- assert event_response["details"]["replies"][0]["role"] == "tool"
- assert '"tool": "handoff_note"' in event_response["details"]["replies"][0]["content"]
- assert event_response["details"]["raw_model_chunks"][0]["event_name"] == "handoff_note"
- assert event_response["details"]["raw_model_chunks"][0]["chunks"][0]["choices"][0]["delta"] == {
- "tool_calls": [
- {
- "index": 0,
- "function": {"name": "handoff_note"},
- }
- ]
- }
- @pytest.mark.asyncio
- async def test_runtime_round_started_separates_available_events_from_round_budget():
- request = DebugRunRequest(
- user_message="debug this",
- system_prompts=[],
- pre_messages=[],
- chat_agent=AgentParams(model="chat-model", temperature=0.1, max_tokens=200),
- event_agent=EventAgentParams(
- enabled_tools=["handoff_note"],
- max_event_loops=1,
- ),
- )
- runtime = DebugRuntime(EventRoundStatsChatClient())
- outputs = [message async for message in runtime.run(request)]
- round_starts = [
- message
- for message in outputs
- if message.get("event") == "chat_round_started"
- ]
- assert round_starts[0]["details"]["events_enabled"] == ["handoff_note"]
- assert round_starts[0]["details"]["configured_events"] == ["handoff_note"]
- assert round_starts[0]["details"]["event_generation_enabled"] is True
- assert round_starts[0]["details"]["event_prompt_events"] == ["handoff_note"]
- assert round_starts[1]["details"]["events_enabled"] == []
- assert round_starts[1]["details"]["configured_events"] == ["handoff_note"]
- assert round_starts[1]["details"]["event_generation_enabled"] is False
- assert round_starts[1]["details"]["event_prompt_events"] == []
- @pytest.mark.asyncio
- async def test_runtime_event_agent_history_excludes_chat_agent_system_context():
- request = DebugRunRequest(
- user_message="debug this",
- system_prompts=["ChatAgent root prompt."],
- pre_messages=[
- ChatMessage(role="system", content="ChatAgent pre system."),
- ChatMessage(role="user", content="earlier user"),
- ChatMessage(role="assistant", content="earlier assistant"),
- ],
- chat_agent=AgentParams(model="chat-model", temperature=0.1, max_tokens=200),
- event_agent=EventAgentParams(
- model="event-model",
- enabled_tools=["handoff_note"],
- max_event_loops=1,
- ),
- )
- client = ContextBoundaryChatClient()
- runtime = DebugRuntime(client)
- outputs = [message async for message in runtime.run(request)]
- event_request = next(
- message for message in outputs if message.get("event") == "event_agent_request"
- )
- assert [
- (message["role"], message["content"])
- for message in event_request["details"]["history"]
- ] == [
- ("user", "earlier user"),
- ("assistant", "earlier assistant"),
- ("user", "debug this"),
- ("assistant", "I will check."),
- ]
- assert not any(
- message["role"] == "system"
- for message in event_request["details"]["history"]
- )
- assert client.event_agent_messages
- assert [
- (message.role, message.content)
- for message in client.event_agent_messages[0]
- if message.content in {"ChatAgent root prompt.", "ChatAgent pre system."}
- ] == []
- @pytest.mark.asyncio
- async def test_runtime_session_event_agent_history_excludes_internal_replies():
- request = DebugRunRequest(
- user_message="debug this",
- system_prompts=["ChatAgent root prompt."],
- pre_messages=[ChatMessage(role="assistant", content="prior answer")],
- chat_agent=AgentParams(model="chat-model", temperature=0.1, max_tokens=200),
- event_agent=EventAgentParams(
- model="event-model",
- enabled_tools=["handoff_note"],
- max_event_loops=1,
- ),
- )
- runtime = DebugRuntime(ContextBoundaryChatClient())
- queues = runtime.start_session(request)
- outputs: list[dict[str, Any]] = []
- while not any(message["type"] == "turn_completed" for message in outputs):
- outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
- await runtime.aclose()
- event_request = next(
- message for message in outputs if message.get("event") == "event_agent_request"
- )
- assert [
- (message["role"], message["content"], message["name"])
- for message in event_request["details"]["history"]
- ] == [
- ("assistant", "prior answer", None),
- ("user", "debug this", None),
- ("assistant", "I will check.", None),
- ]
- assert not any(
- message["name"] == "event_agent"
- for message in event_request["details"]["history"]
- )
- @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(
- [
- 0.9,
- 0.91,
- 1.0,
- 1.01,
- 1.02,
- 1.123,
- 1.2,
- 1.25,
- 1.3,
- 1.456,
- 1.7,
- 1.8,
- ]
- )
- 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(
- [
- 1.9,
- 1.91,
- 2.0,
- 2.01,
- 2.02,
- 2.03,
- 2.04,
- 2.05,
- 2.06,
- 2.07,
- 2.25,
- 2.26,
- 3.0,
- 3.01,
- 3.02,
- 3.05,
- 3.1,
- 3.15,
- 3.18,
- 3.2,
- 3.23,
- 3.24,
- 3.25,
- 3.26,
- ]
- )
- 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
- @pytest.mark.asyncio
- async def test_runtime_persists_session_turn_messages_audit_and_usage(tmp_path):
- store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
- request = DebugRunRequest(
- session_id="session-1",
- user_message="persist this",
- system_prompts=["You are a debugger."],
- pre_messages=[],
- chat_agent=AgentParams(model="chat-model", temperature=0.1, max_tokens=200),
- event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
- )
- runtime = DebugRuntime(RoundStatsChatClient(), session_store=store)
- queues = runtime.start_session(request)
- outputs: list[dict[str, Any]] = []
- while not any(message["type"] == "turn_completed" for message in outputs):
- outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
- await runtime.aclose()
- assert _without_audit(outputs)[0] == {
- "type": "session_started",
- "session_id": "session-1",
- }
- assert store.get_session("session-1")["config"]["chat_agent"]["model"] == "chat-model"
- assert [
- (message["turn_index"], message["role"], message["content"])
- for message in store.list_messages("session-1")
- ] == [
- (1, "user", "persist this"),
- (1, "assistant", "hello"),
- ]
- audit_events = [
- audit["event"]
- for audit in store.list_audit_logs("session-1")
- ]
- assert "session_started" in audit_events
- assert "chat_agent_request" in audit_events
- assert "turn_completed" in audit_events
- usage = store.usage_summary("session-1")
- assert usage["calls"][0]["total_tokens"] == 30
- assert usage["turns"][0]["turn_index"] == 1
- assert usage["session"]["total_tokens"] == 30
- @pytest.mark.asyncio
- async def test_runtime_continues_persisted_turn_indexes_for_existing_session(tmp_path):
- store = SQLiteSessionStore(tmp_path / "agent_lab.sqlite3")
- session_id = store.create_session(title="existing", config={})
- store.start_turn(session_id, turn_index=1, user_message="old")
- request = DebugRunRequest(
- session_id=session_id,
- user_message="new",
- system_prompts=[],
- pre_messages=[],
- chat_agent=AgentParams(model="chat-model", temperature=0.1, max_tokens=200),
- event_agent=EventAgentParams(enabled_tools=[], max_event_loops=1),
- )
- runtime = DebugRuntime(RoundStatsChatClient(), session_store=store)
- queues = runtime.start_session(request)
- outputs: list[dict[str, Any]] = []
- while not any(message["type"] == "turn_completed" for message in outputs):
- outputs.append(await asyncio.wait_for(queues.output.get(), timeout=1))
- await runtime.aclose()
- assert [
- (message["turn_index"], message["content"])
- for message in store.list_messages(session_id)
- ] == [
- (2, "new"),
- (2, "hello"),
- ]
|