fix(webui): make mutations reconnect-safe

This commit is contained in:
chengyongru
2026-08-16 21:40:14 +08:00
committed by chengyongru
parent dec89a49a3
commit e51ffc8978
4 changed files with 325 additions and 47 deletions
+123 -36
View File
@@ -3,14 +3,17 @@
from __future__ import annotations
import asyncio
import hashlib
import hmac
import ipaddress
import json
import re
import ssl
import time
import uuid
from collections.abc import Callable
from contextlib import suppress
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Self, TypeGuard, cast
from urllib.parse import urlsplit, urlunsplit
@@ -92,6 +95,8 @@ from nanobot.webui.websocket_logging import websockets_server_logger
# Plain HTTP WebUI routes also run through websockets.process_request.
_WEBUI_HTTP_OPEN_TIMEOUT_S = 360.0
_WEBUI_REQUEST_CACHE_TTL_S = 5 * 60.0
_WEBUI_REQUEST_CACHE_MAX = 256
_ROUTING_ASSERTION_HEADERS = frozenset(
@@ -348,6 +353,21 @@ def _is_websocket_upgrade(request: WsRequest) -> bool:
return True
@dataclass(frozen=True)
class _WebUIRequestResult:
result: Any = None
status: int | None = None
message: str | None = None
@dataclass
class _WebUIRequestOperation:
action: str
payload_digest: bytes
task: asyncio.Task[_WebUIRequestResult]
completed_at: float | None = None
class WebSocketChannel(BaseChannel):
"""Run a local WebSocket server; forward text/JSON messages to the message bus."""
@@ -373,14 +393,14 @@ class WebSocketChannel(BaseChannel):
self._conn_default: dict[ServerConnection, str] = {}
# Connections authenticated with a one-time token from /webui/bootstrap.
self._webui_connections: set[ServerConnection] = set()
# Request/reply mutations aren't replayed across reconnects. Tasks may
# finish after a client-side deadline so an already-started mutation
# isn't ambiguously cancelled halfway through.
# Delivery tasks are connection-bound, while operations are keyed only
# by request_id so reconnect retries join or replay the original work.
self._webui_request_tasks: dict[
tuple[ServerConnection, str],
asyncio.Task[None],
] = {}
# Preserve request/response order for non-replayable mutations from one
self._webui_request_operations: dict[str, _WebUIRequestOperation] = {}
# Preserve request/response order for mutations from one
# UI. Without this, an earlier slow settings response can overwrite a
# newer settings snapshot in the client.
self._webui_request_locks: dict[ServerConnection, asyncio.Lock] = {}
@@ -1182,32 +1202,107 @@ class WebSocketChannel(BaseChannel):
)
return
key = (connection, request_id)
if key in self._webui_request_tasks:
payload_digest = hashlib.sha256(
json.dumps(
payload,
ensure_ascii=False,
sort_keys=True,
separators=(",", ":"),
).encode("utf-8")
).digest()
self._prune_webui_request_operations()
operation = self._webui_request_operations.get(request_id)
if operation is not None and (
operation.action != action or operation.payload_digest != payload_digest
):
await self._send_webui_response(
connection,
request_id,
status=409,
message="duplicate WebUI request_id",
message="request_id was already used for a different WebUI mutation",
)
return
task = asyncio.create_task(
self._complete_webui_request(
if operation is None:
operation_task = asyncio.create_task(
self._execute_webui_request(
connection,
action,
cast(dict[str, Any], payload),
)
)
new_operation = _WebUIRequestOperation(
action=action,
payload_digest=payload_digest,
task=operation_task,
)
operation = new_operation
self._webui_request_operations[request_id] = new_operation
def mark_complete(_task: asyncio.Task[_WebUIRequestResult]) -> None:
current = self._webui_request_operations.get(request_id)
if current is not new_operation:
return
new_operation.completed_at = time.monotonic()
self._prune_webui_request_operations()
operation_task.add_done_callback(mark_complete)
key = (connection, request_id)
if key in self._webui_request_tasks:
return
delivery_task = asyncio.create_task(
self._deliver_webui_request(
connection,
request_id,
action,
cast(dict[str, Any], payload),
operation.task,
)
)
self._webui_request_tasks[key] = task
self._webui_request_tasks[key] = delivery_task
async def _complete_webui_request(
def _prune_webui_request_operations(self) -> None:
now = time.monotonic()
for request_id, operation in tuple(self._webui_request_operations.items()):
if (
operation.completed_at is not None
and now - operation.completed_at >= _WEBUI_REQUEST_CACHE_TTL_S
):
self._webui_request_operations.pop(request_id, None)
completed = sorted(
(
(operation.completed_at, request_id)
for request_id, operation in self._webui_request_operations.items()
if operation.completed_at is not None
),
key=lambda item: item[0],
)
for _, request_id in completed[:-_WEBUI_REQUEST_CACHE_MAX]:
self._webui_request_operations.pop(request_id, None)
async def _deliver_webui_request(
self,
connection: ServerConnection,
request_id: str,
operation_task: asyncio.Task[_WebUIRequestResult],
) -> None:
try:
result = await asyncio.shield(operation_task)
await self._send_webui_response(
connection,
request_id,
result=result.result,
status=result.status,
message=result.message,
)
finally:
self._webui_request_tasks.pop((connection, request_id), None)
async def _execute_webui_request(
self,
connection: ServerConnection,
action: str,
payload: dict[str, Any],
) -> None:
) -> _WebUIRequestResult:
try:
lock = self._webui_request_locks.setdefault(connection, asyncio.Lock())
async with lock:
@@ -1222,27 +1317,17 @@ class WebSocketChannel(BaseChannel):
try:
result = json.loads(body)
except json.JSONDecodeError:
await self._send_webui_response(
connection,
request_id,
return _WebUIRequestResult(
status=502,
message="WebUI mutation returned an invalid response",
)
return
if action == "sidebar.update" and isinstance(result, dict):
await self._broadcast_webui_event(
"sidebar_state_updated",
state=result,
)
await self._send_webui_response(
connection,
request_id,
result=result,
)
return
await self._send_webui_response(
connection,
request_id,
return _WebUIRequestResult(result=result)
return _WebUIRequestResult(
status=status,
message=body or response.reason_phrase,
)
@@ -1250,14 +1335,10 @@ class WebSocketChannel(BaseChannel):
raise
except Exception:
self.logger.exception("WebUI mutation '{}' failed", action)
await self._send_webui_response(
connection,
request_id,
return _WebUIRequestResult(
status=500,
message="WebUI mutation failed",
)
finally:
self._webui_request_tasks.pop((connection, request_id), None)
async def _send_webui_response(
self,
@@ -1328,13 +1409,19 @@ class WebSocketChannel(BaseChannel):
except Exception as e:
self.logger.warning("server task error during shutdown: {}", e)
self._server_task = None
mutation_tasks = tuple(self._webui_request_tasks.values())
for task in mutation_tasks:
delivery_tasks = tuple(self._webui_request_tasks.values())
operation_tasks = tuple(
operation.task for operation in self._webui_request_operations.values()
)
for task in (*delivery_tasks, *operation_tasks):
task.cancel()
if mutation_tasks:
await asyncio.gather(*mutation_tasks, return_exceptions=True)
if delivery_tasks:
await asyncio.gather(*delivery_tasks, return_exceptions=True)
if operation_tasks:
await asyncio.gather(*operation_tasks, return_exceptions=True)
self._webui_request_tasks.clear()
self._webui_request_locks.clear()
self._webui_request_operations.clear()
self._subs.clear()
self._conn_chats.clear()
self._conn_default.clear()