fix(webui): make mutations reconnect-safe
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user