diff --git a/.gitignore b/.gitignore index 44cc045..3b4f2e9 100644 --- a/.gitignore +++ b/.gitignore @@ -14,6 +14,10 @@ backend/.env # 运行期生成的 SQLite 索引(vault 下的 Markdown 测试数据需提交) backend/data/*.db* backend/data/credentials/ +# 本机 MCP 配置、授权状态及服务器工作目录不得提交。 +backend/data/mcp/ +server.json +servers.json # Editors and operating systems .idea/ diff --git a/README.md b/README.md index cd85754..3a96a48 100644 --- a/README.md +++ b/README.md @@ -2,7 +2,7 @@ > 本文件用于团队开发期间快速配置环境和启动项目,不是正式的项目 README。 -> 当前基线:2026-09-02。第一阶段 Web 联调前后端已经完成;第二阶段已完成 Workspace 去 Mock、Agent Trace 持久化与 SSE 恢复、stdio MCP Bridge、隔离 Plugin Host,以及 Plugin Command/Settings 后端 Contract 和前端 Service。真实音频、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/container.py b/backend/app/container.py index 62b21be..ee05cd8 100644 --- a/backend/app/container.py +++ b/backend/app/container.py @@ -5,6 +5,7 @@ from app.agent.builtin_tools import register_builtin_tools from app.contracts import ModelCapability, ProviderConfig, ProviderType from app.config import BACKEND_DIR, get_settings from app.extensions import PluginRuntime, SkillRuntime +from app.extensions.mcp_registry import McpServerRegistry from app.providers import MockProvider, ProviderFactory, ProviderRegistry from app.providers.credentials import ( ChainedCredentialResolver, @@ -22,6 +23,7 @@ class ApplicationContainer: permissions: PermissionManager skills: SkillRuntime plugins: PluginRuntime + mcp_servers: McpServerRegistry agent: AgentRuntime @@ -61,6 +63,14 @@ def build_container() -> ApplicationContainer: plugins.install(BACKEND_DIR / "extensions" / "plugins" / "text-tools") plugins.enable("text-tools") + mcp_servers = McpServerRegistry( + tools, + credentials, + settings.data_dir, + allow_process_launch=settings.environment == "development", + ) + mcp_servers.restore_enabled() + skills = SkillRuntime(tools) skills.install(BACKEND_DIR / "extensions" / "skills" / "knowledge-assistant") skills.enable("knowledge-assistant") @@ -81,6 +91,7 @@ def build_container() -> ApplicationContainer: permissions=permissions, skills=skills, plugins=plugins, + mcp_servers=mcp_servers, agent=agent, ) diff --git a/backend/app/contracts.py b/backend/app/contracts.py index f95e890..9840f45 100644 --- a/backend/app/contracts.py +++ b/backend/app/contracts.py @@ -199,7 +199,7 @@ class ToolDefinition(Contract): description: str parameters: dict[str, Any] = Field(default_factory=dict) permission: str | None = None - source: Literal["builtin", "plugin"] = "builtin" + source: Literal["builtin", "plugin", "mcp_server"] = "builtin" class ToolCall(Contract): @@ -486,6 +486,97 @@ class PluginHostStatus(Contract): error: str | None = None +# Independent user-managed MCP Server Registry. This is deliberately separate +# from Plugin manifests: a server can contribute tools without being a Plugin. +class McpServerTransport(str, Enum): + stdio = "stdio" + streamable_http = "streamable_http" + sse = "sse" + + +class McpServerConfig(Contract): + name: str = Field(min_length=1, max_length=80) + transport: McpServerTransport = McpServerTransport.stdio + 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 McpServerCreateRequest(McpServerConfig): + pass + + +class McpServerUpdateRequest(McpServerConfig): + version: int = Field(ge=1) + + +class McpServerSecretWriteRequest(Contract): + secret: SecretStr = Field(min_length=1, max_length=32768) + + +class McpServerSecretStatus(Contract): + key: str + configured: bool + + +class McpServerTrustRequest(Contract): + command_digest: str = Field(min_length=64, max_length=64) + + +class McpServerStatus(Contract): + enabled: bool = False + status: PluginHostState = PluginHostState.stopped + tools_count: int = 0 + protocol_version: str | None = None + remote_server_name: str | None = None + remote_server_version: str | None = None + error: str | None = None + last_tested_at: datetime | None = None + 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 dc3e62f..8247edd 100644 --- a/backend/app/extensions/mcp.py +++ b/backend/app/extensions/mcp.py @@ -10,14 +10,19 @@ import asyncio import json import os import queue +import re +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 @@ -45,6 +50,7 @@ MAX_MCP_MESSAGE_BYTES = 2 * 1024 * 1024 MAX_MCP_TOOL_RESULT_BYTES = 256 * 1024 MAX_MCP_TOOLS = 500 MAX_MCP_LIST_PAGES = 100 +_SSE_NEWLINE = re.compile(rb"\r\n?|\n") class McpBridgeError(RuntimeError): @@ -74,12 +80,14 @@ class McpStdioClient: command: list[str], *, cwd: Path, + environment: dict[str, str] | None = None, on_seen: Callable[[], None], on_broken: Callable[[str], None], on_tools_changed: Callable[[], None], ) -> None: self.command = command self.cwd = cwd + self.environment = environment or {} self.on_seen = on_seen self.on_broken = on_broken self.on_tools_changed = on_tools_changed @@ -97,8 +105,14 @@ 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") try: self.process = subprocess.Popen( @@ -114,6 +128,7 @@ class McpStdioClient: shell=False, env=environment, creationflags=creation_flags, + start_new_session=os.name != "nt", ) except OSError as exc: raise McpBridgeError( @@ -133,7 +148,7 @@ class McpStdioClient: timeout_code: str, response_error_code: str = "MCP_TOOL_CALL_FAILED", ) -> dict[str, Any]: - request_id, pending = self.begin_request(method, params) + request_id, pending = self.begin_request(method, params, timeout=timeout) return self.wait_response( request_id, pending, @@ -143,7 +158,7 @@ class McpStdioClient: ) def begin_request( - self, method: str, params: dict[str, Any] + self, method: str, params: dict[str, Any], *, timeout: float | None = None ) -> tuple[int, _PendingRequest]: self._ensure_running() with self._pending_lock: @@ -180,7 +195,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: @@ -213,9 +230,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,避免线程 @@ -240,15 +255,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 @@ -310,14 +327,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: @@ -325,7 +345,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 ) @@ -360,13 +382,498 @@ 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], + startup_timeout_seconds: float = 15, + 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() + self._startup_timeout_seconds = startup_timeout_seconds + + 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, timeout=timeout) + 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], *, timeout: float | None = None + ) -> 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, timeout), + 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(), + timeout=min(self._startup_timeout_seconds, 5), + ) + 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], + timeout: float | None, + ) -> None: + try: + response = self._post(message, timeout=timeout) + 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: + timeout = ( + self._startup_timeout_seconds + if message.get("method") == "notifications/initialized" + else 10 + ) + response = self._post(message, timeout=timeout) + 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(), + timeout=timeout, + ) + 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=self._startup_timeout_seconds) + 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], + timeout: float | None, + ) -> None: + try: + response = self._post(message, timeout=timeout) + 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=(",", ":")) + 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._endpoint, + content=encoded.encode("utf-8"), + headers={ + **self.headers, + "Accept": "application/json, text/event-stream", + "Content-Type": "application/json", + }, + timeout=timeout, + ) + 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)) + if not self._stopping: + self.on_broken("Legacy MCP SSE stream ended unexpectedly.") + 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], *, timeout: float | None = None + ) -> 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 转换。""" @@ -383,19 +890,31 @@ class McpBridge: package_path: Path, declared_permissions: list[str], on_unavailable: Callable[[str, str], None], + *, + 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 = 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, @@ -405,7 +924,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") @@ -415,15 +934,36 @@ 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, - 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 {}, + startup_timeout_seconds=backend.startup_timeout_seconds, + 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: @@ -466,15 +1006,28 @@ 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 + plugin_id, client, backend, declared_permissions, tool_source ) + if status.status == PluginHostState.unhealthy: + raise McpBridgeError( + "PLUGIN_HOST_UNAVAILABLE", + status.error + or "MCP event stream became unavailable during startup.", + status_code=503, + ) 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: @@ -502,7 +1055,9 @@ class McpBridge: ) -> Any: host = self._host(plugin_id) rpc_id, pending = host.client.begin_request( - "tools/call", {"name": remote_name, "arguments": arguments} + "tools/call", + {"name": remote_name, "arguments": arguments}, + timeout=host.backend.tool_timeout_seconds, ) call_key = (plugin_id, request_id) with self._lock: @@ -519,9 +1074,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: @@ -531,7 +1084,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( @@ -598,9 +1153,10 @@ class McpBridge: def _discover_tools( self, plugin_id: str, - client: McpStdioClient, + client: _McpClient, backend: PluginBackend, declared_permissions: list[str], + tool_source: str, ) -> list[McpDiscoveredTool]: discovered: list[McpDiscoveredTool] = [] cursor: str | None = None @@ -616,11 +1172,12 @@ 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( - self._map_tool(plugin_id, raw, declared_permissions) + self._map_tool(plugin_id, raw, declared_permissions, tool_source) ) if len(discovered) > MAX_MCP_TOOLS: raise McpBridgeError( @@ -632,7 +1189,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: @@ -648,7 +1206,10 @@ class McpBridge: @staticmethod def _map_tool( - plugin_id: str, raw: Any, declared_permissions: list[str] + plugin_id: str, + raw: Any, + declared_permissions: list[str], + tool_source: str = "plugin", ) -> McpDiscoveredTool: if not isinstance(raw, dict): raise McpBridgeError( @@ -663,9 +1224,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 ) ): @@ -690,7 +1249,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 @@ -709,10 +1270,12 @@ 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="plugin", + source=tool_source, ), ) @@ -789,3 +1352,179 @@ 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 _bounded_sse_lines(response: httpx.Response): + """Split UTF-8 lines without httpx.iter_lines()'s unbounded line buffer. + + Check each segment before appending it, including partial/no-newline input. + SSE allows LF, CR and CRLF; a CRLF pair can span network chunks. + """ + + pending = bytearray() + event_size = 0 + skip_lf = False + first_line = True + for chunk in response.iter_bytes(): + offset = 0 + if skip_lf and chunk: + offset = int(chunk.startswith(b"\n")) + skip_lf = False + for match in _SSE_NEWLINE.finditer(chunk, offset): + start, end = match.span() + segment = memoryview(chunk)[offset:start] + if event_size + len(pending) + len(segment) + 1 > MAX_MCP_MESSAGE_BYTES: + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", "MCP SSE event is too large." + ) + pending.extend(segment) + line = pending.decode("utf-8", errors="replace") + event_size += len(pending) + 1 + pending.clear() + if first_line: + line = line.removeprefix("\ufeff") + first_line = False + if not line: + event_size = 0 + yield line + skip_lf = chunk[end - 1 : end] == b"\r" and end == len(chunk) + offset = end + tail = memoryview(chunk)[offset:] + if event_size + len(pending) + len(tail) > MAX_MCP_MESSAGE_BYTES: + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", "MCP SSE event is too large." + ) + pending.extend(tail) + if pending: + line = pending.decode("utf-8", errors="replace") + yield line.removeprefix("\ufeff") if first_line else line + + +def _iter_sse(response: httpx.Response): + event = "message" + event_id: str | None = None + data_lines: list[str] = [] + size = 0 + for line in _bounded_sse_lines(response): + 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 new file mode 100644 index 0000000..a81aced --- /dev/null +++ b/backend/app/extensions/mcp_registry.py @@ -0,0 +1,1010 @@ +"""Independent, user-managed MCP server registry for development builds.""" + +from __future__ import annotations + +import hashlib +import json +import re +import threading +from datetime import UTC, datetime +from functools import wraps +from pathlib import Path +from typing import Any, Literal +from urllib.parse import urlsplit +from uuid import uuid4 + +from pydantic import BaseModel, ConfigDict, Field, ValidationError, create_model + +from app.agent.permissions import KNOWN_PERMISSIONS +from app.agent.tools import ToolExecutionContext, ToolRegistry +from app.contracts import ( + McpServer, + McpServerConfig, + McpServerCreateRequest, + McpServerSecretStatus, + McpServerTransport, + McpServerUpdateRequest, + McpToolSummary, + PluginBackend, + PluginHostState, +) +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", +} +_SERVER_ID = re.compile(r"^[A-Za-z0-9][A-Za-z0-9_-]{0,63}$") +_MAX_MCP_SERVERS = 256 + + +class _McpConnectionBackend(PluginBackend): + """Bridge adapter for the independent server's float timeout contract. + + Plugin manifests retain their integer/60-second startup restrictions. + Reusing that validation here used to reject valid 120-second server configs. + """ + + startup_timeout_seconds: float = Field(default=15, ge=1, le=120) + tool_timeout_seconds: float = Field(default=30, ge=1, le=300) + + +class _McpServerRecord(McpServerConfig): + """Validated on-disk representation with defaults for older C.1 records.""" + + version: int = Field(default=1, ge=1) + secret_environment_version: Literal[1, 2] = 1 + enabled: bool = False + approved_digest: str | None = Field( + default=None, min_length=64, max_length=64, pattern=r"^[0-9a-f]{64}$" + ) + tested_digest: str | None = Field( + default=None, min_length=64, max_length=64, pattern=r"^[0-9a-f]{64}$" + ) + last_tested_at: datetime | None = None + last_test_succeeded: bool | None = None + + +class McpRegistryError(RuntimeError): + def __init__(self, code: str, message: str, *, status_code: int = 422) -> None: + super().__init__(message) + self.code = code + self.message = message + 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.""" + + def __init__( + self, + registry: ToolRegistry, + credentials: EncryptedCredentialStore, + data_dir: Path, + *, + allow_process_launch: bool, + bridge: McpBridge | None = None, + ) -> None: + self.tools = registry + self.credentials = credentials + self.data_dir = data_dir + 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]] = {} + self._generations: dict[str, object] = {} + self._migrate_environment_secrets() + + def list(self) -> list[McpServer]: + with self._lock: + return [ + self._public(server_id, record) + for server_id, record in self._records.items() + ] + + def get(self, server_id: str) -> McpServer: + 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) + with self._lock: + if len(self._records) >= _MAX_MCP_SERVERS: + raise McpRegistryError( + "MCP_SERVER_LIMIT_REACHED", + f"At most {_MAX_MCP_SERVERS} MCP servers can be configured.", + status_code=409, + ) + server_id = uuid4().hex[:12] + record = request.model_dump(mode="json") + record["name"] = request.name.strip() + record["command"] = request.command.strip() if request.command else None + record["url"] = request.url.strip() if request.url else None + record.update( + version=1, + secret_environment_version=2, + 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_secret_ids = [ + secret_id + 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 secret_id in self._secret_ids(server_id, old_keys, kind) + - self._secret_ids(server_id, new_keys, kind) + ] + try: + self.credentials.delete_many(removed_secret_ids) + except CredentialStoreError as exc: + raise McpRegistryError( + "MCP_SECRET_STORE_ERROR", str(exc), status_code=500 + ) from exc + with self._lock: + record = request.model_dump(mode="json", exclude={"version"}) + record["name"] = request.name.strip() + 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, + secret_environment_version=2, + 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) + self._summaries.pop(server_id, None) + 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 = [ + secret_id + for kind, keys in ( + ("environment", record.get("secret_environment_keys", [])), + ("header", record.get("secret_header_keys", [])), + ) + for secret_id in self._secret_ids(server_id, keys, kind) + ] + try: + self.credentials.delete_many(secret_ids) + except CredentialStoreError as exc: + raise McpRegistryError( + "MCP_SECRET_STORE_ERROR", str(exc), status_code=500 + ) from exc + with self._lock: + 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) + 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) + current = self._digest(record) + if command_digest != current: + raise McpRegistryError( + "MCP_TRUST_DIGEST_STALE", + "MCP server configuration changed; review it again.", + status_code=409, + ) + approved = {**record, "approved_digest": current} + updated = {**self._records, server_id: approved} + self._write(updated) + self._records = updated + return self.get(server_id) + + @_serialized_lifecycle + def put_secret( + self, server_id: str, key: str, secret: str, *, kind: str = "environment" + ) -> McpServerSecretStatus: + with self._lock: + record = self._record(server_id) + 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.", + ) + if record.get("enabled"): + self.disable(server_id) + self._invalidate_test(server_id) + try: + 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) + + @_serialized_lifecycle + def delete_secret( + self, server_id: str, key: str, *, kind: str = "environment" + ) -> McpServerSecretStatus: + record = self._record(server_id) + 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.", + ) + if record.get("enabled"): + self.disable(server_id) + self._invalidate_test(server_id) + try: + 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"): + raise McpRegistryError( + "MCP_SERVER_ALREADY_ENABLED", + "Disable the MCP server before running an isolated connection test.", + status_code=409, + ) + self._require_launch_allowed(record, require_test=False) + try: + discovered = self._start(server_id, record) + except Exception as exc: + self._generations.pop(server_id, None) + self.bridge.remove(self._host_id(server_id)) + tested_at = datetime.now(UTC) + failure = { + "status": PluginHostState.error, + "error": str(exc), + "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), + "protocol_version": status.protocol_version, + "remote_server_name": status.server_name, + "remote_server_version": status.server_version, + "error": None, + "last_tested_at": tested_at, + "last_test_succeeded": True, + } + self._summaries[server_id] = self._tool_summaries(discovered) + self._generations.pop(server_id, None) + 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, 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: + self._register(server_id, item) + registered.append(item.definition.name) + except Exception: + for name in registered: + self.tools.unregister(name) + self._generations.pop(server_id, None) + self.bridge.stop(self._host_id(server_id)) + raise + try: + with self._lock: + enabled_record = {**record, "enabled": True} + updated = {**self._records, server_id: enabled_record} + self._write(updated) + self._records = updated + self._registered[server_id] = registered + except McpRegistryError: + for name in registered: + self.tools.unregister(name) + self._generations.pop(server_id, None) + self.bridge.stop(self._host_id(server_id)) + raise + return self.get(server_id) + + @_serialized_lifecycle + def disable(self, server_id: str) -> McpServer: + with self._lock: + record = self._record(server_id) + disabled_record = {**record, "enabled": False} + updated = {**self._records, server_id: disabled_record} + self._write(updated) + self._records = updated + for name in self._registered.pop(server_id, []): + self.tools.unregister(name) + self._generations.pop(server_id, None) + 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 + for server_id, record in list(self._records.items()): + if record.get("enabled"): + try: + self.enable(server_id) + except (McpRegistryError, ValueError, OSError) as exc: + self._records[server_id] = {**record, "enabled": False} + self._last_status[server_id] = { + "status": PluginHostState.error, + "error": str(exc), + } + self._write() + + @_serialized_lifecycle + def shutdown(self) -> None: + for server_id in list(self._records): + self._generations.pop(server_id, None) + for name in self._registered.pop(server_id, []): + self.tools.unregister(name) + self.bridge.stop(self._host_id(server_id)) + + def _start(self, server_id: str, record: dict[str, Any]) -> list[McpDiscoveredTool]: + environment = dict(record.get("environment", {})) + for key in record.get("secret_environment_keys", []): + try: + 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 + ) from exc + if value is None: + raise McpRegistryError( + "MCP_SECRET_REQUIRED", + f"Secret environment variable is not configured: {key}", + 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) + # A queued callback from the previous process must not affect its replacement. + generation = object() + self._generations[server_id] = generation + self.bridge.remove(host_id) + try: + return self.bridge.start( + host_id, + self._backend(record), + self._server_dir(server_id), + list(record.get("permissions", [])), + lambda _host, message: self._unavailable( + server_id, generation, message + ), + 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: + self._generations.pop(server_id, None) + self.bridge.remove(host_id) + raise McpRegistryError( + exc.code, exc.message, status_code=exc.status_code + ) from exc + + def _register(self, server_id: str, discovered: McpDiscoveredTool) -> None: + definition = discovered.definition + model_name = "McpArgs_" + re.sub(r"\W+", "_", definition.name) + arguments_model = create_model(model_name, __config__=ConfigDict(extra="allow")) + + async def executor(arguments: BaseModel, context: ToolExecutionContext) -> Any: + return await self.bridge.call_tool( + self._host_id(server_id), + discovered.remote_name, + arguments.model_dump(exclude_unset=True), + request_id=context.tool_call_id + or f"{context.run_id}:{definition.name}", + ) + + self.tools.register(definition, arguments_model, executor) + + def _unavailable(self, server_id: str, generation: object, message: str) -> None: + # A failure may race with enable(). Waiting for the lifecycle mutation makes + # sure tools registered immediately before the callback are also removed. + with self._lifecycle_lock: + if self._generations.get(server_id) is not generation: + return + self._generations.pop(server_id, None) + try: + with self._lock: + record = self._records.get(server_id) + registered = self._registered.pop(server_id, []) + for name in registered: + self.tools.unregister(name) + if record is not None and (record.get("enabled") or registered): + self._records[server_id] = {**record, "enabled": False} + self._last_status[server_id] = { + "status": PluginHostState.unhealthy, + "error": message, + } + self._write() + finally: + # broken() can run on the client's reader/event thread. stop() does + # not join that thread, and setting _stopping before closing the + # transport prevents the close itself from reporting another failure. + self.bridge.remove(self._host_id(server_id)) + + 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") != digest: + raise McpRegistryError( + "MCP_TRUST_APPROVAL_REQUIRED", + "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, + ) + + def _public(self, server_id: str, record: dict[str, Any]) -> McpServer: + digest = self._digest(record) + backend = self._backend(record) + status = self.bridge.status(self._host_id(server_id), backend) + 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.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, + enabled=bool(record.get("enabled")), + trusted=record.get("approved_digest") == digest, + command_digest=digest, + command_summary=self._summary(record), + status=status.status + if record.get("enabled") + else cached.get("status", PluginHostState.stopped), + tools_count=status.tools_count + if record.get("enabled") + else cached.get("tools_count", 0), + protocol_version=status.protocol_version + if record.get("enabled") + else cached.get("protocol_version"), + remote_server_name=status.server_name + if record.get("enabled") + else cached.get("remote_server_name"), + remote_server_version=status.server_version + if record.get("enabled") + else cached.get("remote_server_version"), + error=status.error if record.get("enabled") else cached.get("error"), + 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: + if not request.name.strip(): + raise McpRegistryError( + "MCP_SERVER_NAME_INVALID", "MCP server name cannot be blank." + ) + 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): + raise McpRegistryError( + "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( + "MCP_PERMISSION_INVALID", + f"Unknown MCP permission: {min(unknown_permissions)}", + ) + + @staticmethod + def _validate_environment_key(key: str) -> None: + if not _ENVIRONMENT_KEY.fullmatch(key): + raise McpRegistryError( + "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 _McpConnectionBackend( + type="mcp", + transport="stdio", + 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), + ) + + @staticmethod + def _host_id(server_id: str) -> str: + return f"mcp.{server_id}" + + def _server_dir(self, server_id: str) -> Path: + path = self.data_dir / "mcp" / "workdirs" / server_id + path.mkdir(parents=True, exist_ok=True) + return path + + @staticmethod + def _digest(record: dict[str, Any]) -> str: + executable = { + key: record.get(key) + for key in ( + "transport", + "command", + "args", + "environment", + "secret_environment_keys", + "url", + "headers", + "secret_header_keys", + "permissions", + ) + } + return hashlib.sha256( + json.dumps( + executable, sort_keys=True, ensure_ascii=False, separators=(",", ":") + ).encode() + ).hexdigest() + + @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.get("command") or "", + *[ + json.dumps(arg, ensure_ascii=False) + for arg in record.get("args", []) + ], + ] + ) + + @staticmethod + def _secret_id(server_id: str, key: str, kind: str = "environment") -> str: + identity = ( + f"environment-v2\0{key}" + if kind == "environment" + else f"{kind}\0{key.casefold()}" + ) + suffix = hashlib.sha256(identity.encode()).hexdigest()[:20] + return f"mcp.{server_id}.{suffix}" + + @staticmethod + def _legacy_environment_secret_id(server_id: str, key: str) -> str: + suffix = hashlib.sha256(f"environment\0{key.casefold()}".encode()).hexdigest()[ + :20 + ] + return f"mcp.{server_id}.{suffix}" + + def _secret_ids(self, server_id: str, keys: list[str], kind: str) -> set[str]: + ids = {self._secret_id(server_id, key, kind) for key in keys} + if kind == "environment": + # Include retained ambiguous legacy ciphertext when its last declaration is removed. + ids.update( + self._legacy_environment_secret_id(server_id, key) for key in keys + ) + return ids + + def _migrate_environment_secrets(self) -> None: + """迁移旧的大小写折叠 ID;已碰撞的值无法恢复,保留原密文并要求重新录入。""" + + replacements: dict[str, str] = {} + ambiguous: dict[str, list[str]] = {} + legacy_records = { + server_id: record + for server_id, record in self._records.items() + if record.get("secret_environment_version", 1) == 1 + } + if not legacy_records: + return + for server_id, record in legacy_records.items(): + groups: dict[str, set[str]] = {} + for key in record.get("secret_environment_keys", []): + groups.setdefault(key.casefold(), set()).add(key) + for keys in groups.values(): + key = next(iter(keys)) + legacy_id = self._legacy_environment_secret_id(server_id, key) + if len(keys) == 1: + replacements[legacy_id] = self._secret_id(server_id, key) + else: + ambiguous.setdefault(server_id, []).extend(keys) + try: + self.credentials.move_many(replacements) + for server_id, keys in ambiguous.items(): + if not any( + self.credentials.has( + self._legacy_environment_secret_id(server_id, key) + ) + for key in keys + ): + continue + if all( + self.credentials.has(self._secret_id(server_id, key)) + for key in keys + ): + continue + self._records[server_id].update( + enabled=False, + tested_digest=None, + last_test_succeeded=None, + last_tested_at=None, + ) + self._last_status[server_id] = { + "status": PluginHostState.error, + "error": "环境变量密钥名称曾发生大小写冲突,请分别重新录入密钥并测试连接。", + } + # Persist a migration marker even when legacy values were ambiguous. + # Otherwise a later key removal could make that old shared value look + # unambiguous and resurrect a deleted credential on the next restart. + for server_id in legacy_records: + self._records[server_id]["secret_environment_version"] = 2 + self._write() + except CredentialStoreError as exc: + raise McpRegistryError( + "MCP_SECRET_STORE_ERROR", str(exc), status_code=500 + ) from exc + + def _secret_configured( + self, server_id: str, key: str, kind: str = "environment" + ) -> bool: + try: + 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] + except KeyError as exc: + raise McpRegistryError( + "MCP_SERVER_NOT_FOUND", + "MCP server configuration was not found.", + status_code=404, + ) from exc + + def _invalidate_test(self, server_id: str) -> None: + """Make credential changes safe before touching the encrypted store.""" + + with self._lock: + record = self._record(server_id) + invalidated = { + **record, + "tested_digest": None, + "last_tested_at": None, + "last_test_succeeded": None, + } + updated = {**self._records, server_id: invalidated} + self._write(updated) + self._records = updated + self._last_status.pop(server_id, None) + self._summaries.pop(server_id, None) + + @property + def _path(self) -> Path: + return self.data_dir / "mcp" / "servers.json" + + def _read(self) -> dict[str, dict[str, Any]]: + if not self._path.exists(): + return {} + try: + value = json.loads(self._path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError) as exc: + raise McpRegistryError( + "MCP_REGISTRY_INVALID", + "MCP server registry cannot be loaded.", + status_code=500, + ) from exc + if not isinstance(value, dict): + raise McpRegistryError( + "MCP_REGISTRY_INVALID", + "MCP server registry has an invalid format.", + status_code=500, + ) + if len(value) > _MAX_MCP_SERVERS: + raise McpRegistryError( + "MCP_REGISTRY_INVALID", + "MCP server registry contains too many records.", + status_code=500, + ) + normalized: dict[str, dict[str, Any]] = {} + config_fields = set(McpServerConfig.model_fields) + try: + for server_id, raw in value.items(): + if not isinstance(server_id, str) or not _SERVER_ID.fullmatch( + server_id + ): + raise ValueError("invalid server id") + record = _McpServerRecord.model_validate(raw) + config = record.model_dump(mode="json", include=config_fields) + self._validate(McpServerCreateRequest.model_validate(config)) + normalized[server_id] = record.model_dump(mode="json") + except (McpRegistryError, ValidationError, ValueError, TypeError) as exc: + raise McpRegistryError( + "MCP_REGISTRY_INVALID", + "MCP server registry contains an invalid record.", + status_code=500, + ) from exc + return normalized + + def _write(self, records: dict[str, dict[str, Any]] | None = None) -> None: + temporary = self._path.with_suffix(".tmp") + try: + self._path.parent.mkdir(parents=True, exist_ok=True) + temporary.write_text( + json.dumps( + records if records is not None else self._records, + ensure_ascii=False, + indent=2, + sort_keys=True, + ), + encoding="utf-8", + ) + temporary.replace(self._path) + except OSError as exc: + temporary.unlink(missing_ok=True) + raise McpRegistryError( + "MCP_REGISTRY_WRITE_FAILED", + "MCP server registry cannot be written.", + status_code=500, + ) from exc diff --git a/backend/app/main.py b/backend/app/main.py index 8a3aae5..bfaaf91 100644 --- a/backend/app/main.py +++ b/backend/app/main.py @@ -19,6 +19,7 @@ async def lifespan(_: FastAPI): yield # 第三方 MCP Server 必须跟随 AI Core 退出,不能遗留孤儿进程。 container.plugins.shutdown() + container.mcp_servers.shutdown() app = FastAPI( diff --git a/backend/app/providers/credentials.py b/backend/app/providers/credentials.py index 88d9485..8eec9a2 100644 --- a/backend/app/providers/credentials.py +++ b/backend/app/providers/credentials.py @@ -5,15 +5,15 @@ import os import re import threading from pathlib import Path -from typing import Protocol +from typing import ClassVar, Protocol from cryptography.fernet import Fernet, InvalidToken from app.config import get_settings - _CREDENTIAL_ID = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._-]{0,127}$") _PLUGIN_CREDENTIAL_PREFIX = "plugin." +_MCP_CREDENTIAL_PREFIX = "mcp." class CredentialStoreError(RuntimeError): @@ -27,16 +27,18 @@ class CredentialResolver(Protocol): def validate_provider_credential_id(credential_id: str | None) -> None: """阻止 Provider 和通用凭据 API 跨入 Plugin 私有命名空间。""" - if credential_id and credential_id.casefold().startswith( - _PLUGIN_CREDENTIAL_PREFIX - ): - raise CredentialStoreError("Credential namespace is reserved for Plugin settings.") + if credential_id and credential_id.casefold().startswith(_PLUGIN_CREDENTIAL_PREFIX): + raise CredentialStoreError( + "Credential namespace is reserved for Plugin settings." + ) + if credential_id and credential_id.casefold().startswith(_MCP_CREDENTIAL_PREFIX): + raise CredentialStoreError("Credential namespace is reserved for MCP settings.") class EnvironmentCredentialResolver: """解析由桌面 Host 注入 Sidecar 进程的临时凭证上下文。""" - _development_aliases = { + _development_aliases: ClassVar[dict[str, str]] = { "openai": "OPENAI_API_KEY", "deepseek": "DEEPSEEK_API_KEY", } @@ -84,7 +86,9 @@ class EncryptedCredentialStore: try: return Fernet(environment_key.encode("ascii")) except (ValueError, UnicodeEncodeError) as exc: - raise CredentialStoreError("APP_CREDENTIAL_MASTER_KEY is invalid.") from exc + raise CredentialStoreError( + "APP_CREDENTIAL_MASTER_KEY is invalid." + ) from exc key_path.parent.mkdir(parents=True, exist_ok=True) self._restrict(key_path.parent, 0o700) @@ -101,7 +105,9 @@ class EncryptedCredentialStore: try: return Fernet(key_path.read_bytes().strip()) except (OSError, ValueError) as exc: - raise CredentialStoreError("Credential master key cannot be loaded.") from exc + raise CredentialStoreError( + "Credential master key cannot be loaded." + ) from exc def _read_tokens(self) -> dict[str, str]: _, store_path = self._paths() @@ -110,11 +116,16 @@ class EncryptedCredentialStore: try: data = json.loads(store_path.read_text(encoding="utf-8")) except (OSError, json.JSONDecodeError) as exc: - raise CredentialStoreError("Encrypted credential store cannot be loaded.") from exc + raise CredentialStoreError( + "Encrypted credential store cannot be loaded." + ) from exc if not isinstance(data, dict) or not all( - isinstance(key, str) and isinstance(value, str) for key, value in data.items() + isinstance(key, str) and isinstance(value, str) + for key, value in data.items() ): - raise CredentialStoreError("Encrypted credential store has an invalid format.") + raise CredentialStoreError( + "Encrypted credential store has an invalid format." + ) return data def _write_tokens(self, tokens: dict[str, str]) -> None: @@ -195,6 +206,22 @@ class EncryptedCredentialStore: self._write_tokens(tokens) return removed + def move_many(self, replacements: dict[str, str]) -> None: + """原子迁移凭据 ID,直接移动密文且不覆盖已经写入的新凭据。""" + + for old_id, new_id in replacements.items(): + self._validate_id(old_id) + self._validate_id(new_id) + with self._lock: + tokens = self._read_tokens() + changed = False + for old_id, new_id in replacements.items(): + if old_id != new_id and old_id in tokens: + tokens.setdefault(new_id, tokens.pop(old_id)) + changed = True + if changed: + self._write_tokens(tokens) + class ChainedCredentialResolver: def __init__(self, *resolvers: CredentialResolver) -> None: diff --git a/backend/app/routes.py b/backend/app/routes.py index 962bec7..ec8908f 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, @@ -21,6 +23,14 @@ from app.contracts import ( IndexJob, IndexRebuildRequest, IndexStatus, + McpServer, + McpServerCreateRequest, + McpServerListResponse, + McpServerSecretStatus, + McpServerSecretWriteRequest, + McpServerTrustRequest, + McpServerUpdateRequest, + McpToolSummaryListResponse, ModelEvent, ModelEventType, Note, @@ -68,17 +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.providers.registry import ProviderNotFoundError -from app.providers.factory import UnsupportedProviderError +from app.extensions.mcp_registry import McpRegistryError 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, @@ -91,6 +100,14 @@ from app.services import ( router = APIRouter(prefix="/api") +async def mcp_call_async(operation): + """Even registry reads can wait on lifecycle locks; keep all MCP work off the event loop.""" + try: + return await asyncio.to_thread(operation) + except McpRegistryError as exc: + raise ApiError(exc.status_code, exc.code, exc.message) from exc + + def utc_now() -> datetime: return datetime.now(timezone.utc) @@ -202,14 +219,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, ) @@ -217,7 +241,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 @@ -231,7 +257,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") @@ -275,7 +303,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()) @@ -436,9 +466,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)) @@ -478,7 +506,120 @@ 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 +@router.get("/mcp/servers", response_model=McpServerListResponse, tags=["MCP Servers"]) +async def list_mcp_servers() -> McpServerListResponse: + return McpServerListResponse(items=await mcp_call_async(container.mcp_servers.list)) + + +@router.post( + "/mcp/servers", response_model=McpServer, status_code=201, tags=["MCP Servers"] +) +async def create_mcp_server(request: McpServerCreateRequest) -> McpServer: + return await mcp_call_async(lambda: container.mcp_servers.create(request)) + + +@router.get("/mcp/servers/{server_id}", response_model=McpServer, tags=["MCP Servers"]) +async def get_mcp_server(server_id: str) -> McpServer: + return await mcp_call_async(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=await mcp_call_async(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) + ) + + +@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" + ) + + +@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 await mcp_call_async( + lambda: container.mcp_servers.trust(server_id, request.command_digest) + ) + + +@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"] +) +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"] +) +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, + kind: str = Query(default="environment", pattern="^(environment|header)$"), +) -> McpServerSecretStatus: + return await mcp_call_async( + 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, + kind: str = Query(default="environment", pattern="^(environment|header)$"), +) -> McpServerSecretStatus: + return await mcp_call_async( + lambda: container.mcp_servers.delete_secret(server_id, key, kind=kind) + ) # Plugins @@ -570,11 +711,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 @@ -649,9 +794,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) ) @@ -764,7 +907,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 @@ -793,7 +938,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) @@ -871,7 +1018,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 @@ -887,7 +1036,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") @@ -937,5 +1088,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_api.py b/backend/tests/test_api.py index 3d5d2d5..97d7fee 100644 --- a/backend/tests/test_api.py +++ b/backend/tests/test_api.py @@ -1,26 +1,12 @@ import asyncio +import threading +from types import SimpleNamespace + +import pytest -from app.main import health, service_status -from app.routes import ( - get_index_status, - list_notes, - list_plugins, - list_provider_presets, - list_providers, - list_skills, -) -from app.routes import ( - create_provider, - create_task, - delete_provider, - delete_task, - get_provider, - get_task, - list_tasks, - update_provider, - update_task, -) from app.contracts import ( + McpServerSecretStatus, + McpServerSecretWriteRequest, ProviderCreateRequest, ProviderType, ProviderUpdateRequest, @@ -28,6 +14,61 @@ from app.contracts import ( TaskStatus, TaskUpdateRequest, ) +from app.main import health, service_status +from app.routes import ( + create_provider, + create_task, + delete_provider, + delete_task, + get_index_status, + get_provider, + get_task, + list_notes, + list_plugins, + list_provider_presets, + list_providers, + list_skills, + list_tasks, + update_provider, + update_task, +) + + +def test_mcp_secret_routes_offload_blocking_lifecycle_work(monkeypatch) -> None: + from app import routes + + caller_thread = threading.get_ident() + worker_threads: list[int] = [] + + class FakeMcpRegistry: + def put_secret(self, server_id, key, secret, *, kind): + worker_threads.append(threading.get_ident()) + return McpServerSecretStatus(key=key, configured=True) + + def delete_secret(self, server_id, key, *, kind): + worker_threads.append(threading.get_ident()) + return McpServerSecretStatus(key=key, configured=False) + + monkeypatch.setattr( + routes, + "container", + SimpleNamespace(mcp_servers=FakeMcpRegistry()), + ) + written = asyncio.run( + routes.put_mcp_server_secret( + "server-1", + "TOKEN", + McpServerSecretWriteRequest(secret="hidden"), + kind="environment", + ) + ) + deleted = asyncio.run( + routes.delete_mcp_server_secret("server-1", "TOKEN", kind="environment") + ) + + assert written.configured is True + assert deleted.configured is False + assert worker_threads and all(item != caller_thread for item in worker_threads) def test_health() -> None: @@ -36,6 +77,172 @@ def test_health() -> None: assert response.model_dump() == {"status": "ok"} +def test_mcp_create_and_trust_are_not_executed_on_event_loop(monkeypatch) -> None: + from app import routes + from app.contracts import McpServerCreateRequest, McpServerTrustRequest + + caller = threading.get_ident() + workers = [] + + class Registry: + def create(self, request): + workers.append(threading.get_ident()) + return "created" + + def trust(self, server_id, digest): + workers.append(threading.get_ident()) + return "trusted" + + monkeypatch.setattr(routes, "container", SimpleNamespace(mcp_servers=Registry())) + assert ( + asyncio.run( + routes.create_mcp_server(McpServerCreateRequest(name="test", command="uvx")) + ) + == "created" + ) + assert ( + asyncio.run( + routes.trust_mcp_server( + "test", McpServerTrustRequest(command_digest="a" * 64) + ) + ) + == "trusted" + ) + assert len(workers) == 2 + assert all(worker != caller for worker in workers) + + +def test_mcp_split_config_and_secret_requests_persist_without_plaintext( + monkeypatch, +) -> None: + from fastapi.testclient import TestClient + + from app import routes + from app.agent.tools import ToolRegistry + from app.config import get_settings + from app.extensions.mcp_registry import McpServerRegistry + from app.main import app + from app.providers.credentials import EncryptedCredentialStore + + service = McpServerRegistry( + ToolRegistry(), + EncryptedCredentialStore(), + get_settings().data_dir, + allow_process_launch=True, + ) + monkeypatch.setattr(routes, "container", SimpleNamespace(mcp_servers=service)) + client = TestClient(app) + config = { + "name": "MiniMax configuration test", + "command": "uvx", + "environment": {"MINIMAX_API_HOST": "https://api.minimaxi.com"}, + "secret_environment_keys": ["MINIMAX_API_KEY"], + "startup_timeout_seconds": 120, + "tool_timeout_seconds": 300, + } + # Reproduce the old frontend payload. The backend still enforces separation. + invalid = client.post( + "/api/mcp/servers", + json={ + **config, + "environment": { + **config["environment"], + "MINIMAX_API_KEY": "synthetic-only", + }, + }, + ) + assert invalid.status_code == 422 + assert invalid.json()["error"]["code"] == "MCP_ENVIRONMENT_INVALID" + created = client.post("/api/mcp/servers", json=config) + assert created.status_code == 201 + server_id = created.json()["server_id"] + saved = client.put( + f"/api/mcp/servers/{server_id}/secrets/MINIMAX_API_KEY", + json={"secret": "synthetic-only"}, + ) + assert saved.status_code == 200 + current = client.get(f"/api/mcp/servers/{server_id}") + assert current.json()["secret_environment"] == {"MINIMAX_API_KEY": True} + assert "synthetic-only" not in current.text + assert "synthetic-only" not in service._path.read_text(encoding="utf-8") + _, credentials_path = service.credentials._paths() + assert "synthetic-only" not in credentials_path.read_text(encoding="utf-8") + assert not current.json()["enabled"] # Saving never starts a third-party process. + client.close() + + +@pytest.mark.parametrize("operation", ["create", "trust"]) +def test_mcp_lifecycle_lock_contention_keeps_event_loop_responsive( + monkeypatch, + operation, +) -> None: + from app import routes + from app.agent.tools import ToolRegistry + from app.config import get_settings + from app.contracts import McpServerCreateRequest, McpServerTrustRequest + from app.extensions.mcp_registry import McpServerRegistry + from app.providers.credentials import EncryptedCredentialStore + + service = McpServerRegistry( + ToolRegistry(), + EncryptedCredentialStore(), + get_settings().data_dir, + allow_process_launch=True, + ) + request = McpServerCreateRequest( + name="Lock contention fixture", command="not-executed" + ) + server = service.create(request) + monkeypatch.setattr(routes, "container", SimpleNamespace(mcp_servers=service)) + entered = threading.Event() + locked = threading.Event() + release = threading.Event() + original = getattr(service, operation) + + def observed(*args): + entered.set() + return original(*args) + + def hold_lifecycle_lock(): + with service._lifecycle_lock: + locked.set() + release.wait(timeout=5) + + monkeypatch.setattr(service, operation, observed) + holder = threading.Thread(target=hold_lifecycle_lock, daemon=True) + holder.start() + # An independent watchdog lets the test fail rather than hang if a regression + # blocks the event loop itself (an asyncio timeout alone cannot catch that). + watchdog = threading.Timer(5, release.set) + watchdog.start() + + async def exercise(): + pending = asyncio.create_task( + routes.create_mcp_server(request) + if operation == "create" + else routes.trust_mcp_server( + server.server_id, + McpServerTrustRequest(command_digest=server.command_digest), + ) + ) + try: + assert await asyncio.to_thread(entered.wait, 2) + assert not pending.done() + assert not release.is_set() + assert (await health()).status == "ok" + finally: + release.set() + await pending + + try: + assert locked.wait(timeout=2) + asyncio.run(exercise()) + finally: + release.set() + watchdog.cancel() + holder.join(timeout=2) + + def test_service_status() -> None: response = asyncio.run(service_status()) @@ -52,7 +259,9 @@ def test_core_collections_are_typed() -> None: assert notes.items == [] assert notes.page.limit == 20 - assert [skill.manifest.skill_id for skill in skills.items] == ["knowledge-assistant"] + assert [skill.manifest.skill_id for skill in skills.items] == [ + "knowledge-assistant" + ] assert skills.items[0].status == "ready" assert [plugin.manifest.plugin_id for plugin in plugins.items] == ["text-tools"] assert plugins.items[0].status == "ready" @@ -73,9 +282,15 @@ def test_provider_presets_include_openai_and_deepseek() -> None: def test_provider_presets_static_route_precedes_provider_id_route() -> None: from app.routes import router - get_paths = [route.path for route in router.routes if "GET" in getattr(route, "methods", set())] + get_paths = [ + route.path + for route in router.routes + if "GET" in getattr(route, "methods", set()) + ] - assert get_paths.index("/api/providers/presets") < get_paths.index("/api/providers/{provider_id}") + assert get_paths.index("/api/providers/presets") < get_paths.index( + "/api/providers/{provider_id}" + ) def test_openapi_contains_documented_frontend_interfaces() -> None: @@ -101,6 +316,14 @@ def test_openapi_contains_documented_frontend_interfaces() -> None: "/api/plugins/{plugin_id}/settings/{key}/secret", "/api/plugins/{plugin_id}/enable", "/api/plugins/{plugin_id}/disable", + "/api/mcp/servers", + "/api/mcp/servers/{server_id}", + "/api/mcp/servers/{server_id}/tools", + "/api/mcp/servers/{server_id}/trust", + "/api/mcp/servers/{server_id}/test", + "/api/mcp/servers/{server_id}/enable", + "/api/mcp/servers/{server_id}/disable", + "/api/mcp/servers/{server_id}/secrets/{key}", "/api/providers/test", "/api/providers/presets", "/api/credentials/{credential_id}", diff --git a/backend/tests/test_mcp_registry.py b/backend/tests/test_mcp_registry.py new file mode 100644 index 0000000..2d141c0 --- /dev/null +++ b/backend/tests/test_mcp_registry.py @@ -0,0 +1,769 @@ +import asyncio +import hashlib +import json +import sys +import threading +import time +from concurrent.futures import ThreadPoolExecutor + +import httpx +import pytest + +from app.agent.tools import ToolExecutionContext, ToolRegistry +from app.config import BACKEND_DIR, get_settings +from app.contracts import McpServerCreateRequest, McpServerUpdateRequest, ToolCall +from app.extensions.mcp import McpLegacySseClient +from app.extensions.mcp_registry import McpRegistryError, McpServerRegistry +from app.providers.credentials import CredentialStoreError, EncryptedCredentialStore + +SERVER = BACKEND_DIR / "extensions" / "fixtures" / "mcp-echo" / "server.py" + + +def request(**overrides) -> McpServerCreateRequest: + values = { + "name": "Echo MCP", + "command": sys.executable, + "args": [str(SERVER)], + "permissions": ["notes.read", "secrets.use"], + "secret_environment_keys": ["TEST_MCP_SECRET"], + } + values.update(overrides) + return McpServerCreateRequest(**values) + + +def registry(*, launch: bool = True) -> McpServerRegistry: + return McpServerRegistry( + ToolRegistry(), + EncryptedCredentialStore(), + get_settings().data_dir, + allow_process_launch=launch, + ) + + +def test_registry_requires_current_trust_and_never_returns_secret() -> None: + service = registry() + created = service.create(request()) + assert created.trusted is False + assert created.secret_environment == {"TEST_MCP_SECRET": False} + + service.put_secret(created.server_id, "TEST_MCP_SECRET", "do-not-return") + configured = service.get(created.server_id) + assert configured.secret_environment == {"TEST_MCP_SECRET": True} + assert "do-not-return" not in configured.model_dump_json() + + with pytest.raises(McpRegistryError, match="approve"): + service.test(created.server_id) + + service.trust(created.server_id, created.command_digest) + tested = service.test(created.server_id) + assert tested.status == "stopped" + assert tested.last_test_succeeded is True + assert tested.tools_count > 0 + service.shutdown() + + +def test_secret_change_disables_server_and_requires_a_new_connection_test() -> None: + service = registry() + created = service.create(request()) + service.put_secret(created.server_id, "TEST_MCP_SECRET", "first") + service.trust(created.server_id, created.command_digest) + service.test(created.server_id) + service.enable(created.server_id) + + service.put_secret(created.server_id, "TEST_MCP_SECRET", "second") + current = service.get(created.server_id) + assert current.enabled is False + assert current.last_test_succeeded is None + assert not any( + item.name.startswith(f"mcp.{created.server_id}.") + for item in service.tools.definitions() + ) + with pytest.raises(McpRegistryError) as error: + service.enable(created.server_id) + assert error.value.code == "MCP_CONNECTION_TEST_REQUIRED" + service.shutdown() + + +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( + item.name.startswith(f"mcp.{created.server_id}.") + for item in service.tools.definitions() + ) + + updated = service.update( + created.server_id, + McpServerUpdateRequest( + **request(name="Changed", secret_environment_keys=[]).model_dump(), + version=enabled.version, + ), + ) + assert updated.enabled is False + assert updated.trusted is False + assert not any( + item.name.startswith(f"mcp.{created.server_id}.") + for item in service.tools.definitions() + ) + service.shutdown() + + +def test_update_remains_retryable_when_removed_secret_cleanup_fails( + monkeypatch, +) -> None: + service = registry() + created = service.create(request()) + service.put_secret(created.server_id, "TEST_MCP_SECRET", "keep-until-retry") + + def fail_delete_many(_secret_ids: list[str]) -> set[str]: + raise CredentialStoreError("credential store unavailable") + + monkeypatch.setattr(service.credentials, "delete_many", fail_delete_many) + with pytest.raises(McpRegistryError) as error: + service.update( + created.server_id, + McpServerUpdateRequest( + **request(secret_environment_keys=[]).model_dump(), + version=created.version, + ), + ) + + current = service.get(created.server_id) + assert error.value.code == "MCP_SECRET_STORE_ERROR" + assert current.version == created.version + assert current.secret_environment == {"TEST_MCP_SECRET": True} + service.shutdown() + + +def test_delete_keeps_server_retryable_when_secret_cleanup_fails(monkeypatch) -> None: + service = registry() + created = service.create(request()) + service.put_secret(created.server_id, "TEST_MCP_SECRET", "keep-until-retry") + + def fail_delete_many(_secret_ids: list[str]) -> set[str]: + raise CredentialStoreError("credential store unavailable") + + monkeypatch.setattr(service.credentials, "delete_many", fail_delete_many) + with pytest.raises(McpRegistryError) as error: + service.delete(created.server_id) + + current = service.get(created.server_id) + assert error.value.code == "MCP_SECRET_STORE_ERROR" + assert current.server_id == created.server_id + assert current.secret_environment == {"TEST_MCP_SECRET": True} + service.shutdown() + + +def test_unavailable_server_removes_bridge_host(monkeypatch) -> None: + service = registry() + created = service.create(request(secret_environment_keys=[])) + with service._lock: + service._records[created.server_id] = { + **service._records[created.server_id], + "enabled": True, + } + removed: list[str] = [] + monkeypatch.setattr(service.bridge, "remove", removed.append) + + generation = object() + service._generations[created.server_id] = generation + service._unavailable(created.server_id, generation, "connection lost") + + current = service.get(created.server_id) + assert removed == [f"mcp.{created.server_id}"] + assert current.enabled is False + assert current.status == "unhealthy" + service.shutdown() + + +def test_old_failure_callback_cannot_stop_replacement_host(monkeypatch) -> None: + service = registry() + callbacks = [] + original_start = service.bridge.start + + def capture_callback(*args, **kwargs): + callbacks.append(args[4]) + return original_start(*args, **kwargs) + + monkeypatch.setattr(service.bridge, "start", capture_callback) + created = service.create(request(secret_environment_keys=[])) + service.trust(created.server_id, created.command_digest) + callback_thread = None + try: + service.test(created.server_id) + service.enable(created.server_id) + old_callback = callbacks[-1] + callback_started = threading.Event() + callback_finished = threading.Event() + + def delayed_failure(): + callback_started.set() + old_callback(f"mcp.{created.server_id}", "delayed old failure") + callback_finished.set() + + # Queue the old callback while a replacement owns the lifecycle lock. + with service._lifecycle_lock: + callback_thread = threading.Thread(target=delayed_failure, daemon=True) + callback_thread.start() + assert callback_started.wait(timeout=2) + service.disable(created.server_id) + service.enable(created.server_id) + assert callback_finished.wait(timeout=2) + assert service.get(created.server_id).enabled is True + assert service.get(created.server_id).status == "ready" + assert service.tools.definitions() + callbacks[-1](f"mcp.{created.server_id}", "current failure") + assert service.get(created.server_id).enabled is False + assert service.get(created.server_id).status == "unhealthy" + finally: + service.shutdown() + if callback_thread is not None: + callback_thread.join(timeout=2) + + +def test_header_case_only_rename_preserves_secret() -> None: + service = registry() + config = { + "name": "HTTP", + "transport": "streamable_http", + "url": "https://example.test/mcp", + "secret_header_keys": ["Authorization"], + } + created = service.create(McpServerCreateRequest(**config)) + service.put_secret(created.server_id, "Authorization", "synthetic", kind="header") + config["secret_header_keys"] = ["authorization"] + updated = service.update( + created.server_id, McpServerUpdateRequest(**config, version=created.version) + ) + assert updated.secret_headers == {"authorization": True} + assert ( + service.credentials.resolve( + service._secret_id(created.server_id, "authorization", "header") + ) + == "synthetic" + ) + + +def test_environment_secrets_are_case_sensitive_and_delete_independently() -> None: + service = registry() + created = service.create(request(secret_environment_keys=["TOKEN", "token"])) + service.put_secret(created.server_id, "TOKEN", "upper") + service.put_secret(created.server_id, "token", "lower") + assert ( + service.credentials.resolve(service._secret_id(created.server_id, "TOKEN")) + == "upper" + ) + assert ( + service.credentials.resolve(service._secret_id(created.server_id, "token")) + == "lower" + ) + service.delete_secret(created.server_id, "TOKEN") + assert service.get(created.server_id).secret_environment == { + "TOKEN": False, + "token": True, + } + + +def test_legacy_environment_credential_migration_is_idempotent() -> None: + service = registry() + created = service.create(request(secret_environment_keys=["TOKEN"])) + suffix = hashlib.sha256(b"environment\0token").hexdigest()[:20] + legacy_id = f"mcp.{created.server_id}.{suffix}" + service.credentials.put(legacy_id, "legacy-value") + service._records[created.server_id]["secret_environment_version"] = 1 + service._write() + migrated = registry() + assert migrated.get(created.server_id).secret_environment == {"TOKEN": True} + assert ( + migrated.credentials.resolve(migrated._secret_id(created.server_id, "TOKEN")) + == "legacy-value" + ) + assert not migrated.credentials.has(legacy_id) + migrated.put_secret(created.server_id, "TOKEN", "new-value") + assert ( + registry().credentials.resolve(migrated._secret_id(created.server_id, "TOKEN")) + == "new-value" + ) + + +def test_ambiguous_legacy_credentials_are_not_assigned_to_two_variables() -> None: + service = registry() + created = service.create(request(secret_environment_keys=["TOKEN", "token"])) + suffix = hashlib.sha256(b"environment\0token").hexdigest()[:20] + legacy_id = f"mcp.{created.server_id}.{suffix}" + service.credentials.put(legacy_id, "cannot-reconstruct-originals") + service._records[created.server_id]["secret_environment_version"] = 1 + service._write() + migrated = registry() + current = migrated.get(created.server_id) + assert current.secret_environment == {"TOKEN": False, "token": False} + assert current.enabled is False + assert current.last_test_succeeded is None + assert migrated.credentials.has( + legacy_id + ) # Keep the original ciphertext recoverable. + migrated.put_secret(created.server_id, "TOKEN", "upper") + migrated.put_secret(created.server_id, "token", "lower") + assert registry().get(created.server_id).secret_environment == { + "TOKEN": True, + "token": True, + } + migrated.delete(created.server_id) + assert not migrated.credentials.has(legacy_id) + + +def test_credential_id_migration_keeps_new_values_and_is_atomic(monkeypatch) -> None: + credentials = EncryptedCredentialStore() + credentials.put("mcp.old", "old-value") + credentials.put("mcp.new", "new-value") + original_write = credentials._write_tokens + + def fail_write(_tokens): + raise CredentialStoreError("synthetic failure") + + monkeypatch.setattr(credentials, "_write_tokens", fail_write) + with pytest.raises(CredentialStoreError): + credentials.move_many({"mcp.old": "mcp.new"}) + assert credentials.resolve("mcp.old") == "old-value" + assert credentials.resolve("mcp.new") == "new-value" + monkeypatch.setattr(credentials, "_write_tokens", original_write) + credentials.move_many({"mcp.old": "mcp.new"}) + assert credentials.resolve("mcp.old") is None + assert credentials.resolve("mcp.new") == "new-value" + + +def test_ambiguous_legacy_secret_is_not_resurrected_after_removing_a_key() -> None: + service = registry() + created = service.create(request(secret_environment_keys=["TOKEN", "token"])) + legacy_id = service._legacy_environment_secret_id(created.server_id, "TOKEN") + service.credentials.put(legacy_id, "ambiguous-old-value") + service._records[created.server_id]["secret_environment_version"] = 1 + service._write() + migrated = registry() + migrated.update( + created.server_id, + McpServerUpdateRequest( + **request(secret_environment_keys=["token"]).model_dump(), + version=created.version, + ), + ) + assert registry().get(created.server_id).secret_environment == {"token": False} + + +def test_production_rejects_process_launch_even_after_approval() -> None: + service = registry(launch=False) + 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_SANDBOX_REQUIRED" + + +@pytest.mark.parametrize("startup,tool", [(120, 300), (1.5, 2.5)]) +def test_server_timeouts_survive_bridge_adaptation_and_reload(startup, tool) -> None: + service = registry() + created = service.create( + request( + secret_environment_keys=[], + startup_timeout_seconds=startup, + tool_timeout_seconds=tool, + ) + ) + service.trust(created.server_id, created.command_digest) + try: + tested = service.test(created.server_id) + assert tested.last_test_succeeded is True + assert tested.startup_timeout_seconds == startup + assert tested.tool_timeout_seconds == tool + restored = registry().get(created.server_id) + assert restored.startup_timeout_seconds == startup + assert restored.tool_timeout_seconds == tool + finally: + service.shutdown() + + +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_registry_rejects_structurally_invalid_record(tmp_path) -> None: + path = tmp_path / "mcp" + path.mkdir() + (path / "servers.json").write_text( + json.dumps({"server-1": {"name": "Broken", "transport": "stdio"}}), + 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_registry_rejects_create_before_exceeding_persisted_limit( + monkeypatch, +) -> None: + service = registry() + service.create(request(name="Only server")) + monkeypatch.setattr("app.extensions.mcp_registry._MAX_MCP_SERVERS", 1) + with pytest.raises(McpRegistryError) as error: + service.create(request(name="One too many")) + assert error.value.code == "MCP_SERVER_LIMIT_REACHED" + assert len(service.list()) == 1 + service.shutdown() + + +def test_stdio_command_is_not_parsed_as_a_shell_string() -> None: + service = registry() + created = service.create( + request( + 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 == "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() + + restored = registry() + restored.restore_enabled() + current = restored.get(created.server_id) + assert current.enabled is True + assert current.status == "ready" + assert any( + item.name.startswith(f"mcp.{created.server_id}.") + 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] = [] + request_timeouts: dict[str, float] = {} + + 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) + timeout = request_value.extensions.get("timeout", {}).get("read") + if isinstance(timeout, (int, float)): + request_timeouts[payload.get("method", "notification")] = float(timeout) + 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 + ) + assert request_timeouts["initialize"] == 15 + assert request_timeouts["notifications/initialized"] == 15 + assert request_timeouts["tools/list"] == 15 + 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"} + assert request_timeouts["tools/call"] == 30 + service.disable(created.server_id) + service.shutdown() + + +class _LegacyEventStream(httpx.SyncByteStream): + def __init__(self) -> None: + self.closed = threading.Event() + + 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() + self.closed.wait() + + def close(self) -> None: + self.closed.set() + + +def test_legacy_sse_uses_same_origin_endpoint(monkeypatch) -> None: + posted_urls: list[str] = [] + event_stream = _LegacyEventStream() + + def handler(request_value: httpx.Request) -> httpx.Response: + if request_value.method == "GET": + return httpx.Response( + 200, + headers={"content-type": "text/event-stream"}, + stream=event_stream, + ) + 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() + event_stream.close() + + +class _EndingLegacyEventStream(httpx.SyncByteStream): + def __iter__(self): + yield b"event: endpoint\ndata: /messages\n\n" + + +def test_legacy_sse_eof_marks_client_unavailable(monkeypatch) -> None: + def handler(_request_value: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + headers={"content-type": "text/event-stream"}, + stream=_EndingLegacyEventStream(), + ) + + real_client = httpx.Client + monkeypatch.setattr( + "app.extensions.mcp.httpx.Client", + lambda **kwargs: real_client(transport=httpx.MockTransport(handler), **kwargs), + ) + broken = threading.Event() + client = McpLegacySseClient( + "https://legacy.example.test/sse", + headers={}, + startup_timeout_seconds=1, + on_seen=lambda: None, + on_broken=lambda _message: broken.set(), + on_tools_changed=lambda: None, + ) + client.start() + assert broken.wait(timeout=1) + client.stop() + + +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/backend/tests/test_mcp_sse_limits.py b/backend/tests/test_mcp_sse_limits.py new file mode 100644 index 0000000..6c0a091 --- /dev/null +++ b/backend/tests/test_mcp_sse_limits.py @@ -0,0 +1,83 @@ +import json +from contextlib import closing + +import httpx +import pytest + +from app.extensions import mcp + + +class ChunkStream(httpx.SyncByteStream): + def __init__(self, chunks): + self.chunks = chunks + self.bytes_read = 0 + + def __iter__(self): + for chunk in self.chunks: + self.bytes_read += len(chunk) + yield chunk + + +def test_sse_rejects_unterminated_line_before_reading_entire_stream(monkeypatch): + monkeypatch.setattr(mcp, "MAX_MCP_MESSAGE_BYTES", 1024) + stream = ChunkStream([b"x" * 256] * 256) + with ( + closing(httpx.Response(200, stream=stream)) as response, + pytest.raises(mcp.McpBridgeError, match="too large"), + ): + list(mcp._iter_sse(response)) + assert stream.bytes_read == 1280 + + +def test_sse_limits_combined_event_before_partial_line_is_complete(monkeypatch): + monkeypatch.setattr(mcp, "MAX_MCP_MESSAGE_BYTES", 32) + stream = ChunkStream( + [b"data: 123456789\n", b"data: 123456789\n", b"x", b"not-read"] + ) + with ( + closing(httpx.Response(200, stream=stream)) as response, + pytest.raises(mcp.McpBridgeError, match="too large"), + ): + list(mcp._iter_sse(response)) + assert stream.bytes_read == 33 + + +@pytest.mark.parametrize("separator", [b"\n", b"\r", b"\r\n"]) +@pytest.mark.parametrize("chunk_size", [1, 2, 7, 1024]) +def test_sse_preserves_utf8_and_line_endings_across_chunks(separator, chunk_size): + payload = json.dumps( + {"jsonrpc": "2.0", "id": 1, "result": {"text": "中文"}}, ensure_ascii=False + ) + wire = b"\xef\xbb\xbf" + separator.join( + [ + b": heartbeat", + b"event: message", + b"id: replay-1", + ("data: " + payload).encode(), + b"", + b"", + ] + ) + stream = ChunkStream( + [wire[index : index + chunk_size] for index in range(0, len(wire), chunk_size)] + ) + with closing(httpx.Response(200, stream=stream)) as response: + assert list(mcp._iter_sse(response)) == [("message", "replay-1", payload)] + + +def test_sse_event_limit_resets_between_events(monkeypatch): + monkeypatch.setattr(mcp, "MAX_MCP_MESSAGE_BYTES", 16) + with closing( + httpx.Response(200, stream=ChunkStream([b"data: one\n\ndata: two\r\r"])) + ) as response: + assert list(mcp._iter_sse(response)) == [ + ("message", None, "one"), + ("message", None, "two"), + ] + + +def test_sse_preserves_multiline_data_and_final_unterminated_line(): + with closing( + httpx.Response(200, stream=ChunkStream([b"data: first\ndata: last"])) + ) as response: + assert list(mcp._iter_sse(response)) == [("message", None, "first\nlast")] diff --git a/docs/README.md b/docs/README.md index de9938b..eb9ff2e 100644 --- a/docs/README.md +++ b/docs/README.md @@ -32,6 +32,7 @@ - [Knowledge 与 Retrieval Core 开发说明](development/Knowledge与Retrieval-Core开发说明.md) - [模型提供商与模型发现开发说明](development/模型提供商与模型发现开发说明.md) - [MCP Bridge 与 Plugin Host 开发说明](development/MCP-Bridge与Plugin-Host开发说明.md) +- [独立 MCP Server 配置中心开发说明](development/独立MCP-Server配置中心开发说明.md) - [Plugin Command 与 Settings 开发说明](development/Plugin-Command与Settings开发说明.md) - [前端壳子与接口层开发说明](development/前端壳子与接口层开发说明.md) - [前端写作体验优化开发说明](development/前端写作体验优化开发说明.md) 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 87aceb4..f04a7a2 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,6 +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) | 列出、创建独立 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 与非敏感配置 | @@ -647,6 +654,29 @@ MCP_TOOL_SCHEMA_INVALID MCP_TOOL_CALL_FAILED MCP_TOOL_RESULT_TOO_LARGE MCP_TRUST_APPROVAL_REQUIRED +MCP_TRUST_DIGEST_STALE +MCP_SANDBOX_REQUIRED +MCP_TRANSPORT_UNSUPPORTED +MCP_SERVER_NOT_FOUND +MCP_SERVER_NAME_INVALID +MCP_SERVER_ALREADY_ENABLED +MCP_SERVER_VERSION_CONFLICT +MCP_SERVER_LIMIT_REACHED +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 PLUGIN_COMMAND_NOT_FOUND PLUGIN_COMMAND_CONFLICT PLUGIN_COMMAND_INVALID @@ -670,6 +700,20 @@ PLUGIN_STORAGE_ERROR CREDENTIAL_NAMESPACE_RESERVED ``` +### 7.8 独立 MCP Server Registry(C.1) + +独立 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。 + +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、连接参数、环境/Header 及权限完全一致的摘要;配置变化会撤销旧信任和测试结果。只有当前摘要通过 `/test`,才能调用 `/enable`。测试失败也会持久化时间和失败状态。 + +Secret 明文变化无法进入摘要,因此 Secret 写入和删除采用更严格规则:若 Server 已启用则先停用并注销 Tool,随后清除 `tested_digest` 和最近测试状态。调用方必须用新 Secret 再次执行 `/test`,不能沿用旧凭据的测试结果。 + +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 隔离约束。 + --- ## 8. Provider Adapter 扩展 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 new file mode 100644 index 0000000..e0d0a28 --- /dev/null +++ b/docs/development/独立MCP-Server配置中心开发说明.md @@ -0,0 +1,141 @@ +# 独立 MCP Server 配置中心开发说明 + +> 更新日期:2026-09-03。本文记录第二阶段 C.1 的完整实现;独立 MCP Server Registry 与 Plugin 自带 MCP Host 是两个并列入口。 + +## 1. 已实现范围 + +- 独立 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 输入。 + +Streamable HTTP 按 MCP 当前规范实现;SSE 仅用于兼容旧 Server,不应作为新部署首选。 + +## 2. 配置、版本与 Secret + +普通配置原子写入 `APP_DATA_DIR/mcp/servers.json`。更新请求必须携带读取到的 `version`;版本过期返回 `409 MCP_SERVER_VERSION_CONFLICT`,避免多个页面互相覆盖。改变 Transport、命令、URL、Header、环境变量或权限后,旧授权和测试结果立即失效。 + +`backend/data/mcp/` 是本机运行数据,包含连接配置、授权状态和第三方进程工作目录,不属于团队共享配置。`.gitignore` 忽略整个目录以及 `server.json`、`servers.json` 文件名;不得强制添加到 Git。提交前检查暂存文件清单,不要将本地密钥、连接配置或运行数据推送到远程。 + +Secret 使用带类型和键名哈希的 `mcp.*` 内部 ID 写入 Fernet 凭据存储。查询 API 只返回环境变量或 Header 是否配置,不返回明文。删除 Server 或移除 Secret 键会同步清理密文。前端密码框提交后立即清空;用户主动粘贴到 JSON 的密钥仅在当前编辑会话中暂存,解析后从 JSON 中移除,不写入 localStorage、普通配置或日志。 + +环境变量的凭据 ID 使用区分大小写的 v2 名称规则,Header ID 保持大小写不敏感。更新配置按凭据 ID 的差集删除密文,因此 `Authorization` 改为 `authorization` 不会丢失认证信息。启动时对无歧义的旧环境变量凭据原子迁移密文,不覆盖新 ID 已有的值;若旧配置把 `TOKEN`、`token` 合并存到了同一个 ID,无法推断原来的两个值,会保留旧密文、停用连接并要求重新录入和测试。删除服务器时也会清理这些保留的旧密文。 + +### 2.1 JSON 导入与 API Key 填写 + +前端支持 NotesAgent 完整/精简配置、单个 `command / args / env` 配置,以及只含一个服务器的 `mcpServers` 包装。JSON 与表单之间切换会补齐数组、对象及超时默认值,并校验字段类型。批量导入暂不支持;后端配置接口仍只接收 NotesAgent DTO,兼容转换发生在前端。 + +可以先声明 `secret_environment_keys`,保存后在服务器卡片的密码框填写密钥;也可以把密钥放进 JSON 的 `environment` 或通用配置的 `env`。前端会将已声明的敏感变量,以及名称含 API Key、Token、Secret、Password、Authorization、Cookie、Credential 的常见字段拆出:普通配置请求只包含键名,密钥另经 Secret API 加密保存。其他敏感字段必须显式声明,不能只依赖名称识别;命令与参数中不要携带密钥。 + +例如 MiniMax 的输入结构如下,占位值需在自己的本地页面替换,不要把真实密钥贴进聊天或提交到 Git: + +```json +{ + "name": "MiniMax Coding Plan", + "command": "uvx", + "args": ["--index-url", "https://pypi.tuna.tsinghua.edu.cn/simple", "--with", "mcp<2", "minimax-coding-plan-mcp", "-y"], + "environment": { + "MINIMAX_API_HOST": "https://api.minimaxi.com", + "MINIMAX_API_KEY": "<在本地填入新密钥>" + }, + "secret_environment_keys": ["MINIMAX_API_KEY"], + "startup_timeout_seconds": 120, + "tool_timeout_seconds": 300 +} +``` + +旧版前端将 `environment.MINIMAX_API_KEY` 与 `secret_environment_keys` 原样一起发送,触发后端“普通与敏感变量不可同名”的校验。这是配置保存失败,不是模型服务返回的鉴权失败。现在在前端拆分两类请求,后端仍保留互斥校验。 + +导入兼容规则:`env` 转为 `environment`;`timeout` 作为启动超时;`sse_read_timeout` 作为工具等待上限,不保留其原客户端 SSE 读取超时语义。启动超时范围为 1–120 秒,工具超时为 1–300 秒。URL 必须是纯地址,不能使用 Markdown 链接,JSON 中不能包含 `\_` 这样的非法转义。 + +另一个已修复的失败原因是运行时适配层复用了 `PluginBackend` 的整数超时与 60 秒启动上限,导致合法的 120 秒或小数超时配置在保存返回、读取或测试时失败。独立 Server 现在使用专门的 Bridge 适配模型,保留自己的浮点超时范围,不改变 Plugin 清单原有约束。已有的 120 秒记录可直接读取,无需删库重建。 + +配置保存成功但后续 Secret 写入失败时,窗口保留服务器 ID、新版本和未写入的密钥。点击保存会更新同一服务器并重试,不重复创建记录;取消会清除未保存密钥,已经保存的服务器和凭据不会回滚。错误信息显示在配置窗口内。保存配置不会自动运行第三方进程,仍需确认、测试和启用。 + +暂存的 Header Secret 与已保存凭据使用一致的大小写规则:将 `Authorization` 改为 `authorization` 不会丢弃尚未保存的值,提交时采用当前声明名。重新输入同一 Header 的值会覆盖旧草稿;真正删除声明才清除草稿。环境变量仍区分大小写,不会把 `TOKEN` 的草稿转交给 `token`。 + +跨 Registry 与凭据存储的删除以“失败后可重试”为顺序约束:先原子清理密文,再提交新版本或删除 Registry 记录。凭据存储失败时保留原版本和 Server 记录,避免出现返回 500 但配置已提交、版本无法重试或密文失去清理入口的状态。 + +写入或删除 Secret 会先停用正在运行的连接、注销动态 Tool,并撤销当前配置的测试通过状态;必须使用新凭据重新测试后才能启用。这样页面展示的凭据状态不会与运行中进程实际持有的旧凭据不一致。 + +## 3. 启用与运行时规则 + +一次连接按以下顺序执行: + +1. 用户检查服务端生成的连接摘要并确认当前摘要; +2. 后端临时连接,完成 initialize 和 `tools/list` 后关闭连接; +3. 只有当前摘要测试成功,启用操作才会启动长期连接并注册 Tool; +4. 停用、删除、超时或异常退出会注销 Tool 并关闭连接。 + +运行期连接异常或 Tool 列表变化时,Registry 会在同一生命周期临界区内标记不可用、注销 Tool,并从 Bridge 移除 Host。每次启动分配独立的连接代次;失败回调取得锁后先核对代次,旧连接延迟到达的回调不能停用新连接。停用、测试结束及关闭服务时撤销对应代次。 + +所有独立 MCP API 都通过工作线程执行,包括新增、确认授权和读取接口。虽然部分操作不直接访问网络,但仍可能等待正在测试或启动的连接持有的锁,不能在 FastAPI 事件循环上同步等待。 + +HTTP Header 中 `Host`、`Content-Type`、`MCP-Session-Id` 等协议保留项不可由配置覆盖。URL 不允许内嵌凭据或 Fragment。旧 SSE 返回的 POST endpoint 必须与初始 URL 同源,防止认证 Header 被转发到其他站点。 + +启动与 Tool 请求超时会同时应用于业务等待和底层 HTTP 请求;旧 SSE 的 endpoint 等待也使用启动超时。非主动结束的旧 SSE 事件流视为 Host 不可用,宿主随后注销 Tool。注册表加载时逐条校验 Server ID、配置字段、Transport 组合和摘要格式,损坏记录统一返回 `MCP_REGISTRY_INVALID`。 + +两种 HTTP Transport 共用有界 SSE 行解析器:按响应字节块检查未完成行及当前事件的累计大小,再扩展缓冲区,不依赖 `iter_lines()` 先缓存完整行。持续无换行输入也会及时触发上限;解析兼容跨块 UTF-8、首行 BOM、LF/CR/CRLF、多行 data 和事件间计数重置。`tests/test_mcp_sse_limits.py` 覆盖这些边界,防止仅在完整行生成后检查大小。 + +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 tests/test_extension_core.py + +cd ../frontend +npm run type-check +npm test +npm run build +``` + +后端测试使用无需网络或真实密钥的 stdio Fixture,以及 `httpx.MockTransport` 驱动的确定性 HTTP/SSE Server Fixture。覆盖摘要授权、Secret 不回显、版本冲突、生产门禁、重启恢复、Streamable HTTP Session/Header/工具调用及旧 SSE 同源校验;新增覆盖旧回调隔离、路由线程卸载、凭据大小写差异与旧密文迁移。前端覆盖模板切换、JSON 默认值与格式兼容、明文拆分、模式切换、部分保存失败重试、取消清理、测试失败和删除确认。此处的 MiniMax 配置转换测试使用假密钥,不等同于真实 MiniMax 网络调用验证。 + +## 6. 后续边界 + +### 本轮 P1/P2 修复验收 + +| 审阅问题 | 修复方式 | 回归验证 | +| --- | --- | --- | +| P1:新增或授权等待生命周期锁时阻塞事件循环 | 独立 MCP 路由统一交给工作线程 | `test_mcp_lifecycle_lock_contention_keeps_event_loop_responsive`:分别阻塞 create/trust,在锁释放前仍能执行健康检查 | +| P2:旧失败回调误停新连接 | 回调在锁内核对连接代次 | `test_old_failure_callback_cannot_stop_replacement_host`:旧回调排队期间重启连接,释放锁后新连接仍可用,当前代次的失败仍正确停用 | +| P2:精简 JSON 切换模式或编辑保存报错 | 运行时校验、默认值补全、统一转换 | `configuration.spec.ts` 与 `McpServersView.spec.ts`:精简 JSON、编辑版本、模式切换和 Secret 保存重试 | +| P2:Header 大小写改名删除凭据 | 按规范化凭据 ID 而非原始键名计算差集 | `test_header_case_only_rename_preserves_secret` | +| P2:大小写不同的环境变量覆盖同一凭据 | 区分大小写的 v2 ID,带迁移标记的旧密文迁移 | `test_environment_secrets_are_case_sensitive_and_delete_independently` 及 legacy migration 测试,包含删除后不复活旧密钥 | +| P1:SSE 无换行输入在大小校验前无限缓冲 | 在行拼接前校验字节数与事件累计大小 | `test_mcp_sse_limits.py`,包括小块持续输入和跨块换行 | +| P2:Header 大小写改名丢失未保存密钥 | 草稿使用规范化名称匹配,并重新绑定当前声明名 | `configuration.spec.ts` 与页面保存回归测试 | + +这些修复不放宽 stdio 的 JSON-RPC 校验。第三方程序向 stdout 打印普通日志造成的握手失败,应由服务端调整输出或使用不打印日志的启动入口处理。 + +### 后续工作 + +- 增加真实第三方 Server 的兼容矩阵;确定性 Fixture 只能证明宿主协议行为,不能代表所有实现兼容; +- C.5 在第二阶段开发与测试完成后、第三阶段桌面端实现前冻结沙箱 Contract; +- 第三阶段将 stdio 进程创建和 Secret 托管迁移至 Tauri/Rust Host 与 Stronghold/系统 Keychain。 diff --git a/frontend/src/components/common/PrimarySidebar.vue b/frontend/src/components/common/PrimarySidebar.vue index 052f0b4..9561166 100644 --- a/frontend/src/components/common/PrimarySidebar.vue +++ b/frontend/src/components/common/PrimarySidebar.vue @@ -1,7 +1,7 @@ + + + + diff --git a/frontend/src/features/mcp/configuration.spec.ts b/frontend/src/features/mcp/configuration.spec.ts new file mode 100644 index 0000000..6a23e1f --- /dev/null +++ b/frontend/src/features/mcp/configuration.spec.ts @@ -0,0 +1,63 @@ +import { describe, expect, it } from 'vitest' +import { emptyMcpConfig, mergeImportedSecrets, normalizeMcpConfig, parseMcpJson } from './configuration' + +describe('MCP configuration normalization', () => { + it('retains renamed HTTP drafts with the latest spelling and value', () => { + const config = { ...emptyMcpConfig(), secret_header_keys: ['authorization'] } + const previous = [{ kind: 'header' as const, key: 'Authorization', value: 'old-value' }] + expect(mergeImportedSecrets(config, previous, [])).toEqual([{ kind: 'header', key: 'authorization', value: 'old-value' }]) + expect(mergeImportedSecrets(config, previous, [{ kind: 'header', key: 'AUTHORIZATION', value: 'new-value' }])).toEqual([{ kind: 'header', key: 'authorization', value: 'new-value' }]) + expect(mergeImportedSecrets(emptyMcpConfig(), previous, [])).toEqual([]) + }) + + it('does not transfer an environment draft across a case-only rename', () => { + const config = { ...emptyMcpConfig(), secret_environment_keys: ['TOKEN', 'token'] } + const previous = [{ kind: 'environment' as const, key: 'TOKEN', value: 'upper' }, { kind: 'environment' as const, key: 'token', value: 'lower' }] + expect(mergeImportedSecrets(config, previous, [])).toEqual(previous) + expect(mergeImportedSecrets({ ...config, secret_environment_keys: ['token'] }, [previous[0]!], [])).toEqual([]) + }) + it('fills backend defaults for minimal JSON', () => { + const { config } = parseMcpJson('{"name":"demo","command":"uvx"}') + expect(config).toMatchObject({ transport: 'stdio', args: [], headers: {}, environment: {}, permissions: [], secret_header_keys: [] }) + }) + + it('extracts a key pasted into environment despite its existing secret declaration', () => { + const { config, secrets } = normalizeMcpConfig({ + name: 'MiniMax', command: 'uvx', secret_environment_keys: ['MINIMAX_API_KEY'], + environment: { MINIMAX_API_KEY: 'synthetic-key', MINIMAX_API_HOST: 'https://api.minimaxi.com' }, + }) + expect(config.environment).toEqual({ MINIMAX_API_HOST: 'https://api.minimaxi.com' }) + expect(config.secret_environment_keys).toEqual(['MINIMAX_API_KEY']) + expect(JSON.stringify(config)).not.toContain('synthetic-key') + expect(secrets).toEqual([{ kind: 'environment', key: 'MINIMAX_API_KEY', value: 'synthetic-key' }]) + }) + + it('imports a standard single-server wrapper and legacy timeouts', () => { + const { config, secrets } = normalizeMcpConfig({ mcpServers: { MiniMax: { + command: 'uvx', args: ['--with', 'mcp<2', 'minimax-coding-plan-mcp', '-y'], + env: { MINIMAX_API_KEY: 'synthetic-key' }, timeout: 120, sse_read_timeout: 300, + } } }) + expect(config).toMatchObject({ name: 'MiniMax', transport: 'stdio', environment: {}, startup_timeout_seconds: 120, tool_timeout_seconds: 300 }) + expect(secrets).toHaveLength(1) + }) + + it('extracts case-insensitive HTTP credentials without duplicate declarations', () => { + const { config, secrets } = normalizeMcpConfig({ url: 'https://example.test/mcp', headers: { authorization: 'synthetic' }, secret_header_keys: ['Authorization'] }) + expect(config.headers).toEqual({}) + expect(config.secret_header_keys).toEqual(['Authorization']) + expect(secrets[0]?.key).toBe('Authorization') + }) + + it.each([ + [{ command: 'uvx', args: 'not-array' }, 'args'], + [{ command: 'uvx', environment: [] }, 'environment'], + [{ command: 'uvx', timeout: 121 }, '启动超时'], + [{ command: 'uvx', args: ['[https://example.test](https://example.test)'] }, '纯 URL'], + [{ command: 'uvx', api_key: 'do-not-echo' }, '顶层'], + [{ command: 'uvx', env: {}, environment: {} }, '只保留一个'], + [{ mcpServers: { one: {}, two: {} } }, '一次导入一个'], + ])('rejects invalid fields without leaking their values', (input, hint) => { + expect(() => normalizeMcpConfig(input)).toThrow(hint) + try { normalizeMcpConfig(input) } catch (error) { expect(String(error)).not.toContain('do-not-echo') } + }) +}) diff --git a/frontend/src/features/mcp/configuration.ts b/frontend/src/features/mcp/configuration.ts new file mode 100644 index 0000000..4afbab7 --- /dev/null +++ b/frontend/src/features/mcp/configuration.ts @@ -0,0 +1,139 @@ +import type { McpServerInput } from '@/contracts' + +export type SecretKind = 'environment' | 'header' +export interface ImportedSecret { kind: SecretKind; key: string; value: string } + +export function mergeImportedSecrets(config: McpServerInput, previous: ImportedSecret[], incoming: ImportedSecret[]): ImportedSecret[] { + const merged = new Map() + for (const item of [...previous, ...incoming]) { + const normalize = (key: string) => item.kind === 'header' ? key.toLowerCase() : key + const keys = item.kind === 'header' ? config.secret_header_keys : config.secret_environment_keys + const declared = keys.find(key => normalize(key) === normalize(item.key)) + if (declared === undefined) continue + // HTTP identity is case-insensitive, but the Secret API requires the current + // declared spelling. New inline values replace older drafts of that identity. + merged.set(`${item.kind}:${normalize(declared)}`, { ...item, key: declared }) + } + return [...merged.values()] +} + +export function emptyMcpConfig(): 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, + } +} + +function object(value: unknown, label: string): Record { + if (!value || Array.isArray(value) || typeof value !== 'object') throw new Error(`${label}必须是 JSON 对象`) + return value as Record +} + +function strings(value: unknown, label: string): string[] { + if (value === undefined) return [] + if (!Array.isArray(value) || value.some(item => typeof item !== 'string')) throw new Error(`${label}必须是字符串数组`) + return [...value] +} + +function entries(value: unknown, label: string): Record { + if (value === undefined) return {} + const result = object(value, label) + if (Object.values(result).some(item => typeof item !== 'string')) throw new Error(`${label}必须是字符串键值 JSON 对象`) + return { ...result } as Record +} + +function timeout(value: unknown, fallback: number, max: number, label: string): number { + if (value === undefined) return fallback + if (typeof value !== 'number' || !Number.isFinite(value) || value < 1 || value > max) throw new Error(`${label}必须是 1–${max} 秒之间的数字`) + return value +} + +// Do not silently rewrite executable arguments or secret values copied from chat. +function checkUrl(value: string, label: string) { + if (/^\[https?:\/\//i.test(value)) throw new Error(`${label}请填写纯 URL,不要粘贴 Markdown 链接`) +} + +export function parseMcpJson(raw: string, fallbackName = '', requireConnection = true) { + let parsed: unknown + try { parsed = JSON.parse(raw) } + catch { throw new Error('服务器配置不是有效 JSON;请检查逗号、引号和无效的 \\_ 转义') } + return normalizeMcpConfig(parsed, fallbackName, requireConnection) +} + +/** Normalize external client JSON before it reaches either the form or the API. + * Inline secrets leave the public config here and are sent only to the Secret API. + */ +export function normalizeMcpConfig(parsed: unknown, fallbackName = '', requireConnection = true) { + let raw = object(parsed, '服务器配置') + if ('mcpServers' in raw) { + const servers = Object.entries(object(raw.mcpServers, 'mcpServers')) + if (servers.length !== 1) throw new Error('请一次导入一个 MCP 服务器') + fallbackName = servers[0]![0] + raw = object(servers[0]![1], '服务器配置') + } + const allowed = new Set([...Object.keys(emptyMcpConfig()), 'version', 'env', 'type', 'timeout', 'sse_read_timeout']) + if (Object.keys(raw).some(key => !allowed.has(key))) { + // Never echo arbitrary unknown keys: pasted secrets sometimes become JSON keys. + throw new Error('服务器配置含不支持的字段;API Key 请放在 env/environment 的对应变量中,不要放在顶层') + } + if (raw.env !== undefined && raw.environment !== undefined) throw new Error('env 与 environment 请只保留一个,避免覆盖配置') + const transport = raw.transport ?? raw.type ?? (raw.url ? 'streamable_http' : 'stdio') + if (!['stdio', 'streamable_http', 'sse'].includes(transport as string)) throw new Error('transport 必须是 stdio、streamable_http 或 sse') + const config = emptyMcpConfig() + config.transport = transport as McpServerInput['transport'] + const name = raw.name ?? (fallbackName || (typeof raw.command === 'string' ? raw.command : 'MCP 服务器')) + if (typeof name !== 'string' || (requireConnection && !name.trim()) || name.trim().length > 80) throw new Error('服务器名称必须为 1–80 个字符') + config.name = name.trim() + for (const key of ['command', 'url'] as const) { + const value = raw[key] + if (value !== undefined && value !== null && typeof value !== 'string') throw new Error(`${key}必须是字符串`) + config[key] = typeof value === 'string' ? value.trim() : null + } + config.args = strings(raw.args, 'args') + if (config.args.length > 64) throw new Error('args 最多允许 64 项') + for (const value of config.args) checkUrl(value, 'args 中的地址') + config.environment = entries(raw.environment ?? raw.env, 'environment/env') + config.headers = entries(raw.headers, 'headers') + config.secret_environment_keys = [...new Set(strings(raw.secret_environment_keys, 'secret_environment_keys'))] + config.secret_header_keys = [...new Set(strings(raw.secret_header_keys, 'secret_header_keys'))] + config.permissions = strings(raw.permissions, 'permissions') + config.startup_timeout_seconds = timeout(raw.startup_timeout_seconds ?? raw.timeout, 15, 120, '启动超时') + // Compatibility policy: legacy read timeout becomes the tool wait budget, not an SSE transport setting. + config.tool_timeout_seconds = timeout(raw.tool_timeout_seconds ?? raw.sse_read_timeout, 30, 300, '工具超时') + if (config.transport === 'stdio') { + if (requireConnection && !config.command) throw new Error('stdio 配置必须填写 command') + if (config.url || Object.keys(config.headers).length || config.secret_header_keys.length) throw new Error('stdio 配置不能包含 URL 或 HTTP Header') + } else { + if (requireConnection && !config.url) throw new Error('HTTP/SSE 配置必须填写 url') + if (config.url) { + checkUrl(config.url, 'url') + let url: URL + try { url = new URL(config.url) } catch { throw new Error('url 必须是有效的 HTTP(S) 地址') } + if (!['http:', 'https:'].includes(url.protocol) || url.username || url.password || url.hash) throw new Error('url 必须为不含账号密码或片段的 HTTP(S) 地址') + } + if (config.command || config.args.length || Object.keys(config.environment).length || config.secret_environment_keys.length) throw new Error('HTTP/SSE 配置不能包含 command、args 或环境变量') + } + const secrets: ImportedSecret[] = [] + for (const kind of ['environment', 'header'] as const) { + const values = kind === 'environment' ? config.environment : config.headers + const keys = kind === 'environment' ? config.secret_environment_keys : config.secret_header_keys + const identity = (key: string) => kind === 'header' ? key.toLowerCase() : key + const allKeys = [...Object.keys(values), ...keys] + if (kind === 'header' && (new Set(keys.map(identity)).size !== keys.length || new Set(Object.keys(values).map(identity)).size !== Object.keys(values).length)) throw new Error('HTTP Header 名称不能仅大小写不同而重复声明') + const validKey = kind === 'environment' ? /^[A-Za-z_][A-Za-z0-9_]{0,127}$/ : /^[!#$%&'*+.^_`|~0-9A-Za-z-]{1,128}$/ + if (allKeys.some(key => !validKey.test(key))) throw new Error(`${kind === 'environment' ? '环境变量' : 'Header'}名称无效;敏感变量名只能填名称,不能填密钥值`) + for (const [key, value] of Object.entries(values)) { + const declared = keys.find(item => identity(item) === identity(key)) + const sensitive = /api[_-]?key|token|secret|password|authorization|cookie|credential/i.test(key) + if (declared || sensitive) { + if (!value || value.length > 32768) throw new Error('密钥值必须为 1–32768 个字符') + const secretKey = declared ?? key + if (!declared) keys.push(key) + secrets.push({ kind, key: secretKey, value }) + delete values[key] + } else if (/host|url|endpoint/i.test(key)) checkUrl(value, '环境变量或 Header 地址') + } + } + return { config, secrets } +} diff --git a/frontend/src/features/workspace/FileTreePanel.spec.ts b/frontend/src/features/workspace/FileTreePanel.spec.ts index 4a5933e..ce20566 100644 --- a/frontend/src/features/workspace/FileTreePanel.spec.ts +++ b/frontend/src/features/workspace/FileTreePanel.spec.ts @@ -74,4 +74,33 @@ describe('FileTreePanel file switching', () => { expect(editorStore.content).toContain('# 二叉搜索树') expect(editorStore.currentNoteId).toBe('note-bst') }) + + it('creates a Markdown note inside the selected folder', async () => { + const router = createRouter({ + history: createMemoryHistory(), + routes: [{ path: '/workspace', component: { template: '
' } }], + }) + await router.push('/workspace') + await router.isReady() + + const workspaceStore = useWorkspaceStore() + await workspaceStore.openVault('C:/vault') + const createFile = vi.spyOn(workspaceService, 'createFile').mockResolvedValue({ + id: 'note-new', note_id: 'note-new', name: '新笔记.md', + path: '/数据结构/新笔记.md', type: 'file', + }) + wrapper = mount(FileTreePanel, { attachTo: document.body, global: { plugins: [router] } }) + + await wrapper.findAll('.tree-node').find((node) => node.text().includes('数据结构'))!.trigger('click') + await wrapper.get('button[aria-label="新建笔记"]').trigger('click') + await wrapper.get('.new-item input').setValue('新笔记') + await wrapper.get('.new-item').trigger('submit') + await waitForPath('/数据结构/新笔记.md') + await vi.waitFor(() => { + expect(workspaceStore.activeFilePath).toBe('/数据结构/新笔记.md') + }) + + expect(createFile).toHaveBeenCalledWith('/数据结构', '新笔记.md', '# 新笔记\n\n') + expect(wrapper.findAll('.tree-node').some((node) => node.classes().includes('active') && node.text().includes('新笔记.md'))).toBe(true) + }) }) diff --git a/frontend/src/features/workspace/FileTreePanel.vue b/frontend/src/features/workspace/FileTreePanel.vue index 38c496b..f877bf0 100644 --- a/frontend/src/features/workspace/FileTreePanel.vue +++ b/frontend/src/features/workspace/FileTreePanel.vue @@ -1,5 +1,5 @@