|
@@ -1,6 +1,7 @@
|
|
|
import asyncio
|
|
import asyncio
|
|
|
import importlib
|
|
import importlib
|
|
|
import json
|
|
import json
|
|
|
|
|
+import logging
|
|
|
from collections.abc import AsyncIterator
|
|
from collections.abc import AsyncIterator
|
|
|
from typing import Any
|
|
from typing import Any
|
|
|
|
|
|
|
@@ -23,6 +24,23 @@ async def _collect_outputs(stream: AsyncIterator[dict[str, Any]]) -> list[dict[s
|
|
|
return [message async for message in stream]
|
|
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)]
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+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
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
class RecordingQueue(asyncio.Queue):
|
|
class RecordingQueue(asyncio.Queue):
|
|
|
def __init__(self, name: str, log: list[tuple[str, str, str]]) -> None:
|
|
def __init__(self, name: str, log: list[tuple[str, str, str]]) -> None:
|
|
|
super().__init__()
|
|
super().__init__()
|
|
@@ -258,9 +276,10 @@ async def test_runtime_routes_chat_events_through_event_agent_then_continues_cha
|
|
|
runtime = DebugRuntime(client)
|
|
runtime = DebugRuntime(client)
|
|
|
|
|
|
|
|
outputs = [message async for message in runtime.run(request)]
|
|
outputs = [message async for message in runtime.run(request)]
|
|
|
|
|
+ business_outputs = _without_audit(outputs)
|
|
|
|
|
|
|
|
assert client.calls == 2
|
|
assert client.calls == 2
|
|
|
- assert [message["type"] for message in outputs] == [
|
|
|
|
|
|
|
+ assert _message_types(outputs) == [
|
|
|
"session_started",
|
|
"session_started",
|
|
|
"event",
|
|
"event",
|
|
|
"tool_result",
|
|
"tool_result",
|
|
@@ -269,8 +288,41 @@ async def test_runtime_routes_chat_events_through_event_agent_then_continues_cha
|
|
|
"round_stats",
|
|
"round_stats",
|
|
|
"done",
|
|
"done",
|
|
|
]
|
|
]
|
|
|
- assert outputs[1]["event"]["name"] == "handoff_note"
|
|
|
|
|
- assert outputs[4]["content"] == "final answer"
|
|
|
|
|
|
|
+ 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
|
|
@pytest.mark.asyncio
|
|
@@ -328,7 +380,7 @@ async def test_runtime_batches_round_events_before_continuing_chat_agent():
|
|
|
|
|
|
|
|
outputs = [message async for message in runtime.run(request)]
|
|
outputs = [message async for message in runtime.run(request)]
|
|
|
|
|
|
|
|
- assert [message["type"] for message in outputs] == [
|
|
|
|
|
|
|
+ assert _message_types(outputs) == [
|
|
|
"session_started",
|
|
"session_started",
|
|
|
"message_delta",
|
|
"message_delta",
|
|
|
"event",
|
|
"event",
|
|
@@ -390,7 +442,7 @@ async def test_runtime_start_returns_queues_for_downstream_output_consumer():
|
|
|
outputs.append(message)
|
|
outputs.append(message)
|
|
|
if message["type"] == "done":
|
|
if message["type"] == "done":
|
|
|
break
|
|
break
|
|
|
- assert [message["type"] for message in outputs] == [
|
|
|
|
|
|
|
+ assert _message_types(outputs) == [
|
|
|
"session_started",
|
|
"session_started",
|
|
|
"message_delta",
|
|
"message_delta",
|
|
|
"usage",
|
|
"usage",
|
|
@@ -412,11 +464,12 @@ async def test_runtime_finalizes_chat_after_reaching_event_loop_limit():
|
|
|
runtime = DebugRuntime(client)
|
|
runtime = DebugRuntime(client)
|
|
|
|
|
|
|
|
outputs = [message async for message in runtime.run(request)]
|
|
outputs = [message async for message in runtime.run(request)]
|
|
|
|
|
+ business_outputs = _without_audit(outputs)
|
|
|
|
|
|
|
|
assert client.calls == 2
|
|
assert client.calls == 2
|
|
|
assert client.tools_by_call[0][0]["function"]["name"] == "handoff_note"
|
|
assert client.tools_by_call[0][0]["function"]["name"] == "handoff_note"
|
|
|
assert client.tools_by_call[1] == []
|
|
assert client.tools_by_call[1] == []
|
|
|
- assert [message["type"] for message in outputs] == [
|
|
|
|
|
|
|
+ assert _message_types(outputs) == [
|
|
|
"session_started",
|
|
"session_started",
|
|
|
"event",
|
|
"event",
|
|
|
"tool_result",
|
|
"tool_result",
|
|
@@ -425,7 +478,7 @@ async def test_runtime_finalizes_chat_after_reaching_event_loop_limit():
|
|
|
"round_stats",
|
|
"round_stats",
|
|
|
"done",
|
|
"done",
|
|
|
]
|
|
]
|
|
|
- assert outputs[4]["content"] == "final after event limit"
|
|
|
|
|
|
|
+ assert business_outputs[4]["content"] == "final after event limit"
|
|
|
|
|
|
|
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.asyncio
|
|
@@ -443,8 +496,8 @@ async def test_runtime_buffers_upstream_user_input_until_after_matching_tool_rep
|
|
|
runtime = DebugRuntime(client, queues=queues)
|
|
runtime = DebugRuntime(client, queues=queues)
|
|
|
stream = runtime.run(request)
|
|
stream = runtime.run(request)
|
|
|
|
|
|
|
|
- assert await anext(stream) == {"type": "session_started"}
|
|
|
|
|
- event_message = await anext(stream)
|
|
|
|
|
|
|
+ assert await _next_non_audit(stream) == {"type": "session_started"}
|
|
|
|
|
+ event_message = await _next_non_audit(stream)
|
|
|
assert event_message["type"] == "event"
|
|
assert event_message["type"] == "event"
|
|
|
await queues.input.put(ChatMessage(role="user", content="follow-up while tool runs"))
|
|
await queues.input.put(ChatMessage(role="user", content="follow-up while tool runs"))
|
|
|
remaining = [message async for message in stream]
|
|
remaining = [message async for message in stream]
|
|
@@ -488,8 +541,9 @@ async def test_runtime_continues_when_event_agent_tool_handler_raises():
|
|
|
_collect_outputs(runtime.run(request)),
|
|
_collect_outputs(runtime.run(request)),
|
|
|
timeout=1,
|
|
timeout=1,
|
|
|
)
|
|
)
|
|
|
|
|
+ business_outputs = _without_audit(outputs)
|
|
|
|
|
|
|
|
- assert [message["type"] for message in outputs] == [
|
|
|
|
|
|
|
+ assert _message_types(outputs) == [
|
|
|
"session_started",
|
|
"session_started",
|
|
|
"event",
|
|
"event",
|
|
|
"tool_result",
|
|
"tool_result",
|
|
@@ -498,7 +552,7 @@ async def test_runtime_continues_when_event_agent_tool_handler_raises():
|
|
|
"round_stats",
|
|
"round_stats",
|
|
|
"done",
|
|
"done",
|
|
|
]
|
|
]
|
|
|
- assert json.loads(outputs[2]["message"]["content"]) == {
|
|
|
|
|
|
|
+ assert json.loads(business_outputs[2]["message"]["content"]) == {
|
|
|
"tool": "handoff_note",
|
|
"tool": "handoff_note",
|
|
|
"error": "tool handler failed: boom",
|
|
"error": "tool handler failed: boom",
|
|
|
}
|
|
}
|
|
@@ -524,7 +578,7 @@ async def test_runtime_uses_event_and_input_queues_for_event_agent_handoff():
|
|
|
|
|
|
|
|
outputs = [message async for message in runtime.run(request)]
|
|
outputs = [message async for message in runtime.run(request)]
|
|
|
|
|
|
|
|
- assert [message["type"] for message in outputs] == [
|
|
|
|
|
|
|
+ assert _message_types(outputs) == [
|
|
|
"session_started",
|
|
"session_started",
|
|
|
"event",
|
|
"event",
|
|
|
"tool_result",
|
|
"tool_result",
|
|
@@ -567,7 +621,7 @@ async def test_runtime_run_consumes_output_queue_in_stream_order():
|
|
|
|
|
|
|
|
outputs = [message async for message in runtime.run(request)]
|
|
outputs = [message async for message in runtime.run(request)]
|
|
|
|
|
|
|
|
- assert [message["type"] for message in outputs] == [
|
|
|
|
|
|
|
+ assert _message_types(outputs) == [
|
|
|
"session_started",
|
|
"session_started",
|
|
|
"event",
|
|
"event",
|
|
|
"tool_result",
|
|
"tool_result",
|
|
@@ -582,7 +636,7 @@ async def test_runtime_run_consumes_output_queue_in_stream_order():
|
|
|
output_gets = [
|
|
output_gets = [
|
|
|
entry[2] for entry in queue_log if entry[0] == "output" and entry[1] == "get"
|
|
entry[2] for entry in queue_log if entry[0] == "output" and entry[1] == "get"
|
|
|
]
|
|
]
|
|
|
- assert output_puts == [
|
|
|
|
|
|
|
+ assert [message for message in output_puts if message != "output:audit"] == [
|
|
|
"output:session_started",
|
|
"output:session_started",
|
|
|
"output:event",
|
|
"output:event",
|
|
|
"output:tool_result",
|
|
"output:tool_result",
|
|
@@ -693,15 +747,16 @@ async def test_runtime_emits_round_stats_with_clock_and_usage_after_model_turn()
|
|
|
runtime = DebugRuntime(RoundStatsChatClient(), clock=lambda: next(ticks))
|
|
runtime = DebugRuntime(RoundStatsChatClient(), clock=lambda: next(ticks))
|
|
|
|
|
|
|
|
outputs = [message async for message in runtime.run(request)]
|
|
outputs = [message async for message in runtime.run(request)]
|
|
|
|
|
+ business_outputs = _without_audit(outputs)
|
|
|
|
|
|
|
|
- assert [message["type"] for message in outputs] == [
|
|
|
|
|
|
|
+ assert _message_types(outputs) == [
|
|
|
"session_started",
|
|
"session_started",
|
|
|
"message_delta",
|
|
"message_delta",
|
|
|
"usage",
|
|
"usage",
|
|
|
"round_stats",
|
|
"round_stats",
|
|
|
"done",
|
|
"done",
|
|
|
]
|
|
]
|
|
|
- assert outputs[3] == {
|
|
|
|
|
|
|
+ assert business_outputs[3] == {
|
|
|
"type": "round_stats",
|
|
"type": "round_stats",
|
|
|
"round_index": 1,
|
|
"round_index": 1,
|
|
|
"ttft_ms": 123,
|
|
"ttft_ms": 123,
|