From d4ffbdcadde80a44f0300e41eb1d7a40fe961a09 Mon Sep 17 00:00:00 2001 From: KiriAky 107 Date: Thu, 3 Sep 2026 14:49:49 +0800 Subject: [PATCH 1/4] feat(mcp): add standalone server registry Implement C.1 stdio MCP server CRUD, encrypted environment secrets, command digest approval, connection tests, lifecycle recovery, and dynamic tool registration. Add the standalone frontend configuration center, contracts, regression tests, and development documentation. --- README.md | 2 +- backend/app/container.py | 11 + backend/app/contracts.py | 68 ++- backend/app/extensions/mcp.py | 22 +- backend/app/extensions/mcp_registry.py | 554 ++++++++++++++++++ backend/app/main.py | 1 + backend/app/providers/credentials.py | 7 +- backend/app/routes.py | 80 +++ backend/tests/test_mcp_registry.py | 123 ++++ .../src/components/common/PrimarySidebar.vue | 3 +- frontend/src/contracts/index.ts | 34 +- frontend/src/features/mcp/McpServersView.vue | 169 ++++++ frontend/src/router/index.ts | 6 + frontend/src/services/index.ts | 1 + frontend/src/services/mcpServerService.ts | 17 + 15 files changed, 1086 insertions(+), 12 deletions(-) create mode 100644 backend/app/extensions/mcp_registry.py create mode 100644 backend/tests/test_mcp_registry.py create mode 100644 frontend/src/features/mcp/McpServersView.vue create mode 100644 frontend/src/services/mcpServerService.ts diff --git a/README.md b/README.md index cd85754..6f2c586 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/P0。Streamable HTTP MCP、真实音频、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..ee8f3ae 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,72 @@ 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 McpServerCreateRequest(Contract): + name: str = Field(min_length=1, max_length=80) + transport: McpServerTransport = McpServerTransport.stdio + command: str = Field(min_length=1, max_length=1024) + args: list[str] = Field(default_factory=list, max_length=64) + environment: dict[str, str] = Field(default_factory=dict) + secret_environment_keys: list[str] = Field(default_factory=list) + permissions: list[str] = Field(default_factory=list) + startup_timeout_seconds: float = Field(default=15, ge=1, le=120) + tool_timeout_seconds: float = Field(default=30, ge=1, le=300) + + +class McpServerUpdateRequest(McpServerCreateRequest): + pass + + +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 McpServer(Contract): + server_id: str + name: str + transport: McpServerTransport + command: str + args: list[str] = Field(default_factory=list) + environment: dict[str, str] = Field(default_factory=dict) + secret_environment: dict[str, bool] = Field(default_factory=dict) + permissions: list[str] = Field(default_factory=list) + startup_timeout_seconds: float + tool_timeout_seconds: float + enabled: bool = False + trusted: bool = False + command_digest: str + command_summary: str + status: PluginHostState = PluginHostState.stopped + tools_count: int = 0 + protocol_version: str | None = None + 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 McpServerListResponse(Contract): + items: list[McpServer] = 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..834ef65 100644 --- a/backend/app/extensions/mcp.py +++ b/backend/app/extensions/mcp.py @@ -74,12 +74,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 @@ -99,6 +101,7 @@ class McpStdioClient: # 平台级沙箱启动器;uvx 只隔离 Python 依赖,不能替代系统权限限制。 creation_flags = getattr(subprocess, "CREATE_NO_WINDOW", 0) if os.name == "nt" else 0 environment = _subprocess_environment() + environment.update(self.environment) environment.setdefault("PYTHONUNBUFFERED", "1") try: self.process = subprocess.Popen( @@ -383,6 +386,10 @@ 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", ) -> list[McpDiscoveredTool]: if backend.transport != "stdio": raise McpBridgeError( @@ -390,7 +397,7 @@ class McpBridge: "Phase C only supports the MCP stdio transport.", status_code=501, ) - command = self._resolve_command(package_path, backend) + command = command_override or self._resolve_command(package_path, backend) now = datetime.now(timezone.utc) status = PluginHostStatus( plugin_id=plugin_id, @@ -420,6 +427,7 @@ class McpBridge: client = McpStdioClient( command, cwd=package_path, + environment=environment, on_seen=seen, on_broken=broken, on_tools_changed=tools_changed, @@ -470,7 +478,7 @@ class McpBridge: status.server_version = _optional_string(server_info.get("version")) client.notify("notifications/initialized") discovered = self._discover_tools( - plugin_id, client, backend, declared_permissions + plugin_id, client, backend, declared_permissions, tool_source ) status.status = PluginHostState.ready status.tools_count = len(discovered) @@ -601,6 +609,7 @@ class McpBridge: client: McpStdioClient, backend: PluginBackend, declared_permissions: list[str], + tool_source: str, ) -> list[McpDiscoveredTool]: discovered: list[McpDiscoveredTool] = [] cursor: str | None = None @@ -620,7 +629,7 @@ class McpBridge: ) 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( @@ -648,7 +657,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( @@ -712,7 +724,7 @@ class McpBridge: description=description if isinstance(description, str) else remote_name, parameters=schema, permission=permission, - source="plugin", + source=tool_source, ), ) diff --git a/backend/app/extensions/mcp_registry.py b/backend/app/extensions/mcp_registry.py new file mode 100644 index 0000000..e576f7b --- /dev/null +++ b/backend/app/extensions/mcp_registry.py @@ -0,0 +1,554 @@ +"""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 pathlib import Path +from typing import Any +from uuid import uuid4 + +from pydantic import BaseModel, ConfigDict, create_model + +from app.agent.permissions import KNOWN_PERMISSIONS +from app.agent.tools import ToolExecutionContext, ToolRegistry +from app.contracts import ( + McpServer, + McpServerCreateRequest, + McpServerSecretStatus, + McpServerTransport, + McpServerUpdateRequest, + 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}$") + + +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 + + +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._records = self._read() + self._registered: dict[str, list[str]] = {} + self._last_status: dict[str, dict[str, Any]] = {} + + 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 create(self, request: McpServerCreateRequest) -> McpServer: + self._validate(request) + server_id = uuid4().hex[:12] + record = request.model_dump(mode="json") + record["name"] = request.name.strip() + record["command"] = request.command.strip() + record.update(enabled=False, approved_digest=None) + with self._lock: + updated = {**self._records, server_id: record} + self._write(updated) + self._records = updated + return self.get(server_id) + + def update(self, server_id: str, request: McpServerUpdateRequest) -> McpServer: + self._validate(request) + self.disable(server_id) + with self._lock: + previous = self._record(server_id) + removed = set(previous.get("secret_environment_keys", [])) - set( + request.secret_environment_keys + ) + record = request.model_dump(mode="json") + record["name"] = request.name.strip() + record["command"] = request.command.strip() + record.update(enabled=False, approved_digest=None) + updated = {**self._records, server_id: record} + self._write(updated) + self._records = updated + self._last_status.pop(server_id, None) + for key in removed: + try: + self.credentials.delete(self._secret_id(server_id, key)) + except CredentialStoreError as exc: + raise McpRegistryError( + "MCP_SECRET_STORE_ERROR", str(exc), status_code=500 + ) from exc + return self.get(server_id) + + def delete(self, server_id: str) -> None: + self.disable(server_id) + with self._lock: + record = self._record(server_id) + secret_ids = [ + self._secret_id(server_id, key) + for key in record.get("secret_environment_keys", []) + ] + updated = dict(self._records) + del updated[server_id] + self._write(updated) + self._records = updated + self._last_status.pop(server_id, None) + try: + self.credentials.delete_many(secret_ids) + except CredentialStoreError as exc: + raise McpRegistryError( + "MCP_SECRET_STORE_ERROR", str(exc), status_code=500 + ) from exc + self.bridge.remove(self._host_id(server_id)) + + 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) + + def put_secret( + self, server_id: str, key: str, secret: str + ) -> McpServerSecretStatus: + with self._lock: + record = self._record(server_id) + self._validate_environment_key(key) + if key not in record.get("secret_environment_keys", []): + raise McpRegistryError( + "MCP_SECRET_NOT_DECLARED", + "Secret environment key is not declared in this server configuration.", + ) + try: + self.credentials.put(self._secret_id(server_id, key), secret) + except CredentialStoreError as exc: + raise McpRegistryError( + "MCP_SECRET_STORE_ERROR", str(exc), status_code=500 + ) from exc + return McpServerSecretStatus(key=key, configured=True) + + def delete_secret(self, server_id: str, key: str) -> McpServerSecretStatus: + record = self._record(server_id) + if key not in record.get("secret_environment_keys", []): + raise McpRegistryError( + "MCP_SECRET_NOT_DECLARED", + "Secret environment key is not declared in this server configuration.", + ) + try: + self.credentials.delete(self._secret_id(server_id, key)) + except CredentialStoreError as exc: + raise McpRegistryError( + "MCP_SECRET_STORE_ERROR", str(exc), status_code=500 + ) from exc + return McpServerSecretStatus(key=key, configured=False) + + 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) + try: + discovered = self._start(server_id, record) + except Exception as exc: + self._last_status[server_id] = { + "status": PluginHostState.error, + "error": str(exc), + "last_tested_at": datetime.now(UTC), + "last_test_succeeded": False, + } + raise + status = self.bridge.status(self._host_id(server_id), self._backend(record)) + 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": datetime.now(UTC), + "last_test_succeeded": True, + } + self.bridge.stop(self._host_id(server_id)) + return self.get(server_id) + + 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) + discovered = self._start(server_id, record) + 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.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.bridge.stop(self._host_id(server_id)) + raise + return self.get(server_id) + + 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.bridge.stop(self._host_id(server_id)) + return self.get(server_id) + + 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() + + def shutdown(self) -> None: + for server_id in list(self._records): + 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)) + 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 + host_id = self._host_id(server_id) + 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, message), + command_override=[record["command"], *record.get("args", [])], + environment=environment, + tool_source="mcp_server", + ) + except McpBridgeError as exc: + 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, message: str) -> None: + with self._lock: + for name in self._registered.pop(server_id, []): + self.tools.unregister(name) + record = self._records.get(server_id) + if record is not None: + self._records[server_id] = {**record, "enabled": False} + self._last_status[server_id] = { + "status": PluginHostState.unhealthy, + "error": message, + } + self._write() + + def _require_launch_allowed(self, record: dict[str, Any]) -> None: + if record.get("transport") != McpServerTransport.stdio.value: + raise McpRegistryError( + "MCP_TRANSPORT_UNSUPPORTED", + "C.1 currently supports stdio; Streamable HTTP and SSE are reserved for a later increment.", + status_code=501, + ) + if not self.allow_process_launch: + raise McpRegistryError( + "MCP_SANDBOX_REQUIRED", + "Python process launch is disabled outside development until the desktop sandbox is available.", + status_code=403, + ) + if record.get("approved_digest") != self._digest(record): + raise McpRegistryError( + "MCP_TRUST_APPROVAL_REQUIRED", + "Review and approve the current MCP command before testing or 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, + name=record["name"], + transport=record["transport"], + command=record["command"], + args=list(record.get("args", [])), + environment=dict(record.get("environment", {})), + secret_environment={ + key: self._secret_configured(server_id, key) + for key in record.get("secret_environment_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=cached.get("last_tested_at"), + last_test_succeeded=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 not request.command.strip() or "\x00" in request.command: + raise McpRegistryError("MCP_COMMAND_INVALID", "MCP executable is invalid.") + if any("\x00" in arg for arg in request.args): + raise McpRegistryError( + "MCP_COMMAND_INVALID", "MCP argument contains a null byte." + ) + 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.", + ) + 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 _backend(record: dict[str, Any]) -> PluginBackend: + return PluginBackend( + type="mcp", + transport="stdio", + command=record["command"], + 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", + "permissions", + ) + } + return hashlib.sha256( + json.dumps( + executable, sort_keys=True, ensure_ascii=False, separators=(",", ":") + ).encode() + ).hexdigest() + + @staticmethod + def _summary(record: dict[str, Any]) -> str: + return " ".join( + [ + record["command"], + *[ + json.dumps(arg, ensure_ascii=False) + for arg in record.get("args", []) + ], + ] + ) + + @staticmethod + def _secret_id(server_id: str, key: str) -> str: + suffix = hashlib.sha256(key.encode()).hexdigest()[:20] + return f"mcp.{server_id}.{suffix}" + + def _secret_configured(self, server_id: str, key: str) -> bool: + try: + return self.credentials.has(self._secret_id(server_id, key)) + except CredentialStoreError as exc: + raise McpRegistryError( + "MCP_SECRET_STORE_ERROR", str(exc), status_code=500 + ) from exc + + 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 + + @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, + ) + return value + + 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..accdacb 100644 --- a/backend/app/providers/credentials.py +++ b/backend/app/providers/credentials.py @@ -14,6 +14,7 @@ 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,10 +28,10 @@ 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 - ): + 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: diff --git a/backend/app/routes.py b/backend/app/routes.py index 962bec7..438d41b 100644 --- a/backend/app/routes.py +++ b/backend/app/routes.py @@ -21,6 +21,13 @@ from app.contracts import ( IndexJob, IndexRebuildRequest, IndexStatus, + McpServer, + McpServerCreateRequest, + McpServerListResponse, + McpServerSecretStatus, + McpServerSecretWriteRequest, + McpServerTrustRequest, + McpServerUpdateRequest, ModelEvent, ModelEventType, Note, @@ -72,6 +79,7 @@ from app.agent import AgentCapacityError, AgentRunNotFoundError from app.container import container from app.errors import ApiError from app.extensions import ExtensionError +from app.extensions.mcp_registry import McpRegistryError from app.providers.registry import ProviderNotFoundError from app.providers.factory import UnsupportedProviderError from app.providers.base import ProviderError @@ -91,6 +99,21 @@ from app.services import ( router = APIRouter(prefix="/api") +def mcp_call(operation): + try: + return operation() + except McpRegistryError as exc: + raise ApiError(exc.status_code, exc.code, exc.message) from exc + + +async def mcp_call_async(operation): + """MCP process operations wait on stdio and must not block the API 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) @@ -481,6 +504,63 @@ async def uninstall_skill(skill_id: str) -> OperationResponse: 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=mcp_call(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 mcp_call(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 mcp_call(lambda: container.mcp_servers.get(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 mcp_call(lambda: container.mcp_servers.trust(server_id, request.command_digest)) + + +@router.post("/mcp/servers/{server_id}/test", response_model=McpServer, tags=["MCP Servers"]) +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) -> McpServerSecretStatus: + return mcp_call(lambda: container.mcp_servers.put_secret(server_id, key, request.secret.get_secret_value())) + + +@router.delete("/mcp/servers/{server_id}/secrets/{key}", response_model=McpServerSecretStatus, tags=["MCP Servers"]) +async def delete_mcp_server_secret(server_id: str, key: str) -> McpServerSecretStatus: + return mcp_call(lambda: container.mcp_servers.delete_secret(server_id, key)) + + # Plugins @router.get("/plugins", response_model=PluginListResponse, tags=["Plugins"]) async def list_plugins() -> PluginListResponse: diff --git a/backend/tests/test_mcp_registry.py b/backend/tests/test_mcp_registry.py new file mode 100644 index 0000000..a4eb658 --- /dev/null +++ b/backend/tests/test_mcp_registry.py @@ -0,0 +1,123 @@ +import sys + +import pytest + +from app.agent.tools import ToolRegistry +from app.config import BACKEND_DIR, get_settings +from app.contracts import McpServerCreateRequest, McpServerUpdateRequest +from app.extensions.mcp_registry import McpRegistryError, McpServerRegistry +from app.providers.credentials import 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_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) + 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() + ), + ) + 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_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" + + +def test_non_stdio_transport_is_explicitly_reserved() -> None: + service = registry() + created = service.create( + request( + transport="streamable_http", + command="https://example.invalid/mcp", + secret_environment_keys=[], + ) + ) + service.trust(created.server_id, created.command_digest) + with pytest.raises(McpRegistryError) as error: + service.test(created.server_id) + assert error.value.code == "MCP_TRANSPORT_UNSUPPORTED" + + +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.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() 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/router/index.ts b/frontend/src/router/index.ts index 2116443..4655ccb 100644 --- a/frontend/src/router/index.ts +++ b/frontend/src/router/index.ts @@ -44,6 +44,12 @@ const routes = [ component: () => import('@/features/skills/SkillsView.vue'), meta: { title: 'Skill 管理', requiresVault: true }, }, + { + path: '/extensions/mcp', + name: 'mcp-servers', + component: () => import('@/features/mcp/McpServersView.vue'), + meta: { title: 'MCP 服务器', requiresVault: true }, + }, { path: '/extensions/plugins', name: 'plugins', diff --git a/frontend/src/services/index.ts b/frontend/src/services/index.ts index 968f5b3..849e204 100644 --- a/frontend/src/services/index.ts +++ b/frontend/src/services/index.ts @@ -8,6 +8,7 @@ export * as chatService from './chatService' export * as agentService from './agentService' export * as skillService from './skillService' export * as pluginService from './pluginService' +export * as mcpServerService from './mcpServerService' export * as providerService from './providerService' export * as taskService from './taskService' export * as indexService from './indexService' diff --git a/frontend/src/services/mcpServerService.ts b/frontend/src/services/mcpServerService.ts new file mode 100644 index 0000000..a146a6b --- /dev/null +++ b/frontend/src/services/mcpServerService.ts @@ -0,0 +1,17 @@ +import apiClient from './apiClient' +import type { McpServer, McpServerInput, OperationResponse } from '@/contracts' + +const base = '/api/mcp/servers' + +export async function listMcpServers(): Promise { + return (await apiClient.get<{ items: McpServer[] }>(base)).items +} +export const createMcpServer = (input: McpServerInput) => apiClient.post(base, input) +export const updateMcpServer = (id: string, input: McpServerInput) => apiClient.put(`${base}/${id}`, input) +export const deleteMcpServer = (id: string) => apiClient.delete(`${base}/${id}`) +export const trustMcpServer = (server: McpServer) => apiClient.post(`${base}/${server.server_id}/trust`, { command_digest: server.command_digest }) +export const testMcpServer = (id: string) => apiClient.post(`${base}/${id}/test`) +export const enableMcpServer = (id: string) => apiClient.post(`${base}/${id}/enable`) +export const disableMcpServer = (id: string) => apiClient.post(`${base}/${id}/disable`) +export const putMcpServerSecret = (id: string, key: string, secret: string) => apiClient.put(`${base}/${id}/secrets/${encodeURIComponent(key)}`, { secret }) +export const deleteMcpServerSecret = (id: string, key: string) => apiClient.delete(`${base}/${id}/secrets/${encodeURIComponent(key)}`) From 9f2c46ab39aae920b516c6d15b05c19d518c8174 Mon Sep 17 00:00:00 2001 From: KiriAky 107 Date: Thu, 3 Sep 2026 15:25:41 +0800 Subject: [PATCH 2/4] feat(mcp): complete remote transports and configuration workflow --- README.md | 2 +- backend/app/contracts.py | 59 +- backend/app/extensions/mcp.py | 727 +++++++++++++++++- backend/app/extensions/mcp_registry.py | 361 +++++++-- backend/app/routes.py | 160 +++- backend/tests/test_mcp_registry.py | 288 ++++++- frontend/src/contracts/index.ts | 17 +- .../src/features/mcp/McpServersView.spec.ts | 84 ++ frontend/src/features/mcp/McpServersView.vue | 191 +++-- frontend/src/services/mcpServerService.ts | 7 +- 10 files changed, 1667 insertions(+), 229 deletions(-) create mode 100644 frontend/src/features/mcp/McpServersView.spec.ts diff --git a/README.md b/README.md index 6f2c586..3a96a48 100644 --- a/README.md +++ b/README.md @@ -2,7 +2,7 @@ > 本文件用于团队开发期间快速配置环境和启动项目,不是正式的项目 README。 -> 当前基线:2026-09-03。第一阶段 Web 联调前后端已经完成;第二阶段已完成 Workspace 去 Mock、Agent Trace 持久化与 SSE 恢复、stdio MCP Bridge、隔离 Plugin Host、Plugin Command/Settings,以及独立 MCP Server 配置中心 C.1/P0。Streamable HTTP MCP、真实音频、Provider 协议增强、Benchmark、导出、主题包、Trace 可视化、Mermaid 与函数图像仍在后续开发;Tauri Host、Stronghold、原生多 Vault 文件系统和 Sync Server 尚未接入。 +> 当前基线:2026-09-03。第一阶段 Web 联调前后端已经完成;第二阶段已完成 Workspace 去 Mock、Agent Trace 持久化与 SSE 恢复、stdio MCP Bridge、隔离 Plugin Host、Plugin Command/Settings,以及独立 MCP Server 配置中心 C.1(stdio、Streamable HTTP 与旧 SSE 兼容)。真实音频、Provider 协议增强、Benchmark、导出、主题包、Trace 可视化、Mermaid 与函数图像仍在后续开发;Tauri Host、Stronghold、原生多 Vault 文件系统和 Sync Server 尚未接入。 ## 当前目录 diff --git a/backend/app/contracts.py b/backend/app/contracts.py index ee8f3ae..9840f45 100644 --- a/backend/app/contracts.py +++ b/backend/app/contracts.py @@ -494,22 +494,29 @@ class McpServerTransport(str, Enum): sse = "sse" -class McpServerCreateRequest(Contract): +class McpServerConfig(Contract): name: str = Field(min_length=1, max_length=80) transport: McpServerTransport = McpServerTransport.stdio - command: str = Field(min_length=1, max_length=1024) + command: str | None = Field(default=None, max_length=1024) args: list[str] = Field(default_factory=list, max_length=64) + url: str | None = Field(default=None, max_length=4096) + headers: dict[str, str] = Field(default_factory=dict) environment: dict[str, str] = Field(default_factory=dict) secret_environment_keys: list[str] = Field(default_factory=list) + secret_header_keys: list[str] = Field(default_factory=list) permissions: list[str] = Field(default_factory=list) startup_timeout_seconds: float = Field(default=15, ge=1, le=120) tool_timeout_seconds: float = Field(default=30, ge=1, le=300) -class McpServerUpdateRequest(McpServerCreateRequest): +class McpServerCreateRequest(McpServerConfig): pass +class McpServerUpdateRequest(McpServerConfig): + version: int = Field(ge=1) + + class McpServerSecretWriteRequest(Contract): secret: SecretStr = Field(min_length=1, max_length=32768) @@ -523,21 +530,8 @@ class McpServerTrustRequest(Contract): command_digest: str = Field(min_length=64, max_length=64) -class McpServer(Contract): - server_id: str - name: str - transport: McpServerTransport - command: str - args: list[str] = Field(default_factory=list) - environment: dict[str, str] = Field(default_factory=dict) - secret_environment: dict[str, bool] = Field(default_factory=dict) - permissions: list[str] = Field(default_factory=list) - startup_timeout_seconds: float - tool_timeout_seconds: float +class McpServerStatus(Contract): enabled: bool = False - trusted: bool = False - command_digest: str - command_summary: str status: PluginHostState = PluginHostState.stopped tools_count: int = 0 protocol_version: str | None = None @@ -548,10 +542,41 @@ class McpServer(Contract): last_test_succeeded: bool | None = None +class McpServer(McpServerStatus): + server_id: str + version: int + name: str + transport: McpServerTransport + command: str | None = None + args: list[str] = Field(default_factory=list) + url: str | None = None + headers: dict[str, str] = Field(default_factory=dict) + environment: dict[str, str] = Field(default_factory=dict) + permissions: list[str] = Field(default_factory=list) + startup_timeout_seconds: float + tool_timeout_seconds: float + secret_environment: dict[str, bool] = Field(default_factory=dict) + secret_headers: dict[str, bool] = Field(default_factory=dict) + trusted: bool = False + command_digest: str + command_summary: str + + class McpServerListResponse(Contract): items: list[McpServer] = Field(default_factory=list) +class McpToolSummary(Contract): + name: str + remote_name: str + description: str + permission: str | None = None + + +class McpToolSummaryListResponse(Contract): + items: list[McpToolSummary] = Field(default_factory=list) + + class PluginCommandLocation(str, Enum): command_palette = "command_palette" context_menu = "context_menu" diff --git a/backend/app/extensions/mcp.py b/backend/app/extensions/mcp.py index 834ef65..61c2d5f 100644 --- a/backend/app/extensions/mcp.py +++ b/backend/app/extensions/mcp.py @@ -10,14 +10,18 @@ import asyncio import json import os import queue +import signal import subprocess import threading from collections import deque +from collections.abc import Callable from dataclasses import dataclass -from datetime import datetime, timezone +from datetime import UTC, datetime from pathlib import Path -from typing import Any, Callable +from typing import Any, Protocol +from urllib.parse import urljoin, urlsplit +import httpx from jsonschema import Draft202012Validator from jsonschema.exceptions import SchemaError @@ -99,7 +103,12 @@ class McpStdioClient: return # TODO(extension-security): 社区 Plugin 开放前迁移到 Tauri/Rust Host 的 # 平台级沙箱启动器;uvx 只隔离 Python 依赖,不能替代系统权限限制。 - creation_flags = getattr(subprocess, "CREATE_NO_WINDOW", 0) if os.name == "nt" else 0 + creation_flags = ( + getattr(subprocess, "CREATE_NO_WINDOW", 0) + | getattr(subprocess, "CREATE_NEW_PROCESS_GROUP", 0) + if os.name == "nt" + else 0 + ) environment = _subprocess_environment() environment.update(self.environment) environment.setdefault("PYTHONUNBUFFERED", "1") @@ -117,6 +126,7 @@ class McpStdioClient: shell=False, env=environment, creationflags=creation_flags, + start_new_session=os.name != "nt", ) except OSError as exc: raise McpBridgeError( @@ -183,7 +193,9 @@ class McpStdioClient: except queue.Empty as exc: self.cancel(request_id, "Request timed out.") self.abandon(request_id) - raise McpBridgeError(timeout_code, "MCP request timed out.", status_code=504) from exc + raise McpBridgeError( + timeout_code, "MCP request timed out.", status_code=504 + ) from exc if isinstance(response, BaseException): raise response if "error" in response: @@ -216,9 +228,7 @@ class McpStdioClient: except McpBridgeError: pass - def abandon( - self, request_id: int, wake_error: BaseException | None = None - ) -> None: + def abandon(self, request_id: int, wake_error: BaseException | None = None) -> None: with self._pending_lock: pending = self._pending.pop(request_id, None) # asyncio.to_thread 被取消时不会停止底层线程;主动唤醒 Queue,避免线程 @@ -243,15 +253,17 @@ class McpStdioClient: try: process.wait(timeout=2) except subprocess.TimeoutExpired: - process.terminate() + _terminate_process_tree(process) try: process.wait(timeout=2) except subprocess.TimeoutExpired: - process.kill() + _kill_process_tree(process) process.wait(timeout=2) finally: self._fail_pending( - McpBridgeError("PLUGIN_HOST_UNAVAILABLE", "MCP host stopped.", status_code=503) + McpBridgeError( + "PLUGIN_HOST_UNAVAILABLE", "MCP host stopped.", status_code=503 + ) ) self.process = None @@ -313,14 +325,17 @@ class McpStdioClient: { "jsonrpc": "2.0", "id": message["id"], - "error": {"code": -32601, "message": "Method not supported."}, + "error": { + "code": -32601, + "message": "Method not supported.", + }, } ) except (McpBridgeError, OSError, ValueError) as exc: failure = f"MCP stdout closed unexpectedly: {type(exc).__name__}." finally: if failure and process.poll() is None: - process.terminate() + _terminate_process_tree(process) exit_code = process.poll() if exit_code is None: try: @@ -328,7 +343,9 @@ class McpStdioClient: except subprocess.TimeoutExpired: exit_code = None if not self._stopping: - message = failure or f"MCP host exited unexpectedly with code {exit_code}." + message = ( + failure or f"MCP host exited unexpectedly with code {exit_code}." + ) error = McpBridgeError( "PLUGIN_HOST_UNAVAILABLE", message, status_code=503 ) @@ -363,13 +380,472 @@ class McpStdioClient: item.response.put(error) +class McpHttpClient: + """MCP Streamable HTTP client supporting JSON and SSE POST responses.""" + + def __init__( + self, + url: str, + *, + headers: dict[str, str], + on_seen: Callable[[], None], + on_broken: Callable[[str], None], + on_tools_changed: Callable[[], None], + ) -> None: + self.url = url + self.headers = headers + self.on_seen = on_seen + self.on_broken = on_broken + self.on_tools_changed = on_tools_changed + self._client = httpx.Client(follow_redirects=False, timeout=30) + self._pending_lock = threading.Lock() + self._pending: dict[int, _PendingRequest] = {} + self._next_id = 1 + self._session_id: str | None = None + self._protocol_version: str | None = None + self._stopping = False + self._stream_started = False + self._last_event_id: str | None = None + self._stop_event = threading.Event() + + def start(self) -> None: + return + + def set_protocol_version(self, version: str) -> None: + self._protocol_version = version + + def start_event_stream(self) -> None: + if self._stream_started: + return + self._stream_started = True + threading.Thread(target=self._event_stream_loop, daemon=True).start() + + def request( + self, + method: str, + params: dict[str, Any], + *, + timeout: float, + timeout_code: str, + response_error_code: str = "MCP_TOOL_CALL_FAILED", + ) -> dict[str, Any]: + request_id, pending = self.begin_request(method, params) + return self.wait_response( + request_id, + pending, + timeout=timeout, + timeout_code=timeout_code, + response_error_code=response_error_code, + ) + + def begin_request( + self, method: str, params: dict[str, Any] + ) -> tuple[int, _PendingRequest]: + with self._pending_lock: + request_id = self._next_id + self._next_id += 1 + pending = _PendingRequest(response=queue.Queue(maxsize=1)) + self._pending[request_id] = pending + message = { + "jsonrpc": "2.0", + "id": request_id, + "method": method, + "params": params, + } + threading.Thread( + target=self._dispatch_request, + args=(request_id, message), + daemon=True, + ).start() + return request_id, pending + + def wait_response( + self, + request_id: int, + pending: _PendingRequest, + *, + timeout: float, + timeout_code: str, + response_error_code: str = "MCP_TOOL_CALL_FAILED", + ) -> dict[str, Any]: + try: + response = pending.response.get(timeout=timeout) + except queue.Empty as exc: + self.cancel(request_id, "Request timed out.") + self.abandon(request_id) + raise McpBridgeError( + timeout_code, "MCP request timed out.", status_code=504 + ) from exc + if isinstance(response, BaseException): + raise response + if "error" in response: + error = response.get("error") + message = ( + str(error.get("message", "MCP JSON-RPC error.")) + if isinstance(error, dict) + else "MCP JSON-RPC error." + ) + raise McpBridgeError(response_error_code, message) + result = response.get("result") + if not isinstance(result, dict): + raise McpBridgeError( + response_error_code, "MCP response result must be an object." + ) + return result + + def notify(self, method: str, params: dict[str, Any] | None = None) -> None: + message: dict[str, Any] = {"jsonrpc": "2.0", "method": method} + if params is not None: + message["params"] = params + self._post_notification(message) + + def cancel(self, request_id: int, reason: str = "Cancelled by host.") -> None: + def send() -> None: + try: + self.notify( + "notifications/cancelled", + {"requestId": request_id, "reason": reason}, + ) + except McpBridgeError: + pass + + threading.Thread(target=send, daemon=True).start() + + def abandon(self, request_id: int, wake_error: BaseException | None = None) -> None: + with self._pending_lock: + pending = self._pending.pop(request_id, None) + if pending is not None and wake_error is not None: + try: + pending.response.put_nowait(wake_error) + except queue.Full: + pass + + def stop(self) -> None: + self._stopping = True + self._stop_event.set() + if self._session_id: + try: + request = self._client.build_request( + "DELETE", self.url, headers=self._request_headers() + ) + response = self._client.send(request, stream=True) + response.close() + except httpx.HTTPError: + pass + self._client.close() + self._fail_pending( + McpBridgeError( + "PLUGIN_HOST_UNAVAILABLE", "MCP HTTP client stopped.", status_code=503 + ) + ) + + def _dispatch_request(self, request_id: int, message: dict[str, Any]) -> None: + try: + response = self._post(message, timeout=None) + try: + self._capture_session(response) + content_type = response.headers.get("content-type", "").lower() + if response.status_code >= 400: + raise McpBridgeError( + "MCP_HTTP_REQUEST_FAILED", + f"MCP HTTP server returned status {response.status_code}.", + status_code=502, + ) + if "application/json" in content_type: + payload = _bounded_json_response(response) + self._deliver(payload) + elif "text/event-stream" in content_type: + delivered = False + for _event, _event_id, data in _iter_sse(response): + payload = _json_rpc_message(data) + self._handle_message(payload) + if payload.get("id") == request_id: + delivered = True + break + if not delivered: + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", + "MCP SSE response ended before the matching JSON-RPC response.", + ) + else: + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", + "MCP HTTP response has an unsupported Content-Type.", + ) + finally: + response.close() + except (McpBridgeError, httpx.HTTPError) as exc: + error = ( + exc + if isinstance(exc, McpBridgeError) + else McpBridgeError( + "MCP_HTTP_REQUEST_FAILED", + f"MCP HTTP request failed: {type(exc).__name__}.", + status_code=503, + ) + ) + self.abandon(request_id, error) + + def _post_notification(self, message: dict[str, Any]) -> None: + try: + response = self._post(message, timeout=10) + except httpx.HTTPError as exc: + raise McpBridgeError( + "MCP_HTTP_REQUEST_FAILED", + f"MCP HTTP notification failed: {type(exc).__name__}.", + status_code=503, + ) from exc + try: + self._capture_session(response) + if response.status_code not in {200, 202, 204}: + raise McpBridgeError( + "MCP_HTTP_REQUEST_FAILED", + f"MCP HTTP server rejected a notification with status {response.status_code}.", + ) + finally: + response.close() + + def _post( + self, message: dict[str, Any], *, timeout: float | None + ) -> httpx.Response: + encoded = json.dumps(message, ensure_ascii=False, separators=(",", ":")) + if len(encoded.encode("utf-8")) > MAX_MCP_MESSAGE_BYTES: + raise McpBridgeError("MCP_TOOL_CALL_FAILED", "MCP request is too large.") + request = self._client.build_request( + "POST", + self.url, + content=encoded.encode("utf-8"), + headers=self._request_headers(), + ) + return self._client.send(request, stream=True) + + def _request_headers(self) -> dict[str, str]: + headers = { + **self.headers, + "Accept": "application/json, text/event-stream", + "Content-Type": "application/json", + } + if self._session_id: + headers["MCP-Session-Id"] = self._session_id + if self._protocol_version: + headers["MCP-Protocol-Version"] = self._protocol_version + return headers + + def _capture_session(self, response: httpx.Response) -> None: + session_id = response.headers.get("mcp-session-id") + if session_id is not None: + if ( + not session_id.isascii() + or not session_id.isprintable() + or len(session_id) > 1024 + ): + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", "MCP session id is invalid." + ) + self._session_id = session_id + + def _handle_message(self, message: dict[str, Any]) -> None: + self.on_seen() + if "id" in message and ("result" in message or "error" in message): + self._deliver(message) + elif message.get("method") == "notifications/tools/list_changed": + self.on_tools_changed() + + def _deliver(self, message: dict[str, Any]) -> None: + request_id = message.get("id") + if not isinstance(request_id, int): + return + with self._pending_lock: + pending = self._pending.pop(request_id, None) + if pending: + pending.response.put(message) + + def _fail_pending(self, error: BaseException) -> None: + with self._pending_lock: + pending = list(self._pending.values()) + self._pending.clear() + for item in pending: + item.response.put(error) + + def _event_stream_loop(self) -> None: + while not self._stop_event.is_set(): + headers = {**self._request_headers(), "Accept": "text/event-stream"} + headers.pop("Content-Type", None) + if self._last_event_id: + headers["Last-Event-ID"] = self._last_event_id + try: + with self._client.stream( + "GET", self.url, headers=headers, timeout=None + ) as response: + if response.status_code == 405: + return + if response.status_code >= 400: + self.on_broken( + f"MCP HTTP event stream returned status {response.status_code}." + ) + return + if ( + "text/event-stream" + not in response.headers.get("content-type", "").lower() + ): + self.on_broken("MCP HTTP GET response is not an event stream.") + return + self._capture_session(response) + for _event, event_id, data in _iter_sse(response): + if event_id: + self._last_event_id = event_id + self._handle_message(_json_rpc_message(data)) + if self._stop_event.is_set(): + return + except (McpBridgeError, httpx.HTTPError): + if self._stopping: + return + self._stop_event.wait(0.25) + + +class McpLegacySseClient(McpHttpClient): + """Compatibility client for the deprecated 2024-11-05 HTTP+SSE transport.""" + + def __init__(self, *args: Any, **kwargs: Any) -> None: + super().__init__(*args, **kwargs) + self._endpoint: str | None = None + self._endpoint_ready: queue.Queue[str | BaseException] = queue.Queue(maxsize=1) + + def start(self) -> None: + threading.Thread(target=self._event_loop, daemon=True).start() + try: + endpoint = self._endpoint_ready.get(timeout=15) + except queue.Empty as exc: + raise McpBridgeError( + "MCP_INITIALIZE_FAILED", + "Legacy MCP SSE endpoint event timed out.", + status_code=504, + ) from exc + if isinstance(endpoint, BaseException): + raise endpoint + self._endpoint = endpoint + + def start_event_stream(self) -> None: + """The legacy client already owns its single GET event stream.""" + + return + + def _dispatch_request(self, request_id: int, message: dict[str, Any]) -> None: + try: + response = self._post(message, timeout=10) + try: + if response.status_code not in {200, 202, 204}: + raise McpBridgeError( + "MCP_HTTP_REQUEST_FAILED", + f"Legacy MCP endpoint returned status {response.status_code}.", + ) + finally: + response.close() + except (McpBridgeError, httpx.HTTPError) as exc: + error = ( + exc + if isinstance(exc, McpBridgeError) + else McpBridgeError( + "MCP_HTTP_REQUEST_FAILED", + f"Legacy MCP request failed: {type(exc).__name__}.", + status_code=503, + ) + ) + self.abandon(request_id, error) + + def _post( + self, message: dict[str, Any], *, timeout: float | None + ) -> httpx.Response: + if self._endpoint is None: + raise McpBridgeError( + "MCP_INITIALIZE_FAILED", "Legacy MCP endpoint is not ready." + ) + encoded = json.dumps(message, ensure_ascii=False, separators=(",", ":")) + request = self._client.build_request( + "POST", + self._endpoint, + content=encoded.encode("utf-8"), + headers={ + **self.headers, + "Accept": "application/json, text/event-stream", + "Content-Type": "application/json", + }, + ) + return self._client.send(request, stream=True) + + def _event_loop(self) -> None: + try: + with self._client.stream( + "GET", + self.url, + headers={**self.headers, "Accept": "text/event-stream"}, + timeout=None, + ) as response: + if response.status_code >= 400: + raise McpBridgeError( + "MCP_HTTP_REQUEST_FAILED", + f"Legacy MCP SSE server returned status {response.status_code}.", + ) + if ( + "text/event-stream" + not in response.headers.get("content-type", "").lower() + ): + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", + "Legacy MCP GET response is not an event stream.", + ) + for event, _event_id, data in _iter_sse(response): + if self._endpoint is None and event == "endpoint": + endpoint = _legacy_endpoint_url(self.url, data) + self._endpoint_ready.put(endpoint) + self._endpoint = endpoint + continue + self._handle_message(_json_rpc_message(data)) + except (McpBridgeError, httpx.HTTPError) as exc: + if self._endpoint is None: + self._endpoint_ready.put(exc) + elif not self._stopping: + self.on_broken(f"Legacy MCP SSE stream failed: {type(exc).__name__}.") + + @dataclass(slots=True) class _McpHost: backend: PluginBackend - client: McpStdioClient + client: _McpClient status: PluginHostStatus +class _McpClient(Protocol): + def start(self) -> None: ... + def request( + self, + method: str, + params: dict[str, Any], + *, + timeout: float, + timeout_code: str, + response_error_code: str = "MCP_TOOL_CALL_FAILED", + ) -> dict[str, Any]: ... + def begin_request( + self, method: str, params: dict[str, Any] + ) -> tuple[int, _PendingRequest]: ... + def wait_response( + self, + request_id: int, + pending: _PendingRequest, + *, + timeout: float, + timeout_code: str, + response_error_code: str = "MCP_TOOL_CALL_FAILED", + ) -> dict[str, Any]: ... + def notify(self, method: str, params: dict[str, Any] | None = None) -> None: ... + def cancel(self, request_id: int, reason: str = "Cancelled by host.") -> None: ... + def abandon( + self, request_id: int, wake_error: BaseException | None = None + ) -> None: ... + def stop(self) -> None: ... + + class McpBridge: """管理每个 Plugin 的独立 MCP Client,并执行 Contract 转换。""" @@ -390,19 +866,27 @@ class McpBridge: command_override: list[str] | None = None, environment: dict[str, str] | None = None, tool_source: str = "plugin", + transport_kind: str | None = None, + url: str | None = None, + headers: dict[str, str] | None = None, ) -> list[McpDiscoveredTool]: - if backend.transport != "stdio": + transport = transport_kind or backend.transport + if transport not in {"stdio", "streamable_http", "sse"}: raise McpBridgeError( "MCP_CAPABILITY_UNSUPPORTED", - "Phase C only supports the MCP stdio transport.", + f"Unsupported MCP transport: {transport}", status_code=501, ) - command = command_override or self._resolve_command(package_path, backend) - now = datetime.now(timezone.utc) + command = ( + command_override or self._resolve_command(package_path, backend) + if transport == "stdio" + else None + ) + now = datetime.now(UTC) status = PluginHostStatus( plugin_id=plugin_id, backend_type="mcp", - transport="stdio", + transport="stdio" if transport == "stdio" else "http", status=PluginHostState.starting, started_at=now, last_seen_at=now, @@ -412,7 +896,7 @@ class McpBridge: def seen() -> None: host = host_ref.get("host") if host: - host.status.last_seen_at = datetime.now(timezone.utc) + host.status.last_seen_at = datetime.now(UTC) def broken(message: str) -> None: host = host_ref.get("host") @@ -422,16 +906,35 @@ class McpBridge: on_unavailable(plugin_id, message) def tools_changed() -> None: - broken("MCP tool list changed; restart the Plugin Host to revalidate tools.") + broken( + "MCP tool list changed; restart the Plugin Host to revalidate tools." + ) - client = McpStdioClient( - command, - cwd=package_path, - environment=environment, - on_seen=seen, - on_broken=broken, - on_tools_changed=tools_changed, - ) + if transport == "stdio": + assert command is not None + client: _McpClient = McpStdioClient( + command, + cwd=package_path, + environment=environment, + on_seen=seen, + on_broken=broken, + on_tools_changed=tools_changed, + ) + else: + if not url: + raise McpBridgeError( + "MCP_HOST_START_FAILED", "MCP HTTP transport requires a URL." + ) + client_type = ( + McpHttpClient if transport == "streamable_http" else McpLegacySseClient + ) + client = client_type( + url, + headers=headers or {}, + on_seen=seen, + on_broken=broken, + on_tools_changed=tools_changed, + ) host = _McpHost(backend=backend, client=client, status=status) host_ref["host"] = host with self._lock: @@ -474,15 +977,21 @@ class McpBridge: if not isinstance(server_info, dict): server_info = {} status.protocol_version = str(version) + set_protocol_version = getattr(client, "set_protocol_version", None) + if callable(set_protocol_version): + set_protocol_version(str(version)) status.server_name = _optional_string(server_info.get("name")) status.server_version = _optional_string(server_info.get("version")) client.notify("notifications/initialized") + start_event_stream = getattr(client, "start_event_stream", None) + if callable(start_event_stream): + start_event_stream() discovered = self._discover_tools( plugin_id, client, backend, declared_permissions, tool_source ) status.status = PluginHostState.ready status.tools_count = len(discovered) - status.last_seen_at = datetime.now(timezone.utc) + status.last_seen_at = datetime.now(UTC) status.error = None return discovered except McpBridgeError as exc: @@ -527,9 +1036,7 @@ class McpBridge: host.client.cancel(rpc_id) host.client.abandon( rpc_id, - McpBridgeError( - "MCP_TOOL_CALL_FAILED", "MCP request was cancelled." - ), + McpBridgeError("MCP_TOOL_CALL_FAILED", "MCP request was cancelled."), ) raise except McpBridgeError as exc: @@ -539,7 +1046,9 @@ class McpBridge: self._calls.pop(call_key, None) encoded_size = len( - json.dumps(result, ensure_ascii=False, separators=(",", ":")).encode("utf-8") + json.dumps(result, ensure_ascii=False, separators=(",", ":")).encode( + "utf-8" + ) ) if encoded_size > MAX_MCP_TOOL_RESULT_BYTES: raise ToolExecutionError( @@ -606,7 +1115,7 @@ class McpBridge: def _discover_tools( self, plugin_id: str, - client: McpStdioClient, + client: _McpClient, backend: PluginBackend, declared_permissions: list[str], tool_source: str, @@ -625,7 +1134,8 @@ class McpBridge: raw_tools = result.get("tools") if not isinstance(raw_tools, list): raise McpBridgeError( - "MCP_TOOL_SCHEMA_INVALID", "MCP tools/list must return a tools array." + "MCP_TOOL_SCHEMA_INVALID", + "MCP tools/list must return a tools array.", ) for raw in raw_tools: discovered.append( @@ -641,7 +1151,8 @@ class McpBridge: break if not isinstance(next_cursor, str) or not next_cursor: raise McpBridgeError( - "MCP_TOOL_SCHEMA_INVALID", "MCP nextCursor must be a non-empty string." + "MCP_TOOL_SCHEMA_INVALID", + "MCP nextCursor must be a non-empty string.", ) cursor = next_cursor else: @@ -675,9 +1186,7 @@ class McpBridge: len(remote_name) > 128 or not remote_name[0].isalnum() or not all( - character.islower() - or character.isdigit() - or character in "._-" + character.islower() or character.isdigit() or character in "._-" for character in remote_name ) ): @@ -702,7 +1211,9 @@ class McpBridge: ) from exc metadata = raw.get("_meta") permission = ( - metadata.get("notesagent/permission") if isinstance(metadata, dict) else None + metadata.get("notesagent/permission") + if isinstance(metadata, dict) + else None ) if permission is not None and ( not isinstance(permission, str) or permission not in KNOWN_PERMISSIONS @@ -721,7 +1232,9 @@ class McpBridge: remote_name=remote_name, definition=ToolDefinition( name=f"{plugin_id}.{remote_name}", - description=description if isinstance(description, str) else remote_name, + description=description + if isinstance(description, str) + else remote_name, parameters=schema, permission=permission, source=tool_source, @@ -801,3 +1314,133 @@ def _subprocess_environment() -> dict[str, str]: environment["PYTHONUNBUFFERED"] = "1" environment["PYTHONIOENCODING"] = "utf-8" return environment + + +def _bounded_json_response(response: httpx.Response) -> dict[str, Any]: + content_length = response.headers.get("content-length") + if ( + content_length + and content_length.isdigit() + and int(content_length) > MAX_MCP_MESSAGE_BYTES + ): + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", "MCP HTTP response is too large." + ) + chunks: list[bytes] = [] + size = 0 + for chunk in response.iter_bytes(): + size += len(chunk) + if size > MAX_MCP_MESSAGE_BYTES: + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", "MCP HTTP response is too large." + ) + chunks.append(chunk) + try: + payload = json.loads(b"".join(chunks)) + except (UnicodeDecodeError, json.JSONDecodeError) as exc: + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", "MCP HTTP response is not valid JSON." + ) from exc + if not isinstance(payload, dict) or payload.get("jsonrpc") != "2.0": + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", "MCP HTTP response is not a JSON-RPC message." + ) + return payload + + +def _iter_sse(response: httpx.Response): + event = "message" + event_id: str | None = None + data_lines: list[str] = [] + size = 0 + for line in response.iter_lines(): + size += len(line.encode("utf-8")) + 1 + if size > MAX_MCP_MESSAGE_BYTES: + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", "MCP SSE event is too large." + ) + if line == "": + if data_lines: + yield event, event_id, "\n".join(data_lines) + event, event_id, data_lines, size = "message", None, [], 0 + continue + if line.startswith(":"): + continue + field, _, value = line.partition(":") + value = value.removeprefix(" ") + if field == "event": + event = value + elif field == "id" and "\x00" not in value: + event_id = value + elif field == "data": + data_lines.append(value) + if data_lines: + yield event, event_id, "\n".join(data_lines) + + +def _json_rpc_message(data: str) -> dict[str, Any]: + try: + message = json.loads(data) + except json.JSONDecodeError as exc: + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", "MCP SSE data is not valid JSON." + ) from exc + if not isinstance(message, dict) or message.get("jsonrpc") != "2.0": + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", "MCP SSE data is not a JSON-RPC message." + ) + return message + + +def _legacy_endpoint_url(source_url: str, endpoint: str) -> str: + target = urljoin(source_url, endpoint.strip()) + source_parts = urlsplit(source_url) + target_parts = urlsplit(target) + if ( + target_parts.scheme not in {"http", "https"} + or target_parts.username is not None + or target_parts.password is not None + or (source_parts.scheme, source_parts.hostname, source_parts.port) + != (target_parts.scheme, target_parts.hostname, target_parts.port) + ): + raise McpBridgeError( + "MCP_HTTP_RESPONSE_INVALID", + "Legacy MCP endpoint must use the same origin as the configured SSE URL.", + ) + return target + + +def _terminate_process_tree(process: subprocess.Popen[str]) -> None: + if process.poll() is not None: + return + try: + if os.name == "nt": + subprocess.run( + ["taskkill.exe", "/PID", str(process.pid), "/T"], + check=False, + capture_output=True, + creationflags=getattr(subprocess, "CREATE_NO_WINDOW", 0), + timeout=2, + ) + else: + os.killpg(process.pid, signal.SIGTERM) + except (OSError, subprocess.SubprocessError): + process.terminate() + + +def _kill_process_tree(process: subprocess.Popen[str]) -> None: + if process.poll() is not None: + return + try: + if os.name == "nt": + subprocess.run( + ["taskkill.exe", "/PID", str(process.pid), "/T", "/F"], + check=False, + capture_output=True, + creationflags=getattr(subprocess, "CREATE_NO_WINDOW", 0), + timeout=2, + ) + else: + os.killpg(process.pid, signal.SIGKILL) + except (OSError, subprocess.SubprocessError): + process.kill() diff --git a/backend/app/extensions/mcp_registry.py b/backend/app/extensions/mcp_registry.py index e576f7b..6808769 100644 --- a/backend/app/extensions/mcp_registry.py +++ b/backend/app/extensions/mcp_registry.py @@ -7,8 +7,10 @@ import json import re import threading from datetime import UTC, datetime +from functools import wraps from pathlib import Path from typing import Any +from urllib.parse import urlsplit from uuid import uuid4 from pydantic import BaseModel, ConfigDict, create_model @@ -21,6 +23,7 @@ from app.contracts import ( McpServerSecretStatus, McpServerTransport, McpServerUpdateRequest, + McpToolSummary, PluginBackend, PluginHostState, ) @@ -28,6 +31,15 @@ from app.extensions.mcp import McpBridge, McpBridgeError, McpDiscoveredTool from app.providers.credentials import CredentialStoreError, EncryptedCredentialStore _ENVIRONMENT_KEY = re.compile(r"^[A-Za-z_][A-Za-z0-9_]{0,127}$") +_HEADER_KEY = re.compile(r"^[!#$%&'*+.^_`|~0-9A-Za-z-]{1,128}$") +_RESERVED_HEADERS = { + "accept", + "content-length", + "content-type", + "host", + "mcp-protocol-version", + "mcp-session-id", +} class McpRegistryError(RuntimeError): @@ -38,6 +50,17 @@ class McpRegistryError(RuntimeError): self.status_code = status_code +def _serialized_lifecycle(method): + """Serialize lifecycle mutations without blocking MCP failure callbacks.""" + + @wraps(method) + def wrapped(self, *args, **kwargs): + with self._lifecycle_lock: + return method(self, *args, **kwargs) + + return wrapped + + class McpServerRegistry: """Persists configuration and owns stdio host/tool lifecycles.""" @@ -56,8 +79,10 @@ class McpServerRegistry: self.allow_process_launch = allow_process_launch self.bridge = bridge or McpBridge() self._lock = threading.RLock() + self._lifecycle_lock = threading.RLock() self._records = self._read() self._registered: dict[str, list[str]] = {} + self._summaries: dict[str, list[McpToolSummary]] = {} self._last_status: dict[str, dict[str, Any]] = {} def list(self) -> list[McpServer]: @@ -71,57 +96,108 @@ class McpServerRegistry: with self._lock: return self._public(server_id, self._record(server_id)) + def list_tools(self, server_id: str) -> list[McpToolSummary]: + self._record(server_id) + return [ + item.model_copy(deep=True) for item in self._summaries.get(server_id, []) + ] + + @_serialized_lifecycle def create(self, request: McpServerCreateRequest) -> McpServer: self._validate(request) server_id = uuid4().hex[:12] record = request.model_dump(mode="json") record["name"] = request.name.strip() - record["command"] = request.command.strip() - record.update(enabled=False, approved_digest=None) + record["command"] = request.command.strip() if request.command else None + record["url"] = request.url.strip() if request.url else None + record.update( + version=1, + enabled=False, + approved_digest=None, + tested_digest=None, + last_tested_at=None, + last_test_succeeded=None, + ) with self._lock: updated = {**self._records, server_id: record} self._write(updated) self._records = updated return self.get(server_id) + @_serialized_lifecycle def update(self, server_id: str, request: McpServerUpdateRequest) -> McpServer: self._validate(request) + current = self._record(server_id) + if request.version != current.get("version", 1): + raise McpRegistryError( + "MCP_SERVER_VERSION_CONFLICT", + "MCP server configuration version is stale.", + status_code=409, + ) self.disable(server_id) with self._lock: previous = self._record(server_id) - removed = set(previous.get("secret_environment_keys", [])) - set( - request.secret_environment_keys - ) - record = request.model_dump(mode="json") + removed = [ + (kind, key) + for kind, old_keys, new_keys in ( + ( + "environment", + previous.get("secret_environment_keys", []), + request.secret_environment_keys, + ), + ( + "header", + previous.get("secret_header_keys", []), + request.secret_header_keys, + ), + ) + for key in set(old_keys) - set(new_keys) + ] + record = request.model_dump(mode="json", exclude={"version"}) record["name"] = request.name.strip() - record["command"] = request.command.strip() - record.update(enabled=False, approved_digest=None) + record["command"] = request.command.strip() if request.command else None + record["url"] = request.url.strip() if request.url else None + record.update( + version=request.version + 1, + enabled=False, + approved_digest=None, + tested_digest=None, + last_tested_at=None, + last_test_succeeded=None, + ) updated = {**self._records, server_id: record} self._write(updated) self._records = updated self._last_status.pop(server_id, None) - for key in removed: + self._summaries.pop(server_id, None) + for kind, key in removed: try: - self.credentials.delete(self._secret_id(server_id, key)) + self.credentials.delete(self._secret_id(server_id, key, kind)) except CredentialStoreError as exc: raise McpRegistryError( "MCP_SECRET_STORE_ERROR", str(exc), status_code=500 ) from exc return self.get(server_id) + @_serialized_lifecycle def delete(self, server_id: str) -> None: self.disable(server_id) with self._lock: record = self._record(server_id) secret_ids = [ - self._secret_id(server_id, key) - for key in record.get("secret_environment_keys", []) + self._secret_id(server_id, key, kind) + for kind, keys in ( + ("environment", record.get("secret_environment_keys", [])), + ("header", record.get("secret_header_keys", [])), + ) + for key in keys ] updated = dict(self._records) del updated[server_id] self._write(updated) self._records = updated self._last_status.pop(server_id, None) + self._summaries.pop(server_id, None) try: self.credentials.delete_many(secret_ids) except CredentialStoreError as exc: @@ -130,6 +206,7 @@ class McpServerRegistry: ) from exc self.bridge.remove(self._host_id(server_id)) + @_serialized_lifecycle def trust(self, server_id: str, command_digest: str) -> McpServer: with self._lock: record = self._record(server_id) @@ -146,40 +223,46 @@ class McpServerRegistry: self._records = updated return self.get(server_id) + @_serialized_lifecycle def put_secret( - self, server_id: str, key: str, secret: str + self, server_id: str, key: str, secret: str, *, kind: str = "environment" ) -> McpServerSecretStatus: with self._lock: record = self._record(server_id) - self._validate_environment_key(key) - if key not in record.get("secret_environment_keys", []): + declared = self._secret_keys(record, kind) + self._validate_secret_key(key, kind) + if key not in declared: raise McpRegistryError( "MCP_SECRET_NOT_DECLARED", "Secret environment key is not declared in this server configuration.", ) try: - self.credentials.put(self._secret_id(server_id, key), secret) + self.credentials.put(self._secret_id(server_id, key, kind), secret) except CredentialStoreError as exc: raise McpRegistryError( "MCP_SECRET_STORE_ERROR", str(exc), status_code=500 ) from exc return McpServerSecretStatus(key=key, configured=True) - def delete_secret(self, server_id: str, key: str) -> McpServerSecretStatus: + @_serialized_lifecycle + def delete_secret( + self, server_id: str, key: str, *, kind: str = "environment" + ) -> McpServerSecretStatus: record = self._record(server_id) - if key not in record.get("secret_environment_keys", []): + if key not in self._secret_keys(record, kind): raise McpRegistryError( "MCP_SECRET_NOT_DECLARED", "Secret environment key is not declared in this server configuration.", ) try: - self.credentials.delete(self._secret_id(server_id, key)) + self.credentials.delete(self._secret_id(server_id, key, kind)) except CredentialStoreError as exc: raise McpRegistryError( "MCP_SECRET_STORE_ERROR", str(exc), status_code=500 ) from exc return McpServerSecretStatus(key=key, configured=False) + @_serialized_lifecycle def test(self, server_id: str) -> McpServer: record = self._record(server_id) if record.get("enabled"): @@ -188,18 +271,31 @@ class McpServerRegistry: "Disable the MCP server before running an isolated connection test.", status_code=409, ) - self._require_launch_allowed(record) + self._require_launch_allowed(record, require_test=False) try: discovered = self._start(server_id, record) except Exception as exc: - self._last_status[server_id] = { + tested_at = datetime.now(UTC) + failure = { "status": PluginHostState.error, "error": str(exc), - "last_tested_at": datetime.now(UTC), + "last_tested_at": tested_at, "last_test_succeeded": False, } + self._last_status[server_id] = failure + with self._lock: + failed_record = { + **record, + "tested_digest": None, + "last_tested_at": tested_at.isoformat(), + "last_test_succeeded": False, + } + updated = {**self._records, server_id: failed_record} + self._write(updated) + self._records = updated raise status = self.bridge.status(self._host_id(server_id), self._backend(record)) + tested_at = datetime.now(UTC) self._last_status[server_id] = { "status": PluginHostState.stopped, "tools_count": len(discovered), @@ -207,18 +303,31 @@ class McpServerRegistry: "remote_server_name": status.server_name, "remote_server_version": status.server_version, "error": None, - "last_tested_at": datetime.now(UTC), + "last_tested_at": tested_at, "last_test_succeeded": True, } + self._summaries[server_id] = self._tool_summaries(discovered) self.bridge.stop(self._host_id(server_id)) + with self._lock: + tested_record = { + **record, + "tested_digest": self._digest(record), + "last_tested_at": tested_at.isoformat(), + "last_test_succeeded": True, + } + updated = {**self._records, server_id: tested_record} + self._write(updated) + self._records = updated return self.get(server_id) + @_serialized_lifecycle def enable(self, server_id: str) -> McpServer: record = self._record(server_id) if server_id in self._registered: return self.get(server_id) - self._require_launch_allowed(record) + self._require_launch_allowed(record, require_test=True) discovered = self._start(server_id, record) + self._summaries[server_id] = self._tool_summaries(discovered) registered: list[str] = [] try: for item in discovered: @@ -243,6 +352,7 @@ class McpServerRegistry: raise return self.get(server_id) + @_serialized_lifecycle def disable(self, server_id: str) -> McpServer: with self._lock: record = self._record(server_id) @@ -255,6 +365,7 @@ class McpServerRegistry: self.bridge.stop(self._host_id(server_id)) return self.get(server_id) + @_serialized_lifecycle def restore_enabled(self) -> None: if not self._records: return @@ -270,6 +381,7 @@ class McpServerRegistry: } self._write() + @_serialized_lifecycle def shutdown(self) -> None: for server_id in list(self._records): for name in self._registered.pop(server_id, []): @@ -280,7 +392,9 @@ class McpServerRegistry: environment = dict(record.get("environment", {})) for key in record.get("secret_environment_keys", []): try: - value = self.credentials.resolve(self._secret_id(server_id, key)) + value = self.credentials.resolve( + self._secret_id(server_id, key, "environment") + ) except CredentialStoreError as exc: raise McpRegistryError( "MCP_SECRET_STORE_ERROR", str(exc), status_code=500 @@ -292,6 +406,23 @@ class McpServerRegistry: status_code=409, ) environment[key] = value + headers = dict(record.get("headers", {})) + for key in record.get("secret_header_keys", []): + try: + value = self.credentials.resolve( + self._secret_id(server_id, key, "header") + ) + except CredentialStoreError as exc: + raise McpRegistryError( + "MCP_SECRET_STORE_ERROR", str(exc), status_code=500 + ) from exc + if value is None: + raise McpRegistryError( + "MCP_SECRET_REQUIRED", + f"Secret HTTP header is not configured: {key}", + status_code=409, + ) + headers[key] = value host_id = self._host_id(server_id) self.bridge.remove(host_id) try: @@ -301,9 +432,16 @@ class McpServerRegistry: self._server_dir(server_id), list(record.get("permissions", [])), lambda _host, message: self._unavailable(server_id, message), - command_override=[record["command"], *record.get("args", [])], + command_override=( + [record["command"], *record.get("args", [])] + if record.get("command") + else None + ), environment=environment, tool_source="mcp_server", + transport_kind=record["transport"], + url=record.get("url"), + headers=headers, ) except McpBridgeError as exc: raise McpRegistryError( @@ -339,23 +477,29 @@ class McpServerRegistry: } self._write() - def _require_launch_allowed(self, record: dict[str, Any]) -> None: - if record.get("transport") != McpServerTransport.stdio.value: - raise McpRegistryError( - "MCP_TRANSPORT_UNSUPPORTED", - "C.1 currently supports stdio; Streamable HTTP and SSE are reserved for a later increment.", - status_code=501, - ) - if not self.allow_process_launch: + def _require_launch_allowed( + self, record: dict[str, Any], *, require_test: bool + ) -> None: + digest = self._digest(record) + if ( + record.get("transport") == McpServerTransport.stdio.value + and not self.allow_process_launch + ): raise McpRegistryError( "MCP_SANDBOX_REQUIRED", "Python process launch is disabled outside development until the desktop sandbox is available.", status_code=403, ) - if record.get("approved_digest") != self._digest(record): + if record.get("approved_digest") != digest: raise McpRegistryError( "MCP_TRUST_APPROVAL_REQUIRED", - "Review and approve the current MCP command before testing or enabling it.", + "Review and approve the current MCP connection before testing or enabling it.", + status_code=409, + ) + if require_test and record.get("tested_digest") != digest: + raise McpRegistryError( + "MCP_CONNECTION_TEST_REQUIRED", + "Test the current MCP configuration successfully before enabling it.", status_code=409, ) @@ -366,15 +510,22 @@ class McpServerRegistry: cached = self._last_status.get(server_id, {}) return McpServer( server_id=server_id, + version=record.get("version", 1), name=record["name"], transport=record["transport"], - command=record["command"], + command=record.get("command"), args=list(record.get("args", [])), + url=record.get("url"), + headers=dict(record.get("headers", {})), environment=dict(record.get("environment", {})), secret_environment={ key: self._secret_configured(server_id, key) for key in record.get("secret_environment_keys", []) }, + secret_headers={ + key: self._secret_configured(server_id, key, "header") + for key in record.get("secret_header_keys", []) + }, permissions=list(record.get("permissions", [])), startup_timeout_seconds=backend.startup_timeout_seconds, tool_timeout_seconds=backend.tool_timeout_seconds, @@ -398,8 +549,10 @@ class McpServerRegistry: if record.get("enabled") else cached.get("remote_server_version"), error=status.error if record.get("enabled") else cached.get("error"), - last_tested_at=cached.get("last_tested_at"), - last_test_succeeded=cached.get("last_test_succeeded"), + last_tested_at=record.get("last_tested_at") or cached.get("last_tested_at"), + last_test_succeeded=record.get("last_test_succeeded") + if record.get("last_test_succeeded") is not None + else cached.get("last_test_succeeded"), ) def _validate(self, request: McpServerCreateRequest) -> None: @@ -407,12 +560,36 @@ class McpServerRegistry: raise McpRegistryError( "MCP_SERVER_NAME_INVALID", "MCP server name cannot be blank." ) - if not request.command.strip() or "\x00" in request.command: - raise McpRegistryError("MCP_COMMAND_INVALID", "MCP executable is invalid.") - if any("\x00" in arg for arg in request.args): - raise McpRegistryError( - "MCP_COMMAND_INVALID", "MCP argument contains a null byte." - ) + if request.transport == McpServerTransport.stdio: + if ( + not request.command + or not request.command.strip() + or "\x00" in request.command + ): + raise McpRegistryError( + "MCP_COMMAND_INVALID", "MCP executable is invalid." + ) + if request.url or request.headers or request.secret_header_keys: + raise McpRegistryError( + "MCP_CONFIG_INVALID", + "stdio configuration cannot contain HTTP fields.", + ) + if any("\x00" in arg for arg in request.args): + raise McpRegistryError( + "MCP_COMMAND_INVALID", "MCP argument contains a null byte." + ) + else: + self._validate_http_url(request.url) + if ( + request.command + or request.args + or request.environment + or request.secret_environment_keys + ): + raise McpRegistryError( + "MCP_CONFIG_INVALID", + "HTTP configuration cannot contain stdio fields.", + ) for key in [*request.environment, *request.secret_environment_keys]: self._validate_environment_key(key) if set(request.environment) & set(request.secret_environment_keys): @@ -420,6 +597,21 @@ class McpServerRegistry: "MCP_ENVIRONMENT_INVALID", "An environment key cannot be both plain and secret.", ) + plain_headers = {key.casefold() for key in request.headers} + secret_headers = {key.casefold() for key in request.secret_header_keys} + for key in [*request.headers, *request.secret_header_keys]: + self._validate_header_key(key) + if any( + "\r" in value or "\n" in value or "\x00" in value + for value in request.headers.values() + ): + raise McpRegistryError( + "MCP_HEADER_INVALID", "HTTP header value contains control characters." + ) + if plain_headers & secret_headers: + raise McpRegistryError( + "MCP_HEADER_INVALID", "An HTTP header cannot be both plain and secret." + ) unknown_permissions = set(request.permissions) - KNOWN_PERMISSIONS if unknown_permissions: raise McpRegistryError( @@ -434,12 +626,36 @@ class McpServerRegistry: "MCP_ENVIRONMENT_INVALID", f"Invalid environment variable name: {key}" ) + @staticmethod + def _validate_header_key(key: str) -> None: + if not _HEADER_KEY.fullmatch(key) or key.casefold() in _RESERVED_HEADERS: + raise McpRegistryError( + "MCP_HEADER_INVALID", f"Invalid or reserved HTTP header: {key}" + ) + + @staticmethod + def _validate_http_url(url: str | None) -> None: + if not url: + raise McpRegistryError("MCP_URL_INVALID", "MCP HTTP URL is required.") + parts = urlsplit(url.strip()) + if ( + parts.scheme not in {"http", "https"} + or not parts.hostname + or parts.username is not None + or parts.password is not None + or parts.fragment + ): + raise McpRegistryError( + "MCP_URL_INVALID", + "MCP URL must be an HTTP(S) URL without credentials or fragments.", + ) + @staticmethod def _backend(record: dict[str, Any]) -> PluginBackend: return PluginBackend( type="mcp", transport="stdio", - command=record["command"], + command=record.get("command") or "http", args=record.get("args", []), startup_timeout_seconds=record.get("startup_timeout_seconds", 15), tool_timeout_seconds=record.get("tool_timeout_seconds", 30), @@ -464,6 +680,9 @@ class McpServerRegistry: "args", "environment", "secret_environment_keys", + "url", + "headers", + "secret_header_keys", "permissions", ) } @@ -475,9 +694,16 @@ class McpServerRegistry: @staticmethod def _summary(record: dict[str, Any]) -> str: + if record.get("transport") != McpServerTransport.stdio.value: + header_names = sorted( + [*record.get("headers", {}), *record.get("secret_header_keys", [])], + key=str.casefold, + ) + suffix = f" headers={','.join(header_names)}" if header_names else "" + return f"{record.get('transport')} {record.get('url') or ''}{suffix}" return " ".join( [ - record["command"], + record.get("command") or "", *[ json.dumps(arg, ensure_ascii=False) for arg in record.get("args", []) @@ -486,18 +712,51 @@ class McpServerRegistry: ) @staticmethod - def _secret_id(server_id: str, key: str) -> str: - suffix = hashlib.sha256(key.encode()).hexdigest()[:20] + def _secret_id(server_id: str, key: str, kind: str = "environment") -> str: + suffix = hashlib.sha256(f"{kind}\0{key.casefold()}".encode()).hexdigest()[:20] return f"mcp.{server_id}.{suffix}" - def _secret_configured(self, server_id: str, key: str) -> bool: + def _secret_configured( + self, server_id: str, key: str, kind: str = "environment" + ) -> bool: try: - return self.credentials.has(self._secret_id(server_id, key)) + return self.credentials.has(self._secret_id(server_id, key, kind)) except CredentialStoreError as exc: raise McpRegistryError( "MCP_SECRET_STORE_ERROR", str(exc), status_code=500 ) from exc + @staticmethod + def _secret_keys(record: dict[str, Any], kind: str) -> list[str]: + if kind == "environment": + return list(record.get("secret_environment_keys", [])) + if kind == "header": + return list(record.get("secret_header_keys", [])) + raise McpRegistryError("MCP_SECRET_KIND_INVALID", "Unknown MCP secret kind.") + + @staticmethod + def _validate_secret_key(key: str, kind: str) -> None: + if kind == "environment": + McpServerRegistry._validate_environment_key(key) + elif kind == "header": + McpServerRegistry._validate_header_key(key) + else: + raise McpRegistryError( + "MCP_SECRET_KIND_INVALID", "Unknown MCP secret kind." + ) + + @staticmethod + def _tool_summaries(discovered: list[McpDiscoveredTool]) -> list[McpToolSummary]: + return [ + McpToolSummary( + name=item.definition.name, + remote_name=item.remote_name, + description=item.definition.description, + permission=item.definition.permission, + ) + for item in discovered + ] + def _record(self, server_id: str) -> dict[str, Any]: try: return self._records[server_id] diff --git a/backend/app/routes.py b/backend/app/routes.py index 438d41b..a1bd49a 100644 --- a/backend/app/routes.py +++ b/backend/app/routes.py @@ -6,6 +6,8 @@ from uuid import uuid4 from fastapi import APIRouter, Header, Query from fastapi.responses import StreamingResponse +from app.agent import AgentCapacityError, AgentRunNotFoundError +from app.container import container from app.contracts import ( AgentRun, AgentRunCreateRequest, @@ -28,6 +30,7 @@ from app.contracts import ( McpServerSecretWriteRequest, McpServerTrustRequest, McpServerUpdateRequest, + McpToolSummaryListResponse, ModelEvent, ModelEventType, Note, @@ -75,18 +78,16 @@ from app.contracts import ( WorkspaceOpenRequest, WorkspaceSnapshot, ) -from app.agent import AgentCapacityError, AgentRunNotFoundError -from app.container import container from app.errors import ApiError from app.extensions import ExtensionError from app.extensions.mcp_registry import McpRegistryError -from app.providers.registry import ProviderNotFoundError -from app.providers.factory import UnsupportedProviderError from app.providers.base import ProviderError from app.providers.credentials import ( CredentialStoreError, validate_provider_credential_id, ) +from app.providers.factory import UnsupportedProviderError +from app.providers.registry import ProviderNotFoundError from app.retrieval.engine import engine from app.services import ( index_service, @@ -225,14 +226,21 @@ async def list_notes( folder: str | None = None, tag: str | None = None, ) -> NoteListResponse: - items, total = note_service.list_notes(limit=limit, offset=offset, folder=folder, tag=tag) - return NoteListResponse(items=items, page=PageMeta(total=total, limit=limit, offset=offset)) + items, total = note_service.list_notes( + limit=limit, offset=offset, folder=folder, tag=tag + ) + return NoteListResponse( + items=items, page=PageMeta(total=total, limit=limit, offset=offset) + ) @router.post("/notes", response_model=Note, tags=["Notes"]) async def create_note(request: NoteCreateRequest) -> Note: return await note_service.create_note( - title=request.title, markdown=request.markdown, folder=request.folder, tags=request.tags + title=request.title, + markdown=request.markdown, + folder=request.folder, + tags=request.tags, ) @@ -240,7 +248,9 @@ async def create_note(request: NoteCreateRequest) -> Note: async def get_note(note_id: str) -> Note: note = await note_service.get_note(note_id) if note is None: - raise ApiError(404, "RESOURCE_NOT_FOUND", "note not found", {"note_id": note_id}) + raise ApiError( + 404, "RESOURCE_NOT_FOUND", "note not found", {"note_id": note_id} + ) return note @@ -254,7 +264,9 @@ async def update_note(note_id: str, request: NoteUpdateRequest) -> Note: @router.delete("/notes/{note_id}", response_model=OperationResponse, tags=["Notes"]) async def delete_note(note_id: str) -> OperationResponse: if not await note_service.delete_note(note_id): - raise ApiError(404, "RESOURCE_NOT_FOUND", "note not found", {"note_id": note_id}) + raise ApiError( + 404, "RESOURCE_NOT_FOUND", "note not found", {"note_id": note_id} + ) return OperationResponse(status="completed", resource_id=note_id, message="deleted") @@ -298,7 +310,9 @@ async def chat(request: ChatRequest) -> StreamingResponse: data={"code": "PROVIDER_ERROR", "message": str(exc)}, timestamp=utc_now(), ) - done = ModelEvent(event=ModelEventType.done, sequence=1, timestamp=utc_now()) + done = ModelEvent( + event=ModelEventType.done, sequence=1, timestamp=utc_now() + ) yield as_sse(error.event.value, error.model_dump_json()) yield as_sse(done.event.value, done.model_dump_json()) @@ -459,9 +473,7 @@ async def list_skills() -> SkillListResponse: return SkillListResponse(items=container.skills.list()) -@router.get( - "/skills/{skill_id}", response_model=Skill, tags=["Skills"] -) +@router.get("/skills/{skill_id}", response_model=Skill, tags=["Skills"]) async def get_skill(skill_id: str) -> Skill: return extension_call(lambda: container.skills.get(skill_id)) @@ -501,7 +513,9 @@ async def disable_skill(skill_id: str) -> Skill: ) async def uninstall_skill(skill_id: str) -> OperationResponse: extension_call(lambda: container.skills.uninstall(skill_id)) - return OperationResponse(status="completed", resource_id=skill_id, message="uninstalled") + return OperationResponse( + status="completed", resource_id=skill_id, message="uninstalled" + ) # Independent MCP Server Registry @@ -510,7 +524,9 @@ async def list_mcp_servers() -> McpServerListResponse: return McpServerListResponse(items=mcp_call(container.mcp_servers.list)) -@router.post("/mcp/servers", response_model=McpServer, status_code=201, tags=["MCP Servers"]) +@router.post( + "/mcp/servers", response_model=McpServer, status_code=201, tags=["MCP Servers"] +) async def create_mcp_server(request: McpServerCreateRequest) -> McpServer: return mcp_call(lambda: container.mcp_servers.create(request)) @@ -520,45 +536,97 @@ async def get_mcp_server(server_id: str) -> McpServer: return mcp_call(lambda: container.mcp_servers.get(server_id)) +@router.get( + "/mcp/servers/{server_id}/tools", + response_model=McpToolSummaryListResponse, + tags=["MCP Servers"], +) +async def list_mcp_server_tools(server_id: str) -> McpToolSummaryListResponse: + return McpToolSummaryListResponse( + items=mcp_call(lambda: container.mcp_servers.list_tools(server_id)) + ) + + @router.put("/mcp/servers/{server_id}", response_model=McpServer, tags=["MCP Servers"]) -async def update_mcp_server(server_id: str, request: McpServerUpdateRequest) -> McpServer: - return await mcp_call_async(lambda: container.mcp_servers.update(server_id, request)) +async def update_mcp_server( + server_id: str, request: McpServerUpdateRequest +) -> McpServer: + return await mcp_call_async( + lambda: container.mcp_servers.update(server_id, request) + ) -@router.delete("/mcp/servers/{server_id}", response_model=OperationResponse, tags=["MCP Servers"]) +@router.delete( + "/mcp/servers/{server_id}", response_model=OperationResponse, tags=["MCP Servers"] +) async def delete_mcp_server(server_id: str) -> OperationResponse: await mcp_call_async(lambda: container.mcp_servers.delete(server_id)) - return OperationResponse(status="completed", resource_id=server_id, message="deleted") + return OperationResponse( + status="completed", resource_id=server_id, message="deleted" + ) -@router.post("/mcp/servers/{server_id}/trust", response_model=McpServer, tags=["MCP Servers"]) +@router.post( + "/mcp/servers/{server_id}/trust", response_model=McpServer, tags=["MCP Servers"] +) async def trust_mcp_server(server_id: str, request: McpServerTrustRequest) -> McpServer: - return mcp_call(lambda: container.mcp_servers.trust(server_id, request.command_digest)) + return mcp_call( + lambda: container.mcp_servers.trust(server_id, request.command_digest) + ) -@router.post("/mcp/servers/{server_id}/test", response_model=McpServer, tags=["MCP Servers"]) +@router.post( + "/mcp/servers/{server_id}/test", response_model=McpServer, tags=["MCP Servers"] +) async def test_mcp_server(server_id: str) -> McpServer: return await mcp_call_async(lambda: container.mcp_servers.test(server_id)) -@router.post("/mcp/servers/{server_id}/enable", response_model=McpServer, tags=["MCP Servers"]) +@router.post( + "/mcp/servers/{server_id}/enable", response_model=McpServer, tags=["MCP Servers"] +) async def enable_mcp_server(server_id: str) -> McpServer: return await mcp_call_async(lambda: container.mcp_servers.enable(server_id)) -@router.post("/mcp/servers/{server_id}/disable", response_model=McpServer, tags=["MCP Servers"]) +@router.post( + "/mcp/servers/{server_id}/disable", response_model=McpServer, tags=["MCP Servers"] +) async def disable_mcp_server(server_id: str) -> McpServer: return await mcp_call_async(lambda: container.mcp_servers.disable(server_id)) -@router.put("/mcp/servers/{server_id}/secrets/{key}", response_model=McpServerSecretStatus, tags=["MCP Servers"]) -async def put_mcp_server_secret(server_id: str, key: str, request: McpServerSecretWriteRequest) -> McpServerSecretStatus: - return mcp_call(lambda: container.mcp_servers.put_secret(server_id, key, request.secret.get_secret_value())) +@router.put( + "/mcp/servers/{server_id}/secrets/{key}", + response_model=McpServerSecretStatus, + tags=["MCP Servers"], +) +async def put_mcp_server_secret( + server_id: str, + key: str, + request: McpServerSecretWriteRequest, + kind: str = Query(default="environment", pattern="^(environment|header)$"), +) -> McpServerSecretStatus: + return mcp_call( + lambda: container.mcp_servers.put_secret( + server_id, key, request.secret.get_secret_value(), kind=kind + ) + ) -@router.delete("/mcp/servers/{server_id}/secrets/{key}", response_model=McpServerSecretStatus, tags=["MCP Servers"]) -async def delete_mcp_server_secret(server_id: str, key: str) -> McpServerSecretStatus: - return mcp_call(lambda: container.mcp_servers.delete_secret(server_id, key)) +@router.delete( + "/mcp/servers/{server_id}/secrets/{key}", + response_model=McpServerSecretStatus, + tags=["MCP Servers"], +) +async def delete_mcp_server_secret( + server_id: str, + key: str, + kind: str = Query(default="environment", pattern="^(environment|header)$"), +) -> McpServerSecretStatus: + return mcp_call( + lambda: container.mcp_servers.delete_secret(server_id, key, kind=kind) + ) # Plugins @@ -650,11 +718,15 @@ async def restart_plugin_host(plugin_id: str) -> OperationResponse: ) async def uninstall_plugin(plugin_id: str) -> OperationResponse: plugin = extension_call(lambda: container.plugins.get(plugin_id)) - dependent_skills = container.skills.depending_on_tools(plugin.manifest.contributes.tools) + dependent_skills = container.skills.depending_on_tools( + plugin.manifest.contributes.tools + ) await extension_call_async( lambda: container.plugins.uninstall(plugin_id, dependent_skills) ) - return OperationResponse(status="completed", resource_id=plugin_id, message="uninstalled") + return OperationResponse( + status="completed", resource_id=plugin_id, message="uninstalled" + ) # Plugin Command / Settings Contributions @@ -729,9 +801,7 @@ async def put_plugin_setting_secret( response_model=PluginSecretStatus, tags=["Plugins"], ) -async def delete_plugin_setting_secret( - plugin_id: str, key: str -) -> PluginSecretStatus: +async def delete_plugin_setting_secret(plugin_id: str, key: str) -> PluginSecretStatus: return extension_call( lambda: container.plugins.delete_setting_secret(plugin_id, key) ) @@ -844,7 +914,9 @@ async def update_provider( ) -> ProviderConfig: current = configurable_provider_or_404(provider_id).config if provider_id == "mock": - raise ApiError(409, "BUILTIN_PROVIDER_IMMUTABLE", "Mock provider cannot be modified.") + raise ApiError( + 409, "BUILTIN_PROVIDER_IMMUTABLE", "Mock provider cannot be modified." + ) fields = request.model_fields_set if ("name" in fields and request.name is None) or ( "enabled" in fields and request.enabled is None @@ -873,7 +945,9 @@ async def update_provider( async def delete_provider(provider_id: str) -> OperationResponse: configurable_provider_or_404(provider_id) if provider_id == "mock": - raise ApiError(409, "BUILTIN_PROVIDER_IMMUTABLE", "Mock provider cannot be deleted.") + raise ApiError( + 409, "BUILTIN_PROVIDER_IMMUTABLE", "Mock provider cannot be deleted." + ) container.providers.unregister(provider_id) return OperationResponse(status="completed", resource_id=provider_id) @@ -951,7 +1025,9 @@ async def create_task(request: TaskCreateRequest) -> Task: async def get_task(task_id: str) -> Task: task = task_service.get_task(task_id) if task is None: - raise ApiError(404, "RESOURCE_NOT_FOUND", "task not found", {"task_id": task_id}) + raise ApiError( + 404, "RESOURCE_NOT_FOUND", "task not found", {"task_id": task_id} + ) return task @@ -967,7 +1043,9 @@ async def update_task(task_id: str, request: TaskUpdateRequest) -> Task: ) async def delete_task(task_id: str) -> OperationResponse: if not task_service.delete_task(task_id): - raise ApiError(404, "RESOURCE_NOT_FOUND", "task not found", {"task_id": task_id}) + raise ApiError( + 404, "RESOURCE_NOT_FOUND", "task not found", {"task_id": task_id} + ) return OperationResponse(status="completed", resource_id=task_id, message="deleted") @@ -1017,5 +1095,7 @@ async def rebuild_index(request: IndexRebuildRequest) -> IndexJob: async def get_index_job(job_id: str) -> IndexJob: job = index_service.get_job(job_id) if job is None: - raise ApiError(404, "RESOURCE_NOT_FOUND", "index job not found", {"job_id": job_id}) + raise ApiError( + 404, "RESOURCE_NOT_FOUND", "index job not found", {"job_id": job_id} + ) return job diff --git a/backend/tests/test_mcp_registry.py b/backend/tests/test_mcp_registry.py index a4eb658..ae0e9fd 100644 --- a/backend/tests/test_mcp_registry.py +++ b/backend/tests/test_mcp_registry.py @@ -1,10 +1,15 @@ +import asyncio +import json import sys +import time +from concurrent.futures import ThreadPoolExecutor +import httpx import pytest -from app.agent.tools import ToolRegistry +from app.agent.tools import ToolExecutionContext, ToolRegistry from app.config import BACKEND_DIR, get_settings -from app.contracts import McpServerCreateRequest, McpServerUpdateRequest +from app.contracts import McpServerCreateRequest, McpServerUpdateRequest, ToolCall from app.extensions.mcp_registry import McpRegistryError, McpServerRegistry from app.providers.credentials import EncryptedCredentialStore @@ -58,6 +63,7 @@ def test_update_disables_server_and_revokes_command_trust() -> None: service = registry() created = service.create(request(secret_environment_keys=[])) service.trust(created.server_id, created.command_digest) + service.test(created.server_id) enabled = service.enable(created.server_id) assert enabled.enabled is True assert any( @@ -68,7 +74,8 @@ def test_update_disables_server_and_revokes_command_trust() -> None: updated = service.update( created.server_id, McpServerUpdateRequest( - **request(name="Changed", secret_environment_keys=[]).model_dump() + **request(name="Changed", secret_environment_keys=[]).model_dump(), + version=enabled.version, ), ) assert updated.enabled is False @@ -89,25 +96,73 @@ def test_production_rejects_process_launch_even_after_approval() -> None: assert error.value.code == "MCP_SANDBOX_REQUIRED" -def test_non_stdio_transport_is_explicitly_reserved() -> None: +def test_enable_requires_successful_test_and_update_checks_version() -> None: + service = registry() + created = service.create(request(secret_environment_keys=[])) + service.trust(created.server_id, created.command_digest) + with pytest.raises(McpRegistryError) as error: + service.enable(created.server_id) + assert error.value.code == "MCP_CONNECTION_TEST_REQUIRED" + + with pytest.raises(McpRegistryError) as error: + service.update( + created.server_id, + McpServerUpdateRequest( + **request(secret_environment_keys=[]).model_dump(), version=99 + ), + ) + assert error.value.code == "MCP_SERVER_VERSION_CONFLICT" + + +def test_http_transport_rejects_invalid_cross_transport_fields() -> None: + service = registry() + with pytest.raises(McpRegistryError) as error: + service.create( + request( + transport="streamable_http", + url="https://example.invalid/mcp", + secret_environment_keys=[], + ) + ) + assert error.value.code == "MCP_CONFIG_INVALID" + + +def test_registry_rejects_corrupt_persisted_json(tmp_path) -> None: + path = tmp_path / "mcp" + path.mkdir() + (path / "servers.json").write_text("{broken", encoding="utf-8") + with pytest.raises(McpRegistryError) as error: + McpServerRegistry( + ToolRegistry(), + EncryptedCredentialStore(), + tmp_path, + allow_process_launch=True, + ) + assert error.value.code == "MCP_REGISTRY_INVALID" + + +def test_stdio_command_is_not_parsed_as_a_shell_string() -> None: service = registry() created = service.create( request( - transport="streamable_http", - command="https://example.invalid/mcp", + command=f'"{sys.executable}" "{SERVER}"', + args=[], secret_environment_keys=[], ) ) service.trust(created.server_id, created.command_digest) with pytest.raises(McpRegistryError) as error: service.test(created.server_id) - assert error.value.code == "MCP_TRANSPORT_UNSUPPORTED" + assert error.value.code == "PLUGIN_HOST_START_FAILED" + assert service.get(created.server_id).last_test_succeeded is False + service.shutdown() def test_enabled_server_is_restored_from_persisted_registry() -> None: first = registry() created = first.create(request(secret_environment_keys=[])) first.trust(created.server_id, created.command_digest) + first.test(created.server_id) first.enable(created.server_id) first.shutdown() @@ -121,3 +176,222 @@ def test_enabled_server_is_restored_from_persisted_registry() -> None: for item in restored.tools.definitions() ) restored.shutdown() + + +def test_lifecycle_operations_are_serialized_and_tool_names_are_isolated() -> None: + service = registry() + servers = [ + service.create(request(name=f"Echo {index}", secret_environment_keys=[])) + for index in range(2) + ] + for server in servers: + service.trust(server.server_id, server.command_digest) + service.test(server.server_id) + + with ThreadPoolExecutor(max_workers=4) as pool: + enabled = list(pool.map(lambda item: service.enable(item.server_id), servers * 2)) + assert all(item.enabled for item in enabled) + names = [ + item.name + for item in service.tools.definitions() + if item.source == "mcp_server" + ] + assert len(names) == len(set(names)) + assert all(any(name.startswith(f"mcp.{item.server_id}.") for name in names) for item in servers) + + with ThreadPoolExecutor(max_workers=4) as pool: + list(pool.map(lambda item: service.disable(item.server_id), servers * 2)) + assert not any(item.source == "mcp_server" for item in service.tools.definitions()) + service.shutdown() + + +def _http_result(request_id: int, result: dict) -> httpx.Response: + return httpx.Response( + 200, + headers={"content-type": "application/json"}, + json={"jsonrpc": "2.0", "id": request_id, "result": result}, + ) + + +def test_streamable_http_supports_session_headers_secrets_and_tool_summary( + monkeypatch, +) -> None: + requests: list[httpx.Request] = [] + + def handler(request_value: httpx.Request) -> httpx.Response: + requests.append(request_value) + if request_value.method == "GET": + return httpx.Response(405) + if request_value.method == "DELETE": + return httpx.Response(405) + payload = json.loads(request_value.content) + if payload.get("method") == "initialize": + response = _http_result( + payload["id"], + { + "protocolVersion": "2025-11-25", + "capabilities": {"tools": {}}, + "serverInfo": {"name": "HTTP Fixture", "version": "1"}, + }, + ) + response.headers["MCP-Session-Id"] = "session-test" + return response + if payload.get("method") == "tools/list": + return _http_result( + payload["id"], + { + "tools": [ + { + "name": "echo", + "description": "Echo over HTTP", + "inputSchema": {"type": "object", "properties": {}}, + } + ] + }, + ) + if payload.get("method") == "tools/call": + return _http_result( + payload["id"], {"structuredContent": {"transport": "http"}} + ) + return httpx.Response(202) + + real_client = httpx.Client + monkeypatch.setattr( + "app.extensions.mcp.httpx.Client", + lambda **kwargs: real_client(transport=httpx.MockTransport(handler), **kwargs), + ) + service = registry() + created = service.create( + McpServerCreateRequest( + name="Remote MCP", + transport="streamable_http", + url="https://mcp.example.test/mcp", + headers={"X-Client": "NotesAgent"}, + secret_header_keys=["Authorization"], + ) + ) + service.put_secret( + created.server_id, "Authorization", "Bearer hidden", kind="header" + ) + service.trust(created.server_id, created.command_digest) + tested = service.test(created.server_id) + + assert tested.last_test_succeeded is True + assert tested.secret_headers == {"Authorization": True} + assert "Bearer hidden" not in tested.model_dump_json() + assert service.list_tools(created.server_id)[0].remote_name == "echo" + assert any( + request.headers.get("mcp-session-id") == "session-test" for request in requests + ) + assert any( + request.headers.get("mcp-protocol-version") == "2025-11-25" + for request in requests + ) + assert all( + request.headers.get("authorization") == "Bearer hidden" for request in requests + ) + enabled = service.enable(created.server_id) + tool_name = service.list_tools(created.server_id)[0].name + result = asyncio.run( + service.tools.execute( + ToolCall(tool_call_id="call-1", name=tool_name, arguments={}), + ToolExecutionContext(run_id="run-1"), + ) + ) + assert enabled.enabled is True + assert result.success is True + assert result.output == {"transport": "http"} + service.disable(created.server_id) + service.shutdown() + + +class _LegacyEventStream(httpx.SyncByteStream): + def __iter__(self): + yield b"event: endpoint\ndata: /messages\n\n" + time.sleep(0.1) + initialize = { + "jsonrpc": "2.0", + "id": 1, + "result": { + "protocolVersion": "2024-11-05", + "capabilities": {"tools": {}}, + "serverInfo": {"name": "Legacy Fixture"}, + }, + } + yield f"data: {json.dumps(initialize)}\n\n".encode() + time.sleep(0.1) + tools = { + "jsonrpc": "2.0", + "id": 2, + "result": {"tools": []}, + } + yield f"data: {json.dumps(tools)}\n\n".encode() + + +def test_legacy_sse_uses_same_origin_endpoint(monkeypatch) -> None: + posted_urls: list[str] = [] + + def handler(request_value: httpx.Request) -> httpx.Response: + if request_value.method == "GET": + return httpx.Response( + 200, + headers={"content-type": "text/event-stream"}, + stream=_LegacyEventStream(), + ) + posted_urls.append(str(request_value.url)) + return httpx.Response(202) + + real_client = httpx.Client + monkeypatch.setattr( + "app.extensions.mcp.httpx.Client", + lambda **kwargs: real_client(transport=httpx.MockTransport(handler), **kwargs), + ) + service = registry() + created = service.create( + McpServerCreateRequest( + name="Legacy MCP", + transport="sse", + url="https://legacy.example.test/sse", + ) + ) + service.trust(created.server_id, created.command_digest) + tested = service.test(created.server_id) + assert tested.last_test_succeeded is True + assert posted_urls and all( + url == "https://legacy.example.test/messages" for url in posted_urls + ) + service.shutdown() + + +class _CrossOriginLegacyEventStream(httpx.SyncByteStream): + def __iter__(self): + yield b"event: endpoint\ndata: https://attacker.example/messages\n\n" + + +def test_legacy_sse_rejects_cross_origin_message_endpoint(monkeypatch) -> None: + def handler(request_value: httpx.Request) -> httpx.Response: + assert request_value.method == "GET" + return httpx.Response( + 200, + headers={"content-type": "text/event-stream"}, + stream=_CrossOriginLegacyEventStream(), + ) + + real_client = httpx.Client + monkeypatch.setattr( + "app.extensions.mcp.httpx.Client", + lambda **kwargs: real_client(transport=httpx.MockTransport(handler), **kwargs), + ) + service = registry() + created = service.create( + McpServerCreateRequest( + name="Unsafe legacy MCP", + transport="sse", + url="https://legacy.example.test/sse", + ) + ) + service.trust(created.server_id, created.command_digest) + with pytest.raises(McpRegistryError) as error: + service.test(created.server_id) + assert error.value.code == "MCP_HTTP_RESPONSE_INVALID" + service.shutdown() diff --git a/frontend/src/contracts/index.ts b/frontend/src/contracts/index.ts index 05026c5..07f6679 100644 --- a/frontend/src/contracts/index.ts +++ b/frontend/src/contracts/index.ts @@ -539,20 +539,26 @@ export type McpServerTransport = 'stdio' | 'streamable_http' | 'sse' export type McpServerState = 'stopped' | 'starting' | 'ready' | 'unhealthy' | 'error' export interface McpServerInput { + version?: number name: string transport: McpServerTransport - command: string + command?: string | null args: string[] + url?: string | null + headers: Record environment: Record secret_environment_keys: string[] + secret_header_keys: string[] permissions: string[] startup_timeout_seconds: number tool_timeout_seconds: number } -export interface McpServer extends Omit { +export interface McpServer extends Omit { server_id: string + version: number secret_environment: Record + secret_headers: Record enabled: boolean trusted: boolean command_digest: string @@ -567,6 +573,13 @@ export interface McpServer extends Omit ({ + listMcpServers: vi.fn(), createMcpServer: vi.fn(), updateMcpServer: vi.fn(), + deleteMcpServer: vi.fn(), trustMcpServer: vi.fn(), testMcpServer: vi.fn(), + enableMcpServer: vi.fn(), disableMcpServer: vi.fn(), putMcpServerSecret: vi.fn(), +})) + +const server: McpServer = { + server_id: 'server-1', version: 2, name: 'Remote', transport: 'streamable_http', + command: null, args: [], url: 'https://mcp.example.test/mcp', headers: {}, environment: {}, + secret_environment: {}, secret_headers: { Authorization: false }, permissions: [], + startup_timeout_seconds: 15, tool_timeout_seconds: 30, enabled: false, trusted: true, + command_digest: 'a'.repeat(64), command_summary: 'https://mcp.example.test/mcp', + status: 'stopped', tools_count: 1, last_test_succeeded: false, +} + +async function render(items: McpServer[] = []) { + vi.mocked(service.listMcpServers).mockResolvedValue(items) + const wrapper = mount(McpServersView, { global: { stubs: { AppIcon: true } } }) + await flushPromises() + return wrapper +} + +beforeEach(() => { + vi.clearAllMocks() + vi.stubGlobal('confirm', vi.fn(() => true)) +}) + +describe('McpServersView', () => { + it('switches transport templates and round-trips the JSON configuration mode', async () => { + const wrapper = await render() + await wrapper.findAll('button').find(button => button.text() === '新增服务器')!.trigger('click') + await wrapper.findAll('button').find(button => button.text() === 'Streamable HTTP')!.trigger('click') + expect(wrapper.find('input[placeholder="https://example.com/mcp"]').exists()).toBe(true) + + await wrapper.findAll('button').find(button => button.text() === 'JSON 配置')!.trigger('click') + const raw = (wrapper.get('.json-editor').element as HTMLTextAreaElement).value + expect(JSON.parse(raw)).toMatchObject({ transport: 'streamable_http', command: null }) + expect(raw).not.toContain('secret_value') + + await wrapper.findAll('button').find(button => button.text() === '表单配置')!.trigger('click') + expect(wrapper.text()).toContain('MCP URL') + }) + + it('rejects invalid JSON without sending a create request', async () => { + const wrapper = await render() + await wrapper.findAll('button').find(button => button.text() === '新增服务器')!.trigger('click') + await wrapper.findAll('button').find(button => button.text() === 'JSON 配置')!.trigger('click') + await wrapper.get('.json-editor').setValue('{invalid') + await flushPromises() + await wrapper.get('form').trigger('submit') + await flushPromises() + expect(wrapper.text()).toContain('服务器配置不是有效 JSON') + expect(service.createMcpServer).not.toHaveBeenCalled() + }) + + it('keeps secrets request-only, exposes test failures, and confirms deletion', async () => { + const wrapper = await render([server]) + const password = wrapper.get('input[type="password"]') + await password.setValue('request-only-secret') + vi.mocked(service.putMcpServerSecret).mockResolvedValue({} as never) + await wrapper.findAll('button').find(button => button.text() === '保存')!.trigger('click') + await flushPromises() + expect(service.putMcpServerSecret).toHaveBeenCalledWith('server-1', 'Authorization', 'request-only-secret', 'header') + expect((password.element as HTMLInputElement).value).toBe('') + + vi.mocked(service.testMcpServer).mockRejectedValue(new Error('连接失败')) + await wrapper.findAll('button').find(button => button.text().includes('测试连接'))!.trigger('click') + await flushPromises() + expect(wrapper.text()).toContain('连接失败') + + vi.mocked(service.deleteMcpServer).mockResolvedValue({ status: 'completed' }) + await wrapper.findAll('button').find(button => button.text().includes('删除'))!.trigger('click') + await flushPromises() + expect(confirm).toHaveBeenCalled() + expect(service.deleteMcpServer).toHaveBeenCalledWith('server-1') + }) +}) diff --git a/frontend/src/features/mcp/McpServersView.vue b/frontend/src/features/mcp/McpServersView.vue index 11a6957..04f1105 100644 --- a/frontend/src/features/mcp/McpServersView.vue +++ b/frontend/src/features/mcp/McpServersView.vue @@ -5,76 +5,134 @@ import AppIcon from '@/components/common/AppIcon.vue' import type { McpServer, McpServerInput, McpServerTransport } from '@/contracts' import * as service from '@/services/mcpServerService' +type SecretKind = 'environment' | 'header' + const servers = ref([]) const busy = ref('') const error = ref('') const dialogOpen = ref(false) const editingId = ref(null) +const editingOriginal = ref(null) +const editorMode = ref<'form' | 'json'>('form') const argsText = ref('') const environmentText = ref('{}') +const headersText = ref('{}') const secretKeysText = ref('') +const secretHeaderKeysText = ref('') const permissionsText = ref('') +const rawConfig = ref('') const secretDrafts = reactive>({}) -const form = reactive({ - name: '', transport: 'stdio', command: '', args: [], environment: {}, - secret_environment_keys: [], permissions: [], startup_timeout_seconds: 15, tool_timeout_seconds: 30, -}) +const form = reactive(emptyForm()) const dialogTitle = computed(() => editingId.value ? '编辑 MCP 服务器' : '新增 MCP 服务器') +function emptyForm(): McpServerInput { + return { + name: '', transport: 'stdio', command: '', args: [], url: null, headers: {}, + environment: {}, secret_environment_keys: [], secret_header_keys: [], permissions: [], + startup_timeout_seconds: 15, tool_timeout_seconds: 30, + } +} + async function load() { error.value = '' try { servers.value = await service.listMcpServers() } catch (cause) { error.value = message(cause, '读取 MCP 服务器失败') } } +function resetEditor(input: McpServerInput) { + Object.assign(form, input) + argsText.value = input.args.join('\n') + environmentText.value = JSON.stringify(input.environment, null, 2) + headersText.value = JSON.stringify(input.headers, null, 2) + secretKeysText.value = input.secret_environment_keys.join('\n') + secretHeaderKeysText.value = input.secret_header_keys.join('\n') + permissionsText.value = input.permissions.join(', ') + editorMode.value = 'form' + rawConfig.value = '' +} + function openCreate() { editingId.value = null - Object.assign(form, { name: '', transport: 'stdio', command: '', args: [], environment: {}, secret_environment_keys: [], permissions: [], startup_timeout_seconds: 15, tool_timeout_seconds: 30 }) - argsText.value = ''; environmentText.value = '{}'; secretKeysText.value = ''; permissionsText.value = '' + editingOriginal.value = null + resetEditor(emptyForm()) dialogOpen.value = true } function openEdit(server: McpServer) { editingId.value = server.server_id - Object.assign(form, { - name: server.name, transport: server.transport, command: server.command, - args: [...server.args], environment: { ...server.environment }, - secret_environment_keys: Object.keys(server.secret_environment), permissions: [...server.permissions], + editingOriginal.value = server + resetEditor({ + version: server.version, name: server.name, transport: server.transport, + command: server.command, args: [...server.args], url: server.url, + headers: { ...server.headers }, environment: { ...server.environment }, + secret_environment_keys: Object.keys(server.secret_environment), + secret_header_keys: Object.keys(server.secret_headers), permissions: [...server.permissions], startup_timeout_seconds: server.startup_timeout_seconds, tool_timeout_seconds: server.tool_timeout_seconds, }) - argsText.value = server.args.join('\n') - environmentText.value = JSON.stringify(server.environment, null, 2) - secretKeysText.value = Object.keys(server.secret_environment).join('\n') - permissionsText.value = server.permissions.join(', ') dialogOpen.value = true } function applyTemplate(transport: McpServerTransport) { - if (transport !== 'stdio') return - form.transport = 'stdio'; form.command = 'uvx'; argsText.value = 'mcp-server-fetch' + form.transport = transport + if (transport === 'stdio') { + form.command = 'uvx'; form.url = null + argsText.value = '--isolated\n--from\npackage-name==1.0.0\nserver-command' + } else { + form.command = null; argsText.value = ''; form.url = transport === 'sse' ? 'http://127.0.0.1:3000/sse' : 'http://127.0.0.1:3000/mcp' + } +} + +function parseObject(value: string, label: string): Record { + let parsed: unknown + try { parsed = JSON.parse(value || '{}') } catch { throw new Error(`${label}必须是 JSON 对象`) } + if (!parsed || Array.isArray(parsed) || typeof parsed !== 'object' || Object.values(parsed).some(item => typeof item !== 'string')) throw new Error(`${label}必须是字符串键值 JSON 对象`) + return parsed as Record +} + +function formPayload(): McpServerInput { + const stdio = form.transport === 'stdio' + return { + version: form.version, + name: form.name.trim(), transport: form.transport, + command: stdio ? form.command?.trim() : null, + args: stdio ? argsText.value.split('\n').map(value => value.trim()).filter(Boolean) : [], + url: stdio ? null : form.url?.trim(), + headers: stdio ? {} : parseObject(headersText.value, '普通 Header'), + environment: stdio ? parseObject(environmentText.value, '普通环境变量') : {}, + secret_environment_keys: stdio ? splitKeys(secretKeysText.value) : [], + secret_header_keys: stdio ? [] : splitKeys(secretHeaderKeysText.value), + permissions: permissionsText.value.split(',').map(value => value.trim()).filter(Boolean), + startup_timeout_seconds: form.startup_timeout_seconds, + tool_timeout_seconds: form.tool_timeout_seconds, + } } function payload(): McpServerInput { - let environment: Record - try { environment = JSON.parse(environmentText.value || '{}') } - catch { throw new Error('普通环境变量必须是 JSON 对象') } - if (!environment || Array.isArray(environment) || typeof environment !== 'object') throw new Error('普通环境变量必须是 JSON 对象') - return { - ...form, - name: form.name.trim(), command: form.command.trim(), - args: argsText.value.split('\n').map(value => value.trim()).filter(Boolean), - environment, - secret_environment_keys: secretKeysText.value.split(/[\n,]/).map(value => value.trim()).filter(Boolean), - permissions: permissionsText.value.split(',').map(value => value.trim()).filter(Boolean), - } + if (editorMode.value === 'form') return formPayload() + let parsed: unknown + try { parsed = JSON.parse(rawConfig.value) } catch { throw new Error('服务器配置不是有效 JSON') } + if (!parsed || Array.isArray(parsed) || typeof parsed !== 'object') throw new Error('服务器配置必须是 JSON 对象') + const value = parsed as McpServerInput + if (editingId.value) value.version = form.version + return value +} + +function switchMode(mode: 'form' | 'json') { + try { + if (mode === editorMode.value) return + if (mode === 'json') rawConfig.value = JSON.stringify(formPayload(), null, 2) + else resetEditor(payload()) + editorMode.value = mode + } catch (cause) { error.value = message(cause, '配置转换失败') } } async function save() { try { const input = payload() - if (!input.name || !input.command) throw new Error('请填写服务器名称和可执行命令') + if (!input.name || (input.transport === 'stdio' ? !input.command : !input.url)) throw new Error('请填写服务器名称和连接地址') + if (editingOriginal.value && executionChanged(editingOriginal.value, input) && !confirm('连接命令、地址或认证配置已变化,保存后旧测试与授权会失效。是否保存?')) return busy.value = 'save' editingId.value ? await service.updateMcpServer(editingId.value, input) : await service.createMcpServer(input) dialogOpen.value = false @@ -83,15 +141,19 @@ async function save() { finally { busy.value = '' } } +function executionChanged(server: McpServer, input: McpServerInput) { + return JSON.stringify([server.transport, server.command, server.args, server.url, server.headers, Object.keys(server.secret_headers)]) !== JSON.stringify([input.transport, input.command, input.args, input.url, input.headers, input.secret_header_keys]) +} + async function approve(server: McpServer): Promise { if (server.trusted) return server - const accepted = confirm(`即将允许本机启动以下命令:\n\n${server.command_summary}\n\n当前 Python Host 没有系统级沙箱,仅应运行可信服务器。是否继续?`) - if (!accepted) return null + const localWarning = server.transport === 'stdio' ? '\n\n本机进程尚无系统级沙箱,仅应运行可信服务器。' : '\n\n连接可能向该地址发送配置的 Header。' + if (!confirm(`请确认 MCP 连接:\n\n${server.command_summary}${localWarning}\n\n是否继续?`)) return null return service.trustMcpServer(server) } -async function test(server: McpServer) { await act(server, 'test', async current => service.testMcpServer(current.server_id)) } -async function toggle(server: McpServer) { await act(server, 'toggle', async current => current.enabled ? service.disableMcpServer(current.server_id) : service.enableMcpServer(current.server_id)) } +async function test(server: McpServer) { await act(server, 'test', current => service.testMcpServer(current.server_id)) } +async function toggle(server: McpServer) { await act(server, 'toggle', current => current.enabled ? service.disableMcpServer(current.server_id) : service.enableMcpServer(current.server_id)) } async function act(server: McpServer, action: string, operation: (server: McpServer) => Promise) { busy.value = `${action}:${server.server_id}`; error.value = '' try { const current = action === 'toggle' && server.enabled ? server : await approve(server); if (!current) return; await operation(current); await load() } @@ -105,46 +167,51 @@ async function remove(server: McpServer) { catch (cause) { error.value = message(cause, '删除失败') } finally { busy.value = '' } } -async function saveSecret(server: McpServer, key: string) { - const value = secretDrafts[`${server.server_id}:${key}`]?.trim() +async function saveSecret(server: McpServer, key: string, kind: SecretKind) { + const draftKey = `${server.server_id}:${kind}:${key}` + const value = secretDrafts[draftKey]?.trim() if (!value) return - try { busy.value = `secret:${server.server_id}:${key}`; await service.putMcpServerSecret(server.server_id, key, value); secretDrafts[`${server.server_id}:${key}`] = ''; await load() } + try { busy.value = `secret:${draftKey}`; await service.putMcpServerSecret(server.server_id, key, value, kind); secretDrafts[draftKey] = ''; await load() } catch (cause) { error.value = message(cause, '保存密钥失败') } finally { busy.value = '' } } +function splitKeys(value: string) { return value.split(/[\n,]/).map(item => item.trim()).filter(Boolean) } function message(cause: unknown, fallback: string) { return cause instanceof Error ? cause.message : fallback } onMounted(load) diff --git a/frontend/src/services/mcpServerService.ts b/frontend/src/services/mcpServerService.ts index a146a6b..8acc2b1 100644 --- a/frontend/src/services/mcpServerService.ts +++ b/frontend/src/services/mcpServerService.ts @@ -1,5 +1,5 @@ import apiClient from './apiClient' -import type { McpServer, McpServerInput, OperationResponse } from '@/contracts' +import type { McpServer, McpServerInput, McpToolSummary, OperationResponse } from '@/contracts' const base = '/api/mcp/servers' @@ -8,10 +8,11 @@ export async function listMcpServers(): Promise { } export const createMcpServer = (input: McpServerInput) => apiClient.post(base, input) export const updateMcpServer = (id: string, input: McpServerInput) => apiClient.put(`${base}/${id}`, input) +export const listMcpServerTools = async (id: string) => (await apiClient.get<{ items: McpToolSummary[] }>(`${base}/${id}/tools`)).items export const deleteMcpServer = (id: string) => apiClient.delete(`${base}/${id}`) export const trustMcpServer = (server: McpServer) => apiClient.post(`${base}/${server.server_id}/trust`, { command_digest: server.command_digest }) export const testMcpServer = (id: string) => apiClient.post(`${base}/${id}/test`) export const enableMcpServer = (id: string) => apiClient.post(`${base}/${id}/enable`) export const disableMcpServer = (id: string) => apiClient.post(`${base}/${id}/disable`) -export const putMcpServerSecret = (id: string, key: string, secret: string) => apiClient.put(`${base}/${id}/secrets/${encodeURIComponent(key)}`, { secret }) -export const deleteMcpServerSecret = (id: string, key: string) => apiClient.delete(`${base}/${id}/secrets/${encodeURIComponent(key)}`) +export const putMcpServerSecret = (id: string, key: string, secret: string, kind: 'environment' | 'header' = 'environment') => apiClient.put(`${base}/${id}/secrets/${encodeURIComponent(key)}?kind=${kind}`, { secret }) +export const deleteMcpServerSecret = (id: string, key: string, kind: 'environment' | 'header' = 'environment') => apiClient.delete(`${base}/${id}/secrets/${encodeURIComponent(key)}?kind=${kind}`) From 3001d6089a7b6624e58368173f6b1bc60a307a9f Mon Sep 17 00:00:00 2001 From: KiriAky 107 Date: Thu, 3 Sep 2026 16:16:58 +0800 Subject: [PATCH 3/4] =?UTF-8?q?=E6=B7=BB=E5=8A=A0MCP=E5=AE=A2=E6=88=B7?= =?UTF-8?q?=E7=AB=AF=E8=B6=85=E6=97=B6=E9=85=8D=E7=BD=AE=E5=92=8C=E8=BF=9E?= =?UTF-8?q?=E6=8E=A5=E7=AE=A1=E7=90=86=E6=94=B9=E8=BF=9B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 添加了MCP客户端的超时配置功能,包括启动超时和工具调用超时参数。 改进了HTTP客户端和标准IO客户端的超时处理机制,确保请求在指定时间内完成或取消。 增加了对MCP服务器数量的限制,防止配置过多服务器导致系统不稳定。 增强了错误处理机制,当连接异常时能够正确清理资源并移除桥接主机。 添加了对大型MCP消息的大小验证,防止过大的请求导致系统问题。 优化了密钥更改后的处理流程,确保在修改密钥时停用服务器并要求重新测试。 --- backend/app/extensions/mcp.py | 63 +++++-- backend/app/extensions/mcp_registry.py | 137 +++++++++++--- backend/app/routes.py | 4 +- backend/tests/test_api.py | 89 +++++++-- backend/tests/test_mcp_registry.py | 171 +++++++++++++++++- .../src/features/mcp/McpServersView.spec.ts | 11 ++ frontend/src/features/mcp/McpServersView.vue | 15 +- .../features/workspace/FileTreePanel.spec.ts | 29 +++ .../src/features/workspace/FileTreePanel.vue | 46 ++++- 9 files changed, 493 insertions(+), 72 deletions(-) diff --git a/backend/app/extensions/mcp.py b/backend/app/extensions/mcp.py index 61c2d5f..c3852fc 100644 --- a/backend/app/extensions/mcp.py +++ b/backend/app/extensions/mcp.py @@ -146,7 +146,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, @@ -156,7 +156,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: @@ -388,6 +388,7 @@ class McpHttpClient: 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], @@ -407,6 +408,7 @@ class McpHttpClient: 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 @@ -429,7 +431,7 @@ class McpHttpClient: 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, @@ -439,7 +441,7 @@ class McpHttpClient: ) def begin_request( - self, method: str, params: dict[str, Any] + self, method: str, params: dict[str, Any], *, timeout: float | None = None ) -> tuple[int, _PendingRequest]: with self._pending_lock: request_id = self._next_id @@ -454,7 +456,7 @@ class McpHttpClient: } threading.Thread( target=self._dispatch_request, - args=(request_id, message), + args=(request_id, message, timeout), daemon=True, ).start() return request_id, pending @@ -526,7 +528,10 @@ class McpHttpClient: if self._session_id: try: request = self._client.build_request( - "DELETE", self.url, headers=self._request_headers() + "DELETE", + self.url, + headers=self._request_headers(), + timeout=min(self._startup_timeout_seconds, 5), ) response = self._client.send(request, stream=True) response.close() @@ -539,9 +544,14 @@ class McpHttpClient: ) ) - def _dispatch_request(self, request_id: int, message: dict[str, Any]) -> None: + def _dispatch_request( + self, + request_id: int, + message: dict[str, Any], + timeout: float | None, + ) -> None: try: - response = self._post(message, timeout=None) + response = self._post(message, timeout=timeout) try: self._capture_session(response) content_type = response.headers.get("content-type", "").lower() @@ -588,7 +598,12 @@ class McpHttpClient: def _post_notification(self, message: dict[str, Any]) -> None: try: - response = self._post(message, timeout=10) + 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", @@ -616,6 +631,7 @@ class McpHttpClient: self.url, content=encoded.encode("utf-8"), headers=self._request_headers(), + timeout=timeout, ) return self._client.send(request, stream=True) @@ -714,7 +730,7 @@ class McpLegacySseClient(McpHttpClient): def start(self) -> None: threading.Thread(target=self._event_loop, daemon=True).start() try: - endpoint = self._endpoint_ready.get(timeout=15) + endpoint = self._endpoint_ready.get(timeout=self._startup_timeout_seconds) except queue.Empty as exc: raise McpBridgeError( "MCP_INITIALIZE_FAILED", @@ -730,9 +746,14 @@ class McpLegacySseClient(McpHttpClient): return - def _dispatch_request(self, request_id: int, message: dict[str, Any]) -> None: + def _dispatch_request( + self, + request_id: int, + message: dict[str, Any], + timeout: float | None, + ) -> None: try: - response = self._post(message, timeout=10) + response = self._post(message, timeout=timeout) try: if response.status_code not in {200, 202, 204}: raise McpBridgeError( @@ -761,6 +782,8 @@ class McpLegacySseClient(McpHttpClient): "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, @@ -770,6 +793,7 @@ class McpLegacySseClient(McpHttpClient): "Accept": "application/json, text/event-stream", "Content-Type": "application/json", }, + timeout=timeout, ) return self._client.send(request, stream=True) @@ -801,6 +825,8 @@ class McpLegacySseClient(McpHttpClient): 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) @@ -827,7 +853,7 @@ class _McpClient(Protocol): response_error_code: str = "MCP_TOOL_CALL_FAILED", ) -> dict[str, Any]: ... def begin_request( - self, method: str, params: dict[str, Any] + self, method: str, params: dict[str, Any], *, timeout: float | None = None ) -> tuple[int, _PendingRequest]: ... def wait_response( self, @@ -931,6 +957,7 @@ class McpBridge: 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, @@ -989,6 +1016,12 @@ class McpBridge: discovered = self._discover_tools( 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(UTC) @@ -1019,7 +1052,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: diff --git a/backend/app/extensions/mcp_registry.py b/backend/app/extensions/mcp_registry.py index 6808769..7307ede 100644 --- a/backend/app/extensions/mcp_registry.py +++ b/backend/app/extensions/mcp_registry.py @@ -13,12 +13,13 @@ from typing import Any from urllib.parse import urlsplit from uuid import uuid4 -from pydantic import BaseModel, ConfigDict, create_model +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, @@ -40,6 +41,23 @@ _RESERVED_HEADERS = { "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 _McpServerRecord(McpServerConfig): + """Validated on-disk representation with defaults for older C.1 records.""" + + version: int = Field(default=1, ge=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): @@ -105,6 +123,13 @@ class McpServerRegistry: @_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() @@ -137,8 +162,8 @@ class McpServerRegistry: self.disable(server_id) with self._lock: previous = self._record(server_id) - removed = [ - (kind, key) + removed_secret_ids = [ + self._secret_id(server_id, key, kind) for kind, old_keys, new_keys in ( ( "environment", @@ -153,6 +178,13 @@ class McpServerRegistry: ) for key in set(old_keys) - set(new_keys) ] + 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 @@ -170,13 +202,6 @@ class McpServerRegistry: self._records = updated self._last_status.pop(server_id, None) self._summaries.pop(server_id, None) - for kind, key in removed: - 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 self.get(server_id) @_serialized_lifecycle @@ -192,18 +217,19 @@ class McpServerRegistry: ) for key in keys ] - updated = dict(self._records) - del updated[server_id] - self._write(updated) - self._records = updated - self._last_status.pop(server_id, None) - self._summaries.pop(server_id, None) try: self.credentials.delete_many(secret_ids) except CredentialStoreError as exc: 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 @@ -236,6 +262,9 @@ class McpServerRegistry: "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: @@ -254,6 +283,9 @@ class McpServerRegistry: "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: @@ -465,17 +497,27 @@ class McpServerRegistry: self.tools.register(definition, arguments_model, executor) def _unavailable(self, server_id: str, message: str) -> None: - with self._lock: - for name in self._registered.pop(server_id, []): - self.tools.unregister(name) - record = self._records.get(server_id) - if record is not None: - self._records[server_id] = {**record, "enabled": False} - self._last_status[server_id] = { - "status": PluginHostState.unhealthy, - "error": message, - } - self._write() + # 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: + 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 @@ -767,6 +809,23 @@ class McpServerRegistry: 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" @@ -788,7 +847,29 @@ class McpServerRegistry: "MCP server registry has an invalid format.", status_code=500, ) - return value + 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") diff --git a/backend/app/routes.py b/backend/app/routes.py index a1bd49a..27c3586 100644 --- a/backend/app/routes.py +++ b/backend/app/routes.py @@ -607,7 +607,7 @@ async def put_mcp_server_secret( request: McpServerSecretWriteRequest, kind: str = Query(default="environment", pattern="^(environment|header)$"), ) -> McpServerSecretStatus: - return mcp_call( + return await mcp_call_async( lambda: container.mcp_servers.put_secret( server_id, key, request.secret.get_secret_value(), kind=kind ) @@ -624,7 +624,7 @@ async def delete_mcp_server_secret( key: str, kind: str = Query(default="environment", pattern="^(environment|header)$"), ) -> McpServerSecretStatus: - return mcp_call( + return await mcp_call_async( lambda: container.mcp_servers.delete_secret(server_id, key, kind=kind) ) diff --git a/backend/tests/test_api.py b/backend/tests/test_api.py index 3d5d2d5..73901a5 100644 --- a/backend/tests/test_api.py +++ b/backend/tests/test_api.py @@ -1,26 +1,10 @@ import asyncio +import threading +from types import SimpleNamespace -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 +12,63 @@ 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: @@ -101,6 +142,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 index ae0e9fd..1adfa14 100644 --- a/backend/tests/test_mcp_registry.py +++ b/backend/tests/test_mcp_registry.py @@ -1,6 +1,7 @@ import asyncio import json import sys +import threading import time from concurrent.futures import ThreadPoolExecutor @@ -10,8 +11,9 @@ 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 EncryptedCredentialStore +from app.providers.credentials import CredentialStoreError, EncryptedCredentialStore SERVER = BACKEND_DIR / "extensions" / "fixtures" / "mcp-echo" / "server.py" @@ -59,6 +61,28 @@ def test_registry_requires_current_trust_and_never_returns_secret() -> None: 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=[])) @@ -87,6 +111,70 @@ def test_update_disables_server_and_revokes_command_trust() -> None: 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) + + service._unavailable(created.server_id, "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_production_rejects_process_launch_even_after_approval() -> None: service = registry(launch=False) created = service.create(request(secret_environment_keys=[])) @@ -141,6 +229,36 @@ def test_registry_rejects_corrupt_persisted_json(tmp_path) -> None: 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( @@ -217,6 +335,7 @@ 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) @@ -225,6 +344,9 @@ def test_streamable_http_supports_session_headers_secrets_and_tool_summary( 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"], @@ -290,6 +412,9 @@ def test_streamable_http_supports_session_headers_secrets_and_tool_summary( 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( @@ -301,11 +426,15 @@ def test_streamable_http_supports_session_headers_secrets_and_tool_summary( 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) @@ -326,17 +455,22 @@ class _LegacyEventStream(httpx.SyncByteStream): "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=_LegacyEventStream(), + stream=event_stream, ) posted_urls.append(str(request_value.url)) return httpx.Response(202) @@ -361,6 +495,39 @@ def test_legacy_sse_uses_same_origin_endpoint(monkeypatch) -> None: 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): diff --git a/frontend/src/features/mcp/McpServersView.spec.ts b/frontend/src/features/mcp/McpServersView.spec.ts index 921a717..659295f 100644 --- a/frontend/src/features/mcp/McpServersView.spec.ts +++ b/frontend/src/features/mcp/McpServersView.spec.ts @@ -81,4 +81,15 @@ describe('McpServersView', () => { expect(confirm).toHaveBeenCalled() expect(service.deleteMcpServer).toHaveBeenCalledWith('server-1') }) + + it('confirms permission changes before updating an existing server', async () => { + const wrapper = await render([server]) + vi.mocked(service.updateMcpServer).mockResolvedValue(server) + await wrapper.findAll('button').find(button => button.text().includes('编辑'))!.trigger('click') + await wrapper.get('input[placeholder="network.request, notes.read"]').setValue('notes.read') + await wrapper.get('form').trigger('submit') + await flushPromises() + expect(confirm).toHaveBeenCalledWith(expect.stringContaining('旧测试与授权会失效')) + expect(service.updateMcpServer).toHaveBeenCalled() + }) }) diff --git a/frontend/src/features/mcp/McpServersView.vue b/frontend/src/features/mcp/McpServersView.vue index 04f1105..1f9380f 100644 --- a/frontend/src/features/mcp/McpServersView.vue +++ b/frontend/src/features/mcp/McpServersView.vue @@ -142,7 +142,20 @@ async function save() { } function executionChanged(server: McpServer, input: McpServerInput) { - return JSON.stringify([server.transport, server.command, server.args, server.url, server.headers, Object.keys(server.secret_headers)]) !== JSON.stringify([input.transport, input.command, input.args, input.url, input.headers, input.secret_header_keys]) + const sortedEntries = (value: Record) => Object.entries(value).sort(([left], [right]) => left.localeCompare(right)) + const current = [ + server.transport, server.command, server.args, server.url, + sortedEntries(server.headers), sortedEntries(server.environment), + Object.keys(server.secret_headers).sort(), Object.keys(server.secret_environment).sort(), + [...server.permissions].sort(), server.startup_timeout_seconds, server.tool_timeout_seconds, + ] + const next = [ + input.transport, input.command, input.args, input.url, + sortedEntries(input.headers), sortedEntries(input.environment), + [...input.secret_header_keys].sort(), [...input.secret_environment_keys].sort(), + [...input.permissions].sort(), input.startup_timeout_seconds, input.tool_timeout_seconds, + ] + return JSON.stringify(current) !== JSON.stringify(next) } async function approve(server: McpServer): Promise { 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 @@