From 7d5f4023a9c40171d2c2fc547cd98a6e67ad4a93 Mon Sep 17 00:00:00 2001 From: KiriAky 107 Date: Thu, 3 Sep 2026 15:25:41 +0800 Subject: [PATCH] feat(mcp): complete remote transports and configuration workflow --- README.md | 2 +- backend/app/contracts.py | 59 +- backend/app/extensions/mcp.py | 727 +++++++++++++++++- backend/app/extensions/mcp_registry.py | 361 +++++++-- backend/app/routes.py | 160 +++- backend/tests/test_mcp_registry.py | 288 ++++++- .../AI笔记软件技术栈说明-团队版-v2.3.md | 2 +- docs/contracts/第二阶段接口契约-开发版.md | 35 +- .../AI-Core与Agent-Core开发说明.md | 2 +- .../独立MCP-Server配置中心开发说明.md | 77 +- frontend/src/contracts/index.ts | 17 +- .../src/features/mcp/McpServersView.spec.ts | 84 ++ frontend/src/features/mcp/McpServersView.vue | 191 +++-- frontend/src/services/mcpServerService.ts | 7 +- 14 files changed, 1747 insertions(+), 265 deletions(-) create mode 100644 frontend/src/features/mcp/McpServersView.spec.ts diff --git a/README.md b/README.md index 6f2c586..3a96a48 100644 --- a/README.md +++ b/README.md @@ -2,7 +2,7 @@ > 本文件用于团队开发期间快速配置环境和启动项目,不是正式的项目 README。 -> 当前基线:2026-09-03。第一阶段 Web 联调前后端已经完成;第二阶段已完成 Workspace 去 Mock、Agent Trace 持久化与 SSE 恢复、stdio MCP Bridge、隔离 Plugin Host、Plugin Command/Settings,以及独立 MCP Server 配置中心 C.1/P0。Streamable HTTP MCP、真实音频、Provider 协议增强、Benchmark、导出、主题包、Trace 可视化、Mermaid 与函数图像仍在后续开发;Tauri Host、Stronghold、原生多 Vault 文件系统和 Sync Server 尚未接入。 +> 当前基线:2026-09-03。第一阶段 Web 联调前后端已经完成;第二阶段已完成 Workspace 去 Mock、Agent Trace 持久化与 SSE 恢复、stdio MCP Bridge、隔离 Plugin Host、Plugin Command/Settings,以及独立 MCP Server 配置中心 C.1(stdio、Streamable HTTP 与旧 SSE 兼容)。真实音频、Provider 协议增强、Benchmark、导出、主题包、Trace 可视化、Mermaid 与函数图像仍在后续开发;Tauri Host、Stronghold、原生多 Vault 文件系统和 Sync Server 尚未接入。 ## 当前目录 diff --git a/backend/app/contracts.py b/backend/app/contracts.py index ee8f3ae..9840f45 100644 --- a/backend/app/contracts.py +++ b/backend/app/contracts.py @@ -494,22 +494,29 @@ class McpServerTransport(str, Enum): sse = "sse" -class McpServerCreateRequest(Contract): +class McpServerConfig(Contract): name: str = Field(min_length=1, max_length=80) transport: McpServerTransport = McpServerTransport.stdio - command: str = Field(min_length=1, max_length=1024) + command: str | None = Field(default=None, max_length=1024) args: list[str] = Field(default_factory=list, max_length=64) + url: str | None = Field(default=None, max_length=4096) + headers: dict[str, str] = Field(default_factory=dict) environment: dict[str, str] = Field(default_factory=dict) secret_environment_keys: list[str] = Field(default_factory=list) + secret_header_keys: list[str] = Field(default_factory=list) permissions: list[str] = Field(default_factory=list) startup_timeout_seconds: float = Field(default=15, ge=1, le=120) tool_timeout_seconds: float = Field(default=30, ge=1, le=300) -class McpServerUpdateRequest(McpServerCreateRequest): +class McpServerCreateRequest(McpServerConfig): pass +class McpServerUpdateRequest(McpServerConfig): + version: int = Field(ge=1) + + class McpServerSecretWriteRequest(Contract): secret: SecretStr = Field(min_length=1, max_length=32768) @@ -523,21 +530,8 @@ class McpServerTrustRequest(Contract): command_digest: str = Field(min_length=64, max_length=64) -class McpServer(Contract): - server_id: str - name: str - transport: McpServerTransport - command: str - args: list[str] = Field(default_factory=list) - environment: dict[str, str] = Field(default_factory=dict) - secret_environment: dict[str, bool] = Field(default_factory=dict) - permissions: list[str] = Field(default_factory=list) - startup_timeout_seconds: float - tool_timeout_seconds: float +class McpServerStatus(Contract): enabled: bool = False - trusted: bool = False - command_digest: str - command_summary: str status: PluginHostState = PluginHostState.stopped tools_count: int = 0 protocol_version: str | None = None @@ -548,10 +542,41 @@ class McpServer(Contract): last_test_succeeded: bool | None = None +class McpServer(McpServerStatus): + server_id: str + version: int + name: str + transport: McpServerTransport + command: str | None = None + args: list[str] = Field(default_factory=list) + url: str | None = None + headers: dict[str, str] = Field(default_factory=dict) + environment: dict[str, str] = Field(default_factory=dict) + permissions: list[str] = Field(default_factory=list) + startup_timeout_seconds: float + tool_timeout_seconds: float + secret_environment: dict[str, bool] = Field(default_factory=dict) + secret_headers: dict[str, bool] = Field(default_factory=dict) + trusted: bool = False + command_digest: str + command_summary: str + + class McpServerListResponse(Contract): items: list[McpServer] = Field(default_factory=list) +class McpToolSummary(Contract): + name: str + remote_name: str + description: str + permission: str | None = None + + +class McpToolSummaryListResponse(Contract): + items: list[McpToolSummary] = Field(default_factory=list) + + class PluginCommandLocation(str, Enum): command_palette = "command_palette" context_menu = "context_menu" diff --git a/backend/app/extensions/mcp.py b/backend/app/extensions/mcp.py index 834ef65..61c2d5f 100644 --- a/backend/app/extensions/mcp.py +++ b/backend/app/extensions/mcp.py @@ -10,14 +10,18 @@ import asyncio import json import os import queue +import signal import subprocess import threading from collections import deque +from collections.abc import Callable from dataclasses import dataclass -from datetime import datetime, timezone +from datetime import UTC, datetime from pathlib import Path -from typing import Any, Callable +from typing import Any, Protocol +from urllib.parse import urljoin, urlsplit +import httpx from jsonschema import Draft202012Validator from jsonschema.exceptions import SchemaError @@ -99,7 +103,12 @@ class McpStdioClient: return # TODO(extension-security): 社区 Plugin 开放前迁移到 Tauri/Rust Host 的 # 平台级沙箱启动器;uvx 只隔离 Python 依赖,不能替代系统权限限制。 - creation_flags = getattr(subprocess, "CREATE_NO_WINDOW", 0) if os.name == "nt" else 0 + creation_flags = ( + getattr(subprocess, "CREATE_NO_WINDOW", 0) + | getattr(subprocess, "CREATE_NEW_PROCESS_GROUP", 0) + if os.name == "nt" + else 0 + ) environment = _subprocess_environment() environment.update(self.environment) environment.setdefault("PYTHONUNBUFFERED", "1") @@ -117,6 +126,7 @@ class McpStdioClient: shell=False, env=environment, creationflags=creation_flags, + start_new_session=os.name != "nt", ) except OSError as exc: raise McpBridgeError( @@ -183,7 +193,9 @@ class McpStdioClient: except queue.Empty as exc: self.cancel(request_id, "Request timed out.") self.abandon(request_id) - raise McpBridgeError(timeout_code, "MCP request timed out.", status_code=504) from exc + raise McpBridgeError( + timeout_code, "MCP request timed out.", status_code=504 + ) from exc if isinstance(response, BaseException): raise response if "error" in response: @@ -216,9 +228,7 @@ class McpStdioClient: except McpBridgeError: pass - def abandon( - self, request_id: int, wake_error: BaseException | None = None - ) -> None: + def abandon(self, request_id: int, wake_error: BaseException | None = None) -> None: with self._pending_lock: pending = self._pending.pop(request_id, None) # asyncio.to_thread 被取消时不会停止底层线程;主动唤醒 Queue,避免线程 @@ -243,15 +253,17 @@ class McpStdioClient: try: process.wait(timeout=2) except subprocess.TimeoutExpired: - process.terminate() + _terminate_process_tree(process) try: process.wait(timeout=2) except subprocess.TimeoutExpired: - process.kill() + _kill_process_tree(process) process.wait(timeout=2) finally: self._fail_pending( - McpBridgeError("PLUGIN_HOST_UNAVAILABLE", "MCP host stopped.", status_code=503) + McpBridgeError( + "PLUGIN_HOST_UNAVAILABLE", "MCP host stopped.", status_code=503 + ) ) self.process = None @@ -313,14 +325,17 @@ class McpStdioClient: { "jsonrpc": "2.0", "id": message["id"], - "error": {"code": -32601, "message": "Method not supported."}, + "error": { + "code": -32601, + "message": "Method not supported.", + }, } ) except (McpBridgeError, OSError, ValueError) as exc: failure = f"MCP stdout closed unexpectedly: {type(exc).__name__}." finally: if failure and process.poll() is None: - process.terminate() + _terminate_process_tree(process) exit_code = process.poll() if exit_code is None: try: @@ -328,7 +343,9 @@ class McpStdioClient: except subprocess.TimeoutExpired: exit_code = None if not self._stopping: - message = failure or f"MCP host exited unexpectedly with code {exit_code}." + message = ( + failure or f"MCP host exited unexpectedly with code {exit_code}." + ) error = McpBridgeError( "PLUGIN_HOST_UNAVAILABLE", message, status_code=503 ) @@ -363,13 +380,472 @@ class McpStdioClient: item.response.put(error) +class McpHttpClient: + """MCP Streamable HTTP client supporting JSON and SSE POST responses.""" + + def __init__( + self, + url: str, + *, + headers: dict[str, str], + on_seen: Callable[[], None], + on_broken: Callable[[str], None], + on_tools_changed: Callable[[], None], + ) -> None: + self.url = url + self.headers = headers + self.on_seen = on_seen + self.on_broken = on_broken + self.on_tools_changed = on_tools_changed + self._client = httpx.Client(follow_redirects=False, timeout=30) + self._pending_lock = threading.Lock() + self._pending: dict[int, _PendingRequest] = {} + self._next_id = 1 + self._session_id: str | None = None + self._protocol_version: str | None = None + self._stopping = False + self._stream_started = False + self._last_event_id: str | None = None + self._stop_event = threading.Event() + + def start(self) -> None: + return + + def set_protocol_version(self, version: str) -> None: + self._protocol_version = version + + def start_event_stream(self) -> None: + if self._stream_started: + return + self._stream_started = True + threading.Thread(target=self._event_stream_loop, daemon=True).start() + + def request( + self, + method: str, + params: dict[str, Any], + *, + timeout: float, + timeout_code: str, + response_error_code: str = "MCP_TOOL_CALL_FAILED", + ) -> dict[str, Any]: + request_id, pending = self.begin_request(method, params) + return self.wait_response( + request_id, + pending, + timeout=timeout, + timeout_code=timeout_code, + response_error_code=response_error_code, + ) + + def begin_request( + self, method: str, params: dict[str, Any] + ) -> tuple[int, _PendingRequest]: + with self._pending_lock: + request_id = self._next_id + self._next_id += 1 + pending = _PendingRequest(response=queue.Queue(maxsize=1)) + self._pending[request_id] = pending + message = { + "jsonrpc": "2.0", + "id": request_id, + "method": method, + "params": params, + } + threading.Thread( + target=self._dispatch_request, + args=(request_id, message), + daemon=True, + ).start() + return request_id, pending + + def wait_response( + self, + request_id: int, + pending: _PendingRequest, + *, + timeout: float, + timeout_code: str, + response_error_code: str = "MCP_TOOL_CALL_FAILED", + ) -> dict[str, Any]: + try: + response = pending.response.get(timeout=timeout) + except queue.Empty as exc: + self.cancel(request_id, "Request timed out.") + self.abandon(request_id) + raise McpBridgeError( + timeout_code, "MCP request timed out.", status_code=504 + ) from exc + if isinstance(response, BaseException): + raise response + if "error" in response: + error = response.get("error") + message = ( + str(error.get("message", "MCP JSON-RPC error.")) + if isinstance(error, dict) + else "MCP JSON-RPC error." + ) + raise McpBridgeError(response_error_code, message) + result = response.get("result") + if not isinstance(result, dict): + raise McpBridgeError( + response_error_code, "MCP response result must be an object." + ) + return result + + def notify(self, method: str, params: dict[str, Any] | None = None) -> None: + message: dict[str, Any] = {"jsonrpc": "2.0", "method": method} + if params is not None: + message["params"] = params + self._post_notification(message) + + def cancel(self, request_id: int, reason: str = "Cancelled by host.") -> None: + def send() -> None: + try: + self.notify( + "notifications/cancelled", + {"requestId": request_id, "reason": reason}, + ) + except McpBridgeError: + pass + + threading.Thread(target=send, daemon=True).start() + + def abandon(self, request_id: int, wake_error: BaseException | None = None) -> None: + with self._pending_lock: + pending = self._pending.pop(request_id, None) + if pending is not None and wake_error is not None: + try: + pending.response.put_nowait(wake_error) + except queue.Full: + pass + + def stop(self) -> None: + self._stopping = True + self._stop_event.set() + if self._session_id: + try: + request = self._client.build_request( + "DELETE", self.url, headers=self._request_headers() + ) + response = self._client.send(request, stream=True) + response.close() + except httpx.HTTPError: + pass + self._client.close() + self._fail_pending( + McpBridgeError( + "PLUGIN_HOST_UNAVAILABLE", "MCP HTTP client stopped.", status_code=503 + ) + ) + + def _dispatch_request(self, request_id: int, message: dict[str, Any]) -> None: + try: + response = self._post(message, timeout=None) + try: + self._capture_session(response) + content_type = response.headers.get("content-type", "").lower() + if response.status_code >= 400: + raise McpBridgeError( + "MCP_HTTP_REQUEST_FAILED", + f"MCP HTTP server returned status {response.status_code}.", + status_code=502, + ) + if "application/json" in content_type: + payload = _bounded_json_response(response) + self._deliver(payload) + elif "text/event-stream" in content_type: + delivered = False + for _event, _event_id, data in _iter_sse(response): + payload = _json_rpc_message(data) + self._handle_message(payload) + if payload.get("id") == request_id: + delivered = True + break + if not delivered: + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", + "MCP SSE response ended before the matching JSON-RPC response.", + ) + else: + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", + "MCP HTTP response has an unsupported Content-Type.", + ) + finally: + response.close() + except (McpBridgeError, httpx.HTTPError) as exc: + error = ( + exc + if isinstance(exc, McpBridgeError) + else McpBridgeError( + "MCP_HTTP_REQUEST_FAILED", + f"MCP HTTP request failed: {type(exc).__name__}.", + status_code=503, + ) + ) + self.abandon(request_id, error) + + def _post_notification(self, message: dict[str, Any]) -> None: + try: + response = self._post(message, timeout=10) + except httpx.HTTPError as exc: + raise McpBridgeError( + "MCP_HTTP_REQUEST_FAILED", + f"MCP HTTP notification failed: {type(exc).__name__}.", + status_code=503, + ) from exc + try: + self._capture_session(response) + if response.status_code not in {200, 202, 204}: + raise McpBridgeError( + "MCP_HTTP_REQUEST_FAILED", + f"MCP HTTP server rejected a notification with status {response.status_code}.", + ) + finally: + response.close() + + def _post( + self, message: dict[str, Any], *, timeout: float | None + ) -> httpx.Response: + encoded = json.dumps(message, ensure_ascii=False, separators=(",", ":")) + if len(encoded.encode("utf-8")) > MAX_MCP_MESSAGE_BYTES: + raise McpBridgeError("MCP_TOOL_CALL_FAILED", "MCP request is too large.") + request = self._client.build_request( + "POST", + self.url, + content=encoded.encode("utf-8"), + headers=self._request_headers(), + ) + return self._client.send(request, stream=True) + + def _request_headers(self) -> dict[str, str]: + headers = { + **self.headers, + "Accept": "application/json, text/event-stream", + "Content-Type": "application/json", + } + if self._session_id: + headers["MCP-Session-Id"] = self._session_id + if self._protocol_version: + headers["MCP-Protocol-Version"] = self._protocol_version + return headers + + def _capture_session(self, response: httpx.Response) -> None: + session_id = response.headers.get("mcp-session-id") + if session_id is not None: + if ( + not session_id.isascii() + or not session_id.isprintable() + or len(session_id) > 1024 + ): + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", "MCP session id is invalid." + ) + self._session_id = session_id + + def _handle_message(self, message: dict[str, Any]) -> None: + self.on_seen() + if "id" in message and ("result" in message or "error" in message): + self._deliver(message) + elif message.get("method") == "notifications/tools/list_changed": + self.on_tools_changed() + + def _deliver(self, message: dict[str, Any]) -> None: + request_id = message.get("id") + if not isinstance(request_id, int): + return + with self._pending_lock: + pending = self._pending.pop(request_id, None) + if pending: + pending.response.put(message) + + def _fail_pending(self, error: BaseException) -> None: + with self._pending_lock: + pending = list(self._pending.values()) + self._pending.clear() + for item in pending: + item.response.put(error) + + def _event_stream_loop(self) -> None: + while not self._stop_event.is_set(): + headers = {**self._request_headers(), "Accept": "text/event-stream"} + headers.pop("Content-Type", None) + if self._last_event_id: + headers["Last-Event-ID"] = self._last_event_id + try: + with self._client.stream( + "GET", self.url, headers=headers, timeout=None + ) as response: + if response.status_code == 405: + return + if response.status_code >= 400: + self.on_broken( + f"MCP HTTP event stream returned status {response.status_code}." + ) + return + if ( + "text/event-stream" + not in response.headers.get("content-type", "").lower() + ): + self.on_broken("MCP HTTP GET response is not an event stream.") + return + self._capture_session(response) + for _event, event_id, data in _iter_sse(response): + if event_id: + self._last_event_id = event_id + self._handle_message(_json_rpc_message(data)) + if self._stop_event.is_set(): + return + except (McpBridgeError, httpx.HTTPError): + if self._stopping: + return + self._stop_event.wait(0.25) + + +class McpLegacySseClient(McpHttpClient): + """Compatibility client for the deprecated 2024-11-05 HTTP+SSE transport.""" + + def __init__(self, *args: Any, **kwargs: Any) -> None: + super().__init__(*args, **kwargs) + self._endpoint: str | None = None + self._endpoint_ready: queue.Queue[str | BaseException] = queue.Queue(maxsize=1) + + def start(self) -> None: + threading.Thread(target=self._event_loop, daemon=True).start() + try: + endpoint = self._endpoint_ready.get(timeout=15) + except queue.Empty as exc: + raise McpBridgeError( + "MCP_INITIALIZE_FAILED", + "Legacy MCP SSE endpoint event timed out.", + status_code=504, + ) from exc + if isinstance(endpoint, BaseException): + raise endpoint + self._endpoint = endpoint + + def start_event_stream(self) -> None: + """The legacy client already owns its single GET event stream.""" + + return + + def _dispatch_request(self, request_id: int, message: dict[str, Any]) -> None: + try: + response = self._post(message, timeout=10) + try: + if response.status_code not in {200, 202, 204}: + raise McpBridgeError( + "MCP_HTTP_REQUEST_FAILED", + f"Legacy MCP endpoint returned status {response.status_code}.", + ) + finally: + response.close() + except (McpBridgeError, httpx.HTTPError) as exc: + error = ( + exc + if isinstance(exc, McpBridgeError) + else McpBridgeError( + "MCP_HTTP_REQUEST_FAILED", + f"Legacy MCP request failed: {type(exc).__name__}.", + status_code=503, + ) + ) + self.abandon(request_id, error) + + def _post( + self, message: dict[str, Any], *, timeout: float | None + ) -> httpx.Response: + if self._endpoint is None: + raise McpBridgeError( + "MCP_INITIALIZE_FAILED", "Legacy MCP endpoint is not ready." + ) + encoded = json.dumps(message, ensure_ascii=False, separators=(",", ":")) + request = self._client.build_request( + "POST", + self._endpoint, + content=encoded.encode("utf-8"), + headers={ + **self.headers, + "Accept": "application/json, text/event-stream", + "Content-Type": "application/json", + }, + ) + return self._client.send(request, stream=True) + + def _event_loop(self) -> None: + try: + with self._client.stream( + "GET", + self.url, + headers={**self.headers, "Accept": "text/event-stream"}, + timeout=None, + ) as response: + if response.status_code >= 400: + raise McpBridgeError( + "MCP_HTTP_REQUEST_FAILED", + f"Legacy MCP SSE server returned status {response.status_code}.", + ) + if ( + "text/event-stream" + not in response.headers.get("content-type", "").lower() + ): + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", + "Legacy MCP GET response is not an event stream.", + ) + for event, _event_id, data in _iter_sse(response): + if self._endpoint is None and event == "endpoint": + endpoint = _legacy_endpoint_url(self.url, data) + self._endpoint_ready.put(endpoint) + self._endpoint = endpoint + continue + self._handle_message(_json_rpc_message(data)) + except (McpBridgeError, httpx.HTTPError) as exc: + if self._endpoint is None: + self._endpoint_ready.put(exc) + elif not self._stopping: + self.on_broken(f"Legacy MCP SSE stream failed: {type(exc).__name__}.") + + @dataclass(slots=True) class _McpHost: backend: PluginBackend - client: McpStdioClient + client: _McpClient status: PluginHostStatus +class _McpClient(Protocol): + def start(self) -> None: ... + def request( + self, + method: str, + params: dict[str, Any], + *, + timeout: float, + timeout_code: str, + response_error_code: str = "MCP_TOOL_CALL_FAILED", + ) -> dict[str, Any]: ... + def begin_request( + self, method: str, params: dict[str, Any] + ) -> tuple[int, _PendingRequest]: ... + def wait_response( + self, + request_id: int, + pending: _PendingRequest, + *, + timeout: float, + timeout_code: str, + response_error_code: str = "MCP_TOOL_CALL_FAILED", + ) -> dict[str, Any]: ... + def notify(self, method: str, params: dict[str, Any] | None = None) -> None: ... + def cancel(self, request_id: int, reason: str = "Cancelled by host.") -> None: ... + def abandon( + self, request_id: int, wake_error: BaseException | None = None + ) -> None: ... + def stop(self) -> None: ... + + class McpBridge: """管理每个 Plugin 的独立 MCP Client,并执行 Contract 转换。""" @@ -390,19 +866,27 @@ class McpBridge: command_override: list[str] | None = None, environment: dict[str, str] | None = None, tool_source: str = "plugin", + transport_kind: str | None = None, + url: str | None = None, + headers: dict[str, str] | None = None, ) -> list[McpDiscoveredTool]: - if backend.transport != "stdio": + transport = transport_kind or backend.transport + if transport not in {"stdio", "streamable_http", "sse"}: raise McpBridgeError( "MCP_CAPABILITY_UNSUPPORTED", - "Phase C only supports the MCP stdio transport.", + f"Unsupported MCP transport: {transport}", status_code=501, ) - command = command_override or self._resolve_command(package_path, backend) - now = datetime.now(timezone.utc) + command = ( + command_override or self._resolve_command(package_path, backend) + if transport == "stdio" + else None + ) + now = datetime.now(UTC) status = PluginHostStatus( plugin_id=plugin_id, backend_type="mcp", - transport="stdio", + transport="stdio" if transport == "stdio" else "http", status=PluginHostState.starting, started_at=now, last_seen_at=now, @@ -412,7 +896,7 @@ class McpBridge: def seen() -> None: host = host_ref.get("host") if host: - host.status.last_seen_at = datetime.now(timezone.utc) + host.status.last_seen_at = datetime.now(UTC) def broken(message: str) -> None: host = host_ref.get("host") @@ -422,16 +906,35 @@ class McpBridge: on_unavailable(plugin_id, message) def tools_changed() -> None: - broken("MCP tool list changed; restart the Plugin Host to revalidate tools.") + broken( + "MCP tool list changed; restart the Plugin Host to revalidate tools." + ) - client = McpStdioClient( - command, - cwd=package_path, - environment=environment, - on_seen=seen, - on_broken=broken, - on_tools_changed=tools_changed, - ) + if transport == "stdio": + assert command is not None + client: _McpClient = McpStdioClient( + command, + cwd=package_path, + environment=environment, + on_seen=seen, + on_broken=broken, + on_tools_changed=tools_changed, + ) + else: + if not url: + raise McpBridgeError( + "MCP_HOST_START_FAILED", "MCP HTTP transport requires a URL." + ) + client_type = ( + McpHttpClient if transport == "streamable_http" else McpLegacySseClient + ) + client = client_type( + url, + headers=headers or {}, + on_seen=seen, + on_broken=broken, + on_tools_changed=tools_changed, + ) host = _McpHost(backend=backend, client=client, status=status) host_ref["host"] = host with self._lock: @@ -474,15 +977,21 @@ class McpBridge: if not isinstance(server_info, dict): server_info = {} status.protocol_version = str(version) + set_protocol_version = getattr(client, "set_protocol_version", None) + if callable(set_protocol_version): + set_protocol_version(str(version)) status.server_name = _optional_string(server_info.get("name")) status.server_version = _optional_string(server_info.get("version")) client.notify("notifications/initialized") + start_event_stream = getattr(client, "start_event_stream", None) + if callable(start_event_stream): + start_event_stream() discovered = self._discover_tools( plugin_id, client, backend, declared_permissions, tool_source ) status.status = PluginHostState.ready status.tools_count = len(discovered) - status.last_seen_at = datetime.now(timezone.utc) + status.last_seen_at = datetime.now(UTC) status.error = None return discovered except McpBridgeError as exc: @@ -527,9 +1036,7 @@ class McpBridge: host.client.cancel(rpc_id) host.client.abandon( rpc_id, - McpBridgeError( - "MCP_TOOL_CALL_FAILED", "MCP request was cancelled." - ), + McpBridgeError("MCP_TOOL_CALL_FAILED", "MCP request was cancelled."), ) raise except McpBridgeError as exc: @@ -539,7 +1046,9 @@ class McpBridge: self._calls.pop(call_key, None) encoded_size = len( - json.dumps(result, ensure_ascii=False, separators=(",", ":")).encode("utf-8") + json.dumps(result, ensure_ascii=False, separators=(",", ":")).encode( + "utf-8" + ) ) if encoded_size > MAX_MCP_TOOL_RESULT_BYTES: raise ToolExecutionError( @@ -606,7 +1115,7 @@ class McpBridge: def _discover_tools( self, plugin_id: str, - client: McpStdioClient, + client: _McpClient, backend: PluginBackend, declared_permissions: list[str], tool_source: str, @@ -625,7 +1134,8 @@ class McpBridge: raw_tools = result.get("tools") if not isinstance(raw_tools, list): raise McpBridgeError( - "MCP_TOOL_SCHEMA_INVALID", "MCP tools/list must return a tools array." + "MCP_TOOL_SCHEMA_INVALID", + "MCP tools/list must return a tools array.", ) for raw in raw_tools: discovered.append( @@ -641,7 +1151,8 @@ class McpBridge: break if not isinstance(next_cursor, str) or not next_cursor: raise McpBridgeError( - "MCP_TOOL_SCHEMA_INVALID", "MCP nextCursor must be a non-empty string." + "MCP_TOOL_SCHEMA_INVALID", + "MCP nextCursor must be a non-empty string.", ) cursor = next_cursor else: @@ -675,9 +1186,7 @@ class McpBridge: len(remote_name) > 128 or not remote_name[0].isalnum() or not all( - character.islower() - or character.isdigit() - or character in "._-" + character.islower() or character.isdigit() or character in "._-" for character in remote_name ) ): @@ -702,7 +1211,9 @@ class McpBridge: ) from exc metadata = raw.get("_meta") permission = ( - metadata.get("notesagent/permission") if isinstance(metadata, dict) else None + metadata.get("notesagent/permission") + if isinstance(metadata, dict) + else None ) if permission is not None and ( not isinstance(permission, str) or permission not in KNOWN_PERMISSIONS @@ -721,7 +1232,9 @@ class McpBridge: remote_name=remote_name, definition=ToolDefinition( name=f"{plugin_id}.{remote_name}", - description=description if isinstance(description, str) else remote_name, + description=description + if isinstance(description, str) + else remote_name, parameters=schema, permission=permission, source=tool_source, @@ -801,3 +1314,133 @@ def _subprocess_environment() -> dict[str, str]: environment["PYTHONUNBUFFERED"] = "1" environment["PYTHONIOENCODING"] = "utf-8" return environment + + +def _bounded_json_response(response: httpx.Response) -> dict[str, Any]: + content_length = response.headers.get("content-length") + if ( + content_length + and content_length.isdigit() + and int(content_length) > MAX_MCP_MESSAGE_BYTES + ): + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", "MCP HTTP response is too large." + ) + chunks: list[bytes] = [] + size = 0 + for chunk in response.iter_bytes(): + size += len(chunk) + if size > MAX_MCP_MESSAGE_BYTES: + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", "MCP HTTP response is too large." + ) + chunks.append(chunk) + try: + payload = json.loads(b"".join(chunks)) + except (UnicodeDecodeError, json.JSONDecodeError) as exc: + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", "MCP HTTP response is not valid JSON." + ) from exc + if not isinstance(payload, dict) or payload.get("jsonrpc") != "2.0": + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", "MCP HTTP response is not a JSON-RPC message." + ) + return payload + + +def _iter_sse(response: httpx.Response): + event = "message" + event_id: str | None = None + data_lines: list[str] = [] + size = 0 + for line in response.iter_lines(): + size += len(line.encode("utf-8")) + 1 + if size > MAX_MCP_MESSAGE_BYTES: + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", "MCP SSE event is too large." + ) + if line == "": + if data_lines: + yield event, event_id, "\n".join(data_lines) + event, event_id, data_lines, size = "message", None, [], 0 + continue + if line.startswith(":"): + continue + field, _, value = line.partition(":") + value = value.removeprefix(" ") + if field == "event": + event = value + elif field == "id" and "\x00" not in value: + event_id = value + elif field == "data": + data_lines.append(value) + if data_lines: + yield event, event_id, "\n".join(data_lines) + + +def _json_rpc_message(data: str) -> dict[str, Any]: + try: + message = json.loads(data) + except json.JSONDecodeError as exc: + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", "MCP SSE data is not valid JSON." + ) from exc + if not isinstance(message, dict) or message.get("jsonrpc") != "2.0": + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", "MCP SSE data is not a JSON-RPC message." + ) + return message + + +def _legacy_endpoint_url(source_url: str, endpoint: str) -> str: + target = urljoin(source_url, endpoint.strip()) + source_parts = urlsplit(source_url) + target_parts = urlsplit(target) + if ( + target_parts.scheme not in {"http", "https"} + or target_parts.username is not None + or target_parts.password is not None + or (source_parts.scheme, source_parts.hostname, source_parts.port) + != (target_parts.scheme, target_parts.hostname, target_parts.port) + ): + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", + "Legacy MCP endpoint must use the same origin as the configured SSE URL.", + ) + return target + + +def _terminate_process_tree(process: subprocess.Popen[str]) -> None: + if process.poll() is not None: + return + try: + if os.name == "nt": + subprocess.run( + ["taskkill.exe", "/PID", str(process.pid), "/T"], + check=False, + capture_output=True, + creationflags=getattr(subprocess, "CREATE_NO_WINDOW", 0), + timeout=2, + ) + else: + os.killpg(process.pid, signal.SIGTERM) + except (OSError, subprocess.SubprocessError): + process.terminate() + + +def _kill_process_tree(process: subprocess.Popen[str]) -> None: + if process.poll() is not None: + return + try: + if os.name == "nt": + subprocess.run( + ["taskkill.exe", "/PID", str(process.pid), "/T", "/F"], + check=False, + capture_output=True, + creationflags=getattr(subprocess, "CREATE_NO_WINDOW", 0), + timeout=2, + ) + else: + os.killpg(process.pid, signal.SIGKILL) + except (OSError, subprocess.SubprocessError): + process.kill() diff --git a/backend/app/extensions/mcp_registry.py b/backend/app/extensions/mcp_registry.py index e576f7b..6808769 100644 --- a/backend/app/extensions/mcp_registry.py +++ b/backend/app/extensions/mcp_registry.py @@ -7,8 +7,10 @@ import json import re import threading from datetime import UTC, datetime +from functools import wraps from pathlib import Path from typing import Any +from urllib.parse import urlsplit from uuid import uuid4 from pydantic import BaseModel, ConfigDict, create_model @@ -21,6 +23,7 @@ from app.contracts import ( McpServerSecretStatus, McpServerTransport, McpServerUpdateRequest, + McpToolSummary, PluginBackend, PluginHostState, ) @@ -28,6 +31,15 @@ from app.extensions.mcp import McpBridge, McpBridgeError, McpDiscoveredTool from app.providers.credentials import CredentialStoreError, EncryptedCredentialStore _ENVIRONMENT_KEY = re.compile(r"^[A-Za-z_][A-Za-z0-9_]{0,127}$") +_HEADER_KEY = re.compile(r"^[!#$%&'*+.^_`|~0-9A-Za-z-]{1,128}$") +_RESERVED_HEADERS = { + "accept", + "content-length", + "content-type", + "host", + "mcp-protocol-version", + "mcp-session-id", +} class McpRegistryError(RuntimeError): @@ -38,6 +50,17 @@ class McpRegistryError(RuntimeError): self.status_code = status_code +def _serialized_lifecycle(method): + """Serialize lifecycle mutations without blocking MCP failure callbacks.""" + + @wraps(method) + def wrapped(self, *args, **kwargs): + with self._lifecycle_lock: + return method(self, *args, **kwargs) + + return wrapped + + class McpServerRegistry: """Persists configuration and owns stdio host/tool lifecycles.""" @@ -56,8 +79,10 @@ class McpServerRegistry: self.allow_process_launch = allow_process_launch self.bridge = bridge or McpBridge() self._lock = threading.RLock() + self._lifecycle_lock = threading.RLock() self._records = self._read() self._registered: dict[str, list[str]] = {} + self._summaries: dict[str, list[McpToolSummary]] = {} self._last_status: dict[str, dict[str, Any]] = {} def list(self) -> list[McpServer]: @@ -71,57 +96,108 @@ class McpServerRegistry: with self._lock: return self._public(server_id, self._record(server_id)) + def list_tools(self, server_id: str) -> list[McpToolSummary]: + self._record(server_id) + return [ + item.model_copy(deep=True) for item in self._summaries.get(server_id, []) + ] + + @_serialized_lifecycle def create(self, request: McpServerCreateRequest) -> McpServer: self._validate(request) server_id = uuid4().hex[:12] record = request.model_dump(mode="json") record["name"] = request.name.strip() - record["command"] = request.command.strip() - record.update(enabled=False, approved_digest=None) + record["command"] = request.command.strip() if request.command else None + record["url"] = request.url.strip() if request.url else None + record.update( + version=1, + enabled=False, + approved_digest=None, + tested_digest=None, + last_tested_at=None, + last_test_succeeded=None, + ) with self._lock: updated = {**self._records, server_id: record} self._write(updated) self._records = updated return self.get(server_id) + @_serialized_lifecycle def update(self, server_id: str, request: McpServerUpdateRequest) -> McpServer: self._validate(request) + current = self._record(server_id) + if request.version != current.get("version", 1): + raise McpRegistryError( + "MCP_SERVER_VERSION_CONFLICT", + "MCP server configuration version is stale.", + status_code=409, + ) self.disable(server_id) with self._lock: previous = self._record(server_id) - removed = set(previous.get("secret_environment_keys", [])) - set( - request.secret_environment_keys - ) - record = request.model_dump(mode="json") + removed = [ + (kind, key) + for kind, old_keys, new_keys in ( + ( + "environment", + previous.get("secret_environment_keys", []), + request.secret_environment_keys, + ), + ( + "header", + previous.get("secret_header_keys", []), + request.secret_header_keys, + ), + ) + for key in set(old_keys) - set(new_keys) + ] + record = request.model_dump(mode="json", exclude={"version"}) record["name"] = request.name.strip() - record["command"] = request.command.strip() - record.update(enabled=False, approved_digest=None) + record["command"] = request.command.strip() if request.command else None + record["url"] = request.url.strip() if request.url else None + record.update( + version=request.version + 1, + enabled=False, + approved_digest=None, + tested_digest=None, + last_tested_at=None, + last_test_succeeded=None, + ) updated = {**self._records, server_id: record} self._write(updated) self._records = updated self._last_status.pop(server_id, None) - for key in removed: + self._summaries.pop(server_id, None) + for kind, key in removed: try: - self.credentials.delete(self._secret_id(server_id, key)) + self.credentials.delete(self._secret_id(server_id, key, kind)) except CredentialStoreError as exc: raise McpRegistryError( "MCP_SECRET_STORE_ERROR", str(exc), status_code=500 ) from exc return self.get(server_id) + @_serialized_lifecycle def delete(self, server_id: str) -> None: self.disable(server_id) with self._lock: record = self._record(server_id) secret_ids = [ - self._secret_id(server_id, key) - for key in record.get("secret_environment_keys", []) + self._secret_id(server_id, key, kind) + for kind, keys in ( + ("environment", record.get("secret_environment_keys", [])), + ("header", record.get("secret_header_keys", [])), + ) + for key in keys ] updated = dict(self._records) del updated[server_id] self._write(updated) self._records = updated self._last_status.pop(server_id, None) + self._summaries.pop(server_id, None) try: self.credentials.delete_many(secret_ids) except CredentialStoreError as exc: @@ -130,6 +206,7 @@ class McpServerRegistry: ) from exc self.bridge.remove(self._host_id(server_id)) + @_serialized_lifecycle def trust(self, server_id: str, command_digest: str) -> McpServer: with self._lock: record = self._record(server_id) @@ -146,40 +223,46 @@ class McpServerRegistry: self._records = updated return self.get(server_id) + @_serialized_lifecycle def put_secret( - self, server_id: str, key: str, secret: str + self, server_id: str, key: str, secret: str, *, kind: str = "environment" ) -> McpServerSecretStatus: with self._lock: record = self._record(server_id) - self._validate_environment_key(key) - if key not in record.get("secret_environment_keys", []): + declared = self._secret_keys(record, kind) + self._validate_secret_key(key, kind) + if key not in declared: raise McpRegistryError( "MCP_SECRET_NOT_DECLARED", "Secret environment key is not declared in this server configuration.", ) try: - self.credentials.put(self._secret_id(server_id, key), secret) + self.credentials.put(self._secret_id(server_id, key, kind), secret) except CredentialStoreError as exc: raise McpRegistryError( "MCP_SECRET_STORE_ERROR", str(exc), status_code=500 ) from exc return McpServerSecretStatus(key=key, configured=True) - def delete_secret(self, server_id: str, key: str) -> McpServerSecretStatus: + @_serialized_lifecycle + def delete_secret( + self, server_id: str, key: str, *, kind: str = "environment" + ) -> McpServerSecretStatus: record = self._record(server_id) - if key not in record.get("secret_environment_keys", []): + if key not in self._secret_keys(record, kind): raise McpRegistryError( "MCP_SECRET_NOT_DECLARED", "Secret environment key is not declared in this server configuration.", ) try: - self.credentials.delete(self._secret_id(server_id, key)) + self.credentials.delete(self._secret_id(server_id, key, kind)) except CredentialStoreError as exc: raise McpRegistryError( "MCP_SECRET_STORE_ERROR", str(exc), status_code=500 ) from exc return McpServerSecretStatus(key=key, configured=False) + @_serialized_lifecycle def test(self, server_id: str) -> McpServer: record = self._record(server_id) if record.get("enabled"): @@ -188,18 +271,31 @@ class McpServerRegistry: "Disable the MCP server before running an isolated connection test.", status_code=409, ) - self._require_launch_allowed(record) + self._require_launch_allowed(record, require_test=False) try: discovered = self._start(server_id, record) except Exception as exc: - self._last_status[server_id] = { + tested_at = datetime.now(UTC) + failure = { "status": PluginHostState.error, "error": str(exc), - "last_tested_at": datetime.now(UTC), + "last_tested_at": tested_at, "last_test_succeeded": False, } + self._last_status[server_id] = failure + with self._lock: + failed_record = { + **record, + "tested_digest": None, + "last_tested_at": tested_at.isoformat(), + "last_test_succeeded": False, + } + updated = {**self._records, server_id: failed_record} + self._write(updated) + self._records = updated raise status = self.bridge.status(self._host_id(server_id), self._backend(record)) + tested_at = datetime.now(UTC) self._last_status[server_id] = { "status": PluginHostState.stopped, "tools_count": len(discovered), @@ -207,18 +303,31 @@ class McpServerRegistry: "remote_server_name": status.server_name, "remote_server_version": status.server_version, "error": None, - "last_tested_at": datetime.now(UTC), + "last_tested_at": tested_at, "last_test_succeeded": True, } + self._summaries[server_id] = self._tool_summaries(discovered) self.bridge.stop(self._host_id(server_id)) + with self._lock: + tested_record = { + **record, + "tested_digest": self._digest(record), + "last_tested_at": tested_at.isoformat(), + "last_test_succeeded": True, + } + updated = {**self._records, server_id: tested_record} + self._write(updated) + self._records = updated return self.get(server_id) + @_serialized_lifecycle def enable(self, server_id: str) -> McpServer: record = self._record(server_id) if server_id in self._registered: return self.get(server_id) - self._require_launch_allowed(record) + self._require_launch_allowed(record, require_test=True) discovered = self._start(server_id, record) + self._summaries[server_id] = self._tool_summaries(discovered) registered: list[str] = [] try: for item in discovered: @@ -243,6 +352,7 @@ class McpServerRegistry: raise return self.get(server_id) + @_serialized_lifecycle def disable(self, server_id: str) -> McpServer: with self._lock: record = self._record(server_id) @@ -255,6 +365,7 @@ class McpServerRegistry: self.bridge.stop(self._host_id(server_id)) return self.get(server_id) + @_serialized_lifecycle def restore_enabled(self) -> None: if not self._records: return @@ -270,6 +381,7 @@ class McpServerRegistry: } self._write() + @_serialized_lifecycle def shutdown(self) -> None: for server_id in list(self._records): for name in self._registered.pop(server_id, []): @@ -280,7 +392,9 @@ class McpServerRegistry: environment = dict(record.get("environment", {})) for key in record.get("secret_environment_keys", []): try: - value = self.credentials.resolve(self._secret_id(server_id, key)) + value = self.credentials.resolve( + self._secret_id(server_id, key, "environment") + ) except CredentialStoreError as exc: raise McpRegistryError( "MCP_SECRET_STORE_ERROR", str(exc), status_code=500 @@ -292,6 +406,23 @@ class McpServerRegistry: status_code=409, ) environment[key] = value + headers = dict(record.get("headers", {})) + for key in record.get("secret_header_keys", []): + try: + value = self.credentials.resolve( + self._secret_id(server_id, key, "header") + ) + except CredentialStoreError as exc: + raise McpRegistryError( + "MCP_SECRET_STORE_ERROR", str(exc), status_code=500 + ) from exc + if value is None: + raise McpRegistryError( + "MCP_SECRET_REQUIRED", + f"Secret HTTP header is not configured: {key}", + status_code=409, + ) + headers[key] = value host_id = self._host_id(server_id) self.bridge.remove(host_id) try: @@ -301,9 +432,16 @@ class McpServerRegistry: self._server_dir(server_id), list(record.get("permissions", [])), lambda _host, message: self._unavailable(server_id, message), - command_override=[record["command"], *record.get("args", [])], + command_override=( + [record["command"], *record.get("args", [])] + if record.get("command") + else None + ), environment=environment, tool_source="mcp_server", + transport_kind=record["transport"], + url=record.get("url"), + headers=headers, ) except McpBridgeError as exc: raise McpRegistryError( @@ -339,23 +477,29 @@ class McpServerRegistry: } self._write() - def _require_launch_allowed(self, record: dict[str, Any]) -> None: - if record.get("transport") != McpServerTransport.stdio.value: - raise McpRegistryError( - "MCP_TRANSPORT_UNSUPPORTED", - "C.1 currently supports stdio; Streamable HTTP and SSE are reserved for a later increment.", - status_code=501, - ) - if not self.allow_process_launch: + def _require_launch_allowed( + self, record: dict[str, Any], *, require_test: bool + ) -> None: + digest = self._digest(record) + if ( + record.get("transport") == McpServerTransport.stdio.value + and not self.allow_process_launch + ): raise McpRegistryError( "MCP_SANDBOX_REQUIRED", "Python process launch is disabled outside development until the desktop sandbox is available.", status_code=403, ) - if record.get("approved_digest") != self._digest(record): + if record.get("approved_digest") != digest: raise McpRegistryError( "MCP_TRUST_APPROVAL_REQUIRED", - "Review and approve the current MCP command before testing or enabling it.", + "Review and approve the current MCP connection before testing or enabling it.", + status_code=409, + ) + if require_test and record.get("tested_digest") != digest: + raise McpRegistryError( + "MCP_CONNECTION_TEST_REQUIRED", + "Test the current MCP configuration successfully before enabling it.", status_code=409, ) @@ -366,15 +510,22 @@ class McpServerRegistry: cached = self._last_status.get(server_id, {}) return McpServer( server_id=server_id, + version=record.get("version", 1), name=record["name"], transport=record["transport"], - command=record["command"], + command=record.get("command"), args=list(record.get("args", [])), + url=record.get("url"), + headers=dict(record.get("headers", {})), environment=dict(record.get("environment", {})), secret_environment={ key: self._secret_configured(server_id, key) for key in record.get("secret_environment_keys", []) }, + secret_headers={ + key: self._secret_configured(server_id, key, "header") + for key in record.get("secret_header_keys", []) + }, permissions=list(record.get("permissions", [])), startup_timeout_seconds=backend.startup_timeout_seconds, tool_timeout_seconds=backend.tool_timeout_seconds, @@ -398,8 +549,10 @@ class McpServerRegistry: if record.get("enabled") else cached.get("remote_server_version"), error=status.error if record.get("enabled") else cached.get("error"), - last_tested_at=cached.get("last_tested_at"), - last_test_succeeded=cached.get("last_test_succeeded"), + last_tested_at=record.get("last_tested_at") or cached.get("last_tested_at"), + last_test_succeeded=record.get("last_test_succeeded") + if record.get("last_test_succeeded") is not None + else cached.get("last_test_succeeded"), ) def _validate(self, request: McpServerCreateRequest) -> None: @@ -407,12 +560,36 @@ class McpServerRegistry: raise McpRegistryError( "MCP_SERVER_NAME_INVALID", "MCP server name cannot be blank." ) - if not request.command.strip() or "\x00" in request.command: - raise McpRegistryError("MCP_COMMAND_INVALID", "MCP executable is invalid.") - if any("\x00" in arg for arg in request.args): - raise McpRegistryError( - "MCP_COMMAND_INVALID", "MCP argument contains a null byte." - ) + if request.transport == McpServerTransport.stdio: + if ( + not request.command + or not request.command.strip() + or "\x00" in request.command + ): + raise McpRegistryError( + "MCP_COMMAND_INVALID", "MCP executable is invalid." + ) + if request.url or request.headers or request.secret_header_keys: + raise McpRegistryError( + "MCP_CONFIG_INVALID", + "stdio configuration cannot contain HTTP fields.", + ) + if any("\x00" in arg for arg in request.args): + raise McpRegistryError( + "MCP_COMMAND_INVALID", "MCP argument contains a null byte." + ) + else: + self._validate_http_url(request.url) + if ( + request.command + or request.args + or request.environment + or request.secret_environment_keys + ): + raise McpRegistryError( + "MCP_CONFIG_INVALID", + "HTTP configuration cannot contain stdio fields.", + ) for key in [*request.environment, *request.secret_environment_keys]: self._validate_environment_key(key) if set(request.environment) & set(request.secret_environment_keys): @@ -420,6 +597,21 @@ class McpServerRegistry: "MCP_ENVIRONMENT_INVALID", "An environment key cannot be both plain and secret.", ) + plain_headers = {key.casefold() for key in request.headers} + secret_headers = {key.casefold() for key in request.secret_header_keys} + for key in [*request.headers, *request.secret_header_keys]: + self._validate_header_key(key) + if any( + "\r" in value or "\n" in value or "\x00" in value + for value in request.headers.values() + ): + raise McpRegistryError( + "MCP_HEADER_INVALID", "HTTP header value contains control characters." + ) + if plain_headers & secret_headers: + raise McpRegistryError( + "MCP_HEADER_INVALID", "An HTTP header cannot be both plain and secret." + ) unknown_permissions = set(request.permissions) - KNOWN_PERMISSIONS if unknown_permissions: raise McpRegistryError( @@ -434,12 +626,36 @@ class McpServerRegistry: "MCP_ENVIRONMENT_INVALID", f"Invalid environment variable name: {key}" ) + @staticmethod + def _validate_header_key(key: str) -> None: + if not _HEADER_KEY.fullmatch(key) or key.casefold() in _RESERVED_HEADERS: + raise McpRegistryError( + "MCP_HEADER_INVALID", f"Invalid or reserved HTTP header: {key}" + ) + + @staticmethod + def _validate_http_url(url: str | None) -> None: + if not url: + raise McpRegistryError("MCP_URL_INVALID", "MCP HTTP URL is required.") + parts = urlsplit(url.strip()) + if ( + parts.scheme not in {"http", "https"} + or not parts.hostname + or parts.username is not None + or parts.password is not None + or parts.fragment + ): + raise McpRegistryError( + "MCP_URL_INVALID", + "MCP URL must be an HTTP(S) URL without credentials or fragments.", + ) + @staticmethod def _backend(record: dict[str, Any]) -> PluginBackend: return PluginBackend( type="mcp", transport="stdio", - command=record["command"], + command=record.get("command") or "http", args=record.get("args", []), startup_timeout_seconds=record.get("startup_timeout_seconds", 15), tool_timeout_seconds=record.get("tool_timeout_seconds", 30), @@ -464,6 +680,9 @@ class McpServerRegistry: "args", "environment", "secret_environment_keys", + "url", + "headers", + "secret_header_keys", "permissions", ) } @@ -475,9 +694,16 @@ class McpServerRegistry: @staticmethod def _summary(record: dict[str, Any]) -> str: + if record.get("transport") != McpServerTransport.stdio.value: + header_names = sorted( + [*record.get("headers", {}), *record.get("secret_header_keys", [])], + key=str.casefold, + ) + suffix = f" headers={','.join(header_names)}" if header_names else "" + return f"{record.get('transport')} {record.get('url') or ''}{suffix}" return " ".join( [ - record["command"], + record.get("command") or "", *[ json.dumps(arg, ensure_ascii=False) for arg in record.get("args", []) @@ -486,18 +712,51 @@ class McpServerRegistry: ) @staticmethod - def _secret_id(server_id: str, key: str) -> str: - suffix = hashlib.sha256(key.encode()).hexdigest()[:20] + def _secret_id(server_id: str, key: str, kind: str = "environment") -> str: + suffix = hashlib.sha256(f"{kind}\0{key.casefold()}".encode()).hexdigest()[:20] return f"mcp.{server_id}.{suffix}" - def _secret_configured(self, server_id: str, key: str) -> bool: + def _secret_configured( + self, server_id: str, key: str, kind: str = "environment" + ) -> bool: try: - return self.credentials.has(self._secret_id(server_id, key)) + return self.credentials.has(self._secret_id(server_id, key, kind)) except CredentialStoreError as exc: raise McpRegistryError( "MCP_SECRET_STORE_ERROR", str(exc), status_code=500 ) from exc + @staticmethod + def _secret_keys(record: dict[str, Any], kind: str) -> list[str]: + if kind == "environment": + return list(record.get("secret_environment_keys", [])) + if kind == "header": + return list(record.get("secret_header_keys", [])) + raise McpRegistryError("MCP_SECRET_KIND_INVALID", "Unknown MCP secret kind.") + + @staticmethod + def _validate_secret_key(key: str, kind: str) -> None: + if kind == "environment": + McpServerRegistry._validate_environment_key(key) + elif kind == "header": + McpServerRegistry._validate_header_key(key) + else: + raise McpRegistryError( + "MCP_SECRET_KIND_INVALID", "Unknown MCP secret kind." + ) + + @staticmethod + def _tool_summaries(discovered: list[McpDiscoveredTool]) -> list[McpToolSummary]: + return [ + McpToolSummary( + name=item.definition.name, + remote_name=item.remote_name, + description=item.definition.description, + permission=item.definition.permission, + ) + for item in discovered + ] + def _record(self, server_id: str) -> dict[str, Any]: try: return self._records[server_id] diff --git a/backend/app/routes.py b/backend/app/routes.py index 438d41b..a1bd49a 100644 --- a/backend/app/routes.py +++ b/backend/app/routes.py @@ -6,6 +6,8 @@ from uuid import uuid4 from fastapi import APIRouter, Header, Query from fastapi.responses import StreamingResponse +from app.agent import AgentCapacityError, AgentRunNotFoundError +from app.container import container from app.contracts import ( AgentRun, AgentRunCreateRequest, @@ -28,6 +30,7 @@ from app.contracts import ( McpServerSecretWriteRequest, McpServerTrustRequest, McpServerUpdateRequest, + McpToolSummaryListResponse, ModelEvent, ModelEventType, Note, @@ -75,18 +78,16 @@ from app.contracts import ( WorkspaceOpenRequest, WorkspaceSnapshot, ) -from app.agent import AgentCapacityError, AgentRunNotFoundError -from app.container import container from app.errors import ApiError from app.extensions import ExtensionError from app.extensions.mcp_registry import McpRegistryError -from app.providers.registry import ProviderNotFoundError -from app.providers.factory import UnsupportedProviderError from app.providers.base import ProviderError from app.providers.credentials import ( CredentialStoreError, validate_provider_credential_id, ) +from app.providers.factory import UnsupportedProviderError +from app.providers.registry import ProviderNotFoundError from app.retrieval.engine import engine from app.services import ( index_service, @@ -225,14 +226,21 @@ async def list_notes( folder: str | None = None, tag: str | None = None, ) -> NoteListResponse: - items, total = note_service.list_notes(limit=limit, offset=offset, folder=folder, tag=tag) - return NoteListResponse(items=items, page=PageMeta(total=total, limit=limit, offset=offset)) + items, total = note_service.list_notes( + limit=limit, offset=offset, folder=folder, tag=tag + ) + return NoteListResponse( + items=items, page=PageMeta(total=total, limit=limit, offset=offset) + ) @router.post("/notes", response_model=Note, tags=["Notes"]) async def create_note(request: NoteCreateRequest) -> Note: return await note_service.create_note( - title=request.title, markdown=request.markdown, folder=request.folder, tags=request.tags + title=request.title, + markdown=request.markdown, + folder=request.folder, + tags=request.tags, ) @@ -240,7 +248,9 @@ async def create_note(request: NoteCreateRequest) -> Note: async def get_note(note_id: str) -> Note: note = await note_service.get_note(note_id) if note is None: - raise ApiError(404, "RESOURCE_NOT_FOUND", "note not found", {"note_id": note_id}) + raise ApiError( + 404, "RESOURCE_NOT_FOUND", "note not found", {"note_id": note_id} + ) return note @@ -254,7 +264,9 @@ async def update_note(note_id: str, request: NoteUpdateRequest) -> Note: @router.delete("/notes/{note_id}", response_model=OperationResponse, tags=["Notes"]) async def delete_note(note_id: str) -> OperationResponse: if not await note_service.delete_note(note_id): - raise ApiError(404, "RESOURCE_NOT_FOUND", "note not found", {"note_id": note_id}) + raise ApiError( + 404, "RESOURCE_NOT_FOUND", "note not found", {"note_id": note_id} + ) return OperationResponse(status="completed", resource_id=note_id, message="deleted") @@ -298,7 +310,9 @@ async def chat(request: ChatRequest) -> StreamingResponse: data={"code": "PROVIDER_ERROR", "message": str(exc)}, timestamp=utc_now(), ) - done = ModelEvent(event=ModelEventType.done, sequence=1, timestamp=utc_now()) + done = ModelEvent( + event=ModelEventType.done, sequence=1, timestamp=utc_now() + ) yield as_sse(error.event.value, error.model_dump_json()) yield as_sse(done.event.value, done.model_dump_json()) @@ -459,9 +473,7 @@ async def list_skills() -> SkillListResponse: return SkillListResponse(items=container.skills.list()) -@router.get( - "/skills/{skill_id}", response_model=Skill, tags=["Skills"] -) +@router.get("/skills/{skill_id}", response_model=Skill, tags=["Skills"]) async def get_skill(skill_id: str) -> Skill: return extension_call(lambda: container.skills.get(skill_id)) @@ -501,7 +513,9 @@ async def disable_skill(skill_id: str) -> Skill: ) async def uninstall_skill(skill_id: str) -> OperationResponse: extension_call(lambda: container.skills.uninstall(skill_id)) - return OperationResponse(status="completed", resource_id=skill_id, message="uninstalled") + return OperationResponse( + status="completed", resource_id=skill_id, message="uninstalled" + ) # Independent MCP Server Registry @@ -510,7 +524,9 @@ async def list_mcp_servers() -> McpServerListResponse: return McpServerListResponse(items=mcp_call(container.mcp_servers.list)) -@router.post("/mcp/servers", response_model=McpServer, status_code=201, tags=["MCP Servers"]) +@router.post( + "/mcp/servers", response_model=McpServer, status_code=201, tags=["MCP Servers"] +) async def create_mcp_server(request: McpServerCreateRequest) -> McpServer: return mcp_call(lambda: container.mcp_servers.create(request)) @@ -520,45 +536,97 @@ async def get_mcp_server(server_id: str) -> McpServer: return mcp_call(lambda: container.mcp_servers.get(server_id)) +@router.get( + "/mcp/servers/{server_id}/tools", + response_model=McpToolSummaryListResponse, + tags=["MCP Servers"], +) +async def list_mcp_server_tools(server_id: str) -> McpToolSummaryListResponse: + return McpToolSummaryListResponse( + items=mcp_call(lambda: container.mcp_servers.list_tools(server_id)) + ) + + @router.put("/mcp/servers/{server_id}", response_model=McpServer, tags=["MCP Servers"]) -async def update_mcp_server(server_id: str, request: McpServerUpdateRequest) -> McpServer: - return await mcp_call_async(lambda: container.mcp_servers.update(server_id, request)) +async def update_mcp_server( + server_id: str, request: McpServerUpdateRequest +) -> McpServer: + return await mcp_call_async( + lambda: container.mcp_servers.update(server_id, request) + ) -@router.delete("/mcp/servers/{server_id}", response_model=OperationResponse, tags=["MCP Servers"]) +@router.delete( + "/mcp/servers/{server_id}", response_model=OperationResponse, tags=["MCP Servers"] +) async def delete_mcp_server(server_id: str) -> OperationResponse: await mcp_call_async(lambda: container.mcp_servers.delete(server_id)) - return OperationResponse(status="completed", resource_id=server_id, message="deleted") + return OperationResponse( + status="completed", resource_id=server_id, message="deleted" + ) -@router.post("/mcp/servers/{server_id}/trust", response_model=McpServer, tags=["MCP Servers"]) +@router.post( + "/mcp/servers/{server_id}/trust", response_model=McpServer, tags=["MCP Servers"] +) async def trust_mcp_server(server_id: str, request: McpServerTrustRequest) -> McpServer: - return mcp_call(lambda: container.mcp_servers.trust(server_id, request.command_digest)) + return mcp_call( + lambda: container.mcp_servers.trust(server_id, request.command_digest) + ) -@router.post("/mcp/servers/{server_id}/test", response_model=McpServer, tags=["MCP Servers"]) +@router.post( + "/mcp/servers/{server_id}/test", response_model=McpServer, tags=["MCP Servers"] +) async def test_mcp_server(server_id: str) -> McpServer: return await mcp_call_async(lambda: container.mcp_servers.test(server_id)) -@router.post("/mcp/servers/{server_id}/enable", response_model=McpServer, tags=["MCP Servers"]) +@router.post( + "/mcp/servers/{server_id}/enable", response_model=McpServer, tags=["MCP Servers"] +) async def enable_mcp_server(server_id: str) -> McpServer: return await mcp_call_async(lambda: container.mcp_servers.enable(server_id)) -@router.post("/mcp/servers/{server_id}/disable", response_model=McpServer, tags=["MCP Servers"]) +@router.post( + "/mcp/servers/{server_id}/disable", response_model=McpServer, tags=["MCP Servers"] +) async def disable_mcp_server(server_id: str) -> McpServer: return await mcp_call_async(lambda: container.mcp_servers.disable(server_id)) -@router.put("/mcp/servers/{server_id}/secrets/{key}", response_model=McpServerSecretStatus, tags=["MCP Servers"]) -async def put_mcp_server_secret(server_id: str, key: str, request: McpServerSecretWriteRequest) -> McpServerSecretStatus: - return mcp_call(lambda: container.mcp_servers.put_secret(server_id, key, request.secret.get_secret_value())) +@router.put( + "/mcp/servers/{server_id}/secrets/{key}", + response_model=McpServerSecretStatus, + tags=["MCP Servers"], +) +async def put_mcp_server_secret( + server_id: str, + key: str, + request: McpServerSecretWriteRequest, + kind: str = Query(default="environment", pattern="^(environment|header)$"), +) -> McpServerSecretStatus: + return mcp_call( + lambda: container.mcp_servers.put_secret( + server_id, key, request.secret.get_secret_value(), kind=kind + ) + ) -@router.delete("/mcp/servers/{server_id}/secrets/{key}", response_model=McpServerSecretStatus, tags=["MCP Servers"]) -async def delete_mcp_server_secret(server_id: str, key: str) -> McpServerSecretStatus: - return mcp_call(lambda: container.mcp_servers.delete_secret(server_id, key)) +@router.delete( + "/mcp/servers/{server_id}/secrets/{key}", + response_model=McpServerSecretStatus, + tags=["MCP Servers"], +) +async def delete_mcp_server_secret( + server_id: str, + key: str, + kind: str = Query(default="environment", pattern="^(environment|header)$"), +) -> McpServerSecretStatus: + return mcp_call( + lambda: container.mcp_servers.delete_secret(server_id, key, kind=kind) + ) # Plugins @@ -650,11 +718,15 @@ async def restart_plugin_host(plugin_id: str) -> OperationResponse: ) async def uninstall_plugin(plugin_id: str) -> OperationResponse: plugin = extension_call(lambda: container.plugins.get(plugin_id)) - dependent_skills = container.skills.depending_on_tools(plugin.manifest.contributes.tools) + dependent_skills = container.skills.depending_on_tools( + plugin.manifest.contributes.tools + ) await extension_call_async( lambda: container.plugins.uninstall(plugin_id, dependent_skills) ) - return OperationResponse(status="completed", resource_id=plugin_id, message="uninstalled") + return OperationResponse( + status="completed", resource_id=plugin_id, message="uninstalled" + ) # Plugin Command / Settings Contributions @@ -729,9 +801,7 @@ async def put_plugin_setting_secret( response_model=PluginSecretStatus, tags=["Plugins"], ) -async def delete_plugin_setting_secret( - plugin_id: str, key: str -) -> PluginSecretStatus: +async def delete_plugin_setting_secret(plugin_id: str, key: str) -> PluginSecretStatus: return extension_call( lambda: container.plugins.delete_setting_secret(plugin_id, key) ) @@ -844,7 +914,9 @@ async def update_provider( ) -> ProviderConfig: current = configurable_provider_or_404(provider_id).config if provider_id == "mock": - raise ApiError(409, "BUILTIN_PROVIDER_IMMUTABLE", "Mock provider cannot be modified.") + raise ApiError( + 409, "BUILTIN_PROVIDER_IMMUTABLE", "Mock provider cannot be modified." + ) fields = request.model_fields_set if ("name" in fields and request.name is None) or ( "enabled" in fields and request.enabled is None @@ -873,7 +945,9 @@ async def update_provider( async def delete_provider(provider_id: str) -> OperationResponse: configurable_provider_or_404(provider_id) if provider_id == "mock": - raise ApiError(409, "BUILTIN_PROVIDER_IMMUTABLE", "Mock provider cannot be deleted.") + raise ApiError( + 409, "BUILTIN_PROVIDER_IMMUTABLE", "Mock provider cannot be deleted." + ) container.providers.unregister(provider_id) return OperationResponse(status="completed", resource_id=provider_id) @@ -951,7 +1025,9 @@ async def create_task(request: TaskCreateRequest) -> Task: async def get_task(task_id: str) -> Task: task = task_service.get_task(task_id) if task is None: - raise ApiError(404, "RESOURCE_NOT_FOUND", "task not found", {"task_id": task_id}) + raise ApiError( + 404, "RESOURCE_NOT_FOUND", "task not found", {"task_id": task_id} + ) return task @@ -967,7 +1043,9 @@ async def update_task(task_id: str, request: TaskUpdateRequest) -> Task: ) async def delete_task(task_id: str) -> OperationResponse: if not task_service.delete_task(task_id): - raise ApiError(404, "RESOURCE_NOT_FOUND", "task not found", {"task_id": task_id}) + raise ApiError( + 404, "RESOURCE_NOT_FOUND", "task not found", {"task_id": task_id} + ) return OperationResponse(status="completed", resource_id=task_id, message="deleted") @@ -1017,5 +1095,7 @@ async def rebuild_index(request: IndexRebuildRequest) -> IndexJob: async def get_index_job(job_id: str) -> IndexJob: job = index_service.get_job(job_id) if job is None: - raise ApiError(404, "RESOURCE_NOT_FOUND", "index job not found", {"job_id": job_id}) + raise ApiError( + 404, "RESOURCE_NOT_FOUND", "index job not found", {"job_id": job_id} + ) return job diff --git a/backend/tests/test_mcp_registry.py b/backend/tests/test_mcp_registry.py index a4eb658..ae0e9fd 100644 --- a/backend/tests/test_mcp_registry.py +++ b/backend/tests/test_mcp_registry.py @@ -1,10 +1,15 @@ +import asyncio +import json import sys +import time +from concurrent.futures import ThreadPoolExecutor +import httpx import pytest -from app.agent.tools import ToolRegistry +from app.agent.tools import ToolExecutionContext, ToolRegistry from app.config import BACKEND_DIR, get_settings -from app.contracts import McpServerCreateRequest, McpServerUpdateRequest +from app.contracts import McpServerCreateRequest, McpServerUpdateRequest, ToolCall from app.extensions.mcp_registry import McpRegistryError, McpServerRegistry from app.providers.credentials import EncryptedCredentialStore @@ -58,6 +63,7 @@ def test_update_disables_server_and_revokes_command_trust() -> None: service = registry() created = service.create(request(secret_environment_keys=[])) service.trust(created.server_id, created.command_digest) + service.test(created.server_id) enabled = service.enable(created.server_id) assert enabled.enabled is True assert any( @@ -68,7 +74,8 @@ def test_update_disables_server_and_revokes_command_trust() -> None: updated = service.update( created.server_id, McpServerUpdateRequest( - **request(name="Changed", secret_environment_keys=[]).model_dump() + **request(name="Changed", secret_environment_keys=[]).model_dump(), + version=enabled.version, ), ) assert updated.enabled is False @@ -89,25 +96,73 @@ def test_production_rejects_process_launch_even_after_approval() -> None: assert error.value.code == "MCP_SANDBOX_REQUIRED" -def test_non_stdio_transport_is_explicitly_reserved() -> None: +def test_enable_requires_successful_test_and_update_checks_version() -> None: + service = registry() + created = service.create(request(secret_environment_keys=[])) + service.trust(created.server_id, created.command_digest) + with pytest.raises(McpRegistryError) as error: + service.enable(created.server_id) + assert error.value.code == "MCP_CONNECTION_TEST_REQUIRED" + + with pytest.raises(McpRegistryError) as error: + service.update( + created.server_id, + McpServerUpdateRequest( + **request(secret_environment_keys=[]).model_dump(), version=99 + ), + ) + assert error.value.code == "MCP_SERVER_VERSION_CONFLICT" + + +def test_http_transport_rejects_invalid_cross_transport_fields() -> None: + service = registry() + with pytest.raises(McpRegistryError) as error: + service.create( + request( + transport="streamable_http", + url="https://example.invalid/mcp", + secret_environment_keys=[], + ) + ) + assert error.value.code == "MCP_CONFIG_INVALID" + + +def test_registry_rejects_corrupt_persisted_json(tmp_path) -> None: + path = tmp_path / "mcp" + path.mkdir() + (path / "servers.json").write_text("{broken", encoding="utf-8") + with pytest.raises(McpRegistryError) as error: + McpServerRegistry( + ToolRegistry(), + EncryptedCredentialStore(), + tmp_path, + allow_process_launch=True, + ) + assert error.value.code == "MCP_REGISTRY_INVALID" + + +def test_stdio_command_is_not_parsed_as_a_shell_string() -> None: service = registry() created = service.create( request( - transport="streamable_http", - command="https://example.invalid/mcp", + command=f'"{sys.executable}" "{SERVER}"', + args=[], secret_environment_keys=[], ) ) service.trust(created.server_id, created.command_digest) with pytest.raises(McpRegistryError) as error: service.test(created.server_id) - assert error.value.code == "MCP_TRANSPORT_UNSUPPORTED" + assert error.value.code == "PLUGIN_HOST_START_FAILED" + assert service.get(created.server_id).last_test_succeeded is False + service.shutdown() def test_enabled_server_is_restored_from_persisted_registry() -> None: first = registry() created = first.create(request(secret_environment_keys=[])) first.trust(created.server_id, created.command_digest) + first.test(created.server_id) first.enable(created.server_id) first.shutdown() @@ -121,3 +176,222 @@ def test_enabled_server_is_restored_from_persisted_registry() -> None: for item in restored.tools.definitions() ) restored.shutdown() + + +def test_lifecycle_operations_are_serialized_and_tool_names_are_isolated() -> None: + service = registry() + servers = [ + service.create(request(name=f"Echo {index}", secret_environment_keys=[])) + for index in range(2) + ] + for server in servers: + service.trust(server.server_id, server.command_digest) + service.test(server.server_id) + + with ThreadPoolExecutor(max_workers=4) as pool: + enabled = list(pool.map(lambda item: service.enable(item.server_id), servers * 2)) + assert all(item.enabled for item in enabled) + names = [ + item.name + for item in service.tools.definitions() + if item.source == "mcp_server" + ] + assert len(names) == len(set(names)) + assert all(any(name.startswith(f"mcp.{item.server_id}.") for name in names) for item in servers) + + with ThreadPoolExecutor(max_workers=4) as pool: + list(pool.map(lambda item: service.disable(item.server_id), servers * 2)) + assert not any(item.source == "mcp_server" for item in service.tools.definitions()) + service.shutdown() + + +def _http_result(request_id: int, result: dict) -> httpx.Response: + return httpx.Response( + 200, + headers={"content-type": "application/json"}, + json={"jsonrpc": "2.0", "id": request_id, "result": result}, + ) + + +def test_streamable_http_supports_session_headers_secrets_and_tool_summary( + monkeypatch, +) -> None: + requests: list[httpx.Request] = [] + + def handler(request_value: httpx.Request) -> httpx.Response: + requests.append(request_value) + if request_value.method == "GET": + return httpx.Response(405) + if request_value.method == "DELETE": + return httpx.Response(405) + payload = json.loads(request_value.content) + if payload.get("method") == "initialize": + response = _http_result( + payload["id"], + { + "protocolVersion": "2025-11-25", + "capabilities": {"tools": {}}, + "serverInfo": {"name": "HTTP Fixture", "version": "1"}, + }, + ) + response.headers["MCP-Session-Id"] = "session-test" + return response + if payload.get("method") == "tools/list": + return _http_result( + payload["id"], + { + "tools": [ + { + "name": "echo", + "description": "Echo over HTTP", + "inputSchema": {"type": "object", "properties": {}}, + } + ] + }, + ) + if payload.get("method") == "tools/call": + return _http_result( + payload["id"], {"structuredContent": {"transport": "http"}} + ) + return httpx.Response(202) + + real_client = httpx.Client + monkeypatch.setattr( + "app.extensions.mcp.httpx.Client", + lambda **kwargs: real_client(transport=httpx.MockTransport(handler), **kwargs), + ) + service = registry() + created = service.create( + McpServerCreateRequest( + name="Remote MCP", + transport="streamable_http", + url="https://mcp.example.test/mcp", + headers={"X-Client": "NotesAgent"}, + secret_header_keys=["Authorization"], + ) + ) + service.put_secret( + created.server_id, "Authorization", "Bearer hidden", kind="header" + ) + service.trust(created.server_id, created.command_digest) + tested = service.test(created.server_id) + + assert tested.last_test_succeeded is True + assert tested.secret_headers == {"Authorization": True} + assert "Bearer hidden" not in tested.model_dump_json() + assert service.list_tools(created.server_id)[0].remote_name == "echo" + assert any( + request.headers.get("mcp-session-id") == "session-test" for request in requests + ) + assert any( + request.headers.get("mcp-protocol-version") == "2025-11-25" + for request in requests + ) + assert all( + request.headers.get("authorization") == "Bearer hidden" for request in requests + ) + enabled = service.enable(created.server_id) + tool_name = service.list_tools(created.server_id)[0].name + result = asyncio.run( + service.tools.execute( + ToolCall(tool_call_id="call-1", name=tool_name, arguments={}), + ToolExecutionContext(run_id="run-1"), + ) + ) + assert enabled.enabled is True + assert result.success is True + assert result.output == {"transport": "http"} + service.disable(created.server_id) + service.shutdown() + + +class _LegacyEventStream(httpx.SyncByteStream): + def __iter__(self): + yield b"event: endpoint\ndata: /messages\n\n" + time.sleep(0.1) + initialize = { + "jsonrpc": "2.0", + "id": 1, + "result": { + "protocolVersion": "2024-11-05", + "capabilities": {"tools": {}}, + "serverInfo": {"name": "Legacy Fixture"}, + }, + } + yield f"data: {json.dumps(initialize)}\n\n".encode() + time.sleep(0.1) + tools = { + "jsonrpc": "2.0", + "id": 2, + "result": {"tools": []}, + } + yield f"data: {json.dumps(tools)}\n\n".encode() + + +def test_legacy_sse_uses_same_origin_endpoint(monkeypatch) -> None: + posted_urls: list[str] = [] + + def handler(request_value: httpx.Request) -> httpx.Response: + if request_value.method == "GET": + return httpx.Response( + 200, + headers={"content-type": "text/event-stream"}, + stream=_LegacyEventStream(), + ) + posted_urls.append(str(request_value.url)) + return httpx.Response(202) + + real_client = httpx.Client + monkeypatch.setattr( + "app.extensions.mcp.httpx.Client", + lambda **kwargs: real_client(transport=httpx.MockTransport(handler), **kwargs), + ) + service = registry() + created = service.create( + McpServerCreateRequest( + name="Legacy MCP", + transport="sse", + url="https://legacy.example.test/sse", + ) + ) + service.trust(created.server_id, created.command_digest) + tested = service.test(created.server_id) + assert tested.last_test_succeeded is True + assert posted_urls and all( + url == "https://legacy.example.test/messages" for url in posted_urls + ) + service.shutdown() + + +class _CrossOriginLegacyEventStream(httpx.SyncByteStream): + def __iter__(self): + yield b"event: endpoint\ndata: https://attacker.example/messages\n\n" + + +def test_legacy_sse_rejects_cross_origin_message_endpoint(monkeypatch) -> None: + def handler(request_value: httpx.Request) -> httpx.Response: + assert request_value.method == "GET" + return httpx.Response( + 200, + headers={"content-type": "text/event-stream"}, + stream=_CrossOriginLegacyEventStream(), + ) + + real_client = httpx.Client + monkeypatch.setattr( + "app.extensions.mcp.httpx.Client", + lambda **kwargs: real_client(transport=httpx.MockTransport(handler), **kwargs), + ) + service = registry() + created = service.create( + McpServerCreateRequest( + name="Unsafe legacy MCP", + transport="sse", + url="https://legacy.example.test/sse", + ) + ) + service.trust(created.server_id, created.command_digest) + with pytest.raises(McpRegistryError) as error: + service.test(created.server_id) + assert error.value.code == "MCP_HTTP_RESPONSE_INVALID" + service.shutdown() diff --git a/docs/architecture/AI笔记软件技术栈说明-团队版-v2.3.md b/docs/architecture/AI笔记软件技术栈说明-团队版-v2.3.md index bb1ddf9..3b3a9a2 100644 --- a/docs/architecture/AI笔记软件技术栈说明-团队版-v2.3.md +++ b/docs/architecture/AI笔记软件技术栈说明-团队版-v2.3.md @@ -1050,7 +1050,7 @@ Python 包形式的 MCP Server 推荐使用固定版本的 `uvx --isolated --fro MCP Bridge 用于接入具有 MCP Server 接口的插件或外部工具服务。 -当前已实现本地 stdio 首版:Plugin Runtime 在授权后的启用阶段启动独立 Server 进程,完成 `initialize`、capability negotiation、分页 `tools/list`、`tools/call`、取消、超时、异常退出和 Host Restart。实现接受 `2025-11-25`、`2025-06-18`、`2025-03-26` 与 `2024-11-05` 协议版本;Streamable HTTP、Resource、Prompt、Sampling 与操作系统级沙箱仍属于后续范围。 +当前 Plugin Runtime 已实现本地 stdio Host:在授权后的启用阶段启动独立 Server 进程,完成 `initialize`、capability negotiation、分页 `tools/list`、`tools/call`、取消、超时、异常退出和 Host Restart。独立 MCP Server Registry 另行支持 stdio、Streamable HTTP 与旧 HTTP+SSE 兼容,包括 Session、协议 Header、认证 Header Secret、测试门禁和 Tool 动态映射。实现接受 `2025-11-25`、`2025-06-18`、`2025-03-26` 与 `2024-11-05` 协议版本;Plugin Manifest 的 Streamable HTTP、Resource、Prompt、Sampling 与操作系统级沙箱仍属于后续范围。 MCP Tool 进入系统后的调用路径为: diff --git a/docs/contracts/第二阶段接口契约-开发版.md b/docs/contracts/第二阶段接口契约-开发版.md index a600282..fb104dc 100644 --- a/docs/contracts/第二阶段接口契约-开发版.md +++ b/docs/contracts/第二阶段接口契约-开发版.md @@ -2,7 +2,7 @@ > 文档状态:接口冻结草案 > -> 更新日期:2026-09-02 +> 更新日期:2026-09-03 > > 依据:`../architecture/第二阶段团队分工表.md`、`../architecture/AI笔记软件技术栈说明-团队版-v2.3.md`、`后端接口契约-开发版.md` @@ -47,12 +47,13 @@ | Agent Trace | GET | `/api/agent/runs/{run_id}/trace` | 已实现 | 分页读取可回放 Trace 快照 | | Plugin Host | GET | `/api/plugins/{plugin_id}/host` | 已实现 | 获取 MCP Host 健康状态 | | Plugin Host | POST | `/api/plugins/{plugin_id}/host/restart` | 已实现 | 重启异常 Host 并重新发现 Tool | -| MCP Server | GET/POST | `/api/mcp/servers` | 已实现(C.1/P0) | 列出、创建独立 MCP Server 配置 | -| MCP Server | GET/PUT/DELETE | `/api/mcp/servers/{server_id}` | 已实现(C.1/P0) | 读取、修改、删除独立配置 | -| MCP Server | POST | `/api/mcp/servers/{server_id}/trust` | 已实现(C.1/P0) | 确认当前可执行配置摘要 | -| MCP Server | POST | `/api/mcp/servers/{server_id}/test` | 已实现(C.1/P0) | 隔离启动、握手、发现工具后退出 | -| MCP Server | POST | `/api/mcp/servers/{server_id}/enable`、`disable` | 已实现(C.1/P0) | 控制 Host 与动态 Tool 生命周期 | -| MCP Server | PUT/DELETE | `/api/mcp/servers/{server_id}/secrets/{key}` | 已实现(C.1/P0) | 写入或删除加密环境变量 | +| MCP Server | GET/POST | `/api/mcp/servers` | 已实现(C.1) | 列出、创建独立 MCP Server 配置 | +| MCP Server | GET/PUT/DELETE | `/api/mcp/servers/{server_id}` | 已实现(C.1) | 读取、版本化修改、删除独立配置 | +| MCP Server | GET | `/api/mcp/servers/{server_id}/tools` | 已实现(C.1) | 获取映射后的 Tool 摘要 | +| MCP Server | POST | `/api/mcp/servers/{server_id}/trust` | 已实现(C.1) | 确认当前连接配置摘要 | +| MCP Server | POST | `/api/mcp/servers/{server_id}/test` | 已实现(C.1) | 临时连接、握手、发现工具后关闭 | +| MCP Server | POST | `/api/mcp/servers/{server_id}/enable`、`disable` | 已实现(C.1) | 控制连接与动态 Tool 生命周期 | +| MCP Server | PUT/DELETE | `/api/mcp/servers/{server_id}/secrets/{key}` | 已实现(C.1) | 按 `kind` 写入或删除加密环境变量/Header | | Plugin Command | GET | `/api/plugin-contributions/commands` | 已实现 | 获取前端可展示的 Command | | Plugin Command | POST | `/api/plugin-contributions/commands/{command_id}/execute` | 已实现 | 受控执行 Command | | Plugin Settings | GET | `/api/plugins/{plugin_id}/settings` | 已实现 | 获取 Schema 与非敏感配置 | @@ -659,9 +660,19 @@ MCP_TRANSPORT_UNSUPPORTED MCP_SERVER_NOT_FOUND MCP_SERVER_NAME_INVALID MCP_SERVER_ALREADY_ENABLED +MCP_SERVER_VERSION_CONFLICT MCP_REGISTRY_WRITE_FAILED +MCP_REGISTRY_INVALID +MCP_CONNECTION_TEST_REQUIRED +MCP_CONFIG_INVALID +MCP_COMMAND_INVALID +MCP_URL_INVALID +MCP_HEADER_INVALID +MCP_HTTP_REQUEST_FAILED +MCP_HTTP_RESPONSE_INVALID MCP_SECRET_REQUIRED MCP_SECRET_NOT_DECLARED +MCP_SECRET_KIND_INVALID MCP_SECRET_STORE_ERROR MCP_ENVIRONMENT_INVALID MCP_PERMISSION_INVALID @@ -690,11 +701,15 @@ CREDENTIAL_NAMESPACE_RESERVED ### 7.8 独立 MCP Server Registry(C.1) -独立 Server 不依附 Plugin Manifest,配置持久化于 `APP_DATA_DIR/mcp/servers.json`,敏感环境变量只以 `mcp.*` 引用进入加密凭据存储。响应仅返回每个 Secret 是否已配置,不返回明文。动态工具使用 `mcp.{server_id}.{remote_tool}` 命名空间,来源标记为 `mcp_server`,仍通过统一 Tool Registry、Permission Manager 与 Agent Trace。 +独立 Server 不依附 Plugin Manifest,配置持久化于 `APP_DATA_DIR/mcp/servers.json`。`transport` 支持 `stdio`、`streamable_http` 和兼容旧服务的 `sse`。敏感环境变量与认证 Header 只以 `mcp.*` 引用进入加密凭据存储;读取响应以 `secret_environment`、`secret_headers` 的布尔值表示配置状态,不返回明文。动态工具使用 `mcp.{server_id}.{remote_tool}` 命名空间,来源标记为 `mcp_server`,仍通过统一 Tool Registry、Permission Manager 与 Agent Trace。 -P0 只真实支持 `stdio`。`streamable_http` 与 `sse` 已作为后续 Contract 枚举保留,但测试或启用会返回 `501 MCP_TRANSPORT_UNSUPPORTED`,前端不可伪装为可用。命令始终以 executable 与 args 数组通过 `shell=False` 启动;普通环境变量和加密 Secret 显式注入,不继承 Provider Key、数据库或 Vault 路径。 +stdio 配置使用 `command`、`args`、`environment` 和 `secret_environment_keys`;HTTP/SSE 配置使用 `url`、`headers` 和 `secret_header_keys`,两组 Transport 字段不可混用。更新请求必须携带当前 `version`,成功后版本递增;过期版本返回 `409 MCP_SERVER_VERSION_CONFLICT`。`GET /tools` 返回 `name`、`remote_name`、`description` 和可选 `permission`。 -创建或编辑配置后,调用方必须向 `/trust` 回传服务端计算的 `command_digest`。后端只接受与当前 transport、command、args、环境变量键值及权限完全一致的摘要;配置变化会撤销旧信任。测试连接同样会实际启动进程,因此也要求确认。当前 Python Host 仅在 `APP_ENVIRONMENT=development` 时允许启动;其他环境返回 `403 MCP_SANDBOX_REQUIRED`,等待第三阶段桌面端沙箱接管。 +创建或编辑配置后,调用方必须向 `/trust` 回传服务端计算的 `command_digest`。后端只接受与当前 Transport、连接参数、环境/Header 及权限完全一致的摘要;配置变化会撤销旧信任和测试结果。只有当前摘要通过 `/test`,才能调用 `/enable`。测试失败也会持久化时间和失败状态。 + +Streamable HTTP 支持 Session ID、`MCP-Protocol-Version`、JSON 或 SSE POST 响应、可选 GET 事件流及 `Last-Event-ID`;旧 SSE 按 endpoint 事件确定 POST 地址,并要求与初始 URL 同源。Secret 接口用 `?kind=environment` 或 `?kind=header` 区分类型。HTTP URL 不允许内嵌凭据或 Fragment,配置不得覆盖协议保留 Header。 + +stdio 命令始终以 executable 与 args 数组通过 `shell=False` 启动;普通环境变量和加密 Secret 显式注入,不继承 Provider Key、数据库或 Vault 路径。当前 Python Host 仅在 `APP_ENVIRONMENT=development` 时允许启动 stdio;其他环境返回 `403 MCP_SANDBOX_REQUIRED`。远程 HTTP Transport 不创建本机进程,但仍受摘要确认、成功测试、超时、消息限长与 Secret 隔离约束。 --- diff --git a/docs/development/AI-Core与Agent-Core开发说明.md b/docs/development/AI-Core与Agent-Core开发说明.md index 0aadeb4..dac8187 100644 --- a/docs/development/AI-Core与Agent-Core开发说明.md +++ b/docs/development/AI-Core与Agent-Core开发说明.md @@ -349,4 +349,4 @@ Skill Manifest - Task 已持久化到 SQLite;Attachment Tool 读取 Host 管理目录中的 UTF-8 文件。 - `audio.transcribe` 当前消费 Host 预生成的 transcript;faster-whisper 与说话人分离仍待第二阶段后续接入。 - Extension 安装记录暂存内存;后续接入持久化 Registry 与版本升级流程。 -- 当前 Plugin Host 支持内置声明式 handler、本地 stdio MCP Server 以及 Plugin Command/Settings;Streamable HTTP、OS 级沙箱与 UI Contribution 留在后续阶段。 +- 当前 Plugin Host 支持内置声明式 handler、本地 stdio MCP Server 以及 Plugin Command/Settings;独立 MCP Server Registry 另行支持 stdio、Streamable HTTP 与旧 SSE 兼容。OS 级沙箱与 UI Contribution 留在后续阶段。 diff --git a/docs/development/独立MCP-Server配置中心开发说明.md b/docs/development/独立MCP-Server配置中心开发说明.md index 26d68f0..3bb8b8d 100644 --- a/docs/development/独立MCP-Server配置中心开发说明.md +++ b/docs/development/独立MCP-Server配置中心开发说明.md @@ -1,47 +1,76 @@ # 独立 MCP Server 配置中心开发说明 -> 更新日期:2026-09-03。本文记录第二阶段 C.1 的 P0 实现;它与 Plugin 自带的 MCP Host 是两个并列入口。 +> 更新日期:2026-09-03。本文记录第二阶段 C.1 的完整实现;独立 MCP Server Registry 与 Plugin 自带 MCP Host 是两个并列入口。 ## 1. 已实现范围 -- 独立 Server 的创建、读取、编辑和删除; -- stdio 命令、参数、普通环境变量、加密环境变量及超时配置; -- 命令摘要确认、测试连接、启用、停用和异常状态展示; -- MCP initialize、`tools/list` 与动态 Tool 注册,工具命名为 `mcp.{server_id}.{tool}`; -- 启用状态持久化与开发服务重启恢复; -- 前端独立“MCP”导航与配置弹窗,提供 stdio/uvx 模板; -- Streamable HTTP 与旧 SSE 仅作为后续选项展示为禁用,不属于本次完成范围。 +- 独立 Server 的创建、读取、版本化编辑、删除和 Tool 摘要查询; +- `stdio`、Streamable HTTP 和旧版 HTTP+SSE 三种 Transport; +- stdio 可执行文件、参数、普通/加密环境变量,以及 HTTP URL、普通/加密 Header; +- 配置摘要确认、连接测试、启停、异常状态与最近一次测试结果; +- MCP initialize、`tools/list`、`tools/call`、取消与动态 Tool 注册,名称为 `mcp.{server_id}.{tool}`; +- Streamable HTTP Session、协议版本 Header、JSON/SSE POST 响应、可选 GET 事件流和 `Last-Event-ID` 重连; +- 旧 HTTP+SSE 的 endpoint 事件与消息 POST,并强制消息地址和配置地址同源; +- 前端表单/JSON 双模式、三种模板、高风险变更确认及请求期 Secret 输入。 -## 2. 数据与 Secret +Streamable HTTP 按 MCP 当前规范实现;SSE 仅用于兼容旧 Server,不应作为新部署首选。 -普通配置原子写入 `APP_DATA_DIR/mcp/servers.json`。Secret 使用 `mcp.{server_id}.{key_hash}` 作为内部引用写入现有 Fernet 凭据存储;API 和前端只看到 `configured: true/false`。删除 Server 或移除 Secret 键时同步清理密文。 +## 2. 配置、版本与 Secret -前端 Secret 输入使用密码框,提交后立即清空,不写入 localStorage、普通配置 JSON 或日志。`mcp.*` 同 `plugin.*` 一样属于保留凭据命名空间,Provider 配置与通用凭据 API 无权读取。 +普通配置原子写入 `APP_DATA_DIR/mcp/servers.json`。更新请求必须携带读取到的 `version`;版本过期返回 `409 MCP_SERVER_VERSION_CONFLICT`,避免多个页面互相覆盖。改变 Transport、命令、URL、Header、环境变量或权限后,旧授权和测试结果立即失效。 -## 3. 启动安全边界 +Secret 使用带类型和键名哈希的 `mcp.*` 内部 ID 写入 Fernet 凭据存储。API 只返回环境变量或 Header 是否配置,绝不返回明文。删除 Server 或移除 Secret 键会同步清理密文。前端密码框提交后立即清空,不写入 JSON 编辑器、localStorage、普通配置或日志。 -命令不经过 Shell,管道、重定向和拼接字符串不会被解释。Host 只继承启动所需的系统变量,再叠加用户显式配置;`uvx` 可隔离 Python 依赖,但不能限制文件、网络和系统调用。 +## 3. 启用与运行时规则 -后端会对影响执行的配置计算 SHA-256 摘要。测试或启用前,用户必须确认并回传当前摘要;修改配置会立即撤销旧确认。由于 Python Host 尚无 OS 沙箱,非开发环境硬拒绝启动。第三阶段由 Tauri/Rust Host 提供平台级隔离后再替换这道临时门禁。 +一次连接按以下顺序执行: -## 4. 测试与启动 +1. 用户检查服务端生成的连接摘要并确认当前摘要; +2. 后端临时连接,完成 initialize 和 `tools/list` 后关闭连接; +3. 只有当前摘要测试成功,启用操作才会启动长期连接并注册 Tool; +4. 停用、删除、超时或异常退出会注销 Tool 并关闭连接。 + +HTTP Header 中 `Host`、`Content-Type`、`MCP-Session-Id` 等协议保留项不可由配置覆盖。URL 不允许内嵌凭据或 Fragment。旧 SSE 返回的 POST endpoint 必须与初始 URL 同源,防止认证 Header 被转发到其他站点。 + +stdio 命令不经过 Shell,管道、重定向和命令拼接不会被解释。Windows 使用新进程组并通过 `taskkill /T` 回收子树;POSIX 使用独立 session/process group 并向进程组发信号。Python 阶段仍无法提供文件、网络、系统调用或操作系统版本差异下的绝对隔离保证。 + +`uvx` 模板使用 `--isolated`、明确的 `--from` 和固定包版本。它只能隔离依赖,不能替代安全沙箱。非开发环境仍拒绝启动 stdio Server,并返回 `403 MCP_SANDBOX_REQUIRED`;远程 HTTP Transport 不创建本机子进程,但仍要求摘要确认和成功测试。第三阶段前的 C.5 将冻结 Tauri/Rust 沙箱设计。 + +## 4. 接口 + +```text +GET /api/mcp/servers +POST /api/mcp/servers +GET /api/mcp/servers/{server_id} +PUT /api/mcp/servers/{server_id} +DELETE /api/mcp/servers/{server_id} +GET /api/mcp/servers/{server_id}/tools +POST /api/mcp/servers/{server_id}/trust +POST /api/mcp/servers/{server_id}/test +POST /api/mcp/servers/{server_id}/enable +POST /api/mcp/servers/{server_id}/disable +PUT /api/mcp/servers/{server_id}/secrets/{key}?kind=environment|header +DELETE /api/mcp/servers/{server_id}/secrets/{key}?kind=environment|header +``` + +完整字段、状态和错误码见《第二阶段接口契约-开发版》。协议实现参考 MCP 官方的 [Transports](https://modelcontextprotocol.io/specification/2025-11-25/basic/transports) 与 [Lifecycle](https://modelcontextprotocol.io/specification/2025-11-25/basic/lifecycle)。 + +## 5. 验证 ```powershell cd backend -uv run pytest -q tests/test_mcp_registry.py -uv run uvicorn app.main:app --reload +uv run pytest -q tests/test_mcp_registry.py tests/test_extension_core.py cd ../frontend npm run type-check npm test -npm run dev +npm run build ``` -打开知识库后进入左侧“MCP”。保存配置,按提示确认命令,先执行“测试连接”;成功后再启用。默认模板 `uvx mcp-server-fetch` 仅为配置示例,首次下载是否联网由本机 uv 缓存与网络环境决定。 +后端测试使用无需网络或密钥的 stdio Fixture,以及 `httpx.MockTransport` 驱动的确定性 HTTP/SSE Server Fixture。覆盖摘要授权、Secret 不回显、版本冲突、生产门禁、重启恢复、Streamable HTTP Session/Header/工具调用及旧 SSE 同源校验。前端覆盖模板切换、JSON 校验、Secret 请求期输入、测试失败和删除确认。 -## 5. 后续增量 +## 6. 后续边界 -- P1:Streamable HTTP 连接、认证 Header 与重连策略; -- 兼容项:仅在确有旧服务需求时增加 SSE; -- 第三阶段前:把命令确认与进程创建迁移至 Tauri/Rust 沙箱; -- 增加面向真实第三方 Server 的兼容矩阵,不用单一 Fixture 代表协议全兼容。 +- 增加真实第三方 Server 的兼容矩阵;确定性 Fixture 只能证明宿主协议行为,不能代表所有实现兼容; +- C.5 在第二阶段开发与测试完成后、第三阶段桌面端实现前冻结沙箱 Contract; +- 第三阶段将 stdio 进程创建和 Secret 托管迁移至 Tauri/Rust Host 与 Stronghold/系统 Keychain。 diff --git a/frontend/src/contracts/index.ts b/frontend/src/contracts/index.ts index 05026c5..07f6679 100644 --- a/frontend/src/contracts/index.ts +++ b/frontend/src/contracts/index.ts @@ -539,20 +539,26 @@ export type McpServerTransport = 'stdio' | 'streamable_http' | 'sse' export type McpServerState = 'stopped' | 'starting' | 'ready' | 'unhealthy' | 'error' export interface McpServerInput { + version?: number name: string transport: McpServerTransport - command: string + command?: string | null args: string[] + url?: string | null + headers: Record environment: Record secret_environment_keys: string[] + secret_header_keys: string[] permissions: string[] startup_timeout_seconds: number tool_timeout_seconds: number } -export interface McpServer extends Omit { +export interface McpServer extends Omit { server_id: string + version: number secret_environment: Record + secret_headers: Record enabled: boolean trusted: boolean command_digest: string @@ -567,6 +573,13 @@ export interface McpServer extends Omit ({ + listMcpServers: vi.fn(), createMcpServer: vi.fn(), updateMcpServer: vi.fn(), + deleteMcpServer: vi.fn(), trustMcpServer: vi.fn(), testMcpServer: vi.fn(), + enableMcpServer: vi.fn(), disableMcpServer: vi.fn(), putMcpServerSecret: vi.fn(), +})) + +const server: McpServer = { + server_id: 'server-1', version: 2, name: 'Remote', transport: 'streamable_http', + command: null, args: [], url: 'https://mcp.example.test/mcp', headers: {}, environment: {}, + secret_environment: {}, secret_headers: { Authorization: false }, permissions: [], + startup_timeout_seconds: 15, tool_timeout_seconds: 30, enabled: false, trusted: true, + command_digest: 'a'.repeat(64), command_summary: 'https://mcp.example.test/mcp', + status: 'stopped', tools_count: 1, last_test_succeeded: false, +} + +async function render(items: McpServer[] = []) { + vi.mocked(service.listMcpServers).mockResolvedValue(items) + const wrapper = mount(McpServersView, { global: { stubs: { AppIcon: true } } }) + await flushPromises() + return wrapper +} + +beforeEach(() => { + vi.clearAllMocks() + vi.stubGlobal('confirm', vi.fn(() => true)) +}) + +describe('McpServersView', () => { + it('switches transport templates and round-trips the JSON configuration mode', async () => { + const wrapper = await render() + await wrapper.findAll('button').find(button => button.text() === '新增服务器')!.trigger('click') + await wrapper.findAll('button').find(button => button.text() === 'Streamable HTTP')!.trigger('click') + expect(wrapper.find('input[placeholder="https://example.com/mcp"]').exists()).toBe(true) + + await wrapper.findAll('button').find(button => button.text() === 'JSON 配置')!.trigger('click') + const raw = (wrapper.get('.json-editor').element as HTMLTextAreaElement).value + expect(JSON.parse(raw)).toMatchObject({ transport: 'streamable_http', command: null }) + expect(raw).not.toContain('secret_value') + + await wrapper.findAll('button').find(button => button.text() === '表单配置')!.trigger('click') + expect(wrapper.text()).toContain('MCP URL') + }) + + it('rejects invalid JSON without sending a create request', async () => { + const wrapper = await render() + await wrapper.findAll('button').find(button => button.text() === '新增服务器')!.trigger('click') + await wrapper.findAll('button').find(button => button.text() === 'JSON 配置')!.trigger('click') + await wrapper.get('.json-editor').setValue('{invalid') + await flushPromises() + await wrapper.get('form').trigger('submit') + await flushPromises() + expect(wrapper.text()).toContain('服务器配置不是有效 JSON') + expect(service.createMcpServer).not.toHaveBeenCalled() + }) + + it('keeps secrets request-only, exposes test failures, and confirms deletion', async () => { + const wrapper = await render([server]) + const password = wrapper.get('input[type="password"]') + await password.setValue('request-only-secret') + vi.mocked(service.putMcpServerSecret).mockResolvedValue({} as never) + await wrapper.findAll('button').find(button => button.text() === '保存')!.trigger('click') + await flushPromises() + expect(service.putMcpServerSecret).toHaveBeenCalledWith('server-1', 'Authorization', 'request-only-secret', 'header') + expect((password.element as HTMLInputElement).value).toBe('') + + vi.mocked(service.testMcpServer).mockRejectedValue(new Error('连接失败')) + await wrapper.findAll('button').find(button => button.text().includes('测试连接'))!.trigger('click') + await flushPromises() + expect(wrapper.text()).toContain('连接失败') + + vi.mocked(service.deleteMcpServer).mockResolvedValue({ status: 'completed' }) + await wrapper.findAll('button').find(button => button.text().includes('删除'))!.trigger('click') + await flushPromises() + expect(confirm).toHaveBeenCalled() + expect(service.deleteMcpServer).toHaveBeenCalledWith('server-1') + }) +}) diff --git a/frontend/src/features/mcp/McpServersView.vue b/frontend/src/features/mcp/McpServersView.vue index 11a6957..04f1105 100644 --- a/frontend/src/features/mcp/McpServersView.vue +++ b/frontend/src/features/mcp/McpServersView.vue @@ -5,76 +5,134 @@ import AppIcon from '@/components/common/AppIcon.vue' import type { McpServer, McpServerInput, McpServerTransport } from '@/contracts' import * as service from '@/services/mcpServerService' +type SecretKind = 'environment' | 'header' + const servers = ref([]) const busy = ref('') const error = ref('') const dialogOpen = ref(false) const editingId = ref(null) +const editingOriginal = ref(null) +const editorMode = ref<'form' | 'json'>('form') const argsText = ref('') const environmentText = ref('{}') +const headersText = ref('{}') const secretKeysText = ref('') +const secretHeaderKeysText = ref('') const permissionsText = ref('') +const rawConfig = ref('') const secretDrafts = reactive>({}) -const form = reactive({ - name: '', transport: 'stdio', command: '', args: [], environment: {}, - secret_environment_keys: [], permissions: [], startup_timeout_seconds: 15, tool_timeout_seconds: 30, -}) +const form = reactive(emptyForm()) const dialogTitle = computed(() => editingId.value ? '编辑 MCP 服务器' : '新增 MCP 服务器') +function emptyForm(): McpServerInput { + return { + name: '', transport: 'stdio', command: '', args: [], url: null, headers: {}, + environment: {}, secret_environment_keys: [], secret_header_keys: [], permissions: [], + startup_timeout_seconds: 15, tool_timeout_seconds: 30, + } +} + async function load() { error.value = '' try { servers.value = await service.listMcpServers() } catch (cause) { error.value = message(cause, '读取 MCP 服务器失败') } } +function resetEditor(input: McpServerInput) { + Object.assign(form, input) + argsText.value = input.args.join('\n') + environmentText.value = JSON.stringify(input.environment, null, 2) + headersText.value = JSON.stringify(input.headers, null, 2) + secretKeysText.value = input.secret_environment_keys.join('\n') + secretHeaderKeysText.value = input.secret_header_keys.join('\n') + permissionsText.value = input.permissions.join(', ') + editorMode.value = 'form' + rawConfig.value = '' +} + function openCreate() { editingId.value = null - Object.assign(form, { name: '', transport: 'stdio', command: '', args: [], environment: {}, secret_environment_keys: [], permissions: [], startup_timeout_seconds: 15, tool_timeout_seconds: 30 }) - argsText.value = ''; environmentText.value = '{}'; secretKeysText.value = ''; permissionsText.value = '' + editingOriginal.value = null + resetEditor(emptyForm()) dialogOpen.value = true } function openEdit(server: McpServer) { editingId.value = server.server_id - Object.assign(form, { - name: server.name, transport: server.transport, command: server.command, - args: [...server.args], environment: { ...server.environment }, - secret_environment_keys: Object.keys(server.secret_environment), permissions: [...server.permissions], + editingOriginal.value = server + resetEditor({ + version: server.version, name: server.name, transport: server.transport, + command: server.command, args: [...server.args], url: server.url, + headers: { ...server.headers }, environment: { ...server.environment }, + secret_environment_keys: Object.keys(server.secret_environment), + secret_header_keys: Object.keys(server.secret_headers), permissions: [...server.permissions], startup_timeout_seconds: server.startup_timeout_seconds, tool_timeout_seconds: server.tool_timeout_seconds, }) - argsText.value = server.args.join('\n') - environmentText.value = JSON.stringify(server.environment, null, 2) - secretKeysText.value = Object.keys(server.secret_environment).join('\n') - permissionsText.value = server.permissions.join(', ') dialogOpen.value = true } function applyTemplate(transport: McpServerTransport) { - if (transport !== 'stdio') return - form.transport = 'stdio'; form.command = 'uvx'; argsText.value = 'mcp-server-fetch' + form.transport = transport + if (transport === 'stdio') { + form.command = 'uvx'; form.url = null + argsText.value = '--isolated\n--from\npackage-name==1.0.0\nserver-command' + } else { + form.command = null; argsText.value = ''; form.url = transport === 'sse' ? 'http://127.0.0.1:3000/sse' : 'http://127.0.0.1:3000/mcp' + } +} + +function parseObject(value: string, label: string): Record { + let parsed: unknown + try { parsed = JSON.parse(value || '{}') } catch { throw new Error(`${label}必须是 JSON 对象`) } + if (!parsed || Array.isArray(parsed) || typeof parsed !== 'object' || Object.values(parsed).some(item => typeof item !== 'string')) throw new Error(`${label}必须是字符串键值 JSON 对象`) + return parsed as Record +} + +function formPayload(): McpServerInput { + const stdio = form.transport === 'stdio' + return { + version: form.version, + name: form.name.trim(), transport: form.transport, + command: stdio ? form.command?.trim() : null, + args: stdio ? argsText.value.split('\n').map(value => value.trim()).filter(Boolean) : [], + url: stdio ? null : form.url?.trim(), + headers: stdio ? {} : parseObject(headersText.value, '普通 Header'), + environment: stdio ? parseObject(environmentText.value, '普通环境变量') : {}, + secret_environment_keys: stdio ? splitKeys(secretKeysText.value) : [], + secret_header_keys: stdio ? [] : splitKeys(secretHeaderKeysText.value), + permissions: permissionsText.value.split(',').map(value => value.trim()).filter(Boolean), + startup_timeout_seconds: form.startup_timeout_seconds, + tool_timeout_seconds: form.tool_timeout_seconds, + } } function payload(): McpServerInput { - let environment: Record - try { environment = JSON.parse(environmentText.value || '{}') } - catch { throw new Error('普通环境变量必须是 JSON 对象') } - if (!environment || Array.isArray(environment) || typeof environment !== 'object') throw new Error('普通环境变量必须是 JSON 对象') - return { - ...form, - name: form.name.trim(), command: form.command.trim(), - args: argsText.value.split('\n').map(value => value.trim()).filter(Boolean), - environment, - secret_environment_keys: secretKeysText.value.split(/[\n,]/).map(value => value.trim()).filter(Boolean), - permissions: permissionsText.value.split(',').map(value => value.trim()).filter(Boolean), - } + if (editorMode.value === 'form') return formPayload() + let parsed: unknown + try { parsed = JSON.parse(rawConfig.value) } catch { throw new Error('服务器配置不是有效 JSON') } + if (!parsed || Array.isArray(parsed) || typeof parsed !== 'object') throw new Error('服务器配置必须是 JSON 对象') + const value = parsed as McpServerInput + if (editingId.value) value.version = form.version + return value +} + +function switchMode(mode: 'form' | 'json') { + try { + if (mode === editorMode.value) return + if (mode === 'json') rawConfig.value = JSON.stringify(formPayload(), null, 2) + else resetEditor(payload()) + editorMode.value = mode + } catch (cause) { error.value = message(cause, '配置转换失败') } } async function save() { try { const input = payload() - if (!input.name || !input.command) throw new Error('请填写服务器名称和可执行命令') + if (!input.name || (input.transport === 'stdio' ? !input.command : !input.url)) throw new Error('请填写服务器名称和连接地址') + if (editingOriginal.value && executionChanged(editingOriginal.value, input) && !confirm('连接命令、地址或认证配置已变化,保存后旧测试与授权会失效。是否保存?')) return busy.value = 'save' editingId.value ? await service.updateMcpServer(editingId.value, input) : await service.createMcpServer(input) dialogOpen.value = false @@ -83,15 +141,19 @@ async function save() { finally { busy.value = '' } } +function executionChanged(server: McpServer, input: McpServerInput) { + return JSON.stringify([server.transport, server.command, server.args, server.url, server.headers, Object.keys(server.secret_headers)]) !== JSON.stringify([input.transport, input.command, input.args, input.url, input.headers, input.secret_header_keys]) +} + async function approve(server: McpServer): Promise { if (server.trusted) return server - const accepted = confirm(`即将允许本机启动以下命令:\n\n${server.command_summary}\n\n当前 Python Host 没有系统级沙箱,仅应运行可信服务器。是否继续?`) - if (!accepted) return null + const localWarning = server.transport === 'stdio' ? '\n\n本机进程尚无系统级沙箱,仅应运行可信服务器。' : '\n\n连接可能向该地址发送配置的 Header。' + if (!confirm(`请确认 MCP 连接:\n\n${server.command_summary}${localWarning}\n\n是否继续?`)) return null return service.trustMcpServer(server) } -async function test(server: McpServer) { await act(server, 'test', async current => service.testMcpServer(current.server_id)) } -async function toggle(server: McpServer) { await act(server, 'toggle', async current => current.enabled ? service.disableMcpServer(current.server_id) : service.enableMcpServer(current.server_id)) } +async function test(server: McpServer) { await act(server, 'test', current => service.testMcpServer(current.server_id)) } +async function toggle(server: McpServer) { await act(server, 'toggle', current => current.enabled ? service.disableMcpServer(current.server_id) : service.enableMcpServer(current.server_id)) } async function act(server: McpServer, action: string, operation: (server: McpServer) => Promise) { busy.value = `${action}:${server.server_id}`; error.value = '' try { const current = action === 'toggle' && server.enabled ? server : await approve(server); if (!current) return; await operation(current); await load() } @@ -105,46 +167,51 @@ async function remove(server: McpServer) { catch (cause) { error.value = message(cause, '删除失败') } finally { busy.value = '' } } -async function saveSecret(server: McpServer, key: string) { - const value = secretDrafts[`${server.server_id}:${key}`]?.trim() +async function saveSecret(server: McpServer, key: string, kind: SecretKind) { + const draftKey = `${server.server_id}:${kind}:${key}` + const value = secretDrafts[draftKey]?.trim() if (!value) return - try { busy.value = `secret:${server.server_id}:${key}`; await service.putMcpServerSecret(server.server_id, key, value); secretDrafts[`${server.server_id}:${key}`] = ''; await load() } + try { busy.value = `secret:${draftKey}`; await service.putMcpServerSecret(server.server_id, key, value, kind); secretDrafts[draftKey] = ''; await load() } catch (cause) { error.value = message(cause, '保存密钥失败') } finally { busy.value = '' } } +function splitKeys(value: string) { return value.split(/[\n,]/).map(item => item.trim()).filter(Boolean) } function message(cause: unknown, fallback: string) { return cause instanceof Error ? cause.message : fallback } onMounted(load) diff --git a/frontend/src/services/mcpServerService.ts b/frontend/src/services/mcpServerService.ts index a146a6b..8acc2b1 100644 --- a/frontend/src/services/mcpServerService.ts +++ b/frontend/src/services/mcpServerService.ts @@ -1,5 +1,5 @@ import apiClient from './apiClient' -import type { McpServer, McpServerInput, OperationResponse } from '@/contracts' +import type { McpServer, McpServerInput, McpToolSummary, OperationResponse } from '@/contracts' const base = '/api/mcp/servers' @@ -8,10 +8,11 @@ export async function listMcpServers(): Promise { } export const createMcpServer = (input: McpServerInput) => apiClient.post(base, input) export const updateMcpServer = (id: string, input: McpServerInput) => apiClient.put(`${base}/${id}`, input) +export const listMcpServerTools = async (id: string) => (await apiClient.get<{ items: McpToolSummary[] }>(`${base}/${id}/tools`)).items export const deleteMcpServer = (id: string) => apiClient.delete(`${base}/${id}`) export const trustMcpServer = (server: McpServer) => apiClient.post(`${base}/${server.server_id}/trust`, { command_digest: server.command_digest }) export const testMcpServer = (id: string) => apiClient.post(`${base}/${id}/test`) export const enableMcpServer = (id: string) => apiClient.post(`${base}/${id}/enable`) export const disableMcpServer = (id: string) => apiClient.post(`${base}/${id}/disable`) -export const putMcpServerSecret = (id: string, key: string, secret: string) => apiClient.put(`${base}/${id}/secrets/${encodeURIComponent(key)}`, { secret }) -export const deleteMcpServerSecret = (id: string, key: string) => apiClient.delete(`${base}/${id}/secrets/${encodeURIComponent(key)}`) +export const putMcpServerSecret = (id: string, key: string, secret: string, kind: 'environment' | 'header' = 'environment') => apiClient.put(`${base}/${id}/secrets/${encodeURIComponent(key)}?kind=${kind}`, { secret }) +export const deleteMcpServerSecret = (id: string, key: string, kind: 'environment' | 'header' = 'environment') => apiClient.delete(`${base}/${id}/secrets/${encodeURIComponent(key)}?kind=${kind}`)