From 4819c27be07b16a8ce584bc8d709636f611c1eb8 Mon Sep 17 00:00:00 2001 From: Xubin Ren <52506698+Re-bin@users.noreply.github.com> Date: Sat, 23 May 2026 02:16:17 +0800 Subject: [PATCH] fix(webui): scope image rewrites to webui chats --- nanobot/channels/websocket.py | 34 +++++++++++++++++------- tests/channels/test_websocket_channel.py | 27 +++++++++++++++++++ 2 files changed, 51 insertions(+), 10 deletions(-) diff --git a/nanobot/channels/websocket.py b/nanobot/channels/websocket.py index 856274090..bfb4743aa 100644 --- a/nanobot/channels/websocket.py +++ b/nanobot/channels/websocket.py @@ -477,6 +477,10 @@ class WebSocketChannel(BaseChannel): self._conn_chats: dict[Any, set[str]] = {} # connection -> default chat_id for legacy frames that omit routing. self._conn_default: dict[Any, str] = {} + # Chat IDs that opted into WebUI-specific rendering by sending a typed + # envelope with ``webui: true``. Raw WebSocket clients keep the legacy + # wire shape. + self._webui_chats: set[str] = set() # Single-use tokens consumed at WebSocket handshake. self._issued_tokens: dict[str, float] = {} # Multi-use tokens for HTTP routes served beside WS; checked but not consumed. @@ -1528,6 +1532,7 @@ class WebSocketChannel(BaseChannel): metadata: dict[str, Any] = {"remote": getattr(connection, "remote_address", None)} if envelope.get("webui") is True: metadata["webui"] = True + self._webui_chats.add(cid) cli_apps = normalize_cli_app_mentions(envelope.get("cli_apps")) if cli_apps: metadata["cli_apps"] = cli_apps @@ -1564,6 +1569,7 @@ class WebSocketChannel(BaseChannel): self._subs.clear() self._conn_chats.clear() self._conn_default.clear() + self._webui_chats.clear() self._issued_tokens.clear() self._api_tokens.clear() @@ -1642,7 +1648,8 @@ class WebSocketChannel(BaseChannel): await self._safe_send_to(connection, raw, label=" ") return text = msg.content - wire_text = self._rewrite_local_markdown_images(text) + should_rewrite_images = msg.chat_id in self._webui_chats + wire_text = self._rewrite_local_markdown_images(text) if should_rewrite_images else text payload: dict[str, Any] = { "event": "message", "chat_id": msg.chat_id, @@ -1742,25 +1749,32 @@ class WebSocketChannel(BaseChannel): return meta = metadata or {} stream_key = (chat_id, str(meta.get("_stream_id") or "")) + should_rewrite_images = chat_id in self._webui_chats + transcript_body: dict[str, Any] | None = None if meta.get("_stream_end"): body: dict[str, Any] = {"event": "stream_end", "chat_id": chat_id} - buffered = self._stream_text_buffers.pop(stream_key, []) - if delta: - buffered.append(delta) - full_text = "".join(buffered) - rewritten = self._rewrite_local_markdown_images(full_text) - if rewritten != full_text: - body["text"] = rewritten + if should_rewrite_images: + buffered = self._stream_text_buffers.pop(stream_key, []) + if delta: + buffered.append(delta) + full_text = "".join(buffered) + rewritten = self._rewrite_local_markdown_images(full_text) + if rewritten != full_text or delta: + body["text"] = rewritten + transcript_body = {**body, "text": full_text} else: body = { "event": "delta", "chat_id": chat_id, "text": delta, } - self._stream_text_buffers.setdefault(stream_key, []).append(delta) + if should_rewrite_images: + self._stream_text_buffers.setdefault(stream_key, []).append(delta) if meta.get("_stream_id") is not None: body["stream_id"] = meta["_stream_id"] - self._try_append_webui_transcript(chat_id, body) + if transcript_body is not None: + transcript_body["stream_id"] = meta["_stream_id"] + self._try_append_webui_transcript(chat_id, transcript_body or body) raw = json.dumps(body, ensure_ascii=False) for connection in conns: await self._safe_send_to(connection, raw, label=" stream ") diff --git a/tests/channels/test_websocket_channel.py b/tests/channels/test_websocket_channel.py index f40ec4872..12ca2399c 100644 --- a/tests/channels/test_websocket_channel.py +++ b/tests/channels/test_websocket_channel.py @@ -501,6 +501,7 @@ async def test_send_delta_stream_end_rewrites_local_markdown_image(monkeypatch, ) mock_ws = AsyncMock() channel._attach(mock_ws, "chat-1") + channel._webui_chats.add("chat-1") await channel.send_delta("chat-1", "![Diagram](", {"_stream_delta": True, "_stream_id": "sid"}) await channel.send_delta("chat-1", "diagram.png)", {"_stream_delta": True, "_stream_id": "sid"}) @@ -533,6 +534,7 @@ async def test_send_delta_stream_end_rewrites_inline_final_text(monkeypatch, tmp ) mock_ws = AsyncMock() channel._attach(mock_ws, "chat-1") + channel._webui_chats.add("chat-1") await channel.send_delta( "chat-1", @@ -546,6 +548,31 @@ async def test_send_delta_stream_end_rewrites_inline_final_text(monkeypatch, tmp assert final["text"].startswith("![Diagram](/api/media/") +@pytest.mark.asyncio +async def test_send_delta_stream_end_leaves_non_webui_payload_unchanged(tmp_path) -> None: + bus = MagicMock() + workspace = tmp_path / "workspace" + workspace.mkdir() + (workspace / "diagram.png").write_bytes(b"\x89PNG\r\n\x1a\nimage") + channel = WebSocketChannel( + {"enabled": True, "allowFrom": ["*"], "streaming": True}, + bus, + workspace_path=workspace, + ) + mock_ws = AsyncMock() + channel._attach(mock_ws, "chat-1") + + await channel.send_delta( + "chat-1", + "![Diagram](diagram.png)", + {"_stream_delta": True, "_stream_end": True, "_stream_id": "sid"}, + ) + + mock_ws.send.assert_awaited_once() + final = json.loads(mock_ws.send.await_args.args[0]) + assert final == {"event": "stream_end", "chat_id": "chat-1", "stream_id": "sid"} + + @pytest.mark.asyncio async def test_send_reasoning_delta_emits_streaming_frame() -> None: bus = MagicMock()