refactor: split WebUI gateway dependencies
Maintainer edit for PR 4115: rebase onto origin/main and split gateway HTTP routing from token, media, and workspace services so WebSocketChannel depends on explicit gateway services instead of GatewayHTTPHandler internals. Preserve file edit channel capabilities and restore tools.restrict_to_workspace wiring through ChannelManager.
This commit is contained in:
+37
-267
@@ -3,27 +3,20 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import email.utils
|
||||
import hmac
|
||||
import http
|
||||
import json
|
||||
import re
|
||||
import ssl
|
||||
import uuid
|
||||
from collections.abc import Callable
|
||||
from contextlib import suppress
|
||||
from functools import partial
|
||||
from pathlib import Path
|
||||
from typing import Any, Self
|
||||
from urllib.parse import parse_qs, unquote, urlparse
|
||||
|
||||
from loguru import logger
|
||||
from pydantic import Field, field_validator, model_validator
|
||||
from websockets.asyncio.server import ServerConnection, serve, unix_serve
|
||||
from websockets.datastructures import Headers
|
||||
from websockets.exceptions import ConnectionClosed
|
||||
from websockets.http11 import Request as WsRequest
|
||||
from websockets.http11 import Response
|
||||
|
||||
from nanobot.bus.events import OUTBOUND_META_AGENT_UI, OutboundMessage
|
||||
from nanobot.bus.queue import MessageBus
|
||||
@@ -41,53 +34,22 @@ from nanobot.utils.media_decode import (
|
||||
save_base64_data_url,
|
||||
)
|
||||
from nanobot.webui.cli_apps_api import normalize_cli_app_mentions
|
||||
from nanobot.webui.gateway_services import GatewayServices
|
||||
from nanobot.webui.http_utils import (
|
||||
is_localhost as _is_localhost,
|
||||
)
|
||||
from nanobot.webui.http_utils import (
|
||||
normalize_config_path as _normalize_config_path,
|
||||
)
|
||||
from nanobot.webui.http_utils import (
|
||||
parse_request_path as _parse_request_path,
|
||||
)
|
||||
from nanobot.webui.http_utils import (
|
||||
query_first as _query_first,
|
||||
)
|
||||
from nanobot.webui.mcp_presets_api import normalize_mcp_preset_mentions
|
||||
from nanobot.webui.transcript import append_transcript_object
|
||||
|
||||
|
||||
def _strip_trailing_slash(path: str) -> str:
|
||||
if len(path) > 1 and path.endswith("/"):
|
||||
return path.rstrip("/")
|
||||
return path or "/"
|
||||
|
||||
|
||||
def _normalize_config_path(path: str) -> str:
|
||||
return _strip_trailing_slash(path)
|
||||
|
||||
|
||||
def _case_insensitive_header(headers: Any, key: str) -> str:
|
||||
"""Read a header from websockets/http test stubs without assuming casing."""
|
||||
try:
|
||||
value = headers.get(key)
|
||||
except Exception:
|
||||
value = None
|
||||
if value is None:
|
||||
try:
|
||||
value = headers.get(key.lower())
|
||||
except Exception:
|
||||
value = None
|
||||
return str(value or "").strip()
|
||||
|
||||
|
||||
def _safe_host_header(value: str) -> str:
|
||||
"""Return a safe Host header value, or empty when it should not be echoed."""
|
||||
value = value.strip()
|
||||
if not value:
|
||||
return ""
|
||||
if re.fullmatch(r"\[[0-9A-Fa-f:.]+\](?::\d{1,5})?", value):
|
||||
return value
|
||||
if re.fullmatch(r"[A-Za-z0-9.-]+(?::\d{1,5})?", value):
|
||||
return value
|
||||
return ""
|
||||
|
||||
|
||||
def _host_for_url(host: str, port: int) -> str:
|
||||
host = host.strip()
|
||||
if host in ("0.0.0.0", "::"):
|
||||
host = "127.0.0.1"
|
||||
if ":" in host and not host.startswith("["):
|
||||
host = f"[{host}]"
|
||||
return f"{host}:{port}"
|
||||
from nanobot.webui.websocket_logging import websockets_server_logger
|
||||
|
||||
|
||||
class WebSocketConfig(Base):
|
||||
@@ -182,20 +144,6 @@ class WebSocketConfig(Base):
|
||||
)
|
||||
|
||||
|
||||
def _http_json_response(data: dict[str, Any], *, status: int = 200) -> Response:
|
||||
body = json.dumps(data, ensure_ascii=False).encode("utf-8")
|
||||
headers = Headers(
|
||||
[
|
||||
("Date", email.utils.formatdate(usegmt=True)),
|
||||
("Connection", "close"),
|
||||
("Content-Length", str(len(body))),
|
||||
("Content-Type", "application/json; charset=utf-8"),
|
||||
]
|
||||
)
|
||||
reason = http.HTTPStatus(status).phrase
|
||||
return Response(status, reason, headers, body)
|
||||
|
||||
|
||||
def publish_runtime_model_update(
|
||||
bus: MessageBus,
|
||||
model: str,
|
||||
@@ -214,57 +162,6 @@ def publish_runtime_model_update(
|
||||
))
|
||||
|
||||
|
||||
def _default_model_name_from_config() -> str | None:
|
||||
"""Resolved model string from on-disk config (bootstrap fallback)."""
|
||||
try:
|
||||
from nanobot.config.loader import load_config
|
||||
|
||||
model = load_config().resolve_preset().model.strip()
|
||||
return model or None
|
||||
except Exception as e:
|
||||
logger.debug("bootstrap model_name could not load from config: {}", e)
|
||||
return None
|
||||
|
||||
|
||||
def _resolve_bootstrap_model_name(
|
||||
runtime_name: Callable[[], str | None] | None,
|
||||
) -> str | None:
|
||||
"""Prefer an in-process resolver (e.g. AgentLoop); else config-derived default."""
|
||||
if runtime_name is not None:
|
||||
try:
|
||||
raw = runtime_name()
|
||||
except Exception as e:
|
||||
logger.debug("bootstrap runtime model resolver failed: {}", e)
|
||||
else:
|
||||
if isinstance(raw, str):
|
||||
stripped = raw.strip()
|
||||
if stripped:
|
||||
return stripped
|
||||
return _default_model_name_from_config()
|
||||
|
||||
|
||||
def _parse_request_path(path_with_query: str) -> tuple[str, dict[str, list[str]]]:
|
||||
"""Parse normalized path and query parameters in one pass."""
|
||||
parsed = urlparse("ws://x" + path_with_query)
|
||||
path = _strip_trailing_slash(parsed.path or "/")
|
||||
return path, parse_qs(parsed.query, keep_blank_values=True)
|
||||
|
||||
|
||||
def _normalize_http_path(path_with_query: str) -> str:
|
||||
"""Return the path component (no query string), with trailing slash normalized (root stays ``/``)."""
|
||||
return _parse_request_path(path_with_query)[0]
|
||||
|
||||
|
||||
def _parse_query(path_with_query: str) -> dict[str, list[str]]:
|
||||
return _parse_request_path(path_with_query)[1]
|
||||
|
||||
|
||||
def _query_first(query: dict[str, list[str]], key: str) -> str | None:
|
||||
"""Return the first value for *key*, or None."""
|
||||
values = query.get(key)
|
||||
return values[0] if values else None
|
||||
|
||||
|
||||
def _parse_inbound_payload(raw: str) -> str | None:
|
||||
"""Parse a client frame into text; return None for empty or unrecognized content."""
|
||||
text = raw.strip()
|
||||
@@ -355,67 +252,6 @@ def _extract_data_url_mime(url: str) -> str | None:
|
||||
return m.group(1).strip().lower() or None
|
||||
|
||||
|
||||
_LOCALHOSTS = frozenset({"127.0.0.1", "::1", "localhost"})
|
||||
|
||||
# Matches the legacy chat-id pattern but allows file-system-safe stems too,
|
||||
# so the API can address sessions whose keys came from non-WebSocket channels.
|
||||
_API_KEY_RE = re.compile(r"^[A-Za-z0-9_:.-]{1,128}$")
|
||||
|
||||
|
||||
def _decode_api_key(raw_key: str) -> str | None:
|
||||
"""Decode a percent-encoded API path segment, then validate the result."""
|
||||
key = unquote(raw_key)
|
||||
if _API_KEY_RE.match(key) is None:
|
||||
return None
|
||||
return key
|
||||
|
||||
|
||||
def _is_localhost(connection: Any) -> bool:
|
||||
"""Return True if *connection* originated from the loopback interface."""
|
||||
addr = getattr(connection, "remote_address", None)
|
||||
if not addr:
|
||||
return False
|
||||
host = addr[0] if isinstance(addr, tuple) else addr
|
||||
if not isinstance(host, str):
|
||||
return False
|
||||
# ``::ffff:127.0.0.1`` is loopback in IPv6-mapped form.
|
||||
if host.startswith("::ffff:"):
|
||||
host = host[7:]
|
||||
return host in _LOCALHOSTS
|
||||
|
||||
|
||||
def _http_response(
|
||||
body: bytes,
|
||||
*,
|
||||
status: int = 200,
|
||||
content_type: str = "text/plain; charset=utf-8",
|
||||
extra_headers: list[tuple[str, str]] | None = None,
|
||||
) -> Response:
|
||||
headers = [
|
||||
("Date", email.utils.formatdate(usegmt=True)),
|
||||
("Connection", "close"),
|
||||
("Content-Length", str(len(body))),
|
||||
("Content-Type", content_type),
|
||||
]
|
||||
if extra_headers:
|
||||
headers.extend(extra_headers)
|
||||
reason = http.HTTPStatus(status).phrase
|
||||
return Response(status, reason, Headers(headers), body)
|
||||
|
||||
|
||||
def _http_error(status: int, message: str | None = None) -> Response:
|
||||
body = (message or http.HTTPStatus(status).phrase).encode("utf-8")
|
||||
return _http_response(body, status=status)
|
||||
|
||||
|
||||
def _bearer_token(headers: Any) -> str | None:
|
||||
"""Pull a Bearer token out of standard or query-style headers."""
|
||||
auth = headers.get("Authorization") or headers.get("authorization")
|
||||
if auth and auth.lower().startswith("bearer "):
|
||||
return auth[7:].strip() or None
|
||||
return None
|
||||
|
||||
|
||||
def _is_websocket_upgrade(request: WsRequest) -> bool:
|
||||
"""Detect an actual WS upgrade; plain HTTP GETs to the same path should fall through."""
|
||||
upgrade = request.headers.get("Upgrade") or request.headers.get("upgrade")
|
||||
@@ -427,20 +263,6 @@ def _is_websocket_upgrade(request: WsRequest) -> bool:
|
||||
return True
|
||||
|
||||
|
||||
def _issue_route_secret_matches(headers: Any, configured_secret: str) -> bool:
|
||||
"""Return True if the token-issue HTTP request carries credentials matching ``token_issue_secret``."""
|
||||
if not configured_secret:
|
||||
return True
|
||||
authorization = headers.get("Authorization") or headers.get("authorization")
|
||||
if authorization and authorization.lower().startswith("bearer "):
|
||||
supplied = authorization[7:].strip()
|
||||
return hmac.compare_digest(supplied, configured_secret)
|
||||
header_token = headers.get("X-Nanobot-Auth") or headers.get("x-nanobot-auth")
|
||||
if not header_token:
|
||||
return False
|
||||
return hmac.compare_digest(header_token.strip(), configured_secret)
|
||||
|
||||
|
||||
class WebSocketChannel(BaseChannel):
|
||||
"""Run a local WebSocket server; forward text/JSON messages to the message bus."""
|
||||
|
||||
@@ -452,7 +274,7 @@ class WebSocketChannel(BaseChannel):
|
||||
config: Any,
|
||||
bus: MessageBus,
|
||||
*,
|
||||
http_handler: Any | None = None,
|
||||
gateway: GatewayServices,
|
||||
):
|
||||
if isinstance(config, dict):
|
||||
config = WebSocketConfig.model_validate(config)
|
||||
@@ -467,11 +289,11 @@ class WebSocketChannel(BaseChannel):
|
||||
self._stop_event: asyncio.Event | None = None
|
||||
self._server_task: asyncio.Task[None] | None = None
|
||||
|
||||
# HTTP handler injected from outside (ChannelManager / gateway startup).
|
||||
# Owns tokens, sessions, media, settings, static serving.
|
||||
self._http = http_handler
|
||||
# Backwards-compat: workspace controller used in envelope dispatch
|
||||
self._webui_workspaces = http_handler.workspaces if http_handler else None
|
||||
self.gateway = gateway
|
||||
self._http_router = gateway.http
|
||||
self._tokens = gateway.tokens
|
||||
self._media = gateway.media
|
||||
self._workspaces = gateway.workspaces
|
||||
|
||||
self._stream_text_buffers: dict[tuple[str, str], list[str]] = {}
|
||||
|
||||
@@ -501,9 +323,9 @@ class WebSocketChannel(BaseChannel):
|
||||
connected clients normally see it via ``goal_state`` / ``turn_end`` frames.
|
||||
Pushing here makes refresh + reconnect restore the strip without a new model turn.
|
||||
"""
|
||||
if self._http.session_manager is None:
|
||||
if self.gateway.session_manager is None:
|
||||
return
|
||||
row = self._http.session_manager.read_session_file(f"websocket:{chat_id}")
|
||||
row = self.gateway.session_manager.read_session_file(f"websocket:{chat_id}")
|
||||
meta = row.get("metadata", {}) if isinstance(row, dict) else {}
|
||||
if not isinstance(meta, dict):
|
||||
meta = {}
|
||||
@@ -543,57 +365,6 @@ class WebSocketChannel(BaseChannel):
|
||||
def _expected_path(self) -> str:
|
||||
return _normalize_config_path(self.config.path)
|
||||
|
||||
# -- Backwards-compat property aliases (used by tests) ------------------
|
||||
|
||||
@property
|
||||
def _session_manager(self):
|
||||
return self._http.session_manager
|
||||
|
||||
@_session_manager.setter
|
||||
def _session_manager(self, value):
|
||||
self._http.session_manager = value
|
||||
|
||||
@property
|
||||
def _media_secret(self):
|
||||
return self._http.media_secret
|
||||
|
||||
@property
|
||||
def _issued_tokens(self):
|
||||
return self._http.issued_tokens
|
||||
|
||||
@_issued_tokens.setter
|
||||
def _issued_tokens(self, value):
|
||||
self._http.issued_tokens = value
|
||||
|
||||
@property
|
||||
def _api_tokens(self):
|
||||
return self._http.api_tokens
|
||||
|
||||
def _check_api_token(self, request):
|
||||
return self._http.check_api_token(request)
|
||||
|
||||
def _sign_media_path(self, path):
|
||||
return self._http.sign_media_path(path)
|
||||
|
||||
def _sign_or_stage_media_path(self, path):
|
||||
return self._http.sign_or_stage_media_path(path)
|
||||
|
||||
def _rewrite_local_markdown_images(self, text):
|
||||
return self._http.rewrite_local_markdown_images(text)
|
||||
|
||||
def _handle_bootstrap(self, connection, request):
|
||||
return self._http._handle_bootstrap(connection, request)
|
||||
|
||||
def _handle_sessions_list(self, request):
|
||||
return self._http._handle_sessions_list(request)
|
||||
|
||||
def _handle_webui_thread_get(self, request, key):
|
||||
return self._http._handle_webui_thread_get(request, key)
|
||||
|
||||
@property
|
||||
def _workspace_path(self):
|
||||
return self._http.workspace_path
|
||||
|
||||
def _build_ssl_context(self) -> ssl.SSLContext | None:
|
||||
cert = self.config.ssl_certfile.strip()
|
||||
key = self.config.ssl_keyfile.strip()
|
||||
@@ -624,8 +395,8 @@ class WebSocketChannel(BaseChannel):
|
||||
return connection.respond(403, "Forbidden")
|
||||
return self._authorize_websocket_handshake(connection, query)
|
||||
|
||||
<< # Everything else goes to the HTTP handler
|
||||
return await self._http.dispatch(connection, request)
|
||||
# Everything else goes to the HTTP handler
|
||||
return await self._http_router.dispatch(connection, request)
|
||||
|
||||
def _authorize_websocket_handshake(self, connection: Any, query: dict[str, list[str]]) -> Any:
|
||||
supplied = _query_first(query, "token")
|
||||
@@ -634,17 +405,17 @@ class WebSocketChannel(BaseChannel):
|
||||
if static_token:
|
||||
if supplied and hmac.compare_digest(supplied, static_token):
|
||||
return None
|
||||
if supplied and self._http.take_issued_token_if_valid(supplied):
|
||||
if supplied and self._tokens.take_issued_token_if_valid(supplied):
|
||||
return None
|
||||
return connection.respond(401, "Unauthorized")
|
||||
|
||||
if self.config.websocket_requires_token:
|
||||
if supplied and self._http.take_issued_token_if_valid(supplied):
|
||||
if supplied and self._tokens.take_issued_token_if_valid(supplied):
|
||||
return None
|
||||
return connection.respond(401, "Unauthorized")
|
||||
|
||||
if supplied:
|
||||
self._http.take_issued_token_if_valid(supplied)
|
||||
self._tokens.take_issued_token_if_valid(supplied)
|
||||
return None
|
||||
|
||||
# -- Server lifecycle and connection ingress ---------------------------
|
||||
@@ -878,14 +649,14 @@ class WebSocketChannel(BaseChannel):
|
||||
new_id = str(uuid.uuid4())
|
||||
scope = await self._workspace_scope_or_error(
|
||||
connection,
|
||||
lambda: self._webui_workspaces.scope_for_new_chat(
|
||||
lambda: self._workspaces.scope_for_new_chat(
|
||||
envelope,
|
||||
controls_available=_is_localhost(connection),
|
||||
),
|
||||
)
|
||||
if scope is None:
|
||||
return
|
||||
self._webui_workspaces.persist_scope(new_id, scope)
|
||||
self._workspaces.persist_scope(new_id, scope)
|
||||
self._attach(connection, new_id)
|
||||
await self._send_event(connection, "attached", chat_id=new_id)
|
||||
await self._send_event(
|
||||
@@ -913,7 +684,7 @@ class WebSocketChannel(BaseChannel):
|
||||
return
|
||||
scope = await self._workspace_scope_or_error(
|
||||
connection,
|
||||
lambda: self._webui_workspaces.scope_for_set_request(
|
||||
lambda: self._workspaces.scope_for_set_request(
|
||||
envelope,
|
||||
chat_id=cid,
|
||||
chat_running=websocket_turn_wall_started_at(cid) is not None,
|
||||
@@ -923,7 +694,7 @@ class WebSocketChannel(BaseChannel):
|
||||
)
|
||||
if scope is None:
|
||||
return
|
||||
self._webui_workspaces.persist_scope(cid, scope)
|
||||
self._workspaces.persist_scope(cid, scope)
|
||||
await self._send_event(
|
||||
connection,
|
||||
"session_updated",
|
||||
@@ -965,7 +736,7 @@ class WebSocketChannel(BaseChannel):
|
||||
return
|
||||
scope = await self._workspace_scope_or_error(
|
||||
connection,
|
||||
lambda: self._webui_workspaces.scope_for_message(
|
||||
lambda: self._workspaces.scope_for_message(
|
||||
envelope,
|
||||
chat_id=cid,
|
||||
chat_running=websocket_turn_wall_started_at(cid) is not None,
|
||||
@@ -989,7 +760,7 @@ class WebSocketChannel(BaseChannel):
|
||||
if mcp_presets:
|
||||
metadata["mcp_presets"] = mcp_presets
|
||||
metadata[WORKSPACE_SCOPE_METADATA_KEY] = scope.metadata()
|
||||
self._webui_workspaces.persist_scope(cid, scope)
|
||||
self._workspaces.persist_scope(cid, scope)
|
||||
image_generation = envelope.get("image_generation")
|
||||
if isinstance(image_generation, dict) and image_generation.get("enabled") is True:
|
||||
aspect_ratio = image_generation.get("aspect_ratio")
|
||||
@@ -1044,8 +815,7 @@ class WebSocketChannel(BaseChannel):
|
||||
self._subs.clear()
|
||||
self._conn_chats.clear()
|
||||
self._conn_default.clear()
|
||||
self._http.issued_tokens.clear()
|
||||
self._http.api_tokens.clear()
|
||||
self._tokens.clear()
|
||||
|
||||
async def _safe_send_to(self, connection: Any, raw: str, *, label: str = "") -> None:
|
||||
"""Send a raw frame to one connection, cleaning up on ConnectionClosed."""
|
||||
@@ -1127,7 +897,7 @@ class WebSocketChannel(BaseChannel):
|
||||
)
|
||||
return
|
||||
text = msg.content
|
||||
wire_text = self._http.rewrite_local_markdown_images(text)
|
||||
wire_text = self._media.rewrite_local_markdown_images(text)
|
||||
payload: dict[str, Any] = {
|
||||
"event": "message",
|
||||
"chat_id": msg.chat_id,
|
||||
@@ -1137,7 +907,7 @@ class WebSocketChannel(BaseChannel):
|
||||
payload["media"] = msg.media
|
||||
urls: list[dict[str, str]] = []
|
||||
for entry in msg.media:
|
||||
signed = self._http.sign_or_stage_media_path(Path(entry))
|
||||
signed = self._media.sign_or_stage_media_path(Path(entry))
|
||||
if signed is not None:
|
||||
urls.append(signed)
|
||||
if urls:
|
||||
@@ -1252,7 +1022,7 @@ class WebSocketChannel(BaseChannel):
|
||||
if delta:
|
||||
buffered.append(delta)
|
||||
full_text = "".join(buffered)
|
||||
rewritten = self._http.rewrite_local_markdown_images(full_text)
|
||||
rewritten = self._media.rewrite_local_markdown_images(full_text)
|
||||
if rewritten != full_text:
|
||||
body["text"] = rewritten
|
||||
else:
|
||||
|
||||
Reference in New Issue
Block a user