From 4894029a0fbf1081c03cdf289f4eb57fbce9fe73 Mon Sep 17 00:00:00 2001 From: KiriAky 107 Date: Tue, 1 Sep 2026 11:32:05 +0800 Subject: [PATCH 1/5] =?UTF-8?q?feat(extension):=20=E6=8E=A5=E5=85=A5=20std?= =?UTF-8?q?io=20MCP=20Plugin=20Host?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- backend/app/agent/runtime.py | 8 +- backend/app/agent/tools.py | 62 +- backend/app/contracts.py | 26 + backend/app/extensions/__init__.py | 10 +- backend/app/extensions/mcp.py | 746 ++++++++++++++++++ backend/app/extensions/runtime.py | 264 +++++-- backend/app/main.py | 12 + backend/app/routes.py | 47 +- .../extensions/fixtures/mcp-echo/plugin.yaml | 21 + .../extensions/fixtures/mcp-echo/server.py | 183 +++++ backend/tests/test_api.py | 2 + backend/tests/test_extension_core.py | 212 ++++- frontend/src/contracts/index.ts | 25 +- frontend/src/services/pluginService.ts | 10 +- 14 files changed, 1557 insertions(+), 71 deletions(-) create mode 100644 backend/app/extensions/mcp.py create mode 100644 backend/extensions/fixtures/mcp-echo/plugin.yaml create mode 100644 backend/extensions/fixtures/mcp-echo/server.py diff --git a/backend/app/agent/runtime.py b/backend/app/agent/runtime.py index 3daf87d..e573317 100644 --- a/backend/app/agent/runtime.py +++ b/backend/app/agent/runtime.py @@ -491,7 +491,13 @@ class AgentRuntime: async def _invoke_tool(self, record: RunRecord, call: ToolCall) -> ToolResult: try: return await asyncio.wait_for( - self.tools.execute(call, ToolExecutionContext(run_id=record.run.run_id)), + self.tools.execute( + call, + ToolExecutionContext( + run_id=record.run.run_id, + tool_call_id=call.tool_call_id, + ), + ), timeout=record.request.tool_timeout_seconds, ) except TimeoutError: diff --git a/backend/app/agent/tools.py b/backend/app/agent/tools.py index a3724d6..df90158 100644 --- a/backend/app/agent/tools.py +++ b/backend/app/agent/tools.py @@ -1,6 +1,7 @@ """Agent 工具注册与执行边界。""" import inspect +import threading from dataclasses import dataclass from time import perf_counter from typing import Any, Awaitable, Callable @@ -17,6 +18,7 @@ ToolExecutor = Callable[[BaseModel, "ToolExecutionContext"], Any | Awaitable[Any @dataclass(frozen=True, slots=True) class ToolExecutionContext: run_id: str + tool_call_id: str | None = None @dataclass(slots=True) @@ -30,11 +32,21 @@ class ToolNotFoundError(LookupError): pass +class ToolExecutionError(RuntimeError): + """Executor 可预期失败,保留领域错误码而不是折叠成通用异常。""" + + def __init__(self, code: str, message: str) -> None: + super().__init__(message) + self.code = code + self.message = message + + class ToolRegistry: """统一校验工具入参并隔离执行异常,避免单个工具击穿 Agent 主循环。""" def __init__(self) -> None: self._tools: dict[str, RegisteredTool] = {} + self._lock = threading.RLock() def register( self, @@ -42,33 +54,38 @@ class ToolRegistry: arguments_model: type[BaseModel], executor: ToolExecutor, ) -> None: - if definition.name in self._tools: - raise ValueError(f"Tool already registered: {definition.name}") - self._tools[definition.name] = RegisteredTool( - definition=definition, - arguments_model=arguments_model, - executor=executor, - ) + with self._lock: + if definition.name in self._tools: + raise ValueError(f"Tool already registered: {definition.name}") + self._tools[definition.name] = RegisteredTool( + definition=definition, + arguments_model=arguments_model, + executor=executor, + ) def unregister(self, name: str) -> None: - self._tools.pop(name, None) + with self._lock: + self._tools.pop(name, None) def contains(self, name: str) -> bool: - return name in self._tools + with self._lock: + return name in self._tools def get(self, name: str) -> RegisteredTool: - try: - return self._tools[name] - except KeyError as exc: - raise ToolNotFoundError(name) from exc + with self._lock: + try: + return self._tools[name] + except KeyError as exc: + raise ToolNotFoundError(name) from exc def definitions(self, allowed: list[str] | None = None) -> list[ToolDefinition]: names = set(allowed) if allowed is not None else None - return [ - item.definition.model_copy(deep=True) - for name, item in self._tools.items() - if names is None or name in names - ] + with self._lock: + return [ + item.definition.model_copy(deep=True) + for name, item in self._tools.items() + if names is None or name in names + ] async def execute(self, call: ToolCall, context: ToolExecutionContext) -> ToolResult: started = perf_counter() @@ -108,6 +125,15 @@ class ToolRegistry: output=output, duration_ms=round((perf_counter() - started) * 1000), ) + except ToolExecutionError as exc: + return ToolResult( + tool_call_id=call.tool_call_id, + name=call.name, + success=False, + error_code=exc.code, + error_message=exc.message, + duration_ms=round((perf_counter() - started) * 1000), + ) except Exception as exc: # 工具失败转换成结构化结果,由模型决定是否降级或重试。 return ToolResult( tool_call_id=call.tool_call_id, diff --git a/backend/app/contracts.py b/backend/app/contracts.py index bb8bf90..8d1aace 100644 --- a/backend/app/contracts.py +++ b/backend/app/contracts.py @@ -417,6 +417,10 @@ class ExtensionInstallRequest(Contract): class PluginBackend(Contract): type: Literal["mcp", "internal_rpc", "none"] = "none" transport: Literal["stdio", "http", "none"] = "none" + command: str | None = None + args: list[str] = Field(default_factory=list) + startup_timeout_seconds: int = Field(default=10, ge=1, le=60) + tool_timeout_seconds: int = Field(default=30, ge=1, le=600) class PluginContribution(Contract): @@ -460,6 +464,28 @@ class PluginListResponse(Contract): items: list[Plugin] = Field(default_factory=list) +class PluginHostState(str, Enum): + stopped = "stopped" + starting = "starting" + ready = "ready" + unhealthy = "unhealthy" + error = "error" + + +class PluginHostStatus(Contract): + plugin_id: str + backend_type: Literal["mcp", "internal_rpc", "none"] + transport: Literal["stdio", "http", "none"] + status: PluginHostState + tools_count: int = 0 + started_at: datetime | None = None + last_seen_at: datetime | None = None + protocol_version: str | None = None + server_name: str | None = None + server_version: str | None = None + error: str | None = None + + class PluginPermissionGrantRequest(Contract): permissions: list[str] = Field(default_factory=list) diff --git a/backend/app/extensions/__init__.py b/backend/app/extensions/__init__.py index 3495110..ed2e570 100644 --- a/backend/app/extensions/__init__.py +++ b/backend/app/extensions/__init__.py @@ -4,5 +4,13 @@ from app.extensions.runtime import ( PluginRuntime, SkillRuntime, ) +from app.extensions.mcp import McpBridge, McpBridgeError -__all__ = ["AgentConfiguration", "ExtensionError", "PluginRuntime", "SkillRuntime"] +__all__ = [ + "AgentConfiguration", + "ExtensionError", + "McpBridge", + "McpBridgeError", + "PluginRuntime", + "SkillRuntime", +] diff --git a/backend/app/extensions/mcp.py b/backend/app/extensions/mcp.py new file mode 100644 index 0000000..2647c20 --- /dev/null +++ b/backend/app/extensions/mcp.py @@ -0,0 +1,746 @@ +"""本地 stdio MCP Bridge。 + +第三方 Server 始终运行在子进程中。Bridge 只把通过校验的 MCP Tool 转换为项目内部 +ToolDefinition/ToolResult,不把 MCP 原始协议泄露给 Agent Runtime 或前端。 +""" + +from __future__ import annotations + +import asyncio +import json +import os +import queue +import subprocess +import threading +from collections import deque +from dataclasses import dataclass +from datetime import datetime, timezone +from pathlib import Path +from typing import Any, Callable + +from jsonschema import Draft202012Validator +from jsonschema.exceptions import SchemaError + +from app.agent.permissions import KNOWN_PERMISSIONS +from app.agent.tools import ToolExecutionError +from app.contracts import ( + PluginBackend, + PluginHostState, + PluginHostStatus, + ToolDefinition, +) + +MCP_PROTOCOL_VERSION = "2025-11-25" +SUPPORTED_PROTOCOL_VERSIONS = { + MCP_PROTOCOL_VERSION, + "2025-06-18", + "2025-03-26", + "2024-11-05", +} +MAX_MCP_MESSAGE_BYTES = 2 * 1024 * 1024 +MAX_MCP_TOOL_RESULT_BYTES = 256 * 1024 +MAX_MCP_TOOLS = 500 +MAX_MCP_LIST_PAGES = 100 + + +class McpBridgeError(RuntimeError): + def __init__(self, code: str, message: str, *, status_code: int = 502) -> None: + super().__init__(message) + self.code = code + self.message = message + self.status_code = status_code + + +@dataclass(frozen=True, slots=True) +class McpDiscoveredTool: + remote_name: str + definition: ToolDefinition + + +@dataclass(slots=True) +class _PendingRequest: + response: queue.Queue[dict[str, Any] | BaseException] + + +class McpStdioClient: + """线程驱动的换行分隔 JSON-RPC 客户端,避免阻塞 FastAPI 事件循环。""" + + def __init__( + self, + command: list[str], + *, + cwd: Path, + on_seen: Callable[[], None], + on_broken: Callable[[str], None], + on_tools_changed: Callable[[], None], + ) -> None: + self.command = command + self.cwd = cwd + self.on_seen = on_seen + self.on_broken = on_broken + self.on_tools_changed = on_tools_changed + self.process: subprocess.Popen[str] | None = None + self._write_lock = threading.Lock() + self._pending_lock = threading.Lock() + self._pending: dict[int, _PendingRequest] = {} + self._next_id = 1 + self._stopping = False + # stderr 只在 Host 内部保留有限尾部,不进入 API、Trace 或普通日志。 + self._stderr_tail: deque[str] = deque(maxlen=50) + + def start(self) -> None: + if self.process is not None and self.process.poll() is None: + return + creation_flags = getattr(subprocess, "CREATE_NO_WINDOW", 0) if os.name == "nt" else 0 + environment = _subprocess_environment() + environment.setdefault("PYTHONUNBUFFERED", "1") + try: + self.process = subprocess.Popen( + self.command, + cwd=self.cwd, + stdin=subprocess.PIPE, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + text=True, + encoding="utf-8", + errors="replace", + bufsize=1, + shell=False, + env=environment, + creationflags=creation_flags, + ) + except OSError as exc: + raise McpBridgeError( + "PLUGIN_HOST_START_FAILED", + f"Cannot start MCP server process: {exc}", + status_code=503, + ) from exc + threading.Thread(target=self._stdout_loop, daemon=True).start() + threading.Thread(target=self._stderr_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]: + self._ensure_running() + 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 + try: + self._send( + { + "jsonrpc": "2.0", + "id": request_id, + "method": method, + "params": params, + } + ) + except BaseException: + with self._pending_lock: + self._pending.pop(request_id, None) + raise + 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( + "MCP_TOOL_CALL_FAILED", "MCP response result must be an object." + ) + return result + + def notify(self, method: str, params: dict[str, Any] | None = None) -> None: + payload: dict[str, Any] = {"jsonrpc": "2.0", "method": method} + if params is not None: + payload["params"] = params + self._send(payload) + + def cancel(self, request_id: int, reason: str = "Cancelled by host.") -> None: + try: + self.notify( + "notifications/cancelled", + {"requestId": request_id, "reason": reason}, + ) + except McpBridgeError: + pass + + def abandon(self, request_id: int) -> None: + with self._pending_lock: + self._pending.pop(request_id, None) + + def stop(self) -> None: + process = self.process + if process is None: + return + self._stopping = True + try: + if process.stdin: + try: + process.stdin.close() + except (BrokenPipeError, OSError, ValueError): + pass + try: + process.wait(timeout=2) + except subprocess.TimeoutExpired: + process.terminate() + try: + process.wait(timeout=2) + except subprocess.TimeoutExpired: + process.kill() + process.wait(timeout=2) + finally: + self._fail_pending( + McpBridgeError("PLUGIN_HOST_UNAVAILABLE", "MCP host stopped.", status_code=503) + ) + self.process = None + + def _send(self, message: dict[str, Any]) -> None: + self._ensure_running() + 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.") + process = self.process + assert process is not None and process.stdin is not None + try: + with self._write_lock: + process.stdin.write(encoded + "\n") + process.stdin.flush() + except (BrokenPipeError, OSError, ValueError) as exc: + raise McpBridgeError( + "PLUGIN_HOST_UNAVAILABLE", "MCP host input is closed.", status_code=503 + ) from exc + + def _stdout_loop(self) -> None: + process = self.process + assert process is not None and process.stdout is not None + failure: str | None = None + try: + for raw_line in process.stdout: + if len(raw_line.encode("utf-8")) > MAX_MCP_MESSAGE_BYTES: + failure = "MCP server emitted an oversized protocol message." + break + try: + message = json.loads(raw_line) + except json.JSONDecodeError: + failure = "MCP server emitted invalid JSON on stdout." + break + if not isinstance(message, dict) or message.get("jsonrpc") != "2.0": + failure = "MCP server emitted an invalid JSON-RPC message." + break + self.on_seen() + if "id" in message and ("result" in message or "error" in message): + request_id = message.get("id") + if isinstance(request_id, int): + with self._pending_lock: + pending = self._pending.pop(request_id, None) + if pending: + pending.response.put(message) + continue + method = message.get("method") + if method == "notifications/tools/list_changed": + self.on_tools_changed() + elif isinstance(method, str) and "id" in message: + self._send( + { + "jsonrpc": "2.0", + "id": message["id"], + "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() + exit_code = process.poll() + if exit_code is None: + try: + exit_code = process.wait(timeout=1) + except subprocess.TimeoutExpired: + exit_code = None + if not self._stopping: + message = failure or f"MCP host exited unexpectedly with code {exit_code}." + error = McpBridgeError( + "PLUGIN_HOST_UNAVAILABLE", message, status_code=503 + ) + self._fail_pending(error) + self.on_broken(message) + + def _stderr_loop(self) -> None: + process = self.process + assert process is not None and process.stderr is not None + try: + for line in process.stderr: + self._stderr_tail.append(line.rstrip()[:1024]) + except (OSError, ValueError): + return + + def _ensure_running(self) -> None: + if self.process is None or self.process.poll() is not None: + raise McpBridgeError( + "PLUGIN_HOST_UNAVAILABLE", "MCP host is not running.", status_code=503 + ) + + 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) + + +@dataclass(slots=True) +class _McpHost: + backend: PluginBackend + client: McpStdioClient + status: PluginHostStatus + + +class McpBridge: + """管理每个 Plugin 的独立 MCP Client,并执行 Contract 转换。""" + + def __init__(self) -> None: + self._hosts: dict[str, _McpHost] = {} + self._statuses: dict[str, PluginHostStatus] = {} + self._calls: dict[tuple[str, str], int] = {} + self._lock = threading.RLock() + + def start( + self, + plugin_id: str, + backend: PluginBackend, + package_path: Path, + declared_permissions: list[str], + on_unavailable: Callable[[str, str], None], + ) -> list[McpDiscoveredTool]: + if backend.transport != "stdio": + raise McpBridgeError( + "MCP_CAPABILITY_UNSUPPORTED", + "Phase C only supports the MCP stdio transport.", + status_code=501, + ) + command = self._resolve_command(package_path, backend) + now = datetime.now(timezone.utc) + status = PluginHostStatus( + plugin_id=plugin_id, + backend_type="mcp", + transport="stdio", + status=PluginHostState.starting, + started_at=now, + last_seen_at=now, + ) + host_ref: dict[str, _McpHost] = {} + + def seen() -> None: + host = host_ref.get("host") + if host: + host.status.last_seen_at = datetime.now(timezone.utc) + + def broken(message: str) -> None: + host = host_ref.get("host") + if host: + host.status.status = PluginHostState.unhealthy + host.status.error = message + on_unavailable(plugin_id, message) + + def tools_changed() -> None: + broken("MCP tool list changed; restart the Plugin Host to revalidate tools.") + + client = McpStdioClient( + command, + cwd=package_path, + on_seen=seen, + on_broken=broken, + on_tools_changed=tools_changed, + ) + host = _McpHost(backend=backend, client=client, status=status) + host_ref["host"] = host + with self._lock: + if plugin_id in self._hosts: + raise McpBridgeError( + "PLUGIN_HOST_START_FAILED", + f"MCP host is already running: {plugin_id}", + status_code=409, + ) + self._hosts[plugin_id] = host + self._statuses[plugin_id] = status + try: + client.start() + initialize = client.request( + "initialize", + { + "protocolVersion": MCP_PROTOCOL_VERSION, + "capabilities": {}, + "clientInfo": {"name": "NotesAgent", "version": "0.1.0"}, + }, + timeout=backend.startup_timeout_seconds, + timeout_code="MCP_INITIALIZE_FAILED", + response_error_code="MCP_INITIALIZE_FAILED", + ) + version = initialize.get("protocolVersion") + if version not in SUPPORTED_PROTOCOL_VERSIONS: + raise McpBridgeError( + "MCP_INITIALIZE_FAILED", + f"Unsupported MCP protocol version: {version}", + ) + capabilities = initialize.get("capabilities") + if not isinstance(capabilities, dict) or not isinstance( + capabilities.get("tools"), dict + ): + raise McpBridgeError( + "MCP_CAPABILITY_UNSUPPORTED", + "MCP server does not declare the tools capability.", + ) + server_info = initialize.get("serverInfo") + if not isinstance(server_info, dict): + server_info = {} + status.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") + discovered = self._discover_tools( + plugin_id, client, backend, declared_permissions + ) + status.status = PluginHostState.ready + status.tools_count = len(discovered) + status.last_seen_at = datetime.now(timezone.utc) + status.error = None + return discovered + except McpBridgeError as exc: + status.status = PluginHostState.error + status.error = exc.message + client.stop() + with self._lock: + self._hosts.pop(plugin_id, None) + raise + except Exception as exc: + status.status = PluginHostState.error + status.error = f"MCP initialization failed: {type(exc).__name__}." + client.stop() + with self._lock: + self._hosts.pop(plugin_id, None) + raise McpBridgeError("MCP_INITIALIZE_FAILED", status.error) from exc + + async def call_tool( + self, + plugin_id: str, + remote_name: str, + arguments: dict[str, Any], + *, + request_id: str, + ) -> Any: + host = self._host(plugin_id) + rpc_id, pending = host.client.begin_request( + "tools/call", {"name": remote_name, "arguments": arguments} + ) + call_key = (plugin_id, request_id) + with self._lock: + self._calls[call_key] = rpc_id + try: + result = await asyncio.to_thread( + host.client.wait_response, + rpc_id, + pending, + timeout=host.backend.tool_timeout_seconds, + timeout_code="MCP_TOOL_CALL_FAILED", + ) + except asyncio.CancelledError: + host.client.cancel(rpc_id) + host.client.abandon(rpc_id) + raise + except McpBridgeError as exc: + raise ToolExecutionError(exc.code, exc.message) from exc + finally: + with self._lock: + self._calls.pop(call_key, None) + + encoded_size = len( + json.dumps(result, ensure_ascii=False, separators=(",", ":")).encode("utf-8") + ) + if encoded_size > MAX_MCP_TOOL_RESULT_BYTES: + raise ToolExecutionError( + "MCP_TOOL_RESULT_TOO_LARGE", + "MCP tool result exceeds the configured size limit.", + ) + if result.get("isError") is True: + raise ToolExecutionError( + "MCP_TOOL_CALL_FAILED", _mcp_error_message(result.get("content")) + ) + structured = result.get("structuredContent") + if structured is not None: + if not isinstance(structured, dict): + raise ToolExecutionError( + "MCP_TOOL_CALL_FAILED", + "MCP structuredContent must be an object.", + ) + return structured + content = result.get("content", []) + if not isinstance(content, list): + raise ToolExecutionError( + "MCP_TOOL_CALL_FAILED", "MCP tool content must be an array." + ) + return {"content": content} + + def cancel(self, plugin_id: str, request_id: str) -> None: + with self._lock: + rpc_id = self._calls.get((plugin_id, request_id)) + host = self._hosts.get(plugin_id) + if rpc_id is not None and host is not None: + host.client.cancel(rpc_id) + + def stop(self, plugin_id: str) -> None: + with self._lock: + host = self._hosts.pop(plugin_id, None) + if host: + host.client.stop() + host.status.status = PluginHostState.stopped + host.status.tools_count = 0 + host.status.error = None + + def status(self, plugin_id: str, backend: PluginBackend) -> PluginHostStatus: + with self._lock: + status = self._statuses.get(plugin_id) + if status: + return status.model_copy(deep=True) + return PluginHostStatus( + plugin_id=plugin_id, + backend_type=backend.type, + transport=backend.transport, + status=PluginHostState.stopped, + ) + + def _discover_tools( + self, + plugin_id: str, + client: McpStdioClient, + backend: PluginBackend, + declared_permissions: list[str], + ) -> list[McpDiscoveredTool]: + discovered: list[McpDiscoveredTool] = [] + cursor: str | None = None + for _ in range(MAX_MCP_LIST_PAGES): + params = {"cursor": cursor} if cursor else {} + result = client.request( + "tools/list", + params, + timeout=backend.startup_timeout_seconds, + timeout_code="MCP_INITIALIZE_FAILED", + response_error_code="MCP_INITIALIZE_FAILED", + ) + 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." + ) + for raw in raw_tools: + discovered.append( + self._map_tool(plugin_id, raw, declared_permissions) + ) + if len(discovered) > MAX_MCP_TOOLS: + raise McpBridgeError( + "MCP_TOOL_SCHEMA_INVALID", + f"MCP server exposes more than {MAX_MCP_TOOLS} tools.", + ) + next_cursor = result.get("nextCursor") + if next_cursor is None: + 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." + ) + cursor = next_cursor + else: + raise McpBridgeError( + "MCP_TOOL_SCHEMA_INVALID", "MCP tools/list exceeded the page limit." + ) + names = [item.definition.name for item in discovered] + if len(names) != len(set(names)): + raise McpBridgeError( + "MCP_TOOL_SCHEMA_INVALID", "MCP server returned duplicate tool names." + ) + return discovered + + @staticmethod + def _map_tool( + plugin_id: str, raw: Any, declared_permissions: list[str] + ) -> McpDiscoveredTool: + if not isinstance(raw, dict): + raise McpBridgeError( + "MCP_TOOL_SCHEMA_INVALID", "MCP tool definition must be an object." + ) + remote_name = raw.get("name") + if not isinstance(remote_name, str) or not remote_name: + raise McpBridgeError( + "MCP_TOOL_SCHEMA_INVALID", "MCP tool name must be a non-empty string." + ) + if ( + len(remote_name) > 128 + or not remote_name[0].isalnum() + or not all( + character.islower() + or character.isdigit() + or character in "._-" + for character in remote_name + ) + ): + raise McpBridgeError( + "MCP_TOOL_SCHEMA_INVALID", + f"MCP tool name is not a valid NotesAgent id: {remote_name}", + ) + schema = raw.get("inputSchema", {"type": "object", "properties": {}}) + if not isinstance(schema, dict) or schema.get("type", "object") != "object": + raise McpBridgeError( + "MCP_TOOL_SCHEMA_INVALID", + f"MCP tool inputSchema must be an object schema: {remote_name}", + ) + try: + Draft202012Validator.check_schema(schema) + except SchemaError as exc: + raise McpBridgeError( + "MCP_TOOL_SCHEMA_INVALID", + f"Invalid MCP tool schema for {remote_name}: {exc.message}", + ) from exc + metadata = raw.get("_meta") + permission = ( + 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 + ): + raise McpBridgeError( + "MCP_TOOL_SCHEMA_INVALID", + f"MCP tool declares an unknown permission: {remote_name}", + ) + if permission and permission not in declared_permissions: + raise McpBridgeError( + "MCP_TOOL_SCHEMA_INVALID", + f"MCP tool permission is missing from Plugin manifest: {permission}", + ) + description = raw.get("description") + return McpDiscoveredTool( + remote_name=remote_name, + definition=ToolDefinition( + name=f"{plugin_id}.{remote_name}", + description=description if isinstance(description, str) else remote_name, + parameters=schema, + permission=permission, + source="plugin", + ), + ) + + def _host(self, plugin_id: str) -> _McpHost: + with self._lock: + host = self._hosts.get(plugin_id) + if host is None or host.status.status != PluginHostState.ready: + raise ToolExecutionError( + "PLUGIN_HOST_UNAVAILABLE", f"MCP Plugin Host is not ready: {plugin_id}" + ) + return host + + @staticmethod + def _resolve_command(root: Path, backend: PluginBackend) -> list[str]: + if not backend.command or not backend.command.strip(): + raise McpBridgeError( + "PLUGIN_HOST_START_FAILED", "MCP stdio backend requires a command." + ) + command = backend.command.strip() + if Path(command).is_absolute() or "/" in command or "\\" in command: + executable = ( + (root / command).resolve() + if not Path(command).is_absolute() + else Path(command).resolve() + ) + try: + executable.relative_to(root) + except ValueError as exc: + raise McpBridgeError( + "PLUGIN_HOST_START_FAILED", + "MCP executable path must stay inside the Plugin package.", + ) from exc + command = str(executable) + return [command, *backend.args] + + +def _mcp_error_message(content: Any) -> str: + if isinstance(content, list): + texts = [ + item.get("text") + for item in content + if isinstance(item, dict) + and item.get("type") == "text" + and isinstance(item.get("text"), str) + ] + if texts: + return "\n".join(texts)[:4096] + return "MCP tool returned an error result." + + +def _optional_string(value: Any) -> str | None: + return value if isinstance(value, str) else None + + +def _subprocess_environment() -> dict[str, str]: + """只传递启动进程所需的系统变量,隔离 Provider Key、Vault 路径等宿主状态。""" + + allowed = { + "PATH", + "PATHEXT", + "SYSTEMROOT", + "WINDIR", + "COMSPEC", + "TEMP", + "TMP", + "TMPDIR", + "LANG", + "LC_ALL", + "VIRTUAL_ENV", + } + environment = { + key: value for key, value in os.environ.items() if key.upper() in allowed + } + environment["PYTHONUNBUFFERED"] = "1" + environment["PYTHONIOENCODING"] = "utf-8" + return environment diff --git a/backend/app/extensions/runtime.py b/backend/app/extensions/runtime.py index 2f701cb..33d0bbc 100644 --- a/backend/app/extensions/runtime.py +++ b/backend/app/extensions/runtime.py @@ -1,6 +1,7 @@ from __future__ import annotations import re +import threading from dataclasses import dataclass from pathlib import Path from typing import Any, Literal @@ -16,6 +17,7 @@ from app.contracts import ( ModelCapability, Plugin, PluginManifest, + PluginHostStatus, PluginStatus, RetrievalConfig, Skill, @@ -23,6 +25,7 @@ from app.contracts import ( SkillStatus, ToolDefinition, ) +from app.extensions.mcp import McpBridge, McpBridgeError, McpDiscoveredTool _EXTENSION_ID = re.compile(r"^[a-z0-9][a-z0-9._-]*$") @@ -242,18 +245,26 @@ class _PluginRecord: tools: list[DeclarativeToolSpec] package_path: Path registered_tools: list[str] + mcp_remote_names: dict[str, str] class PluginRuntime: """Plugin Manifest、生命周期及 Tool Contribution 注册。""" - def __init__(self, tools: ToolRegistry, host: DeclarativePluginHost | None = None) -> None: + def __init__( + self, + tools: ToolRegistry, + host: DeclarativePluginHost | None = None, + mcp_bridge: McpBridge | None = None, + ) -> None: self.registry = tools self.host = host or DeclarativePluginHost() + self.mcp = mcp_bridge or McpBridge() self._records: dict[str, _PluginRecord] = {} + self._lock = threading.RLock() def install(self, package_path: str | Path) -> Plugin: - # 当前只加载声明式清单,不导入或执行插件包中的任意 Python 代码。 + # 安装阶段只读取清单;MCP 子进程必须在权限授予后的 enable 阶段启动。 root = _package_dir(package_path) raw = _read_yaml(root / "plugin.yaml") if "id" in raw and "plugin_id" not in raw: @@ -271,15 +282,17 @@ class PluginRuntime: status_code=409, ) - specs = self._load_tools(root) - declared = set(manifest.contributes.tools) - actual = {spec.name for spec in specs} - if declared != actual: - raise ExtensionError( - "PLUGIN_CONTRIBUTION_INVALID", - "plugin.yaml tool contributions must exactly match tools.yaml", - details={"declared": sorted(declared), "actual": sorted(actual)}, - ) + _validate_backend(manifest) + specs = [] if manifest.backend.type == "mcp" else self._load_tools(root) + if manifest.backend.type != "mcp": + declared = set(manifest.contributes.tools) + actual = {spec.name for spec in specs} + if declared != actual: + raise ExtensionError( + "PLUGIN_CONTRIBUTION_INVALID", + "plugin.yaml tool contributions must exactly match tools.yaml", + details={"declared": sorted(declared), "actual": sorted(actual)}, + ) for spec in specs: _validate_id("tool", spec.name) _validate_tool_schema(spec) @@ -302,6 +315,7 @@ class PluginRuntime: tools=specs, package_path=root, registered_tools=[], + mcp_remote_names={}, ) self._records[manifest.plugin_id] = record return record.plugin.model_copy(deep=True) @@ -313,18 +327,14 @@ class PluginRuntime: return self._record(plugin_id).plugin.model_copy(deep=True) def enable(self, plugin_id: str) -> Plugin: + # Host 启动和 Tool 批量注册必须串行,避免并发 enable 产生重复进程或半注册状态。 + with self._lock: + return self._enable(plugin_id) + + def _enable(self, plugin_id: str) -> Plugin: record = self._record(plugin_id) if record.plugin.enabled: return record.plugin.model_copy(deep=True) - if record.plugin.manifest.backend.type == "mcp": - # TODO(extension): 第二阶段以隔离进程实现 MCP Host,并补充签名与来源校验。 - record.plugin.status = PluginStatus.dependency_missing - raise ExtensionError( - "PLUGIN_HOST_UNAVAILABLE", - "MCP Plugin Host is reserved for the second development phase.", - status_code=501, - details={"plugin_id": plugin_id, "backend": "mcp"}, - ) missing_grants = sorted( set(record.plugin.manifest.permissions) - set(record.plugin.granted_permissions) ) @@ -336,7 +346,8 @@ class PluginRuntime: status_code=409, details={"plugin_id": plugin_id, "permissions": missing_grants}, ) - conflicts = [spec.name for spec in record.tools if self.registry.contains(spec.name)] + declared_tools = list(record.plugin.manifest.contributes.tools) + conflicts = [name for name in declared_tools if self.registry.contains(name)] if conflicts: raise ExtensionError( "PLUGIN_TOOL_CONFLICT", @@ -346,42 +357,75 @@ class PluginRuntime: ) record.plugin.status = PluginStatus.starting try: - for spec in record.tools: - arguments_model = _arguments_model(spec) + if record.plugin.manifest.backend.type == "mcp": + discovered = self._start_mcp(record) + actual = {item.definition.name for item in discovered} + declared = set(declared_tools) + if actual != declared: + raise ExtensionError( + "PLUGIN_CONTRIBUTION_INVALID", + "Discovered MCP tools must exactly match Plugin contributions.", + details={"declared": sorted(declared), "actual": sorted(actual)}, + ) + for item in discovered: + self._register_mcp_tool(record, item) + else: + for spec in record.tools: + arguments_model = _arguments_model(spec) - async def executor( - arguments: BaseModel, - context: ToolExecutionContext, - _handler: str = spec.handler, - ) -> Any: - return await self.host.execute(_handler, arguments, context) + async def executor( + arguments: BaseModel, + context: ToolExecutionContext, + _handler: str = spec.handler, + ) -> Any: + return await self.host.execute(_handler, arguments, context) - self.registry.register( - ToolDefinition( - name=spec.name, - description=spec.description, - parameters=spec.parameters, - permission=spec.permission, - source="plugin", - ), - arguments_model, - executor, - ) - record.registered_tools.append(spec.name) + self.registry.register( + ToolDefinition( + name=spec.name, + description=spec.description, + parameters=spec.parameters, + permission=spec.permission, + source="plugin", + ), + arguments_model, + executor, + ) + record.registered_tools.append(spec.name) except Exception as exc: # 注册过程必须具备回滚语义,防止半启用插件污染全局工具表。 for name in record.registered_tools: self.registry.unregister(name) record.registered_tools.clear() + record.mcp_remote_names.clear() + self.mcp.stop(plugin_id) record.plugin.status = PluginStatus.error - record.plugin.error_message = str(exc) - raise + record.plugin.error_message = _safe_extension_message(exc) + if isinstance(exc, ExtensionError): + raise + if isinstance(exc, McpBridgeError): + raise ExtensionError( + exc.code, + exc.message, + status_code=exc.status_code, + details={"plugin_id": plugin_id}, + ) from exc + raise ExtensionError( + "PLUGIN_HOST_START_FAILED", + record.plugin.error_message, + status_code=503, + details={"plugin_id": plugin_id}, + ) from exc record.plugin.enabled = True record.plugin.status = PluginStatus.ready record.plugin.error_message = None return record.plugin.model_copy(deep=True) def set_permissions(self, plugin_id: str, permissions: list[str]) -> Plugin: + with self._lock: + return self._set_permissions(plugin_id, permissions) + + def _set_permissions(self, plugin_id: str, permissions: list[str]) -> Plugin: record = self._record(plugin_id) requested = set(permissions) declared = set(record.plugin.manifest.permissions) @@ -403,15 +447,112 @@ class PluginRuntime: return record.plugin.model_copy(deep=True) def disable(self, plugin_id: str) -> Plugin: + with self._lock: + return self._disable(plugin_id) + + def _disable(self, plugin_id: str) -> Plugin: record = self._record(plugin_id) for name in record.registered_tools: self.registry.unregister(name) record.registered_tools.clear() + record.mcp_remote_names.clear() + if record.plugin.manifest.backend.type == "mcp": + self.mcp.stop(plugin_id) record.plugin.enabled = False record.plugin.status = PluginStatus.disabled return record.plugin.model_copy(deep=True) + def get_host_status(self, plugin_id: str) -> PluginHostStatus: + record = self._record(plugin_id) + return self.mcp.status(plugin_id, record.plugin.manifest.backend) + + def restart_host(self, plugin_id: str) -> PluginHostStatus: + with self._lock: + return self._restart_host(plugin_id) + + def _restart_host(self, plugin_id: str) -> PluginHostStatus: + record = self._record(plugin_id) + if record.plugin.manifest.backend.type != "mcp": + raise ExtensionError( + "PLUGIN_HOST_UNAVAILABLE", + "Plugin does not use an MCP Host.", + status_code=409, + details={"plugin_id": plugin_id}, + ) + for name in record.registered_tools: + self.registry.unregister(name) + record.registered_tools.clear() + record.mcp_remote_names.clear() + self.mcp.stop(plugin_id) + record.plugin.enabled = False + record.plugin.status = PluginStatus.installed + record.plugin.error_message = None + self.enable(plugin_id) + return self.get_host_status(plugin_id) + + def shutdown(self) -> None: + """关闭所有隔离 Host;用于 FastAPI lifespan 和测试清理。""" + + with self._lock: + for plugin_id, record in list(self._records.items()): + if record.plugin.manifest.backend.type == "mcp": + self.mcp.stop(plugin_id) + + def _start_mcp(self, record: _PluginRecord) -> list[McpDiscoveredTool]: + manifest = record.plugin.manifest + return self.mcp.start( + manifest.plugin_id, + manifest.backend, + record.package_path, + manifest.permissions, + self._handle_mcp_unavailable, + ) + + def _register_mcp_tool( + self, record: _PluginRecord, discovered: McpDiscoveredTool + ) -> None: + definition = discovered.definition + arguments_model = _arguments_model_from_schema( + definition.name, definition.parameters + ) + plugin_id = record.plugin.manifest.plugin_id + remote_name = discovered.remote_name + + async def executor( + arguments: BaseModel, + context: ToolExecutionContext, + ) -> Any: + return await self.mcp.call_tool( + plugin_id, + remote_name, + arguments.model_dump(), + request_id=context.tool_call_id or f"{context.run_id}:{definition.name}", + ) + + self.registry.register(definition, arguments_model, executor) + record.registered_tools.append(definition.name) + record.mcp_remote_names[definition.name] = remote_name + + def _handle_mcp_unavailable(self, plugin_id: str, message: str) -> None: + with self._lock: + record = self._records.get(plugin_id) + if record is None: + return + for name in record.registered_tools: + self.registry.unregister(name) + record.registered_tools.clear() + record.mcp_remote_names.clear() + record.plugin.enabled = False + record.plugin.status = PluginStatus.error + record.plugin.error_message = message + def uninstall(self, plugin_id: str, dependent_skills: list[str] | None = None) -> None: + with self._lock: + self._uninstall(plugin_id, dependent_skills) + + def _uninstall( + self, plugin_id: str, dependent_skills: list[str] | None = None + ) -> None: record = self._record(plugin_id) if dependent_skills: raise ExtensionError( @@ -422,6 +563,8 @@ class PluginRuntime: ) if record.plugin.enabled: self.disable(plugin_id) + elif record.plugin.manifest.backend.type == "mcp": + self.mcp.stop(plugin_id) del self._records[plugin_id] def _record(self, plugin_id: str) -> _PluginRecord: @@ -498,6 +641,12 @@ def _manifest_error(kind: str, exc: ValidationError) -> ExtensionError: def _arguments_model(spec: DeclarativeToolSpec) -> type[BaseModel]: schema = spec.parameters or {"type": "object", "properties": {}} + return _arguments_model_from_schema(spec.name, schema) + + +def _arguments_model_from_schema( + tool_name: str, schema: dict[str, Any] +) -> type[BaseModel]: if schema.get("type", "object") != "object": raise ExtensionError("PLUGIN_TOOL_SCHEMA_INVALID", "Tool parameters must be an object schema.") properties = schema.get("properties", {}) @@ -514,7 +663,7 @@ def _arguments_model(spec: DeclarativeToolSpec) -> type[BaseModel]: for name, field_schema in properties.items(): annotation = types.get(field_schema.get("type"), Any) fields[name] = (annotation, ... if name in required else None) - model_name = "PluginArgs_" + re.sub(r"\W+", "_", spec.name) + model_name = "PluginArgs_" + re.sub(r"\W+", "_", tool_name) return create_model(model_name, __config__=ConfigDict(extra="forbid"), **fields) @@ -536,3 +685,30 @@ def _validate_tool_schema(spec: DeclarativeToolSpec) -> None: "Tool parameters must be an object schema with object properties.", details={"tool": spec.name}, ) + + +def _validate_backend(manifest: PluginManifest) -> None: + backend = manifest.backend + if backend.type == "mcp": + if backend.transport != "stdio": + raise ExtensionError( + "MCP_CAPABILITY_UNSUPPORTED", + "Phase C MCP Plugins must use stdio transport.", + status_code=501, + ) + if not backend.command or not backend.command.strip(): + raise ExtensionError( + "EXTENSION_MANIFEST_INVALID", + "MCP stdio backend requires a command.", + ) + elif backend.command is not None or backend.args: + raise ExtensionError( + "EXTENSION_MANIFEST_INVALID", + "Only MCP stdio backends may declare command or args.", + ) + + +def _safe_extension_message(exc: Exception) -> str: + if isinstance(exc, (ExtensionError, McpBridgeError)): + return exc.message + return f"Plugin Host operation failed: {type(exc).__name__}." diff --git a/backend/app/main.py b/backend/app/main.py index 510e019..8a3aae5 100644 --- a/backend/app/main.py +++ b/backend/app/main.py @@ -1,19 +1,31 @@ +from contextlib import asynccontextmanager + from fastapi import FastAPI from fastapi.exceptions import RequestValidationError from fastapi.middleware.cors import CORSMiddleware from starlette.exceptions import HTTPException as StarletteHttpException from app.config import get_settings +from app.container import container from app.errors import ApiError, api_error_handler, http_error_handler, validation_error_handler from app.routes import router as api_router from app.schemas import HealthResponse, ServiceStatusResponse settings = get_settings() + +@asynccontextmanager +async def lifespan(_: FastAPI): + yield + # 第三方 MCP Server 必须跟随 AI Core 退出,不能遗留孤儿进程。 + container.plugins.shutdown() + + app = FastAPI( title=settings.name, version=settings.version, description="AI 笔记软件的本地 AI Core 与 Agent Core 服务。", + lifespan=lifespan, ) app.add_middleware( diff --git a/backend/app/routes.py b/backend/app/routes.py index e5e2d07..201973e 100644 --- a/backend/app/routes.py +++ b/backend/app/routes.py @@ -1,3 +1,4 @@ +import asyncio from collections.abc import AsyncIterator from datetime import datetime, timezone from uuid import uuid4 @@ -32,6 +33,7 @@ from app.contracts import ( PageMeta, PermissionDecisionRequest, Plugin, + PluginHostStatus, PluginListResponse, PluginPermissionGrantRequest, ProviderConfig, @@ -130,6 +132,15 @@ def extension_call(operation): raise ApiError(exc.status_code, exc.code, exc.message, exc.details) from exc +async def extension_call_async(operation): + """进程启动/关闭可能等待 stdio Host,移出 FastAPI 事件循环。""" + + try: + return await asyncio.to_thread(operation) + except ExtensionError as exc: + raise ApiError(exc.status_code, exc.code, exc.message, exc.details) from exc + + # Workspace (single configured Vault in Web development mode) @router.get("/workspace", response_model=WorkspaceInfo, tags=["Workspace"]) async def get_workspace() -> WorkspaceInfo: @@ -483,7 +494,7 @@ async def install_plugin(request: ExtensionInstallRequest) -> Plugin: tags=["Plugins"], ) async def enable_plugin(plugin_id: str) -> Plugin: - return extension_call(lambda: container.plugins.enable(plugin_id)) + return await extension_call_async(lambda: container.plugins.enable(plugin_id)) @router.post( @@ -492,7 +503,7 @@ async def enable_plugin(plugin_id: str) -> Plugin: tags=["Plugins"], ) async def disable_plugin(plugin_id: str) -> Plugin: - return extension_call(lambda: container.plugins.disable(plugin_id)) + return await extension_call_async(lambda: container.plugins.disable(plugin_id)) @router.put( @@ -503,11 +514,37 @@ async def disable_plugin(plugin_id: str) -> Plugin: async def set_plugin_permissions( plugin_id: str, request: PluginPermissionGrantRequest ) -> Plugin: - return extension_call( + return await extension_call_async( lambda: container.plugins.set_permissions(plugin_id, request.permissions) ) +@router.get( + "/plugins/{plugin_id}/host", + response_model=PluginHostStatus, + tags=["Plugins"], +) +async def get_plugin_host_status(plugin_id: str) -> PluginHostStatus: + return extension_call(lambda: container.plugins.get_host_status(plugin_id)) + + +@router.post( + "/plugins/{plugin_id}/host/restart", + response_model=OperationResponse, + status_code=202, + tags=["Plugins"], +) +async def restart_plugin_host(plugin_id: str) -> OperationResponse: + status = await extension_call_async( + lambda: container.plugins.restart_host(plugin_id) + ) + return OperationResponse( + status="accepted", + resource_id=plugin_id, + message=f"Plugin Host status: {status.status.value}", + ) + + @router.delete( "/plugins/{plugin_id}", response_model=OperationResponse, @@ -516,7 +553,9 @@ async def set_plugin_permissions( 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) - extension_call(lambda: container.plugins.uninstall(plugin_id, dependent_skills)) + await extension_call_async( + lambda: container.plugins.uninstall(plugin_id, dependent_skills) + ) return OperationResponse(status="completed", resource_id=plugin_id, message="uninstalled") diff --git a/backend/extensions/fixtures/mcp-echo/plugin.yaml b/backend/extensions/fixtures/mcp-echo/plugin.yaml new file mode 100644 index 0000000..34c4312 --- /dev/null +++ b/backend/extensions/fixtures/mcp-echo/plugin.yaml @@ -0,0 +1,21 @@ +id: mcp-fixture +name: MCP Fixture +version: 1.0.0 +description: 阶段 C 离线联调 Fixture,覆盖 MCP Tool 生命周期与错误边界。 +permissions: + - notes.read +contributes: + tools: + - mcp-fixture.echo + - mcp-fixture.fail + - mcp-fixture.sleep + - mcp-fixture.large + - mcp-fixture.environment + - mcp-fixture.exit +backend: + type: mcp + transport: stdio + command: python + args: [server.py] + startup_timeout_seconds: 5 + tool_timeout_seconds: 1 diff --git a/backend/extensions/fixtures/mcp-echo/server.py b/backend/extensions/fixtures/mcp-echo/server.py new file mode 100644 index 0000000..6cf43e1 --- /dev/null +++ b/backend/extensions/fixtures/mcp-echo/server.py @@ -0,0 +1,183 @@ +"""确定性的 MCP stdio 测试 Server;仅使用标准库,不依赖产品代码。""" + +from __future__ import annotations + +import json +import os +import sys +import threading +import time +from typing import Any + +WRITE_LOCK = threading.Lock() +CANCELLED: dict[int, threading.Event] = {} +MODE = sys.argv[1] if len(sys.argv) > 1 else "normal" + + +def send(message: dict[str, Any]) -> None: + with WRITE_LOCK: + sys.stdout.write(json.dumps(message, ensure_ascii=False, separators=(",", ":")) + "\n") + sys.stdout.flush() + + +def respond(request_id: int, result: dict[str, Any]) -> None: + send({"jsonrpc": "2.0", "id": request_id, "result": result}) + + +def tool(name: str, description: str, properties: dict[str, Any] | None = None) -> dict[str, Any]: + return { + "name": name, + "description": description, + "inputSchema": { + "type": "object", + "properties": properties or {}, + "required": list(properties or {}), + "additionalProperties": False, + }, + } + + +TOOLS = { + "echo": { + **tool("echo", "Return the provided text.", {"text": {"type": "string"}}), + "_meta": {"notesagent/permission": "notes.read"}, + }, + "fail": tool("fail", "Return an MCP business error."), + "sleep": tool("sleep", "Wait until completed or cancelled.", {"seconds": {"type": "number"}}), + "large": tool("large", "Return a result larger than the host limit."), + "environment": tool("environment", "Report whether host secrets leaked into the process."), + "exit": tool("exit", "Terminate the fixture process."), +} + + +def call_tool(request_id: int, params: dict[str, Any]) -> None: + name = params.get("name") + arguments = params.get("arguments") or {} + if name == "echo": + text = str(arguments.get("text", "")) + respond( + request_id, + { + "content": [{"type": "text", "text": text}], + "structuredContent": {"echo": text}, + "isError": False, + }, + ) + return + if name == "fail": + respond( + request_id, + { + "content": [{"type": "text", "text": "fixture failure"}], + "isError": True, + }, + ) + return + if name == "large": + respond( + request_id, + { + "content": [{"type": "text", "text": "x" * 300_000}], + "isError": False, + }, + ) + return + if name == "environment": + respond( + request_id, + { + "content": [{"type": "text", "text": "environment checked"}], + "structuredContent": { + "has_openai_key": "OPENAI_API_KEY" in os.environ, + "has_app_db_path": "APP_DB_PATH" in os.environ, + }, + "isError": False, + }, + ) + return + if name == "exit": + os._exit(17) + if name == "sleep": + cancelled = CANCELLED.setdefault(request_id, threading.Event()) + seconds = max(0.0, min(float(arguments.get("seconds", 0)), 30.0)) + if cancelled.wait(seconds): + respond( + request_id, + { + "content": [{"type": "text", "text": "cancelled"}], + "isError": True, + }, + ) + else: + respond( + request_id, + { + "content": [{"type": "text", "text": "completed"}], + "structuredContent": {"slept": seconds}, + "isError": False, + }, + ) + CANCELLED.pop(request_id, None) + return + send( + { + "jsonrpc": "2.0", + "id": request_id, + "error": {"code": -32602, "message": f"Unknown tool: {name}"}, + } + ) + + +def main() -> None: + for line in sys.stdin: + message = json.loads(line) + method = message.get("method") + request_id = message.get("id") + params = message.get("params") or {} + if method == "initialize" and isinstance(request_id, int): + respond( + request_id, + { + "protocolVersion": params.get("protocolVersion"), + "capabilities": ( + {} if MODE == "no-tools" else {"tools": {"listChanged": False}} + ), + "serverInfo": {"name": "notesagent-mcp-fixture", "version": "1.0.0"}, + }, + ) + elif method == "tools/list" and isinstance(request_id, int): + if MODE == "invalid-schema": + respond( + request_id, + { + "tools": [ + { + "name": "broken", + "description": "invalid schema", + "inputSchema": {"type": "string"}, + } + ] + }, + ) + elif params.get("cursor") == "page-2": + respond( + request_id, + {"tools": [TOOLS["large"], TOOLS["environment"], TOOLS["exit"]]}, + ) + else: + respond( + request_id, + {"tools": [TOOLS["echo"], TOOLS["fail"], TOOLS["sleep"]], "nextCursor": "page-2"}, + ) + elif method == "tools/call" and isinstance(request_id, int): + threading.Thread(target=call_tool, args=(request_id, params), daemon=True).start() + elif method == "notifications/cancelled": + cancelled_id = params.get("requestId") + if isinstance(cancelled_id, int): + CANCELLED.setdefault(cancelled_id, threading.Event()).set() + elif method == "ping" and isinstance(request_id, int): + respond(request_id, {}) + + +if __name__ == "__main__": + main() diff --git a/backend/tests/test_api.py b/backend/tests/test_api.py index f62496e..e76a1e8 100644 --- a/backend/tests/test_api.py +++ b/backend/tests/test_api.py @@ -93,6 +93,8 @@ def test_openapi_contains_documented_frontend_interfaces() -> None: "/api/skills", "/api/plugins", "/api/plugins/install", + "/api/plugins/{plugin_id}/host", + "/api/plugins/{plugin_id}/host/restart", "/api/plugins/{plugin_id}/enable", "/api/plugins/{plugin_id}/disable", "/api/providers/test", diff --git a/backend/tests/test_extension_core.py b/backend/tests/test_extension_core.py index 17f9891..1c3274a 100644 --- a/backend/tests/test_extension_core.py +++ b/backend/tests/test_extension_core.py @@ -1,4 +1,6 @@ import asyncio +import shutil +import time import pytest @@ -13,13 +15,28 @@ from app.contracts import ( ) from app.extensions import ExtensionError from app.services import note_service -from app.config import get_settings +from app.config import BACKEND_DIR, get_settings + + +MCP_FIXTURE = BACKEND_DIR / "extensions" / "fixtures" / "mcp-echo" def run(coroutine): return asyncio.run(coroutine) +@pytest.fixture +def mcp_container(): + container = build_container() + installed = container.plugins.install(MCP_FIXTURE) + assert installed.status == "permission_required" + container.plugins.set_permissions("mcp-fixture", ["notes.read"]) + try: + yield container + finally: + container.plugins.shutdown() + + def test_bundled_plugin_registers_tool_and_skill_is_ready() -> None: async def scenario() -> None: container = build_container() @@ -297,3 +314,196 @@ def test_attachment_and_transcription_tools_use_host_storage() -> None: assert transcription.output["text"] == "会议转写内容" run(scenario()) + + +def test_mcp_stdio_host_discovers_namespaced_tools_and_maps_results( + mcp_container, monkeypatch +) -> None: + async def scenario() -> None: + monkeypatch.setenv("OPENAI_API_KEY", "must-not-enter-plugin-host") + enabled = mcp_container.plugins.enable("mcp-fixture") + status = mcp_container.plugins.get_host_status("mcp-fixture") + definition = mcp_container.tools.get("mcp-fixture.echo").definition + result = await mcp_container.tools.execute( + ToolCall( + tool_call_id="call_mcp_echo", + name="mcp-fixture.echo", + arguments={"text": "hello mcp"}, + ), + ToolExecutionContext( + run_id="run_mcp_fixture", tool_call_id="call_mcp_echo" + ), + ) + + assert enabled.status == "ready" and enabled.enabled is True + assert status.status == "ready" + environment = await mcp_container.tools.execute( + ToolCall( + tool_call_id="call_mcp_environment", + name="mcp-fixture.environment", + arguments={}, + ), + ToolExecutionContext(run_id="run_mcp_fixture"), + ) + + assert status.tools_count == 6 + assert status.protocol_version == "2025-11-25" + assert status.server_name == "notesagent-mcp-fixture" + assert definition.permission == "notes.read" + assert result.success is True + assert result.output == {"echo": "hello mcp"} + assert environment.success is True + assert environment.output == { + "has_openai_key": False, + "has_app_db_path": False, + } + + disabled = mcp_container.plugins.disable("mcp-fixture") + assert disabled.status == "disabled" + assert mcp_container.plugins.get_host_status("mcp-fixture").status == "stopped" + assert not mcp_container.tools.contains("mcp-fixture.echo") + + run(scenario()) + + +def test_agent_calls_mcp_tool_through_registry_and_writes_trace(mcp_container) -> None: + async def scenario() -> None: + mcp_container.plugins.enable("mcp-fixture") + created = await mcp_container.agent.create_run( + AgentRunCreateRequest( + input='/tool mcp-fixture.echo {"text":"agent mcp"}', + provider_id="mock", + model="mock-1", + allowed_tools=["mcp-fixture.echo"], + ) + ) + completed = await mcp_container.agent.wait(created.run_id) + trace = mcp_container.agent.get_trace( + created.run_id, after_sequence=-1, limit=100 + ) + + assert completed.status == AgentRunStatus.completed + assert completed.tool_results[0].success is True + assert completed.tool_results[0].output == {"echo": "agent mcp"} + assert any( + item.event == "ToolCall" and item.data.get("name") == "mcp-fixture.echo" + for item in trace.items + ) + + run(scenario()) + + +def test_mcp_business_error_size_limit_and_timeout_are_structured(mcp_container) -> None: + async def scenario() -> None: + mcp_container.plugins.enable("mcp-fixture") + context = ToolExecutionContext(run_id="run_mcp_errors") + + failed = await mcp_container.tools.execute( + ToolCall(tool_call_id="call_fail", name="mcp-fixture.fail", arguments={}), + context, + ) + oversized = await mcp_container.tools.execute( + ToolCall(tool_call_id="call_large", name="mcp-fixture.large", arguments={}), + context, + ) + timed_out = await mcp_container.tools.execute( + ToolCall( + tool_call_id="call_sleep", + name="mcp-fixture.sleep", + arguments={"seconds": 5}, + ), + ToolExecutionContext( + run_id="run_mcp_errors", tool_call_id="call_sleep" + ), + ) + recovered = await mcp_container.tools.execute( + ToolCall( + tool_call_id="call_after_timeout", + name="mcp-fixture.echo", + arguments={"text": "still ready"}, + ), + context, + ) + + assert failed.success is False + assert failed.error_code == "MCP_TOOL_CALL_FAILED" + assert failed.error_message == "fixture failure" + assert oversized.success is False + assert oversized.error_code == "MCP_TOOL_RESULT_TOO_LARGE" + assert timed_out.success is False + assert timed_out.error_code == "MCP_TOOL_CALL_FAILED" + assert recovered.success is True + assert mcp_container.plugins.get_host_status("mcp-fixture").status == "ready" + + run(scenario()) + + +def test_mcp_abnormal_exit_unregisters_tools_and_restart_recovers(mcp_container) -> None: + async def scenario() -> None: + mcp_container.plugins.enable("mcp-fixture") + crashed = await mcp_container.tools.execute( + ToolCall(tool_call_id="call_exit", name="mcp-fixture.exit", arguments={}), + ToolExecutionContext(run_id="run_mcp_exit", tool_call_id="call_exit"), + ) + + deadline = time.monotonic() + 2 + while mcp_container.tools.contains("mcp-fixture.echo") and time.monotonic() < deadline: + await asyncio.sleep(0.02) + + plugin = mcp_container.plugins.get("mcp-fixture") + status = mcp_container.plugins.get_host_status("mcp-fixture") + assert crashed.success is False + assert crashed.error_code == "PLUGIN_HOST_UNAVAILABLE" + assert plugin.status == "error" and plugin.enabled is False + assert status.status == "unhealthy" + assert not mcp_container.tools.contains("mcp-fixture.echo") + + restarted = mcp_container.plugins.restart_host("mcp-fixture") + assert restarted.status == "ready" + assert restarted.tools_count == 6 + assert mcp_container.tools.contains("mcp-fixture.echo") + + run(scenario()) + + +@pytest.mark.parametrize( + ("mode", "contributions", "expected_code"), + [ + ("no-tools", "[]", "MCP_CAPABILITY_UNSUPPORTED"), + ("invalid-schema", "[mcp-invalid.broken]", "MCP_TOOL_SCHEMA_INVALID"), + ], +) +def test_mcp_rejects_missing_capability_and_invalid_discovery( + tmp_path, mode, contributions, expected_code +) -> None: + package = tmp_path / f"mcp-{mode}" + package.mkdir() + shutil.copyfile(MCP_FIXTURE / "server.py", package / "server.py") + (package / "plugin.yaml").write_text( + f""" +id: mcp-invalid +name: Invalid MCP Fixture +version: 1.0.0 +contributes: + tools: {contributions} +backend: + type: mcp + transport: stdio + command: python + args: [server.py, {mode}] + startup_timeout_seconds: 5 + tool_timeout_seconds: 1 +""".strip(), + encoding="utf-8", + ) + container = build_container() + container.plugins.install(package) + try: + with pytest.raises(ExtensionError) as exc: + container.plugins.enable("mcp-invalid") + assert exc.value.code == expected_code + assert container.plugins.get("mcp-invalid").status == "error" + assert container.plugins.get_host_status("mcp-invalid").status == "error" + assert not container.tools.contains("mcp-invalid.broken") + finally: + container.plugins.shutdown() diff --git a/frontend/src/contracts/index.ts b/frontend/src/contracts/index.ts index 0fcfcfd..8891885 100644 --- a/frontend/src/contracts/index.ts +++ b/frontend/src/contracts/index.ts @@ -255,6 +255,22 @@ export type PluginStatus = | 'dependency_missing' | 'permission_required' +export type PluginHostState = 'stopped' | 'starting' | 'ready' | 'unhealthy' | 'error' + +export interface PluginHostStatus { + plugin_id: string + backend_type: 'mcp' | 'internal_rpc' | 'none' + transport: 'stdio' | 'http' | 'none' + status: PluginHostState + tools_count: number + started_at?: string | null + last_seen_at?: string | null + protocol_version?: string | null + server_name?: string | null + server_version?: string | null + error?: string | null +} + export interface PluginContribution { type: 'tool' | 'command' | 'importer' | 'exporter' | 'sidebar_panel' | 'settings_section' id: string @@ -539,7 +555,14 @@ export interface ApiPlugin { panels: string[] settings_sections: string[] } - backend: { type: 'mcp' | 'internal_rpc' | 'none'; transport: 'stdio' | 'http' | 'none' } + backend: { + type: 'mcp' | 'internal_rpc' | 'none' + transport: 'stdio' | 'http' | 'none' + command?: string | null + args?: string[] + startup_timeout_seconds?: number + tool_timeout_seconds?: number + } } status: PluginStatus enabled: boolean diff --git a/frontend/src/services/pluginService.ts b/frontend/src/services/pluginService.ts index 726a7e3..48b6a28 100644 --- a/frontend/src/services/pluginService.ts +++ b/frontend/src/services/pluginService.ts @@ -1,5 +1,5 @@ import apiClient from './apiClient' -import type { ApiPlugin, OperationResponse, Plugin, PluginContribution } from '@/contracts' +import type { ApiPlugin, OperationResponse, Plugin, PluginContribution, PluginHostStatus } from '@/contracts' function toPlugin(plugin: ApiPlugin): Plugin { const { manifest } = plugin @@ -54,6 +54,14 @@ export async function grantPluginPermissions(pluginId: string, permissions: stri return toPlugin(await apiClient.put(`/api/plugins/${pluginId}/permissions`, { permissions })) } +export async function getPluginHostStatus(pluginId: string): Promise { + return apiClient.get(`/api/plugins/${pluginId}/host`) +} + +export async function restartPluginHost(pluginId: string): Promise { + return apiClient.post(`/api/plugins/${pluginId}/host/restart`) +} + export async function uninstallPlugin(pluginId: string): Promise { return apiClient.delete(`/api/plugins/${pluginId}`) } From a8e0fe8ff3d6ac79090168d1d8884ea4d64c62da Mon Sep 17 00:00:00 2001 From: KiriAky 107 Date: Tue, 1 Sep 2026 11:32:47 +0800 Subject: [PATCH 2/5] =?UTF-8?q?docs(extension):=20=E8=A1=A5=E5=85=85?= =?UTF-8?q?=E9=98=B6=E6=AE=B5C=20MCP=E5=BC=80=E5=8F=91=E8=AF=B4=E6=98=8E?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- README.md | 3 ++- backend/README.md | 2 +- 2 files changed, 3 insertions(+), 2 deletions(-) diff --git a/README.md b/README.md index 7bfa242..bfc698e 100644 --- a/README.md +++ b/README.md @@ -118,7 +118,7 @@ cd frontend pnpm test ``` -当前回归基线为后端 81 项测试、前端 27 项测试,且生产构建通过。测试数量会随功能增长,以本地实际输出和 CI 为准。 +当前回归基线为后端 87 项测试、前端 27 项测试,且生产构建通过。测试数量会随功能增长,以本地实际输出和 CI 为准。 构建产物位于 `frontend/dist`,该目录不提交到 Git。 @@ -132,6 +132,7 @@ pnpm test | [后端接口契约](docs/contracts/后端接口契约-开发版.md) | HTTP/SSE 接口、错误和当前实现状态 | | [第二阶段接口契约](docs/contracts/第二阶段接口契约-开发版.md) | 第二阶段公共 DTO、计划接口、SSE、错误码与联调顺序 | | [AI Core 与 Agent Core](docs/development/AI-Core与Agent-Core开发说明.md) | Provider、Agent、Tool、Permission 与 Extension Core | +| [MCP Bridge 与 Plugin Host](docs/development/MCP-Bridge与Plugin-Host开发说明.md) | stdio MCP、隔离进程、Tool 映射、状态与错误边界 | | [Git 使用细则](docs/guides/Git使用细则-团队开发版.md) | 分支、提交、PR、Review 与合并流程 | | [CI/CD 细则](docs/guides/CI-CD细则-团队开发版.md) | Gitea 流水线、质量门禁、产物、发布与回滚规则 | | [Agent Trace 复盘](docs/retrospectives/Agent-Core第二阶段问题与修复复盘.md) | Agent 持久化、SSE 恢复、事件契约与脱敏问题复盘 | diff --git a/backend/README.md b/backend/README.md index c211fc1..3c8429e 100644 --- a/backend/README.md +++ b/backend/README.md @@ -23,7 +23,7 @@ uv run uvicorn app.main:app --reload --host 127.0.0.1 --port 8000 uv run pytest ``` -当前基线为 81 项测试通过。Provider API Key 可通过前端设置页写入,也可用 `OPENAI_API_KEY`、`DEEPSEEK_API_KEY` 或 `AINOTE_CREDENTIAL_` 注入;不要把真实密钥写入仓库。 +当前基线为 87 项测试通过。Provider API Key 可通过前端设置页写入,也可用 `OPENAI_API_KEY`、`DEEPSEEK_API_KEY` 或 `AINOTE_CREDENTIAL_` 注入;不要把真实密钥写入仓库。 团队接口清单见 `../docs/contracts/后端接口契约-开发版.md`,机器可读契约以运行时的 `/openapi.json` 为准。 From 20920b684556f5468bd357f53f4b8fffdfc87357 Mon Sep 17 00:00:00 2001 From: KiriAky 107 Date: Tue, 1 Sep 2026 12:11:30 +0800 Subject: [PATCH 3/5] =?UTF-8?q?fix(extension):=20=E4=BF=AE=E5=A4=8D=20MCP?= =?UTF-8?q?=20Host=20=E8=B5=84=E6=BA=90=E4=B8=8E=E5=8D=8F=E8=AE=AE?= =?UTF-8?q?=E8=BE=B9=E7=95=8C?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- README.md | 2 +- backend/README.md | 2 +- backend/app/extensions/mcp.py | 39 ++++++++++-- backend/app/extensions/runtime.py | 15 ++++- .../extensions/fixtures/mcp-echo/server.py | 9 +++ backend/tests/test_extension_core.py | 62 ++++++++++++++++++- 6 files changed, 119 insertions(+), 10 deletions(-) diff --git a/README.md b/README.md index bfc698e..ac7a7db 100644 --- a/README.md +++ b/README.md @@ -118,7 +118,7 @@ cd frontend pnpm test ``` -当前回归基线为后端 87 项测试、前端 27 项测试,且生产构建通过。测试数量会随功能增长,以本地实际输出和 CI 为准。 +当前回归基线为后端 91 项测试、前端 27 项测试,且生产构建通过。测试数量会随功能增长,以本地实际输出和 CI 为准。 构建产物位于 `frontend/dist`,该目录不提交到 Git。 diff --git a/backend/README.md b/backend/README.md index 3c8429e..6a0930e 100644 --- a/backend/README.md +++ b/backend/README.md @@ -23,7 +23,7 @@ uv run uvicorn app.main:app --reload --host 127.0.0.1 --port 8000 uv run pytest ``` -当前基线为 87 项测试通过。Provider API Key 可通过前端设置页写入,也可用 `OPENAI_API_KEY`、`DEEPSEEK_API_KEY` 或 `AINOTE_CREDENTIAL_` 注入;不要把真实密钥写入仓库。 +当前基线为 91 项测试通过。Provider API Key 可通过前端设置页写入,也可用 `OPENAI_API_KEY`、`DEEPSEEK_API_KEY` 或 `AINOTE_CREDENTIAL_` 注入;不要把真实密钥写入仓库。 团队接口清单见 `../docs/contracts/后端接口契约-开发版.md`,机器可读契约以运行时的 `/openapi.json` 为准。 diff --git a/backend/app/extensions/mcp.py b/backend/app/extensions/mcp.py index 2647c20..d5036a7 100644 --- a/backend/app/extensions/mcp.py +++ b/backend/app/extensions/mcp.py @@ -188,7 +188,7 @@ class McpStdioClient: result = response.get("result") if not isinstance(result, dict): raise McpBridgeError( - "MCP_TOOL_CALL_FAILED", "MCP response result must be an object." + response_error_code, "MCP response result must be an object." ) return result @@ -207,9 +207,18 @@ class McpStdioClient: except McpBridgeError: pass - def abandon(self, request_id: int) -> None: + def abandon( + self, request_id: int, wake_error: BaseException | None = None + ) -> None: with self._pending_lock: - self._pending.pop(request_id, None) + pending = self._pending.pop(request_id, None) + # asyncio.to_thread 被取消时不会停止底层线程;主动唤醒 Queue,避免线程 + # 一直占用默认线程池直至远端超时。 + 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: process = self.process @@ -258,7 +267,15 @@ class McpStdioClient: assert process is not None and process.stdout is not None failure: str | None = None try: - for raw_line in process.stdout: + while True: + # readline(size) 在换行缺失时仍有硬上限,不能先把任意大的 + # 第三方 stdout 行完整读入宿主内存再检查。 + raw_line = process.stdout.readline(MAX_MCP_MESSAGE_BYTES + 1) + if raw_line == "": + break + if not raw_line.endswith("\n"): + failure = "MCP server emitted an oversized or unterminated message." + break if len(raw_line.encode("utf-8")) > MAX_MCP_MESSAGE_BYTES: failure = "MCP server emitted an oversized protocol message." break @@ -313,7 +330,12 @@ class McpStdioClient: process = self.process assert process is not None and process.stderr is not None try: - for line in process.stderr: + while True: + # stderr 不是协议通道,但同样按块读取,避免无换行日志造成 + # 宿主侧的无界字符串分配。 + line = process.stderr.readline(1025) + if line == "": + break self._stderr_tail.append(line.rstrip()[:1024]) except (OSError, ValueError): return @@ -489,7 +511,12 @@ class McpBridge: ) except asyncio.CancelledError: host.client.cancel(rpc_id) - host.client.abandon(rpc_id) + host.client.abandon( + rpc_id, + McpBridgeError( + "MCP_TOOL_CALL_FAILED", "MCP request was cancelled." + ), + ) raise except McpBridgeError as exc: raise ToolExecutionError(exc.code, exc.message) from exc diff --git a/backend/app/extensions/runtime.py b/backend/app/extensions/runtime.py index 33d0bbc..fa6f72f 100644 --- a/backend/app/extensions/runtime.py +++ b/backend/app/extensions/runtime.py @@ -479,6 +479,17 @@ class PluginRuntime: status_code=409, details={"plugin_id": plugin_id}, ) + if record.plugin.status in { + PluginStatus.installed, + PluginStatus.disabled, + PluginStatus.permission_required, + }: + raise ExtensionError( + "PLUGIN_HOST_UNAVAILABLE", + "Disabled or inactive MCP Plugins must be started with Enable.", + status_code=409, + details={"plugin_id": plugin_id, "status": record.plugin.status.value}, + ) for name in record.registered_tools: self.registry.unregister(name) record.registered_tools.clear() @@ -664,7 +675,9 @@ def _arguments_model_from_schema( annotation = types.get(field_schema.get("type"), Any) fields[name] = (annotation, ... if name in required else None) model_name = "PluginArgs_" + re.sub(r"\W+", "_", tool_name) - return create_model(model_name, __config__=ConfigDict(extra="forbid"), **fields) + # 完整 JSON Schema 已在 ToolRegistry 中先行校验。这里允许额外字段,避免 + # Pydantic 再次拒绝 additionalProperties/patternProperties 接受的合法参数。 + return create_model(model_name, __config__=ConfigDict(extra="allow"), **fields) def _validate_tool_schema(spec: DeclarativeToolSpec) -> None: diff --git a/backend/extensions/fixtures/mcp-echo/server.py b/backend/extensions/fixtures/mcp-echo/server.py index 6cf43e1..59f2be9 100644 --- a/backend/extensions/fixtures/mcp-echo/server.py +++ b/backend/extensions/fixtures/mcp-echo/server.py @@ -135,6 +135,15 @@ def main() -> None: request_id = message.get("id") params = message.get("params") or {} if method == "initialize" and isinstance(request_id, int): + if MODE == "invalid-result": + send({"jsonrpc": "2.0", "id": request_id, "result": None}) + continue + if MODE == "oversized-stdout": + # 不带换行,验证 Host 在读取完整内容前执行硬上限。 + sys.stdout.write("x" * (2 * 1024 * 1024 + 1)) + sys.stdout.flush() + time.sleep(10) + return respond( request_id, { diff --git a/backend/tests/test_extension_core.py b/backend/tests/test_extension_core.py index 1c3274a..4e0a9a8 100644 --- a/backend/tests/test_extension_core.py +++ b/backend/tests/test_extension_core.py @@ -1,5 +1,6 @@ import asyncio import shutil +import threading import time import pytest @@ -14,6 +15,8 @@ from app.contracts import ( ToolCall, ) from app.extensions import ExtensionError +from app.extensions.mcp import McpStdioClient +from app.extensions.runtime import _arguments_model_from_schema from app.services import note_service from app.config import BACKEND_DIR, get_settings @@ -362,6 +365,11 @@ def test_mcp_stdio_host_discovers_namespaced_tools_and_maps_results( assert disabled.status == "disabled" assert mcp_container.plugins.get_host_status("mcp-fixture").status == "stopped" assert not mcp_container.tools.contains("mcp-fixture.echo") + with pytest.raises(ExtensionError) as exc: + mcp_container.plugins.restart_host("mcp-fixture") + assert exc.value.code == "PLUGIN_HOST_UNAVAILABLE" + assert mcp_container.plugins.get("mcp-fixture").status == "disabled" + assert not mcp_container.tools.contains("mcp-fixture.echo") run(scenario()) @@ -438,6 +446,56 @@ def test_mcp_business_error_size_limit_and_timeout_are_structured(mcp_container) run(scenario()) +def test_mcp_cancel_releases_blocking_response_thread( + mcp_container, monkeypatch +) -> None: + async def scenario() -> None: + mcp_container.plugins.enable("mcp-fixture") + released = threading.Event() + original_wait = McpStdioClient.wait_response + + def tracked_wait(self, *args, **kwargs): + try: + return original_wait(self, *args, **kwargs) + finally: + released.set() + + monkeypatch.setattr(McpStdioClient, "wait_response", tracked_wait) + task = asyncio.create_task( + mcp_container.plugins.mcp.call_tool( + "mcp-fixture", + "sleep", + {"seconds": 5}, + request_id="call_cancel_release", + ) + ) + await asyncio.sleep(0.05) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + deadline = time.monotonic() + 0.5 + while not released.is_set() and time.monotonic() < deadline: + await asyncio.sleep(0.01) + assert released.is_set(), "cancelled MCP wait must not occupy a worker until timeout" + + run(scenario()) + + +def test_mcp_argument_model_preserves_json_schema_additional_properties() -> None: + arguments_model = _arguments_model_from_schema( + "mcp-fixture.dynamic", + { + "type": "object", + "additionalProperties": {"type": "string"}, + }, + ) + + arguments = arguments_model.model_validate({"dynamic_key": "value"}) + + assert arguments.model_dump() == {"dynamic_key": "value"} + + def test_mcp_abnormal_exit_unregisters_tools_and_restart_recovers(mcp_container) -> None: async def scenario() -> None: mcp_container.plugins.enable("mcp-fixture") @@ -471,9 +529,11 @@ def test_mcp_abnormal_exit_unregisters_tools_and_restart_recovers(mcp_container) [ ("no-tools", "[]", "MCP_CAPABILITY_UNSUPPORTED"), ("invalid-schema", "[mcp-invalid.broken]", "MCP_TOOL_SCHEMA_INVALID"), + ("invalid-result", "[]", "MCP_INITIALIZE_FAILED"), + ("oversized-stdout", "[]", "PLUGIN_HOST_UNAVAILABLE"), ], ) -def test_mcp_rejects_missing_capability_and_invalid_discovery( +def test_mcp_rejects_invalid_initialization_and_discovery( tmp_path, mode, contributions, expected_code ) -> None: package = tmp_path / f"mcp-{mode}" From 199dd25c3e224841aa34f94a9635d418d1226a68 Mon Sep 17 00:00:00 2001 From: KiriAky 107 Date: Tue, 1 Sep 2026 16:06:23 +0800 Subject: [PATCH 4/5] =?UTF-8?q?fix(extension):=20=E5=AE=8C=E5=96=84=20MCP?= =?UTF-8?q?=20=E5=8F=82=E6=95=B0=E4=B8=8E=E8=BF=90=E8=A1=8C=E5=AE=89?= =?UTF-8?q?=E5=85=A8=E8=BE=B9=E7=95=8C?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- backend/app/extensions/mcp.py | 12 ++++++++++ backend/app/extensions/runtime.py | 16 ++++++++++---- .../extensions/fixtures/mcp-echo/server.py | 16 ++++++++++++-- backend/tests/test_extension_core.py | 22 +++++++++++++++++++ 4 files changed, 60 insertions(+), 6 deletions(-) diff --git a/backend/app/extensions/mcp.py b/backend/app/extensions/mcp.py index d5036a7..68f056b 100644 --- a/backend/app/extensions/mcp.py +++ b/backend/app/extensions/mcp.py @@ -91,6 +91,8 @@ class McpStdioClient: def start(self) -> None: if self.process is not None and self.process.poll() is None: return + # TODO(extension-security): 社区 Plugin 开放前迁移到 Tauri/Rust Host 的 + # 平台级沙箱启动器;uvx 只隔离 Python 依赖,不能替代系统权限限制。 creation_flags = getattr(subprocess, "CREATE_NO_WINDOW", 0) if os.name == "nt" else 0 environment = _subprocess_environment() environment.setdefault("PYTHONUNBUFFERED", "1") @@ -567,6 +569,16 @@ class McpBridge: host.status.tools_count = 0 host.status.error = None + def remove(self, plugin_id: str) -> None: + """停止 Host,并清除卸载后不应跨安装保留的状态与调用索引。""" + + self.stop(plugin_id) + with self._lock: + self._statuses.pop(plugin_id, None) + stale_calls = [key for key in self._calls if key[0] == plugin_id] + for key in stale_calls: + self._calls.pop(key, None) + def status(self, plugin_id: str, backend: PluginBackend) -> PluginHostStatus: with self._lock: status = self._statuses.get(plugin_id) diff --git a/backend/app/extensions/runtime.py b/backend/app/extensions/runtime.py index fa6f72f..cacbb12 100644 --- a/backend/app/extensions/runtime.py +++ b/backend/app/extensions/runtime.py @@ -536,7 +536,9 @@ class PluginRuntime: return await self.mcp.call_tool( plugin_id, remote_name, - arguments.model_dump(), + # 省略的可选字段不能被补成 null;显式传入的 null 仍由 + # model_fields_set 保留并交给 MCP Server。 + arguments.model_dump(exclude_unset=True), request_id=context.tool_call_id or f"{context.run_id}:{definition.name}", ) @@ -572,10 +574,13 @@ class PluginRuntime: status_code=409, details={"plugin_id": plugin_id, "skills": dependent_skills}, ) + is_mcp = record.plugin.manifest.backend.type == "mcp" if record.plugin.enabled: self.disable(plugin_id) - elif record.plugin.manifest.backend.type == "mcp": - self.mcp.stop(plugin_id) + if is_mcp: + # stop 只结束本次进程并保留状态供故障诊断;真正卸载时必须连同 + # 历史状态一起遗忘,避免同 ID 重装继承旧协商信息。 + self.mcp.remove(plugin_id) del self._records[plugin_id] def _record(self, plugin_id: str) -> _PluginRecord: @@ -672,7 +677,10 @@ def _arguments_model_from_schema( "object": dict[str, Any], } for name, field_schema in properties.items(): - annotation = types.get(field_schema.get("type"), Any) + schema_type = field_schema.get("type") + # JSON Schema 允许联合类型数组;复杂类型继续由 Draft Validator + # 精确校验,Pydantic 在这里只承担参数载体职责。 + annotation = types.get(schema_type, Any) if isinstance(schema_type, str) else Any fields[name] = (annotation, ... if name in required else None) model_name = "PluginArgs_" + re.sub(r"\W+", "_", tool_name) # 完整 JSON Schema 已在 ToolRegistry 中先行校验。这里允许额外字段,避免 diff --git a/backend/extensions/fixtures/mcp-echo/server.py b/backend/extensions/fixtures/mcp-echo/server.py index 59f2be9..b1fefda 100644 --- a/backend/extensions/fixtures/mcp-echo/server.py +++ b/backend/extensions/fixtures/mcp-echo/server.py @@ -39,7 +39,14 @@ def tool(name: str, description: str, properties: dict[str, Any] | None = None) TOOLS = { "echo": { - **tool("echo", "Return the provided text.", {"text": {"type": "string"}}), + **tool( + "echo", + "Return the provided text.", + { + "text": {"type": "string"}, + "suffix": {"type": ["string", "null"]}, + }, + ), "_meta": {"notesagent/permission": "notes.read"}, }, "fail": tool("fail", "Return an MCP business error."), @@ -48,6 +55,8 @@ TOOLS = { "environment": tool("environment", "Report whether host secrets leaked into the process."), "exit": tool("exit", "Terminate the fixture process."), } +# suffix 是可选字段,用于验证 Host 不会把缺省值擅自补成 null。 +TOOLS["echo"]["inputSchema"]["required"] = ["text"] def call_tool(request_id: int, params: dict[str, Any]) -> None: @@ -55,11 +64,14 @@ def call_tool(request_id: int, params: dict[str, Any]) -> None: arguments = params.get("arguments") or {} if name == "echo": text = str(arguments.get("text", "")) + structured_content = {"echo": text} + if "suffix" in arguments: + structured_content["suffix"] = arguments["suffix"] respond( request_id, { "content": [{"type": "text", "text": text}], - "structuredContent": {"echo": text}, + "structuredContent": structured_content, "isError": False, }, ) diff --git a/backend/tests/test_extension_core.py b/backend/tests/test_extension_core.py index 4e0a9a8..5b368aa 100644 --- a/backend/tests/test_extension_core.py +++ b/backend/tests/test_extension_core.py @@ -355,6 +355,19 @@ def test_mcp_stdio_host_discovers_namespaced_tools_and_maps_results( assert definition.permission == "notes.read" assert result.success is True assert result.output == {"echo": "hello mcp"} + explicit_null = await mcp_container.tools.execute( + ToolCall( + tool_call_id="call_mcp_explicit_null", + name="mcp-fixture.echo", + arguments={"text": "null stays explicit", "suffix": None}, + ), + ToolExecutionContext(run_id="run_mcp_fixture"), + ) + assert explicit_null.success is True + assert explicit_null.output == { + "echo": "null stays explicit", + "suffix": None, + } assert environment.success is True assert environment.output == { "has_openai_key": False, @@ -371,6 +384,15 @@ def test_mcp_stdio_host_discovers_namespaced_tools_and_maps_results( assert mcp_container.plugins.get("mcp-fixture").status == "disabled" assert not mcp_container.tools.contains("mcp-fixture.echo") + mcp_container.plugins.uninstall("mcp-fixture") + reinstalled = mcp_container.plugins.install(MCP_FIXTURE) + fresh_status = mcp_container.plugins.get_host_status("mcp-fixture") + assert reinstalled.status == "permission_required" + assert fresh_status.status == "stopped" + assert fresh_status.started_at is None + assert fresh_status.protocol_version is None + assert fresh_status.server_name is None + run(scenario()) From 605bfc1c1a5dbc692842ac068a5d6f50aac01525 Mon Sep 17 00:00:00 2001 From: KiriAky 107 Date: Tue, 1 Sep 2026 21:25:05 +0800 Subject: [PATCH 5/5] =?UTF-8?q?fix(extension):=20=E5=BC=BA=E5=8C=96=20MCP?= =?UTF-8?q?=20=E5=8F=82=E6=95=B0=E4=B8=8E=E7=94=9F=E4=BA=A7=E8=BF=90?= =?UTF-8?q?=E8=A1=8C=E9=97=A8=E7=A6=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- README.md | 2 +- backend/README.md | 2 +- backend/app/container.py | 10 ++++++-- backend/app/extensions/runtime.py | 36 +++++++++++++--------------- backend/tests/test_extension_core.py | 29 ++++++++++++++++++++-- 5 files changed, 53 insertions(+), 26 deletions(-) diff --git a/README.md b/README.md index ac7a7db..10d735b 100644 --- a/README.md +++ b/README.md @@ -118,7 +118,7 @@ cd frontend pnpm test ``` -当前回归基线为后端 91 项测试、前端 27 项测试,且生产构建通过。测试数量会随功能增长,以本地实际输出和 CI 为准。 +当前回归基线为后端 92 项测试、前端 27 项测试,且生产构建通过。测试数量会随功能增长,以本地实际输出和 CI 为准。 构建产物位于 `frontend/dist`,该目录不提交到 Git。 diff --git a/backend/README.md b/backend/README.md index 6a0930e..3d4fde4 100644 --- a/backend/README.md +++ b/backend/README.md @@ -23,7 +23,7 @@ uv run uvicorn app.main:app --reload --host 127.0.0.1 --port 8000 uv run pytest ``` -当前基线为 91 项测试通过。Provider API Key 可通过前端设置页写入,也可用 `OPENAI_API_KEY`、`DEEPSEEK_API_KEY` 或 `AINOTE_CREDENTIAL_` 注入;不要把真实密钥写入仓库。 +当前基线为 92 项测试通过。Provider API Key 可通过前端设置页写入,也可用 `OPENAI_API_KEY`、`DEEPSEEK_API_KEY` 或 `AINOTE_CREDENTIAL_` 注入;不要把真实密钥写入仓库。 团队接口清单见 `../docs/contracts/后端接口契约-开发版.md`,机器可读契约以运行时的 `/openapi.json` 为准。 diff --git a/backend/app/container.py b/backend/app/container.py index 0cff70a..9263ca9 100644 --- a/backend/app/container.py +++ b/backend/app/container.py @@ -3,7 +3,7 @@ from dataclasses import dataclass from app.agent import AgentRuntime, PermissionManager, PermissionPolicy, ToolRegistry from app.agent.builtin_tools import register_builtin_tools from app.contracts import ModelCapability, ProviderConfig, ProviderType -from app.config import BACKEND_DIR +from app.config import BACKEND_DIR, get_settings from app.extensions import PluginRuntime, SkillRuntime from app.providers import MockProvider, ProviderFactory, ProviderRegistry from app.providers.credentials import ( @@ -26,6 +26,7 @@ class ApplicationContainer: def build_container() -> ApplicationContainer: + settings = get_settings() credentials = EncryptedCredentialStore() provider_factory = ProviderFactory( ChainedCredentialResolver(credentials, EnvironmentCredentialResolver()) @@ -50,7 +51,12 @@ def build_container() -> ApplicationContainer: tools = ToolRegistry() register_builtin_tools(tools) - plugins = PluginRuntime(tools) + plugins = PluginRuntime( + tools, + # 当前 Python Host 尚无 OS 沙箱。生产构建必须保持关闭,直到 + # Tauri/Rust Host 能签发绑定命令摘要的可信启动许可。 + allow_unsandboxed_mcp=settings.environment == "development", + ) plugins.install(BACKEND_DIR / "extensions" / "plugins" / "text-tools") plugins.enable("text-tools") diff --git a/backend/app/extensions/runtime.py b/backend/app/extensions/runtime.py index cacbb12..832b726 100644 --- a/backend/app/extensions/runtime.py +++ b/backend/app/extensions/runtime.py @@ -256,10 +256,13 @@ class PluginRuntime: tools: ToolRegistry, host: DeclarativePluginHost | None = None, mcp_bridge: McpBridge | None = None, + *, + allow_unsandboxed_mcp: bool = False, ) -> None: self.registry = tools self.host = host or DeclarativePluginHost() self.mcp = mcp_bridge or McpBridge() + self.allow_unsandboxed_mcp = allow_unsandboxed_mcp self._records: dict[str, _PluginRecord] = {} self._lock = threading.RLock() @@ -346,6 +349,16 @@ class PluginRuntime: status_code=409, details={"plugin_id": plugin_id, "permissions": missing_grants}, ) + if ( + record.plugin.manifest.backend.type == "mcp" + and not self.allow_unsandboxed_mcp + ): + raise ExtensionError( + "MCP_TRUST_APPROVAL_REQUIRED", + "Unsandboxed MCP Hosts are disabled outside development mode.", + status_code=403, + details={"plugin_id": plugin_id}, + ) declared_tools = list(record.plugin.manifest.contributes.tools) conflicts = [name for name in declared_tools if self.registry.contains(name)] if conflicts: @@ -665,27 +678,10 @@ def _arguments_model_from_schema( ) -> type[BaseModel]: if schema.get("type", "object") != "object": raise ExtensionError("PLUGIN_TOOL_SCHEMA_INVALID", "Tool parameters must be an object schema.") - properties = schema.get("properties", {}) - required = set(schema.get("required", [])) - fields: dict[str, tuple[Any, Any]] = {} - types = { - "string": str, - "number": float, - "integer": int, - "boolean": bool, - "array": list[Any], - "object": dict[str, Any], - } - for name, field_schema in properties.items(): - schema_type = field_schema.get("type") - # JSON Schema 允许联合类型数组;复杂类型继续由 Draft Validator - # 精确校验,Pydantic 在这里只承担参数载体职责。 - annotation = types.get(schema_type, Any) if isinstance(schema_type, str) else Any - fields[name] = (annotation, ... if name in required else None) model_name = "PluginArgs_" + re.sub(r"\W+", "_", tool_name) - # 完整 JSON Schema 已在 ToolRegistry 中先行校验。这里允许额外字段,避免 - # Pydantic 再次拒绝 additionalProperties/patternProperties 接受的合法参数。 - return create_model(model_name, __config__=ConfigDict(extra="allow"), **fields) + # 完整 JSON Schema 已在 ToolRegistry 中先行校验。参数载体不重复声明字段, + # 从而完整保留 model_dump、连字符键、联合类型和动态属性等合法 JSON 键值。 + return create_model(model_name, __config__=ConfigDict(extra="allow")) def _validate_tool_schema(spec: DeclarativeToolSpec) -> None: diff --git a/backend/tests/test_extension_core.py b/backend/tests/test_extension_core.py index 5b368aa..29a2d92 100644 --- a/backend/tests/test_extension_core.py +++ b/backend/tests/test_extension_core.py @@ -509,13 +509,38 @@ def test_mcp_argument_model_preserves_json_schema_additional_properties() -> Non "mcp-fixture.dynamic", { "type": "object", + "properties": {"model_dump": {"type": "string"}}, + "required": ["model_dump"], "additionalProperties": {"type": "string"}, }, ) - arguments = arguments_model.model_validate({"dynamic_key": "value"}) + arguments = arguments_model.model_validate( + {"model_dump": "method name remains data", "dynamic-key": "value"} + ) - assert arguments.model_dump() == {"dynamic_key": "value"} + assert arguments.model_dump() == { + "model_dump": "method name remains data", + "dynamic-key": "value", + } + + +def test_production_rejects_unsandboxed_mcp_host(monkeypatch) -> None: + monkeypatch.setenv("APP_ENVIRONMENT", "production") + get_settings.cache_clear() + container = build_container() + installed = container.plugins.install(MCP_FIXTURE) + assert installed.status == "permission_required" + container.plugins.set_permissions("mcp-fixture", ["notes.read"]) + try: + with pytest.raises(ExtensionError) as exc: + container.plugins.enable("mcp-fixture") + assert exc.value.code == "MCP_TRUST_APPROVAL_REQUIRED" + assert container.plugins.get_host_status("mcp-fixture").status == "stopped" + assert not container.tools.contains("mcp-fixture.echo") + finally: + container.plugins.shutdown() + get_settings.cache_clear() def test_mcp_abnormal_exit_unregisters_tools_and_restart_recovers(mcp_container) -> None: