feat(webui): add project workspaces and access controls (#4007)
* feat(webui): add project workspaces and access controls * feat(webui): add project workspaces and access controls * refactor(tools): centralize workspace access resolution * refactor(webui): remove unused workspace host state * fix(webui): hide estimated file edit label * fix(webui): clarify file edit deletion feedback * fix(webui): label deleted file activity * fix(webui): flatten file edit activity rows * fix(core): remove path-only patch deletion * fix(core): keep apply patch non-destructive * refactor(webui): trim workspace host plumbing * fix(tools): register exec with tools config
This commit is contained in:
+310
-31
@@ -18,19 +18,24 @@ import ssl
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Callable
|
||||
from contextlib import suppress
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, 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
|
||||
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.agent.tools.mcp import request_mcp_reload
|
||||
from nanobot.security.workspace_access import (
|
||||
WORKSPACE_SCOPE_METADATA_KEY,
|
||||
WorkspaceScopeError,
|
||||
)
|
||||
from nanobot.bus.events import OUTBOUND_META_AGENT_UI, OutboundMessage
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.channels.base import BaseChannel
|
||||
@@ -48,9 +53,15 @@ from nanobot.utils.subagent_channel_display import scrub_subagent_messages_for_c
|
||||
from nanobot.webui.settings_api import (
|
||||
WebUISettingsError,
|
||||
create_model_configuration,
|
||||
decorate_settings_payload,
|
||||
login_oauth_provider,
|
||||
logout_oauth_provider,
|
||||
runtime_capabilities,
|
||||
settings_payload,
|
||||
update_agent_settings,
|
||||
update_image_generation_settings,
|
||||
update_model_configuration,
|
||||
update_network_safety_settings,
|
||||
update_provider_settings,
|
||||
update_web_search_settings,
|
||||
)
|
||||
@@ -73,6 +84,9 @@ from nanobot.webui.transcript import (
|
||||
build_webui_thread_response,
|
||||
rewrite_local_markdown_images,
|
||||
)
|
||||
from nanobot.webui.workspaces import (
|
||||
WebUIWorkspaceController,
|
||||
)
|
||||
|
||||
_MCP_PRESET_ACTIONS_BY_PATH = {
|
||||
"/api/settings/mcp-presets/enable": "enable",
|
||||
@@ -100,6 +114,41 @@ 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}"
|
||||
|
||||
|
||||
class WebSocketConfig(Base):
|
||||
"""WebSocket server channel configuration.
|
||||
|
||||
@@ -123,6 +172,7 @@ class WebSocketConfig(Base):
|
||||
enabled: bool = False
|
||||
host: str = "127.0.0.1"
|
||||
port: int = 8765
|
||||
unix_socket_path: str = ""
|
||||
path: str = "/"
|
||||
token: str = ""
|
||||
token_issue_path: str = ""
|
||||
@@ -141,6 +191,19 @@ class WebSocketConfig(Base):
|
||||
ssl_certfile: str = ""
|
||||
ssl_keyfile: str = ""
|
||||
|
||||
@field_validator("unix_socket_path")
|
||||
@classmethod
|
||||
def unix_socket_path_format(cls, value: str) -> str:
|
||||
value = value.strip()
|
||||
if not value:
|
||||
return ""
|
||||
if "\x00" in value:
|
||||
raise ValueError("unix_socket_path must not contain NUL bytes")
|
||||
path = Path(value).expanduser()
|
||||
if not path.is_absolute():
|
||||
raise ValueError("unix_socket_path must be an absolute path")
|
||||
return str(path)
|
||||
|
||||
@field_validator("path")
|
||||
@classmethod
|
||||
def path_must_start_with_slash(cls, value: str) -> str:
|
||||
@@ -503,7 +566,10 @@ class WebSocketChannel(BaseChannel):
|
||||
session_manager: "SessionManager | None" = None,
|
||||
static_dist_path: Path | None = None,
|
||||
workspace_path: Path | None = None,
|
||||
restrict_to_workspace: bool = False,
|
||||
runtime_model_name: Callable[[], str | None] | None = None,
|
||||
runtime_surface: str = "browser",
|
||||
runtime_capabilities_overrides: dict[str, Any] | None = None,
|
||||
):
|
||||
if isinstance(config, dict):
|
||||
config = WebSocketConfig.model_validate(config)
|
||||
@@ -530,7 +596,20 @@ class WebSocketChannel(BaseChannel):
|
||||
if workspace_path is not None
|
||||
else get_workspace_path()
|
||||
).resolve(strict=False)
|
||||
self._default_restrict_to_workspace = restrict_to_workspace
|
||||
self._webui_workspaces = WebUIWorkspaceController(
|
||||
session_manager=self._session_manager,
|
||||
default_workspace=self._workspace_path,
|
||||
default_restrict_to_workspace=self._default_restrict_to_workspace,
|
||||
)
|
||||
self._runtime_model_name = runtime_model_name
|
||||
self._runtime_surface = (
|
||||
"native" if runtime_surface in {"native", "desktop"} else "browser"
|
||||
)
|
||||
self._runtime_capabilities = runtime_capabilities(
|
||||
self._runtime_surface,
|
||||
runtime_capabilities_overrides,
|
||||
)
|
||||
self._settings_restart_sections: set[str] = set()
|
||||
self._stream_text_buffers: dict[tuple[str, str], list[str]] = {}
|
||||
# Process-local secret used to HMAC-sign media URLs. The signed URL is
|
||||
@@ -695,6 +774,9 @@ class WebSocketChannel(BaseChannel):
|
||||
if got == "/api/commands":
|
||||
return self._handle_commands(request)
|
||||
|
||||
if got == "/api/workspaces":
|
||||
return self._handle_workspaces(connection, request)
|
||||
|
||||
if got == "/api/webui/sidebar-state":
|
||||
return self._handle_webui_sidebar_state(request)
|
||||
|
||||
@@ -707,15 +789,27 @@ class WebSocketChannel(BaseChannel):
|
||||
if got == "/api/settings/model-configurations/create":
|
||||
return self._handle_settings_model_configuration_create(request)
|
||||
|
||||
if got == "/api/settings/model-configurations/update":
|
||||
return self._handle_settings_model_configuration_update(request)
|
||||
|
||||
if got == "/api/settings/provider/update":
|
||||
return self._handle_settings_provider_update(request)
|
||||
|
||||
if got == "/api/settings/provider/oauth-login":
|
||||
return await self._handle_settings_provider_oauth(request, "login")
|
||||
|
||||
if got == "/api/settings/provider/oauth-logout":
|
||||
return await self._handle_settings_provider_oauth(request, "logout")
|
||||
|
||||
if got == "/api/settings/web-search/update":
|
||||
return self._handle_settings_web_search_update(request)
|
||||
|
||||
if got == "/api/settings/image-generation/update":
|
||||
return self._handle_settings_image_generation_update(request)
|
||||
|
||||
if got == "/api/settings/network-safety/update":
|
||||
return self._handle_settings_network_safety_update(request)
|
||||
|
||||
if got == "/api/settings/cli-apps":
|
||||
return self._handle_settings_cli_apps(request)
|
||||
|
||||
@@ -773,6 +867,12 @@ class WebSocketChannel(BaseChannel):
|
||||
return connection.respond(403, "Forbidden")
|
||||
return self._authorize_websocket_handshake(connection, query)
|
||||
|
||||
# API clients should never receive the SPA shell for an unknown route.
|
||||
# Returning HTML here makes the WebUI fail with "Unexpected token <"
|
||||
# when a dev server is pointed at an older gateway.
|
||||
if got.startswith("/api/"):
|
||||
return _http_error(404, "API route not found")
|
||||
|
||||
# 5. Static SPA serving (only if a build directory was wired in).
|
||||
if self._static_dist_path is not None:
|
||||
response = self._serve_static(got)
|
||||
@@ -832,15 +932,32 @@ class WebSocketChannel(BaseChannel):
|
||||
# while the REST surface keeps validating the other until TTL expiry.
|
||||
self._issued_tokens[token] = expiry
|
||||
self._api_tokens[token] = expiry
|
||||
ws_url = self._bootstrap_ws_url(request)
|
||||
return _http_json_response(
|
||||
{
|
||||
"token": token,
|
||||
"ws_path": self._expected_path(),
|
||||
"ws_url": ws_url,
|
||||
"expires_in": self.config.token_ttl_s,
|
||||
"model_name": _resolve_bootstrap_model_name(self._runtime_model_name),
|
||||
"runtime_surface": self._runtime_surface,
|
||||
"runtime_capabilities": self._runtime_capabilities,
|
||||
}
|
||||
)
|
||||
|
||||
def _bootstrap_ws_url(self, request: Any) -> str:
|
||||
"""Absolute WS URL clients should prefer over a dev-server proxy."""
|
||||
headers = getattr(request, "headers", {}) or {}
|
||||
host = _safe_host_header(_case_insensitive_header(headers, "Host"))
|
||||
if not host:
|
||||
host = _host_for_url(self.config.host, self.config.port)
|
||||
|
||||
proto = _case_insensitive_header(headers, "X-Forwarded-Proto")
|
||||
proto = proto.split(",", 1)[0].strip().lower()
|
||||
secure = proto in {"https", "wss"} or bool(self.config.ssl_certfile.strip())
|
||||
scheme = "wss" if secure else "ws"
|
||||
return f"{scheme}://{host}{self._expected_path()}"
|
||||
|
||||
def _handle_sessions_list(self, request: WsRequest) -> Response:
|
||||
if not self._check_api_token(request):
|
||||
return _http_error(401, "Unauthorized")
|
||||
@@ -859,13 +976,29 @@ class WebSocketChannel(BaseChannel):
|
||||
started_at = websocket_turn_wall_started_at(chat_id)
|
||||
if started_at is not None:
|
||||
row["run_started_at"] = started_at
|
||||
scope = self._webui_workspaces.scope_for_session_key(key)
|
||||
row["workspace_scope"] = scope.payload()
|
||||
cleaned.append(row)
|
||||
return _http_json_response({"sessions": cleaned})
|
||||
|
||||
def _handle_workspaces(self, connection: Any, request: WsRequest) -> Response:
|
||||
if not self._check_api_token(request):
|
||||
return _http_error(401, "Unauthorized")
|
||||
return _http_json_response(
|
||||
self._webui_workspaces.payload(controls_available=_is_localhost(connection))
|
||||
)
|
||||
|
||||
def _handle_settings(self, request: WsRequest) -> Response:
|
||||
if not self._check_api_token(request):
|
||||
return _http_error(401, "Unauthorized")
|
||||
return _http_json_response(self._with_settings_restart_state(settings_payload()))
|
||||
return _http_json_response(
|
||||
self._with_settings_restart_state(
|
||||
settings_payload(
|
||||
surface=self._runtime_surface,
|
||||
runtime_capability_overrides=self._runtime_capabilities,
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
def _with_settings_restart_state(
|
||||
self,
|
||||
@@ -876,14 +1009,16 @@ class WebSocketChannel(BaseChannel):
|
||||
"""Keep restart-required state alive for this gateway process."""
|
||||
if section and payload.get("requires_restart"):
|
||||
self._settings_restart_sections.add(section)
|
||||
if self._settings_restart_sections:
|
||||
payload = dict(payload)
|
||||
sections = sorted(self._settings_restart_sections)
|
||||
payload = dict(payload)
|
||||
if sections:
|
||||
payload["requires_restart"] = True
|
||||
payload["restart_required_sections"] = sorted(self._settings_restart_sections)
|
||||
else:
|
||||
payload = dict(payload)
|
||||
payload["restart_required_sections"] = []
|
||||
return payload
|
||||
return decorate_settings_payload(
|
||||
payload,
|
||||
surface=self._runtime_surface,
|
||||
runtime_capability_overrides=self._runtime_capabilities,
|
||||
restart_required_sections=sections,
|
||||
)
|
||||
|
||||
def _handle_commands(self, request: WsRequest) -> Response:
|
||||
if not self._check_api_token(request):
|
||||
@@ -939,6 +1074,16 @@ class WebSocketChannel(BaseChannel):
|
||||
return _http_error(e.status, e.message)
|
||||
return _http_json_response(self._with_settings_restart_state(payload))
|
||||
|
||||
def _handle_settings_model_configuration_update(self, request: WsRequest) -> Response:
|
||||
if not self._check_api_token(request):
|
||||
return _http_error(401, "Unauthorized")
|
||||
query = _parse_query(request.path)
|
||||
try:
|
||||
payload = update_model_configuration(query)
|
||||
except WebUISettingsError as e:
|
||||
return _http_error(e.status, e.message)
|
||||
return _http_json_response(self._with_settings_restart_state(payload))
|
||||
|
||||
def _handle_settings_provider_update(self, request: WsRequest) -> Response:
|
||||
if not self._check_api_token(request):
|
||||
return _http_error(401, "Unauthorized")
|
||||
@@ -949,6 +1094,19 @@ class WebSocketChannel(BaseChannel):
|
||||
return _http_error(e.status, e.message)
|
||||
return _http_json_response(self._with_settings_restart_state(payload, section="image"))
|
||||
|
||||
async def _handle_settings_provider_oauth(self, request: WsRequest, action: str) -> Response:
|
||||
if not self._check_api_token(request):
|
||||
return _http_error(401, "Unauthorized")
|
||||
query = _parse_query(request.path)
|
||||
try:
|
||||
if action == "login":
|
||||
payload = await asyncio.to_thread(login_oauth_provider, query)
|
||||
else:
|
||||
payload = await asyncio.to_thread(logout_oauth_provider, query)
|
||||
except WebUISettingsError as e:
|
||||
return _http_error(e.status, e.message)
|
||||
return _http_json_response(self._with_settings_restart_state(payload))
|
||||
|
||||
def _handle_settings_web_search_update(self, request: WsRequest) -> Response:
|
||||
if not self._check_api_token(request):
|
||||
return _http_error(401, "Unauthorized")
|
||||
@@ -957,7 +1115,7 @@ class WebSocketChannel(BaseChannel):
|
||||
payload = update_web_search_settings(query)
|
||||
except WebUISettingsError as e:
|
||||
return _http_error(e.status, e.message)
|
||||
return _http_json_response(self._with_settings_restart_state(payload, section="web"))
|
||||
return _http_json_response(self._with_settings_restart_state(payload, section="browser"))
|
||||
|
||||
def _handle_settings_image_generation_update(self, request: WsRequest) -> Response:
|
||||
if not self._check_api_token(request):
|
||||
@@ -969,6 +1127,16 @@ class WebSocketChannel(BaseChannel):
|
||||
return _http_error(e.status, e.message)
|
||||
return _http_json_response(self._with_settings_restart_state(payload, section="image"))
|
||||
|
||||
def _handle_settings_network_safety_update(self, request: WsRequest) -> Response:
|
||||
if not self._check_api_token(request):
|
||||
return _http_error(401, "Unauthorized")
|
||||
query = _parse_query(request.path)
|
||||
try:
|
||||
payload = update_network_safety_settings(query)
|
||||
except WebUISettingsError as e:
|
||||
return _http_error(e.status, e.message)
|
||||
return _http_json_response(self._with_settings_restart_state(payload, section="runtime"))
|
||||
|
||||
def _handle_settings_cli_apps(self, request: WsRequest) -> Response:
|
||||
if not self._check_api_token(request):
|
||||
return _http_error(401, "Unauthorized")
|
||||
@@ -1058,13 +1226,19 @@ class WebSocketChannel(BaseChannel):
|
||||
return _http_error(400, "invalid session key")
|
||||
if not self._is_websocket_channel_session_key(decoded_key):
|
||||
return _http_error(404, "session not found")
|
||||
scope = self._webui_workspaces.scope_for_session_key(decoded_key)
|
||||
data = build_webui_thread_response(
|
||||
decoded_key,
|
||||
augment_user_media=self._augment_transcript_user_media,
|
||||
augment_assistant_text=self._rewrite_local_markdown_images,
|
||||
augment_assistant_text=lambda text: rewrite_local_markdown_images(
|
||||
text,
|
||||
workspace_path=scope.project_path,
|
||||
sign_path=self._sign_or_stage_media_path,
|
||||
),
|
||||
)
|
||||
if data is None:
|
||||
return _http_error(404, "webui thread not found")
|
||||
data["workspace_scope"] = scope.payload()
|
||||
return _http_json_response(data)
|
||||
|
||||
def _try_append_webui_transcript(self, chat_id: str, wire: dict[str, Any]) -> None:
|
||||
@@ -1359,34 +1533,63 @@ class WebSocketChannel(BaseChannel):
|
||||
await self._connection_loop(connection)
|
||||
|
||||
self.logger.info(
|
||||
"WebSocket server listening on {}://{}:{}{}",
|
||||
scheme,
|
||||
self.config.host,
|
||||
self.config.port,
|
||||
self.config.path,
|
||||
"WebSocket server listening on {}",
|
||||
(
|
||||
f"unix:{self.config.unix_socket_path}{self.config.path}"
|
||||
if self.config.unix_socket_path
|
||||
else f"{scheme}://{self.config.host}:{self.config.port}{self.config.path}"
|
||||
),
|
||||
)
|
||||
if self.config.token_issue_path:
|
||||
self.logger.info(
|
||||
"WebSocket token issue route: {}://{}:{}{}",
|
||||
scheme,
|
||||
self.config.host,
|
||||
self.config.port,
|
||||
_normalize_config_path(self.config.token_issue_path),
|
||||
"WebSocket token issue route: {}",
|
||||
(
|
||||
f"unix:{self.config.unix_socket_path}{_normalize_config_path(self.config.token_issue_path)}"
|
||||
if self.config.unix_socket_path
|
||||
else (
|
||||
f"{scheme}://{self.config.host}:{self.config.port}"
|
||||
f"{_normalize_config_path(self.config.token_issue_path)}"
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
async def runner() -> None:
|
||||
async with serve(
|
||||
handler,
|
||||
self.config.host,
|
||||
self.config.port,
|
||||
process_request=process_request,
|
||||
max_size=self.config.max_message_bytes,
|
||||
ping_interval=self.config.ping_interval_s,
|
||||
ping_timeout=self.config.ping_timeout_s,
|
||||
ssl=ssl_context,
|
||||
):
|
||||
socket_path = self.config.unix_socket_path
|
||||
if socket_path:
|
||||
path_obj = Path(socket_path)
|
||||
path_obj.parent.mkdir(parents=True, exist_ok=True)
|
||||
with suppress(FileNotFoundError):
|
||||
path_obj.unlink()
|
||||
server = await unix_serve(
|
||||
handler,
|
||||
socket_path,
|
||||
process_request=process_request,
|
||||
max_size=self.config.max_message_bytes,
|
||||
ping_interval=self.config.ping_interval_s,
|
||||
ping_timeout=self.config.ping_timeout_s,
|
||||
)
|
||||
with suppress(OSError):
|
||||
path_obj.chmod(0o600)
|
||||
else:
|
||||
server = await serve(
|
||||
handler,
|
||||
self.config.host,
|
||||
self.config.port,
|
||||
process_request=process_request,
|
||||
max_size=self.config.max_message_bytes,
|
||||
ping_interval=self.config.ping_interval_s,
|
||||
ping_timeout=self.config.ping_timeout_s,
|
||||
ssl=ssl_context,
|
||||
)
|
||||
try:
|
||||
assert self._stop_event is not None
|
||||
await self._stop_event.wait()
|
||||
finally:
|
||||
server.close()
|
||||
await server.wait_closed()
|
||||
if socket_path:
|
||||
with suppress(FileNotFoundError):
|
||||
Path(socket_path).unlink()
|
||||
|
||||
self._server_task = asyncio.create_task(runner())
|
||||
await self._server_task
|
||||
@@ -1530,8 +1733,25 @@ class WebSocketChannel(BaseChannel):
|
||||
t = envelope.get("type")
|
||||
if t == "new_chat":
|
||||
new_id = str(uuid.uuid4())
|
||||
scope = await self._workspace_scope_or_error(
|
||||
connection,
|
||||
lambda: self._webui_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._attach(connection, new_id)
|
||||
await self._send_event(connection, "attached", chat_id=new_id)
|
||||
await self._send_event(
|
||||
connection,
|
||||
"session_updated",
|
||||
chat_id=new_id,
|
||||
scope="metadata",
|
||||
workspace_scope=scope.payload(),
|
||||
)
|
||||
await self._hydrate_after_subscribe(new_id)
|
||||
return
|
||||
if t == "attach":
|
||||
@@ -1543,6 +1763,32 @@ class WebSocketChannel(BaseChannel):
|
||||
await self._send_event(connection, "attached", chat_id=cid)
|
||||
await self._hydrate_after_subscribe(cid)
|
||||
return
|
||||
if t == "set_workspace_scope":
|
||||
cid = envelope.get("chat_id")
|
||||
if not _is_valid_chat_id(cid):
|
||||
await self._send_event(connection, "error", detail="invalid chat_id")
|
||||
return
|
||||
scope = await self._workspace_scope_or_error(
|
||||
connection,
|
||||
lambda: self._webui_workspaces.scope_for_set_request(
|
||||
envelope,
|
||||
chat_id=cid,
|
||||
chat_running=websocket_turn_wall_started_at(cid) is not None,
|
||||
controls_available=_is_localhost(connection),
|
||||
),
|
||||
chat_id=cid,
|
||||
)
|
||||
if scope is None:
|
||||
return
|
||||
self._webui_workspaces.persist_scope(cid, scope)
|
||||
await self._send_event(
|
||||
connection,
|
||||
"session_updated",
|
||||
chat_id=cid,
|
||||
scope="metadata",
|
||||
workspace_scope=scope.payload(),
|
||||
)
|
||||
return
|
||||
if t == "message":
|
||||
cid = envelope.get("chat_id")
|
||||
content = envelope.get("content")
|
||||
@@ -1574,6 +1820,18 @@ class WebSocketChannel(BaseChannel):
|
||||
if not content.strip() and not media_paths:
|
||||
await self._send_event(connection, "error", detail="missing content")
|
||||
return
|
||||
scope = await self._workspace_scope_or_error(
|
||||
connection,
|
||||
lambda: self._webui_workspaces.scope_for_message(
|
||||
envelope,
|
||||
chat_id=cid,
|
||||
chat_running=websocket_turn_wall_started_at(cid) is not None,
|
||||
controls_available=_is_localhost(connection),
|
||||
),
|
||||
chat_id=cid,
|
||||
)
|
||||
if scope is None:
|
||||
return
|
||||
|
||||
# Auto-attach on first use so clients can one-shot without a separate attach.
|
||||
self._attach(connection, cid)
|
||||
@@ -1587,6 +1845,8 @@ class WebSocketChannel(BaseChannel):
|
||||
mcp_presets = normalize_mcp_preset_mentions(envelope.get("mcp_presets"))
|
||||
if mcp_presets:
|
||||
metadata["mcp_presets"] = mcp_presets
|
||||
metadata[WORKSPACE_SCOPE_METADATA_KEY] = scope.metadata()
|
||||
self._webui_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")
|
||||
@@ -1605,6 +1865,25 @@ class WebSocketChannel(BaseChannel):
|
||||
return
|
||||
await self._send_event(connection, "error", detail=f"unknown type: {t!r}")
|
||||
|
||||
async def _workspace_scope_or_error(
|
||||
self,
|
||||
connection: Any,
|
||||
resolver: Callable[[], Any],
|
||||
*,
|
||||
chat_id: str | None = None,
|
||||
) -> Any | None:
|
||||
try:
|
||||
return resolver()
|
||||
except WorkspaceScopeError as exc:
|
||||
await self._send_event(
|
||||
connection,
|
||||
"error",
|
||||
detail="workspace_scope_rejected",
|
||||
reason=exc.message,
|
||||
**({"chat_id": chat_id} if chat_id else {}),
|
||||
)
|
||||
return None
|
||||
|
||||
async def stop(self) -> None:
|
||||
if not self._running:
|
||||
return
|
||||
|
||||
Reference in New Issue
Block a user