浏览代码

feat: run websocket sessions over queues

zhenyu.hu 3 周之前
父节点
当前提交
e8eeedb643

+ 13 - 5
docs/plans/todo-13-real-queue-websocket-loop.md

@@ -1,6 +1,6 @@
 # Todo 13 Real Queue WebSocket Loop Plan
 
-**Status:** in_progress
+**Status:** done
 
 ## Goal
 
@@ -90,8 +90,16 @@ uv run pytest
 
 ## Verification
 
-Expected result:
+Result:
 
-- Runtime tests prove output queue and event queue have separate consumers.
-- WebSocket tests prove upstream messages can enter the running input queue.
-- Full suite passes.
+- Red run passed as expected: the new tests failed because `DebugRuntime.start()` did not exist, event-time user input was ordered before the tool reply, and WebSocket still called only `run()`.
+- `uv run pytest tests/test_debug_runtime.py tests/test_websocket_api.py` passed: 26 tests, 1 existing Starlette deprecation warning.
+- `uv run pytest` passed: 31 tests, 1 existing Starlette deprecation warning.
+
+## Evaluation
+
+- `DebugRuntime.start()` now starts a queue-backed session and returns `RuntimeQueues` for downstream consumers.
+- `DebugRuntime.run()` remains as a compatibility wrapper over the output queue.
+- EventAgent now consumes the event queue in a worker task and writes tool replies to the input queue.
+- WebSocket sessions now split downstream output sending from upstream `user_message` receiving.
+- User messages that arrive during a tool handoff are buffered until after the matching tool reply, preserving provider-valid message order.

+ 1 - 1
docs/plans/todos.md

@@ -39,7 +39,7 @@
 | 10 | done | `docs/plans/todo-10-tool-registry-management.md` | Replace the single hardcoded EventAgent tool checkbox with a backend tool registry and UI-managed loaded tools. | Tests prove available tools are listed by API and selected tools control EventAgent behavior. |
 | 11 | done | `docs/plans/todo-11-prompt-workspace.md` | Improve prompt/pre-message editing with browser-side persistence and reusable prompt sets. | UI tests or focused JS tests prove prompts persist and can be restored. |
 | 12 | done | `docs/plans/todo-12-round-stats.md` | Move round statistics to structured backend events for token counts, cached tokens, TTFT, elapsed time, and per-turn summaries. | `uv run pytest tests/test_debug_runtime.py tests/test_websocket_api.py`; `uv run pytest`. |
-| 13 | in_progress | `docs/plans/todo-13-real-queue-websocket-loop.md` | Make runtime queues real producer/consumer boundaries and let WebSocket run separate upstream/downstream tasks. | Tests prove output queue can be consumed directly, WS upstream enqueues user input, and event tool replies keep provider-valid order. |
+| 13 | done | `docs/plans/todo-13-real-queue-websocket-loop.md` | Make runtime queues real producer/consumer boundaries and let WebSocket run separate upstream/downstream tasks. | `uv run pytest tests/test_debug_runtime.py tests/test_websocket_api.py`; `uv run pytest`. |
 | 14 | pending | `docs/plans/todo-14-message-role-validation.md` | Reject provider-invalid API message roles and malformed tool replies before network calls. | Tests prove request validation rejects invalid roles and orphan tool messages. |
 | 15 | pending | `docs/plans/todo-15-static-package-data.md` | Include static UI assets in installed package builds. | Build/package inspection proves `index.html`, JS, and CSS are present. |
 | 16 | pending | `docs/plans/todo-16-tool-error-isolation.md` | Convert EventAgent tool handler exceptions into structured tool-result errors instead of aborting the session. | Tests prove a failing tool returns an error payload and the WebSocket run stays structured. |

+ 1 - 1
src/agent_lab/application/queues.py

@@ -10,4 +10,4 @@ from agent_lab.domain.messages import ChatMessage
 class RuntimeQueues:
     input: asyncio.Queue[ChatMessage] = field(default_factory=asyncio.Queue)
     output: asyncio.Queue[dict[str, Any]] = field(default_factory=asyncio.Queue)
-    events: asyncio.Queue[ToolCallEvent] = field(default_factory=asyncio.Queue)
+    events: asyncio.Queue[ToolCallEvent | None] = field(default_factory=asyncio.Queue)

+ 93 - 31
src/agent_lab/application/runtime.py

@@ -1,3 +1,4 @@
+import asyncio
 import time
 from collections.abc import AsyncIterator, Callable
 from typing import Any, Protocol
@@ -31,18 +32,66 @@ class DebugRuntime:
         self.queues = queues
         self.registry = registry or build_default_tool_registry()
         self.clock = clock
+        self._tasks: list[asyncio.Task[None]] = []
 
-    async def run(self, request: DebugRunRequest) -> AsyncIterator[dict[str, Any]]:
+    def start(self, request: DebugRunRequest) -> RuntimeQueues:
         queues = self.queues or RuntimeQueues()
-        messages = self._build_initial_messages(request)
-        await queues.input.put(ChatMessage(role="user", content=request.user_message))
-        tools = self.registry.chat_tools(request.event_agent.enabled_tools)
+        task = asyncio.create_task(self._produce(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):
+            yield message
+        await self._wait_for_tasks()
+
+    async def output_messages(
+        self,
+        queues: RuntimeQueues,
+    ) -> AsyncIterator[dict[str, Any]]:
+        while True:
+            message = await queues.output.get()
+            yield message
+            if message.get("type") in {"done", "error"}:
+                break
+
+    async def aclose(self) -> None:
+        for task in self._tasks:
+            if not task.done():
+                task.cancel()
+        await self._wait_for_tasks()
+
+    async def _wait_for_tasks(self) -> None:
+        if not self._tasks:
+            return
+        await asyncio.gather(*self._tasks, return_exceptions=True)
+        self._tasks = [task for task in self._tasks if not task.done()]
+
+    async def _produce(self, request: DebugRunRequest, queues: RuntimeQueues) -> None:
         event_agent = EventAgent(
             request.event_agent.enabled_tools,
             registry=self.registry,
         )
+        event_worker = asyncio.create_task(self._consume_events(queues, event_agent))
+        try:
+            await self._run_chat_agent(request, queues)
+        except Exception as exc:
+            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,
+        queues: RuntimeQueues,
+    ) -> None:
+        messages = self._build_initial_messages(request)
+        await queues.input.put(ChatMessage(role="user", content=request.user_message))
+        tools = self.registry.chat_tools(request.event_agent.enabled_tools)
 
-        yield await self._emit(queues, {"type": "session_started"})
+        await queues.output.put({"type": "session_started"})
 
         event_loops = 0
         round_index = 0
@@ -58,6 +107,7 @@ class DebugRuntime:
             saw_event = False
             assistant_content: list[str] = []
             assistant_tool_calls: list[dict[str, Any]] = []
+            tool_replies: list[ChatMessage] = []
 
             async for item in self.chat_client.stream_chat(
                 messages=messages,
@@ -68,9 +118,8 @@ class DebugRuntime:
                     if ttft_ms is None:
                         ttft_ms = self._elapsed_ms(started_at)
                     assistant_content.append(item.content or "")
-                    yield await self._emit(
-                        queues,
-                        {"type": "message_delta", "content": item.content},
+                    await queues.output.put(
+                        {"type": "message_delta", "content": item.content}
                     )
                     continue
 
@@ -79,9 +128,8 @@ class DebugRuntime:
                     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()},
+                    await queues.output.put(
+                        {"type": "usage", "usage": item.usage.model_dump()}
                     )
                     continue
 
@@ -98,16 +146,13 @@ class DebugRuntime:
                         }
                     )
                     await queues.events.put(item.event)
-                    yield await self._emit(
-                        queues,
-                        {"type": "event", "event": item.event.model_dump()},
+                    await queues.output.put(
+                        {"type": "event", "event": item.event.model_dump()}
                     )
-                    event = await queues.events.get()
-                    reply = await event_agent.handle(event)
-                    await queues.input.put(reply)
-                    yield await self._emit(
-                        queues,
-                        {"type": "tool_result", "message": reply.model_dump()},
+                    reply = await self._wait_for_tool_reply(queues, item.event.id)
+                    tool_replies.append(reply)
+                    await queues.output.put(
+                        {"type": "tool_result", "message": reply.model_dump()}
                     )
 
             if assistant_content or assistant_tool_calls:
@@ -118,9 +163,9 @@ class DebugRuntime:
                         tool_calls=assistant_tool_calls or None,
                     )
                 )
+            messages.extend(tool_replies)
 
-            yield await self._emit(
-                queues,
+            await queues.output.put(
                 {
                     "type": "round_stats",
                     "round_index": round_index,
@@ -141,7 +186,32 @@ class DebugRuntime:
             if event_loops >= request.event_agent.max_event_loops:
                 break
 
-        yield await self._emit(queues, {"type": "done"})
+        await queues.output.put({"type": "done"})
+
+    async def _consume_events(
+        self,
+        queues: RuntimeQueues,
+        event_agent: EventAgent,
+    ) -> None:
+        while True:
+            event = await queues.events.get()
+            if event is None:
+                return
+            await queues.input.put(await event_agent.handle(event))
+
+    async def _wait_for_tool_reply(
+        self,
+        queues: RuntimeQueues,
+        tool_call_id: str,
+    ) -> ChatMessage:
+        deferred: list[ChatMessage] = []
+        while True:
+            message = await queues.input.get()
+            if message.role == "tool" and message.tool_call_id == tool_call_id:
+                for deferred_message in deferred:
+                    await queues.input.put(deferred_message)
+                return message
+            deferred.append(message)
 
     def _build_initial_messages(self, request: DebugRunRequest) -> list[ChatMessage]:
         messages = [
@@ -159,13 +229,5 @@ class DebugRuntime:
         while not queues.input.empty():
             messages.append(await queues.input.get())
 
-    async def _emit(
-        self,
-        queues: RuntimeQueues,
-        payload: dict[str, Any],
-    ) -> 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)

+ 58 - 2
src/agent_lab/presentation/web.py

@@ -1,3 +1,4 @@
+import asyncio
 from collections.abc import Callable
 from pathlib import Path
 from typing import Any
@@ -8,8 +9,10 @@ from fastapi.staticfiles import StaticFiles
 from pydantic import ValidationError
 
 from agent_lab.application.contracts import DebugRunRequest
+from agent_lab.application.queues import RuntimeQueues
 from agent_lab.application.runtime import DebugRuntime
 from agent_lab.application.tools import ToolRegistry, build_default_tool_registry
+from agent_lab.domain.messages import ChatMessage
 from agent_lab.infrastructure.chat_client import OpenAICompatibleChatClient
 from agent_lab.settings import Settings
 
@@ -17,6 +20,7 @@ from agent_lab.settings import Settings
 RuntimeFactory = Callable[[], Any]
 
 STATIC_DIR = Path(__file__).parent / "static"
+TERMINAL_MESSAGE_TYPES = {"done", "error"}
 
 
 def create_app(
@@ -53,8 +57,12 @@ def create_app(
         try:
             payload = await websocket.receive_json()
             request = DebugRunRequest.model_validate(payload)
-            async for message in runtime.run(request):
-                await websocket.send_json(message)
+            if hasattr(runtime, "start"):
+                queues = runtime.start(request)
+                await _run_queue_session(websocket, queues)
+            else:
+                async for message in runtime.run(request):
+                    await websocket.send_json(message)
         except WebSocketDisconnect:
             should_close = False
             return
@@ -70,6 +78,50 @@ def create_app(
     return app
 
 
+async def _run_queue_session(websocket: WebSocket, queues: RuntimeQueues) -> None:
+    downstream = asyncio.create_task(_send_downstream(websocket, queues))
+    upstream = asyncio.create_task(_receive_upstream(websocket, queues))
+    done, pending = await asyncio.wait(
+        {downstream, upstream},
+        return_when=asyncio.FIRST_COMPLETED,
+    )
+    for task in done:
+        await task
+    for task in pending:
+        task.cancel()
+    await asyncio.gather(*pending, return_exceptions=True)
+
+
+async def _send_downstream(websocket: WebSocket, queues: RuntimeQueues) -> None:
+    while True:
+        message = await queues.output.get()
+        await websocket.send_json(message)
+        if message.get("type") in TERMINAL_MESSAGE_TYPES:
+            return
+
+
+async def _receive_upstream(websocket: WebSocket, queues: RuntimeQueues) -> None:
+    while True:
+        payload = await websocket.receive_json()
+        try:
+            message = _parse_upstream_message(payload)
+        except (TypeError, ValueError) as exc:
+            await queues.output.put({"type": "error", "message": str(exc)})
+            return
+        await queues.input.put(message)
+
+
+def _parse_upstream_message(payload: Any) -> ChatMessage:
+    if not isinstance(payload, dict):
+        raise TypeError("upstream message must be a JSON object")
+    if payload.get("type") != "user_message":
+        raise ValueError("unsupported upstream message type")
+    content = payload.get("content")
+    if not isinstance(content, str):
+        raise TypeError("user_message content must be a string")
+    return ChatMessage(role="user", content=content)
+
+
 def _runtime_factory(
     settings: Settings,
     tool_registry: ToolRegistry,
@@ -88,6 +140,10 @@ def _runtime_factory(
 
 
 async def _close_runtime(runtime: Any) -> None:
+    runtime_close = getattr(runtime, "aclose", None)
+    if runtime_close is not None:
+        await runtime_close()
+
     chat_client = getattr(runtime, "chat_client", None)
     close = getattr(chat_client, "aclose", None)
     if close is not None:

+ 76 - 20
tests/test_debug_runtime.py

@@ -201,6 +201,66 @@ async def test_runtime_routes_chat_events_through_event_agent_then_continues_cha
     assert outputs[4]["content"] == "final answer"
 
 
+@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["type"] for message in outputs] == [
+        "session_started",
+        "message_delta",
+        "usage",
+        "round_stats",
+        "done",
+    ]
+
+
+@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 anext(stream) == {"type": "session_started"}
+    event_message = await anext(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] == [
+        "user",
+        "assistant",
+        "tool",
+        "user",
+    ]
+    assert client.second_call_messages[2].tool_call_id == "call_1"
+    assert client.second_call_messages[3].content == "follow-up while tool runs"
+
+
 @pytest.mark.asyncio
 async def test_runtime_uses_event_and_input_queues_for_event_agent_handoff():
     RuntimeQueues = _runtime_queues_class()
@@ -245,7 +305,7 @@ async def test_runtime_uses_event_and_input_queues_for_event_agent_handoff():
 
 
 @pytest.mark.asyncio
-async def test_runtime_yields_existing_output_order_from_output_queue():
+async def test_runtime_run_consumes_output_queue_in_stream_order():
     RuntimeQueues = _runtime_queues_class()
     queue_log: list[tuple[str, str, str]] = []
     queues = RuntimeQueues(
@@ -273,26 +333,22 @@ async def test_runtime_yields_existing_output_order_from_output_queue():
         "round_stats",
         "done",
     ]
-    assert [
-        entry
-        for entry in queue_log
-        if entry[0] == "output" and entry[1] in {"put", "get"}
-    ] == [
-        ("output", "put", "output:session_started"),
-        ("output", "get", "output:session_started"),
-        ("output", "put", "output:event"),
-        ("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"),
+    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 output_puts == [
+        "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

+ 48 - 0
tests/test_websocket_api.py

@@ -1,4 +1,5 @@
 import json
+import asyncio
 from collections.abc import AsyncIterator
 from pathlib import Path
 
@@ -7,6 +8,7 @@ import pytest
 from fastapi.testclient import TestClient
 
 from agent_lab.application.contracts import AgentParams, DebugRunRequest, EventAgentParams
+from agent_lab.application.queues import RuntimeQueues
 from agent_lab.application.runtime import DebugRuntime
 from agent_lab.domain.events import ToolCallEvent
 from agent_lab.domain.messages import ChatMessage, StreamItem
@@ -25,6 +27,33 @@ class FakeRuntime:
         yield {"type": "done"}
 
 
+class QueueAwareRuntime:
+    def __init__(self) -> None:
+        self.requests: list[DebugRunRequest] = []
+        self.queues: RuntimeQueues | None = None
+        self.task: asyncio.Task | None = None
+
+    def start(self, request: DebugRunRequest) -> RuntimeQueues:
+        self.requests.append(request)
+        self.queues = RuntimeQueues()
+        self.task = asyncio.create_task(self._run())
+        return self.queues
+
+    async def _run(self) -> None:
+        assert self.queues is not None
+        await self.queues.output.put({"type": "session_started"})
+        message = await self.queues.input.get()
+        await self.queues.output.put(
+            {"type": "message_delta", "content": message.content}
+        )
+        await self.queues.output.put({"type": "done"})
+
+    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",
@@ -93,6 +122,25 @@ def test_websocket_debug_streams_runtime_messages():
     assert runtime.requests[0].pre_messages[0].content == "previous turn"
 
 
+def test_websocket_debug_enqueues_user_messages_during_running_session():
+    runtime = QueueAwareRuntime()
+    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"}
+        websocket.send_json({"type": "user_message", "content": "follow-up"})
+        assert websocket.receive_json() == {
+            "type": "message_delta",
+            "content": "follow-up",
+        }
+        assert websocket.receive_json() == {"type": "done"}
+
+    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)