feat(mcp): complete remote transports and configuration workflow
This commit is contained in:
+42
-17
@@ -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"
|
||||
|
||||
+685
-42
@@ -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()
|
||||
|
||||
@@ -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]
|
||||
|
||||
+120
-40
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user