Sfoglia il codice sorgente

feat: support persistent websocket turns

zhenyu.hu 3 settimane fa
parent
commit
9f11708991

+ 30 - 0
docs/plans/todo-28-session-turns.md

@@ -0,0 +1,30 @@
+# Todo 28: Session Turns
+
+## Status
+
+done
+
+## Goal
+
+Separate session and turn behavior. A WebSocket session is the ongoing conversation. A turn starts with one user message and finishes after ChatAgent has produced the visible reply, including any EventAgent handoff and final ChatAgent response.
+
+## Root Cause
+
+The UI currently closes and recreates the WebSocket for every submit. The backend runtime also treats `done` as the end of the whole run after one user message. `round_index` is a ChatAgent call counter, so the post-event ChatAgent call appears as a second round with `events_enabled=[]`, which is expected for a one-event-loop turn but confusing without turn boundaries.
+
+## Scope
+
+- Add a persistent runtime session entrypoint for WebSocket use.
+- Emit `turn_started` and `turn_completed` events.
+- Reset EventAgent event-loop budget for each new user turn.
+- Keep post-event final ChatAgent calls inside the same turn.
+- Keep existing one-shot `runtime.run()` behavior for tests and smoke helpers.
+- Reuse an open WebSocket from the frontend instead of closing it on every submit.
+
+## Verification
+
+- Runtime test proves a second user turn gets enabled events again.
+- WebSocket test proves one connection can process two user messages.
+- Static UI test proves the browser sends `user_message` on an open socket and handles `turn_completed`.
+- `uv run pytest tests/test_debug_runtime.py tests/test_websocket_api.py -q` passes: 45 tests.
+- `uv run pytest` passes: 56 tests.

+ 1 - 0
docs/plans/todos.md

@@ -54,3 +54,4 @@
 | 25 | done | `docs/plans/todo-25-prompt-list-message-types.md` | Make the ChatAgent prompt list visually compact and show the message type for each prompt item. | `uv run pytest` passes (`51 passed`, one existing Starlette deprecation warning); local service returned `/health` and versioned JS asset. |
 | 26 | done | `docs/plans/todo-26-workspace-snapshot-save.md` | Move save/load/delete out of ChatAgent config and save one workspace snapshot covering prompts, Agent config, and selected tools. | `uv run pytest tests/test_websocket_api.py -q` passes (`27 passed`, one existing Starlette deprecation warning). |
 | 27 | done | `docs/plans/todo-27-event-agent-tool-arguments.md` | Make EventAgent tool argument generation robust for compatible providers that do not emit the expected tool-call finish reason or ignore tool calls. | `uv run pytest` passes (`53 passed`, one existing Starlette deprecation warning). |
+| 28 | done | `docs/plans/todo-28-session-turns.md` | Split WebSocket session lifetime from user turns so one conversation can run multiple user messages while each turn has its own event budget. | `uv run pytest` passes (`56 passed`, one existing Starlette deprecation warning). |

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

@@ -45,6 +45,12 @@ class DebugRuntime:
         self._tasks.append(task)
         return queues
 
+    def start_session(self, request: DebugRunRequest) -> RuntimeQueues:
+        queues = self.queues or RuntimeQueues()
+        task = asyncio.create_task(self._produce_session(request, queues))
+        self._tasks.append(task)
+        return queues
+
     async def run(self, request: DebugRunRequest) -> AsyncIterator[dict[str, Any]]:
         queues = self.start(request)
         async for message in self.output_messages(queues):
@@ -90,6 +96,29 @@ class DebugRuntime:
             await queues.events.put(None)
             await asyncio.gather(event_worker, return_exceptions=True)
 
+    async def _produce_session(
+        self,
+        request: DebugRunRequest,
+        queues: RuntimeQueues,
+    ) -> None:
+        event_agent = EventAgent(
+            request.event_agent.enabled_tools,
+            registry=self.registry,
+            chat_client=self.chat_client,
+            params=request.event_agent,
+        )
+        event_worker = asyncio.create_task(self._consume_events(queues, event_agent))
+        try:
+            await self._run_chat_session(request, queues)
+        except asyncio.CancelledError:
+            raise
+        except Exception as exc:
+            logger.exception("runtime session failed")
+            await queues.output.put({"type": "error", "message": str(exc)})
+        finally:
+            await queues.events.put(None)
+            await asyncio.gather(event_worker, return_exceptions=True)
+
     async def _run_chat_agent(
         self,
         request: DebugRunRequest,
@@ -231,6 +260,196 @@ class DebugRuntime:
         await self._audit(queues, "session_finished", round_count=round_index)
         await queues.output.put({"type": "done"})
 
+    async def _run_chat_session(
+        self,
+        request: DebugRunRequest,
+        queues: RuntimeQueues,
+    ) -> None:
+        messages = self._build_initial_messages(request)
+        await queues.output.put({"type": "session_started"})
+        await self._audit(queues, "session_started")
+        await queues.input.put(ChatMessage(role="user", content=request.user_message))
+
+        turn_index = 0
+        while True:
+            user_message = await queues.input.get()
+            if user_message.role != "user" or user_message.name == "event_agent":
+                messages.append(user_message)
+                continue
+            turn_index += 1
+            await self._run_chat_turn(
+                request,
+                queues,
+                messages,
+                user_message,
+                turn_index,
+            )
+
+    async def _run_chat_turn(
+        self,
+        request: DebugRunRequest,
+        queues: RuntimeQueues,
+        messages: list[ChatMessage],
+        user_message: ChatMessage,
+        turn_index: int,
+    ) -> None:
+        messages.append(user_message)
+        await queues.output.put({"type": "turn_started", "turn_index": turn_index})
+        await self._audit(queues, "turn_started", turn_index=turn_index)
+
+        event_loops = 0
+        round_index = 0
+        deferred_user_messages: list[ChatMessage] = []
+        while True:
+            if round_index:
+                await self._drain_input(
+                    queues,
+                    messages,
+                    deferred_user_messages=deferred_user_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] = []
+            events: list[ToolCallEvent] = []
+            enabled_events = (
+                request.event_agent.enabled_tools
+                if event_loops < request.event_agent.max_event_loops
+                else []
+            )
+            chat_messages = self._chat_messages_for_round(messages, enabled_events)
+            await self._audit(
+                queues,
+                "chat_round_started",
+                turn_index=turn_index,
+                round_index=round_index,
+                events_enabled=enabled_events,
+            )
+
+            async for item in self.chat_client.stream_chat(
+                messages=chat_messages,
+                tools=[],
+                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 "")
+                    await queues.output.put(
+                        {"type": "message_delta", "content": item.content}
+                    )
+                    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
+                    await queues.output.put(
+                        {"type": "usage", "usage": item.usage.model_dump()}
+                    )
+                    continue
+
+                if item.kind == "event" and item.event is not None:
+                    saw_event = True
+                    event = self._event_name_only(item.event)
+                    events.append(event)
+                    await queues.output.put(
+                        {"type": "event", "event": event.model_dump()}
+                    )
+                    await self._audit(
+                        queues,
+                        "chat_event_detected",
+                        turn_index=turn_index,
+                        round_index=round_index,
+                        event_id=event.id,
+                        event_name=event.name,
+                    )
+
+            if assistant_content or events:
+                messages.append(
+                    ChatMessage(
+                        role="assistant",
+                        content="".join(assistant_content),
+                    )
+                )
+
+            if events:
+                await queues.events.put(
+                    EventAgentRequest(
+                        events=events,
+                        history=list(messages),
+                        system_prompt=request.event_agent.system_prompt,
+                        extra_body=request.event_agent.extra_body,
+                    )
+                )
+                tool_replies = await self._wait_for_tool_replies(
+                    queues,
+                    [event.id for event in events],
+                )
+                for reply in tool_replies:
+                    await queues.output.put(
+                        {"type": "tool_result", "message": reply.model_dump()}
+                    )
+                await self._audit(
+                    queues,
+                    "event_agent_completed",
+                    turn_index=turn_index,
+                    round_index=round_index,
+                    event_names=[event.name for event in events],
+                    result_count=len(tool_replies),
+                    result_summary="\n".join(reply.content for reply in tool_replies),
+                )
+
+            elapsed_ms = self._elapsed_ms(started_at)
+            await queues.output.put(
+                {
+                    "type": "round_stats",
+                    "round_index": round_index,
+                    "ttft_ms": ttft_ms,
+                    "elapsed_ms": elapsed_ms,
+                    "prompt_tokens": prompt_tokens,
+                    "completion_tokens": completion_tokens,
+                    "total_tokens": total_tokens,
+                    "cached_tokens": cached_tokens,
+                    "had_event": saw_event,
+                },
+            )
+            await self._audit(
+                queues,
+                "chat_round_finished",
+                turn_index=turn_index,
+                round_index=round_index,
+                had_event=saw_event,
+                elapsed_ms=elapsed_ms,
+            )
+
+            if not saw_event:
+                break
+
+            event_loops += 1
+
+        for deferred_message in deferred_user_messages:
+            await queues.input.put(deferred_message)
+        await self._audit(
+            queues,
+            "turn_completed",
+            turn_index=turn_index,
+            round_count=round_index,
+        )
+        await queues.output.put(
+            {
+                "type": "turn_completed",
+                "turn_index": turn_index,
+                "round_count": round_index,
+            }
+        )
+
     async def _consume_events(
         self,
         queues: RuntimeQueues,
@@ -320,9 +539,18 @@ class DebugRuntime:
         self,
         queues: RuntimeQueues,
         messages: list[ChatMessage],
+        deferred_user_messages: list[ChatMessage] | None = None,
     ) -> None:
         while not queues.input.empty():
-            messages.append(await queues.input.get())
+            message = await queues.input.get()
+            if (
+                deferred_user_messages is not None
+                and message.role == "user"
+                and message.name != "event_agent"
+            ):
+                deferred_user_messages.append(message)
+                continue
+            messages.append(message)
 
     def _elapsed_ms(self, started_at: float) -> int:
         return round((self.clock() - started_at) * 1000)

+ 46 - 6
src/agent_lab/presentation/static/app.js

@@ -611,18 +611,21 @@ function runDebugSession() {
     return;
   }
 
+  if (isSocketOpen()) {
+    beginTurn(submittedMessage, "Waiting for first token");
+    socket.send(JSON.stringify({ type: "user_message", content: submittedMessage }));
+    userMessage.value = "";
+    return;
+  }
+
   closeSocket();
   resetRun();
-  appendLog("user", submittedMessage);
-  setRunState(true);
-  runStartedAt = performance.now();
-  startElapsedTimer();
-  statusEl.textContent = "Connecting";
-
+  beginTurn(submittedMessage, "Connecting");
   socket = new WebSocket(wsUrl());
   socket.addEventListener("open", () => {
     statusEl.textContent = "Waiting for first token";
     socket.send(JSON.stringify(buildRequest(submittedMessage)));
+    userMessage.value = "";
   });
   socket.addEventListener("message", (event) => {
     handleServerMessage(JSON.parse(event.data));
@@ -638,6 +641,22 @@ function runDebugSession() {
   });
 }
 
+function isSocketOpen() {
+  return socket && socket.readyState === WebSocket.OPEN;
+}
+
+function beginTurn(submittedMessage, statusText) {
+  activeAssistant = null;
+  firstTokenAt = 0;
+  hasBackendRoundStats = false;
+  resetTurnStats();
+  appendLog("user", submittedMessage);
+  setRunState(true);
+  runStartedAt = performance.now();
+  startElapsedTimer();
+  statusEl.textContent = statusText;
+}
+
 function buildRequest(submittedMessage = userMessage.value) {
   return {
     user_message: submittedMessage,
@@ -678,6 +697,10 @@ function handleServerMessage(message) {
     appendLog("session", "Session started");
     return;
   }
+  if (message.type === "turn_started") {
+    statusEl.textContent = "Waiting for first token";
+    return;
+  }
   if (message.type === "message_delta") {
     appendAssistantDelta(message.content || "");
     return;
@@ -704,10 +727,23 @@ function handleServerMessage(message) {
   }
   if (message.type === "error") {
     appendLog("error", message.message);
+    stopElapsedTimer();
+    setRunState(false);
+    return;
+  }
+  if (message.type === "turn_completed") {
+    appendLog("session", `Turn ${message.turn_index || ""} completed`.trim());
+    statusEl.textContent = "Idle";
+    stopElapsedTimer();
+    setRunState(false);
+    updateElapsed();
     return;
   }
   if (message.type === "done") {
     appendLog("session", "Done");
+    statusEl.textContent = "Idle";
+    stopElapsedTimer();
+    setRunState(false);
     closeSocket();
   }
 }
@@ -803,6 +839,10 @@ function resetRun() {
   firstTokenAt = 0;
   hasBackendRoundStats = false;
   messagesEl.textContent = "";
+  resetTurnStats();
+}
+
+function resetTurnStats() {
   document.querySelector("#stat-tokens").textContent = "0";
   document.querySelector("#stat-cached").textContent = "0";
   document.querySelector("#stat-ttft").textContent = "-";

+ 8 - 1
src/agent_lab/presentation/web.py

@@ -61,7 +61,10 @@ def create_app(
             payload = await websocket.receive_json()
             request = DebugRunRequest.model_validate(payload)
             logger.info("websocket request accepted")
-            if hasattr(runtime, "start"):
+            if hasattr(runtime, "start_session"):
+                queues = runtime.start_session(request)
+                await _run_queue_session(websocket, queues)
+            elif hasattr(runtime, "start"):
                 queues = runtime.start(request)
                 await _run_queue_session(websocket, queues)
             else:
@@ -71,6 +74,10 @@ def create_app(
             logger.info("websocket session disconnected")
             should_close = False
             return
+        except asyncio.CancelledError:
+            logger.info("websocket session cancelled")
+            should_close = False
+            return
         except (ValidationError, ValueError) as exc:
             logger.warning("websocket request rejected: %s", exc)
             await websocket.send_json({"type": "error", "message": str(exc)})

+ 61 - 0
tests/test_debug_runtime.py

@@ -309,6 +309,36 @@ class EventRoundStatsChatClient:
         )
 
 
+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
@@ -918,3 +948,34 @@ async def test_runtime_emits_round_stats_for_each_chat_call_in_event_handoff():
             "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

+ 73 - 0
tests/test_websocket_api.py

@@ -56,6 +56,42 @@ class QueueAwareRuntime:
             await asyncio.gather(self.task, return_exceptions=True)
 
 
+class PersistentSessionRuntime:
+    def __init__(self) -> None:
+        self.requests: list[DebugRunRequest] = []
+        self.queues: RuntimeQueues | None = None
+        self.task: asyncio.Task | None = None
+
+    def start_session(self, request: DebugRunRequest) -> RuntimeQueues:
+        self.requests.append(request)
+        self.queues = RuntimeQueues()
+        self.task = asyncio.create_task(self._run(request))
+        return self.queues
+
+    async def _run(self, request: DebugRunRequest) -> None:
+        assert self.queues is not None
+        await self.queues.output.put({"type": "session_started"})
+        await self._emit_turn(1, request.user_message)
+        turn_index = 1
+        while True:
+            message = await self.queues.input.get()
+            turn_index += 1
+            await self._emit_turn(turn_index, message.content)
+
+    async def _emit_turn(self, turn_index: int, content: str) -> None:
+        assert self.queues is not None
+        await self.queues.output.put({"type": "turn_started", "turn_index": turn_index})
+        await self.queues.output.put({"type": "message_delta", "content": content})
+        await self.queues.output.put(
+            {"type": "turn_completed", "turn_index": turn_index}
+        )
+
+    async def aclose(self) -> None:
+        if self.task is not None and not self.task.done():
+            self.task.cancel()
+            await asyncio.gather(self.task, return_exceptions=True)
+
+
 def _request_payload() -> dict:
     return {
         "user_message": "debug this",
@@ -203,6 +239,34 @@ def test_websocket_debug_enqueues_user_messages_during_running_session():
     assert runtime.requests[0].user_message == "debug this"
 
 
+def test_websocket_debug_keeps_session_open_for_multiple_turns():
+    runtime = PersistentSessionRuntime()
+    app = create_app(runtime_factory=lambda: runtime)
+    client = TestClient(app)
+
+    with client.websocket_connect("/ws/debug") as websocket:
+        websocket.send_json(_request_payload())
+
+        assert websocket.receive_json() == {"type": "session_started"}
+        assert websocket.receive_json() == {"type": "turn_started", "turn_index": 1}
+        assert websocket.receive_json() == {
+            "type": "message_delta",
+            "content": "debug this",
+        }
+        assert websocket.receive_json() == {"type": "turn_completed", "turn_index": 1}
+
+        websocket.send_json({"type": "user_message", "content": "follow-up"})
+        assert websocket.receive_json() == {"type": "turn_started", "turn_index": 2}
+        assert websocket.receive_json() == {
+            "type": "message_delta",
+            "content": "follow-up",
+        }
+        assert websocket.receive_json() == {"type": "turn_completed", "turn_index": 2}
+        websocket.close()
+
+    assert runtime.requests[0].user_message == "debug this"
+
+
 def test_websocket_debug_sends_error_for_invalid_request():
     app = create_app(runtime_factory=FakeRuntime)
     client = TestClient(app)
@@ -607,6 +671,15 @@ def test_static_app_shows_wait_state_and_immediate_user_echo():
     assert ".user" in css
 
 
+def test_static_app_reuses_websocket_for_session_turns():
+    js = Path("src/agent_lab/presentation/static/app.js").read_text()
+
+    assert "function isSocketOpen()" in js
+    assert "socket.send(JSON.stringify({ type: \"user_message\"" in js
+    assert 'message.type === "turn_completed"' in js
+    assert 'message.type === "done"' in js
+
+
 def test_static_workspace_snapshot_uses_stable_storage_hooks():
     js = Path("src/agent_lab/presentation/static/app.js").read_text()