fix(mcp): 修复配置导入、凭据管理与协议边界
修复生命周期锁阻塞事件循环、旧连接回调误停新连接及超时契约不一致。 补齐 MCP JSON 兼容导入、密钥拆分与失败重试,修复 Header 大小写草稿丢失,迁移大小写敏感的环境变量凭据。 在 SSE 行拼接前限制缓冲大小,增加并发、迁移和流式输入回归测试;忽略本机 MCP 数据及 server.json/servers.json。 验证:后端 185 项、前端 54 项测试通过,前端生产构建、相关文件 Ruff 与暂存差异检查通过。
This commit is contained in:
@@ -10,6 +10,7 @@ import asyncio
|
||||
import json
|
||||
import os
|
||||
import queue
|
||||
import re
|
||||
import signal
|
||||
import subprocess
|
||||
import threading
|
||||
@@ -49,6 +50,7 @@ MAX_MCP_MESSAGE_BYTES = 2 * 1024 * 1024
|
||||
MAX_MCP_TOOL_RESULT_BYTES = 256 * 1024
|
||||
MAX_MCP_TOOLS = 500
|
||||
MAX_MCP_LIST_PAGES = 100
|
||||
_SSE_NEWLINE = re.compile(rb"\r\n?|\n")
|
||||
|
||||
|
||||
class McpBridgeError(RuntimeError):
|
||||
@@ -1019,7 +1021,8 @@ class McpBridge:
|
||||
if status.status == PluginHostState.unhealthy:
|
||||
raise McpBridgeError(
|
||||
"PLUGIN_HOST_UNAVAILABLE",
|
||||
status.error or "MCP event stream became unavailable during startup.",
|
||||
status.error
|
||||
or "MCP event stream became unavailable during startup.",
|
||||
status_code=503,
|
||||
)
|
||||
status.status = PluginHostState.ready
|
||||
@@ -1383,12 +1386,58 @@ def _bounded_json_response(response: httpx.Response) -> dict[str, Any]:
|
||||
return payload
|
||||
|
||||
|
||||
def _bounded_sse_lines(response: httpx.Response):
|
||||
"""Split UTF-8 lines without httpx.iter_lines()'s unbounded line buffer.
|
||||
|
||||
Check each segment before appending it, including partial/no-newline input.
|
||||
SSE allows LF, CR and CRLF; a CRLF pair can span network chunks.
|
||||
"""
|
||||
|
||||
pending = bytearray()
|
||||
event_size = 0
|
||||
skip_lf = False
|
||||
first_line = True
|
||||
for chunk in response.iter_bytes():
|
||||
offset = 0
|
||||
if skip_lf and chunk:
|
||||
offset = int(chunk.startswith(b"\n"))
|
||||
skip_lf = False
|
||||
for match in _SSE_NEWLINE.finditer(chunk, offset):
|
||||
start, end = match.span()
|
||||
segment = memoryview(chunk)[offset:start]
|
||||
if event_size + len(pending) + len(segment) + 1 > MAX_MCP_MESSAGE_BYTES:
|
||||
raise McpBridgeError(
|
||||
"MCP_HTTP_RESPONSE_INVALID", "MCP SSE event is too large."
|
||||
)
|
||||
pending.extend(segment)
|
||||
line = pending.decode("utf-8", errors="replace")
|
||||
event_size += len(pending) + 1
|
||||
pending.clear()
|
||||
if first_line:
|
||||
line = line.removeprefix("\ufeff")
|
||||
first_line = False
|
||||
if not line:
|
||||
event_size = 0
|
||||
yield line
|
||||
skip_lf = chunk[end - 1 : end] == b"\r" and end == len(chunk)
|
||||
offset = end
|
||||
tail = memoryview(chunk)[offset:]
|
||||
if event_size + len(pending) + len(tail) > MAX_MCP_MESSAGE_BYTES:
|
||||
raise McpBridgeError(
|
||||
"MCP_HTTP_RESPONSE_INVALID", "MCP SSE event is too large."
|
||||
)
|
||||
pending.extend(tail)
|
||||
if pending:
|
||||
line = pending.decode("utf-8", errors="replace")
|
||||
yield line.removeprefix("\ufeff") if first_line else line
|
||||
|
||||
|
||||
def _iter_sse(response: httpx.Response):
|
||||
event = "message"
|
||||
event_id: str | None = None
|
||||
data_lines: list[str] = []
|
||||
size = 0
|
||||
for line in response.iter_lines():
|
||||
for line in _bounded_sse_lines(response):
|
||||
size += len(line.encode("utf-8")) + 1
|
||||
if size > MAX_MCP_MESSAGE_BYTES:
|
||||
raise McpBridgeError(
|
||||
|
||||
@@ -9,7 +9,7 @@ import threading
|
||||
from datetime import UTC, datetime
|
||||
from functools import wraps
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, Literal
|
||||
from urllib.parse import urlsplit
|
||||
from uuid import uuid4
|
||||
|
||||
@@ -45,10 +45,22 @@ _SERVER_ID = re.compile(r"^[A-Za-z0-9][A-Za-z0-9_-]{0,63}$")
|
||||
_MAX_MCP_SERVERS = 256
|
||||
|
||||
|
||||
class _McpConnectionBackend(PluginBackend):
|
||||
"""Bridge adapter for the independent server's float timeout contract.
|
||||
|
||||
Plugin manifests retain their integer/60-second startup restrictions.
|
||||
Reusing that validation here used to reject valid 120-second server configs.
|
||||
"""
|
||||
|
||||
startup_timeout_seconds: float = Field(default=15, ge=1, le=120)
|
||||
tool_timeout_seconds: float = Field(default=30, ge=1, le=300)
|
||||
|
||||
|
||||
class _McpServerRecord(McpServerConfig):
|
||||
"""Validated on-disk representation with defaults for older C.1 records."""
|
||||
|
||||
version: int = Field(default=1, ge=1)
|
||||
secret_environment_version: Literal[1, 2] = 1
|
||||
enabled: bool = False
|
||||
approved_digest: str | None = Field(
|
||||
default=None, min_length=64, max_length=64, pattern=r"^[0-9a-f]{64}$"
|
||||
@@ -102,6 +114,8 @@ class McpServerRegistry:
|
||||
self._registered: dict[str, list[str]] = {}
|
||||
self._summaries: dict[str, list[McpToolSummary]] = {}
|
||||
self._last_status: dict[str, dict[str, Any]] = {}
|
||||
self._generations: dict[str, object] = {}
|
||||
self._migrate_environment_secrets()
|
||||
|
||||
def list(self) -> list[McpServer]:
|
||||
with self._lock:
|
||||
@@ -137,6 +151,7 @@ class McpServerRegistry:
|
||||
record["url"] = request.url.strip() if request.url else None
|
||||
record.update(
|
||||
version=1,
|
||||
secret_environment_version=2,
|
||||
enabled=False,
|
||||
approved_digest=None,
|
||||
tested_digest=None,
|
||||
@@ -163,7 +178,7 @@ class McpServerRegistry:
|
||||
with self._lock:
|
||||
previous = self._record(server_id)
|
||||
removed_secret_ids = [
|
||||
self._secret_id(server_id, key, kind)
|
||||
secret_id
|
||||
for kind, old_keys, new_keys in (
|
||||
(
|
||||
"environment",
|
||||
@@ -176,7 +191,8 @@ class McpServerRegistry:
|
||||
request.secret_header_keys,
|
||||
),
|
||||
)
|
||||
for key in set(old_keys) - set(new_keys)
|
||||
for secret_id in self._secret_ids(server_id, old_keys, kind)
|
||||
- self._secret_ids(server_id, new_keys, kind)
|
||||
]
|
||||
try:
|
||||
self.credentials.delete_many(removed_secret_ids)
|
||||
@@ -191,6 +207,7 @@ class McpServerRegistry:
|
||||
record["url"] = request.url.strip() if request.url else None
|
||||
record.update(
|
||||
version=request.version + 1,
|
||||
secret_environment_version=2,
|
||||
enabled=False,
|
||||
approved_digest=None,
|
||||
tested_digest=None,
|
||||
@@ -210,12 +227,12 @@ class McpServerRegistry:
|
||||
with self._lock:
|
||||
record = self._record(server_id)
|
||||
secret_ids = [
|
||||
self._secret_id(server_id, key, kind)
|
||||
secret_id
|
||||
for kind, keys in (
|
||||
("environment", record.get("secret_environment_keys", [])),
|
||||
("header", record.get("secret_header_keys", [])),
|
||||
)
|
||||
for key in keys
|
||||
for secret_id in self._secret_ids(server_id, keys, kind)
|
||||
]
|
||||
try:
|
||||
self.credentials.delete_many(secret_ids)
|
||||
@@ -307,6 +324,8 @@ class McpServerRegistry:
|
||||
try:
|
||||
discovered = self._start(server_id, record)
|
||||
except Exception as exc:
|
||||
self._generations.pop(server_id, None)
|
||||
self.bridge.remove(self._host_id(server_id))
|
||||
tested_at = datetime.now(UTC)
|
||||
failure = {
|
||||
"status": PluginHostState.error,
|
||||
@@ -339,6 +358,7 @@ class McpServerRegistry:
|
||||
"last_test_succeeded": True,
|
||||
}
|
||||
self._summaries[server_id] = self._tool_summaries(discovered)
|
||||
self._generations.pop(server_id, None)
|
||||
self.bridge.stop(self._host_id(server_id))
|
||||
with self._lock:
|
||||
tested_record = {
|
||||
@@ -368,6 +388,7 @@ class McpServerRegistry:
|
||||
except Exception:
|
||||
for name in registered:
|
||||
self.tools.unregister(name)
|
||||
self._generations.pop(server_id, None)
|
||||
self.bridge.stop(self._host_id(server_id))
|
||||
raise
|
||||
try:
|
||||
@@ -380,6 +401,7 @@ class McpServerRegistry:
|
||||
except McpRegistryError:
|
||||
for name in registered:
|
||||
self.tools.unregister(name)
|
||||
self._generations.pop(server_id, None)
|
||||
self.bridge.stop(self._host_id(server_id))
|
||||
raise
|
||||
return self.get(server_id)
|
||||
@@ -394,6 +416,7 @@ class McpServerRegistry:
|
||||
self._records = updated
|
||||
for name in self._registered.pop(server_id, []):
|
||||
self.tools.unregister(name)
|
||||
self._generations.pop(server_id, None)
|
||||
self.bridge.stop(self._host_id(server_id))
|
||||
return self.get(server_id)
|
||||
|
||||
@@ -416,6 +439,7 @@ class McpServerRegistry:
|
||||
@_serialized_lifecycle
|
||||
def shutdown(self) -> None:
|
||||
for server_id in list(self._records):
|
||||
self._generations.pop(server_id, None)
|
||||
for name in self._registered.pop(server_id, []):
|
||||
self.tools.unregister(name)
|
||||
self.bridge.stop(self._host_id(server_id))
|
||||
@@ -456,6 +480,9 @@ class McpServerRegistry:
|
||||
)
|
||||
headers[key] = value
|
||||
host_id = self._host_id(server_id)
|
||||
# A queued callback from the previous process must not affect its replacement.
|
||||
generation = object()
|
||||
self._generations[server_id] = generation
|
||||
self.bridge.remove(host_id)
|
||||
try:
|
||||
return self.bridge.start(
|
||||
@@ -463,7 +490,9 @@ class McpServerRegistry:
|
||||
self._backend(record),
|
||||
self._server_dir(server_id),
|
||||
list(record.get("permissions", [])),
|
||||
lambda _host, message: self._unavailable(server_id, message),
|
||||
lambda _host, message: self._unavailable(
|
||||
server_id, generation, message
|
||||
),
|
||||
command_override=(
|
||||
[record["command"], *record.get("args", [])]
|
||||
if record.get("command")
|
||||
@@ -476,6 +505,8 @@ class McpServerRegistry:
|
||||
headers=headers,
|
||||
)
|
||||
except McpBridgeError as exc:
|
||||
self._generations.pop(server_id, None)
|
||||
self.bridge.remove(host_id)
|
||||
raise McpRegistryError(
|
||||
exc.code, exc.message, status_code=exc.status_code
|
||||
) from exc
|
||||
@@ -496,10 +527,13 @@ class McpServerRegistry:
|
||||
|
||||
self.tools.register(definition, arguments_model, executor)
|
||||
|
||||
def _unavailable(self, server_id: str, message: str) -> None:
|
||||
def _unavailable(self, server_id: str, generation: object, message: str) -> None:
|
||||
# A failure may race with enable(). Waiting for the lifecycle mutation makes
|
||||
# sure tools registered immediately before the callback are also removed.
|
||||
with self._lifecycle_lock:
|
||||
if self._generations.get(server_id) is not generation:
|
||||
return
|
||||
self._generations.pop(server_id, None)
|
||||
try:
|
||||
with self._lock:
|
||||
record = self._records.get(server_id)
|
||||
@@ -694,7 +728,7 @@ class McpServerRegistry:
|
||||
|
||||
@staticmethod
|
||||
def _backend(record: dict[str, Any]) -> PluginBackend:
|
||||
return PluginBackend(
|
||||
return _McpConnectionBackend(
|
||||
type="mcp",
|
||||
transport="stdio",
|
||||
command=record.get("command") or "http",
|
||||
@@ -755,9 +789,89 @@ class McpServerRegistry:
|
||||
|
||||
@staticmethod
|
||||
def _secret_id(server_id: str, key: str, kind: str = "environment") -> str:
|
||||
suffix = hashlib.sha256(f"{kind}\0{key.casefold()}".encode()).hexdigest()[:20]
|
||||
identity = (
|
||||
f"environment-v2\0{key}"
|
||||
if kind == "environment"
|
||||
else f"{kind}\0{key.casefold()}"
|
||||
)
|
||||
suffix = hashlib.sha256(identity.encode()).hexdigest()[:20]
|
||||
return f"mcp.{server_id}.{suffix}"
|
||||
|
||||
@staticmethod
|
||||
def _legacy_environment_secret_id(server_id: str, key: str) -> str:
|
||||
suffix = hashlib.sha256(f"environment\0{key.casefold()}".encode()).hexdigest()[
|
||||
:20
|
||||
]
|
||||
return f"mcp.{server_id}.{suffix}"
|
||||
|
||||
def _secret_ids(self, server_id: str, keys: list[str], kind: str) -> set[str]:
|
||||
ids = {self._secret_id(server_id, key, kind) for key in keys}
|
||||
if kind == "environment":
|
||||
# Include retained ambiguous legacy ciphertext when its last declaration is removed.
|
||||
ids.update(
|
||||
self._legacy_environment_secret_id(server_id, key) for key in keys
|
||||
)
|
||||
return ids
|
||||
|
||||
def _migrate_environment_secrets(self) -> None:
|
||||
"""迁移旧的大小写折叠 ID;已碰撞的值无法恢复,保留原密文并要求重新录入。"""
|
||||
|
||||
replacements: dict[str, str] = {}
|
||||
ambiguous: dict[str, list[str]] = {}
|
||||
legacy_records = {
|
||||
server_id: record
|
||||
for server_id, record in self._records.items()
|
||||
if record.get("secret_environment_version", 1) == 1
|
||||
}
|
||||
if not legacy_records:
|
||||
return
|
||||
for server_id, record in legacy_records.items():
|
||||
groups: dict[str, set[str]] = {}
|
||||
for key in record.get("secret_environment_keys", []):
|
||||
groups.setdefault(key.casefold(), set()).add(key)
|
||||
for keys in groups.values():
|
||||
key = next(iter(keys))
|
||||
legacy_id = self._legacy_environment_secret_id(server_id, key)
|
||||
if len(keys) == 1:
|
||||
replacements[legacy_id] = self._secret_id(server_id, key)
|
||||
else:
|
||||
ambiguous.setdefault(server_id, []).extend(keys)
|
||||
try:
|
||||
self.credentials.move_many(replacements)
|
||||
for server_id, keys in ambiguous.items():
|
||||
if not any(
|
||||
self.credentials.has(
|
||||
self._legacy_environment_secret_id(server_id, key)
|
||||
)
|
||||
for key in keys
|
||||
):
|
||||
continue
|
||||
if all(
|
||||
self.credentials.has(self._secret_id(server_id, key))
|
||||
for key in keys
|
||||
):
|
||||
continue
|
||||
self._records[server_id].update(
|
||||
enabled=False,
|
||||
tested_digest=None,
|
||||
last_test_succeeded=None,
|
||||
last_tested_at=None,
|
||||
)
|
||||
self._last_status[server_id] = {
|
||||
"status": PluginHostState.error,
|
||||
"error": "环境变量密钥名称曾发生大小写冲突,请分别重新录入密钥并测试连接。",
|
||||
}
|
||||
# Persist a migration marker even when legacy values were ambiguous.
|
||||
# Otherwise a later key removal could make that old shared value look
|
||||
# unambiguous and resurrect a deleted credential on the next restart.
|
||||
for server_id in legacy_records:
|
||||
self._records[server_id]["secret_environment_version"] = 2
|
||||
self._write()
|
||||
except CredentialStoreError as exc:
|
||||
raise McpRegistryError(
|
||||
"MCP_SECRET_STORE_ERROR", str(exc), status_code=500
|
||||
) from exc
|
||||
|
||||
def _secret_configured(
|
||||
self, server_id: str, key: str, kind: str = "environment"
|
||||
) -> bool:
|
||||
@@ -857,7 +971,9 @@ class McpServerRegistry:
|
||||
config_fields = set(McpServerConfig.model_fields)
|
||||
try:
|
||||
for server_id, raw in value.items():
|
||||
if not isinstance(server_id, str) or not _SERVER_ID.fullmatch(server_id):
|
||||
if not isinstance(server_id, str) or not _SERVER_ID.fullmatch(
|
||||
server_id
|
||||
):
|
||||
raise ValueError("invalid server id")
|
||||
record = _McpServerRecord.model_validate(raw)
|
||||
config = record.model_dump(mode="json", include=config_fields)
|
||||
|
||||
Reference in New Issue
Block a user