瀏覽代碼

feat: emit structured round stats

zhenyu.hu 3 周之前
父節點
當前提交
295d72fdb9

+ 36 - 1
src/agent_lab/application/runtime.py

@@ -1,4 +1,5 @@
-from collections.abc import AsyncIterator
+import time
+from collections.abc import AsyncIterator, Callable
 from typing import Any, Protocol
 
 from agent_lab.application.contracts import AgentParams, DebugRunRequest
@@ -24,10 +25,12 @@ class DebugRuntime:
         chat_client: ChatClient,
         queues: RuntimeQueues | None = None,
         registry: ToolRegistry | None = None,
+        clock: Callable[[], float] = time.perf_counter,
     ) -> None:
         self.chat_client = chat_client
         self.queues = queues
         self.registry = registry or build_default_tool_registry()
+        self.clock = clock
 
     async def run(self, request: DebugRunRequest) -> AsyncIterator[dict[str, Any]]:
         queues = self.queues or RuntimeQueues()
@@ -42,8 +45,16 @@ class DebugRuntime:
         yield await self._emit(queues, {"type": "session_started"})
 
         event_loops = 0
+        round_index = 0
         while True:
             await self._drain_input(queues, messages)
+            round_index += 1
+            started_at = self.clock()
+            ttft_ms: int | None = None
+            prompt_tokens = 0
+            completion_tokens = 0
+            total_tokens = 0
+            cached_tokens = 0
             saw_event = False
             assistant_content: list[str] = []
             assistant_tool_calls: list[dict[str, Any]] = []
@@ -54,6 +65,8 @@ class DebugRuntime:
                 params=request.chat_agent,
             ):
                 if item.kind == "message_delta":
+                    if ttft_ms is None:
+                        ttft_ms = self._elapsed_ms(started_at)
                     assistant_content.append(item.content or "")
                     yield await self._emit(
                         queues,
@@ -62,6 +75,10 @@ class DebugRuntime:
                     continue
 
                 if item.kind == "usage" and item.usage is not None:
+                    prompt_tokens = item.usage.prompt_tokens
+                    completion_tokens = item.usage.completion_tokens
+                    total_tokens = item.usage.total_tokens
+                    cached_tokens = item.usage.cached_tokens
                     yield await self._emit(
                         queues,
                         {"type": "usage", "usage": item.usage.model_dump()},
@@ -102,6 +119,21 @@ class DebugRuntime:
                     )
                 )
 
+            yield await self._emit(
+                queues,
+                {
+                    "type": "round_stats",
+                    "round_index": round_index,
+                    "ttft_ms": ttft_ms,
+                    "elapsed_ms": self._elapsed_ms(started_at),
+                    "prompt_tokens": prompt_tokens,
+                    "completion_tokens": completion_tokens,
+                    "total_tokens": total_tokens,
+                    "cached_tokens": cached_tokens,
+                    "had_event": saw_event,
+                },
+            )
+
             if not saw_event:
                 break
 
@@ -134,3 +166,6 @@ class DebugRuntime:
     ) -> dict[str, Any]:
         await queues.output.put(payload)
         return await queues.output.get()
+
+    def _elapsed_ms(self, started_at: float) -> int:
+        return round((self.clock() - started_at) * 1000)

+ 31 - 2
src/agent_lab/presentation/static/app.js

@@ -14,6 +14,7 @@ let socket = null;
 let activeAssistant = null;
 let runStartedAt = 0;
 let firstTokenAt = 0;
+let hasBackendRoundStats = false;
 let pendingEnabledTools = null;
 
 document.querySelector("#add-system-prompt").addEventListener("click", () => {
@@ -322,6 +323,10 @@ function handleServerMessage(message) {
     updateUsage(message.usage || {});
     return;
   }
+  if (message.type === "round_stats") {
+    updateRoundStats(message);
+    return;
+  }
   if (message.type === "event") {
     appendLog("event", `${message.event.name}: ${JSON.stringify(message.event.arguments)}`);
     return;
@@ -343,7 +348,9 @@ function handleServerMessage(message) {
 function appendAssistantDelta(content) {
   if (!firstTokenAt) {
     firstTokenAt = performance.now();
-    document.querySelector("#stat-ttft").textContent = `${Math.round(firstTokenAt - runStartedAt)}ms`;
+    if (!hasBackendRoundStats) {
+      document.querySelector("#stat-ttft").textContent = `${Math.round(firstTokenAt - runStartedAt)}ms`;
+    }
   }
   if (!activeAssistant) {
     activeAssistant = document.createElement("div");
@@ -364,12 +371,33 @@ function appendLog(kind, content) {
 }
 
 function updateUsage(usage) {
+  if (hasBackendRoundStats) {
+    return;
+  }
   document.querySelector("#stat-tokens").textContent = usage.total_tokens || 0;
   document.querySelector("#stat-cached").textContent = usage.cached_tokens || 0;
 }
 
+function updateRoundStats(stats) {
+  hasBackendRoundStats = true;
+  document.querySelector("#stat-tokens").textContent = stats.total_tokens || 0;
+  document.querySelector("#stat-cached").textContent = stats.cached_tokens || 0;
+  const ttftText = formatMs(stats.ttft_ms);
+  const elapsedText = formatMs(stats.elapsed_ms);
+  document.querySelector("#stat-ttft").textContent = ttftText;
+  document.querySelector("#stat-elapsed").textContent = elapsedText;
+  appendLog(
+    "session",
+    `Round ${stats.round_index}: tokens ${stats.total_tokens || 0}, cached ${stats.cached_tokens || 0}, TTFT ${ttftText}, elapsed ${elapsedText}`,
+  );
+}
+
+function formatMs(value) {
+  return value === null || value === undefined ? "-" : `${value}ms`;
+}
+
 function updateElapsed() {
-  if (!runStartedAt) {
+  if (!runStartedAt || hasBackendRoundStats) {
     return;
   }
   document.querySelector("#stat-elapsed").textContent = `${Math.round(performance.now() - runStartedAt)}ms`;
@@ -379,6 +407,7 @@ function resetRun() {
   activeAssistant = null;
   runStartedAt = 0;
   firstTokenAt = 0;
+  hasBackendRoundStats = false;
   messagesEl.textContent = "";
   document.querySelector("#stat-tokens").textContent = "0";
   document.querySelector("#stat-cached").textContent = "0";

+ 138 - 2
tests/test_debug_runtime.py

@@ -9,7 +9,7 @@ from agent_lab.application.contracts import AgentParams, DebugRunRequest, EventA
 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
+from agent_lab.domain.messages import ChatMessage, StreamItem, TokenUsage
 
 
 def _runtime_queues_class():
@@ -111,6 +111,55 @@ class ToolCapturingChatClient:
         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()
 
@@ -143,11 +192,13 @@ async def test_runtime_routes_chat_events_through_event_agent_then_continues_cha
         "session_started",
         "event",
         "tool_result",
+        "round_stats",
         "message_delta",
+        "round_stats",
         "done",
     ]
     assert outputs[1]["event"]["name"] == "handoff_note"
-    assert outputs[3]["content"] == "final answer"
+    assert outputs[4]["content"] == "final answer"
 
 
 @pytest.mark.asyncio
@@ -174,7 +225,9 @@ async def test_runtime_uses_event_and_input_queues_for_event_agent_handoff():
         "session_started",
         "event",
         "tool_result",
+        "round_stats",
         "message_delta",
+        "round_stats",
         "done",
     ]
     assert queue_log.index(("input", "put", "user")) < queue_log.index(
@@ -215,7 +268,9 @@ async def test_runtime_yields_existing_output_order_from_output_queue():
         "session_started",
         "event",
         "tool_result",
+        "round_stats",
         "message_delta",
+        "round_stats",
         "done",
     ]
     assert [
@@ -229,8 +284,12 @@ async def test_runtime_yields_existing_output_order_from_output_queue():
         ("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"),
     ]
@@ -323,3 +382,80 @@ async def test_runtime_passes_selected_tool_schema_from_registry_to_chat_agent()
             },
         }
     ]
+
+
+@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,
+        },
+    ]

+ 9 - 0
tests/test_websocket_api.py

@@ -381,6 +381,15 @@ def test_static_prompt_workspace_uses_stable_storage_hooks():
     assert "function deletePromptSet()" in js
 
 
+def test_static_app_handles_backend_round_stats_messages():
+    js = Path("src/agent_lab/presentation/static/app.js").read_text()
+
+    assert 'message.type === "round_stats"' in js
+    assert "function updateRoundStats(" in js
+    assert "#stat-ttft" in js
+    assert "#stat-elapsed" in js
+
+
 @pytest.mark.asyncio
 async def test_openai_chat_client_omits_stream_options_when_usage_disabled():
     captured_payloads: list[dict] = []