fix(mcp): 修复配置导入、凭据管理与协议边界

修复生命周期锁阻塞事件循环、旧连接回调误停新连接及超时契约不一致。

补齐 MCP JSON 兼容导入、密钥拆分与失败重试,修复 Header 大小写草稿丢失,迁移大小写敏感的环境变量凭据。

在 SSE 行拼接前限制缓冲大小,增加并发、迁移和流式输入回归测试;忽略本机 MCP 数据及 server.json/servers.json。

验证:后端 185 项、前端 54 项测试通过,前端生产构建、相关文件 Ruff 与暂存差异检查通过。
This commit is contained in:
2026-09-03 22:38:35 +08:00
parent 2f7066aa92
commit 9b8b10cdb1
13 changed files with 1097 additions and 80 deletions
+51 -2
View File
@@ -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(
+126 -10
View File
@@ -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)
+35 -9
View File
@@ -5,13 +5,12 @@ import os
import re
import threading
from pathlib import Path
from typing import Protocol
from typing import ClassVar, Protocol
from cryptography.fernet import Fernet, InvalidToken
from app.config import get_settings
_CREDENTIAL_ID = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._-]{0,127}$")
_PLUGIN_CREDENTIAL_PREFIX = "plugin."
_MCP_CREDENTIAL_PREFIX = "mcp."
@@ -29,7 +28,9 @@ def validate_provider_credential_id(credential_id: str | None) -> None:
"""阻止 Provider 和通用凭据 API 跨入 Plugin 私有命名空间。"""
if credential_id and credential_id.casefold().startswith(_PLUGIN_CREDENTIAL_PREFIX):
raise CredentialStoreError("Credential namespace is reserved for Plugin settings.")
raise CredentialStoreError(
"Credential namespace is reserved for Plugin settings."
)
if credential_id and credential_id.casefold().startswith(_MCP_CREDENTIAL_PREFIX):
raise CredentialStoreError("Credential namespace is reserved for MCP settings.")
@@ -37,7 +38,7 @@ def validate_provider_credential_id(credential_id: str | None) -> None:
class EnvironmentCredentialResolver:
"""解析由桌面 Host 注入 Sidecar 进程的临时凭证上下文。"""
_development_aliases = {
_development_aliases: ClassVar[dict[str, str]] = {
"openai": "OPENAI_API_KEY",
"deepseek": "DEEPSEEK_API_KEY",
}
@@ -85,7 +86,9 @@ class EncryptedCredentialStore:
try:
return Fernet(environment_key.encode("ascii"))
except (ValueError, UnicodeEncodeError) as exc:
raise CredentialStoreError("APP_CREDENTIAL_MASTER_KEY is invalid.") from exc
raise CredentialStoreError(
"APP_CREDENTIAL_MASTER_KEY is invalid."
) from exc
key_path.parent.mkdir(parents=True, exist_ok=True)
self._restrict(key_path.parent, 0o700)
@@ -102,7 +105,9 @@ class EncryptedCredentialStore:
try:
return Fernet(key_path.read_bytes().strip())
except (OSError, ValueError) as exc:
raise CredentialStoreError("Credential master key cannot be loaded.") from exc
raise CredentialStoreError(
"Credential master key cannot be loaded."
) from exc
def _read_tokens(self) -> dict[str, str]:
_, store_path = self._paths()
@@ -111,11 +116,16 @@ class EncryptedCredentialStore:
try:
data = json.loads(store_path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError) as exc:
raise CredentialStoreError("Encrypted credential store cannot be loaded.") from exc
raise CredentialStoreError(
"Encrypted credential store cannot be loaded."
) from exc
if not isinstance(data, dict) or not all(
isinstance(key, str) and isinstance(value, str) for key, value in data.items()
isinstance(key, str) and isinstance(value, str)
for key, value in data.items()
):
raise CredentialStoreError("Encrypted credential store has an invalid format.")
raise CredentialStoreError(
"Encrypted credential store has an invalid format."
)
return data
def _write_tokens(self, tokens: dict[str, str]) -> None:
@@ -196,6 +206,22 @@ class EncryptedCredentialStore:
self._write_tokens(tokens)
return removed
def move_many(self, replacements: dict[str, str]) -> None:
"""原子迁移凭据 ID,直接移动密文且不覆盖已经写入的新凭据。"""
for old_id, new_id in replacements.items():
self._validate_id(old_id)
self._validate_id(new_id)
with self._lock:
tokens = self._read_tokens()
changed = False
for old_id, new_id in replacements.items():
if old_id != new_id and old_id in tokens:
tokens.setdefault(new_id, tokens.pop(old_id))
changed = True
if changed:
self._write_tokens(tokens)
class ChainedCredentialResolver:
def __init__(self, *resolvers: CredentialResolver) -> None:
+6 -13
View File
@@ -100,15 +100,8 @@ from app.services import (
router = APIRouter(prefix="/api")
def mcp_call(operation):
try:
return operation()
except McpRegistryError as exc:
raise ApiError(exc.status_code, exc.code, exc.message) from exc
async def mcp_call_async(operation):
"""MCP process operations wait on stdio and must not block the API event loop."""
"""Even registry reads can wait on lifecycle locks; keep all MCP work off the event loop."""
try:
return await asyncio.to_thread(operation)
except McpRegistryError as exc:
@@ -521,19 +514,19 @@ async def uninstall_skill(skill_id: str) -> OperationResponse:
# Independent MCP Server Registry
@router.get("/mcp/servers", response_model=McpServerListResponse, tags=["MCP Servers"])
async def list_mcp_servers() -> McpServerListResponse:
return McpServerListResponse(items=mcp_call(container.mcp_servers.list))
return McpServerListResponse(items=await mcp_call_async(container.mcp_servers.list))
@router.post(
"/mcp/servers", response_model=McpServer, status_code=201, tags=["MCP Servers"]
)
async def create_mcp_server(request: McpServerCreateRequest) -> McpServer:
return mcp_call(lambda: container.mcp_servers.create(request))
return await mcp_call_async(lambda: container.mcp_servers.create(request))
@router.get("/mcp/servers/{server_id}", response_model=McpServer, tags=["MCP Servers"])
async def get_mcp_server(server_id: str) -> McpServer:
return mcp_call(lambda: container.mcp_servers.get(server_id))
return await mcp_call_async(lambda: container.mcp_servers.get(server_id))
@router.get(
@@ -543,7 +536,7 @@ async def get_mcp_server(server_id: str) -> McpServer:
)
async def list_mcp_server_tools(server_id: str) -> McpToolSummaryListResponse:
return McpToolSummaryListResponse(
items=mcp_call(lambda: container.mcp_servers.list_tools(server_id))
items=await mcp_call_async(lambda: container.mcp_servers.list_tools(server_id))
)
@@ -570,7 +563,7 @@ async def delete_mcp_server(server_id: str) -> OperationResponse:
"/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(
return await mcp_call_async(
lambda: container.mcp_servers.trust(server_id, request.command_digest)
)
+180 -6
View File
@@ -2,6 +2,8 @@ import asyncio
import threading
from types import SimpleNamespace
import pytest
from app.contracts import (
McpServerSecretStatus,
McpServerSecretWriteRequest,
@@ -61,9 +63,7 @@ def test_mcp_secret_routes_offload_blocking_lifecycle_work(monkeypatch) -> None:
)
)
deleted = asyncio.run(
routes.delete_mcp_server_secret(
"server-1", "TOKEN", kind="environment"
)
routes.delete_mcp_server_secret("server-1", "TOKEN", kind="environment")
)
assert written.configured is True
@@ -77,6 +77,172 @@ def test_health() -> None:
assert response.model_dump() == {"status": "ok"}
def test_mcp_create_and_trust_are_not_executed_on_event_loop(monkeypatch) -> None:
from app import routes
from app.contracts import McpServerCreateRequest, McpServerTrustRequest
caller = threading.get_ident()
workers = []
class Registry:
def create(self, request):
workers.append(threading.get_ident())
return "created"
def trust(self, server_id, digest):
workers.append(threading.get_ident())
return "trusted"
monkeypatch.setattr(routes, "container", SimpleNamespace(mcp_servers=Registry()))
assert (
asyncio.run(
routes.create_mcp_server(McpServerCreateRequest(name="test", command="uvx"))
)
== "created"
)
assert (
asyncio.run(
routes.trust_mcp_server(
"test", McpServerTrustRequest(command_digest="a" * 64)
)
)
== "trusted"
)
assert len(workers) == 2
assert all(worker != caller for worker in workers)
def test_mcp_split_config_and_secret_requests_persist_without_plaintext(
monkeypatch,
) -> None:
from fastapi.testclient import TestClient
from app import routes
from app.agent.tools import ToolRegistry
from app.config import get_settings
from app.extensions.mcp_registry import McpServerRegistry
from app.main import app
from app.providers.credentials import EncryptedCredentialStore
service = McpServerRegistry(
ToolRegistry(),
EncryptedCredentialStore(),
get_settings().data_dir,
allow_process_launch=True,
)
monkeypatch.setattr(routes, "container", SimpleNamespace(mcp_servers=service))
client = TestClient(app)
config = {
"name": "MiniMax configuration test",
"command": "uvx",
"environment": {"MINIMAX_API_HOST": "https://api.minimaxi.com"},
"secret_environment_keys": ["MINIMAX_API_KEY"],
"startup_timeout_seconds": 120,
"tool_timeout_seconds": 300,
}
# Reproduce the old frontend payload. The backend still enforces separation.
invalid = client.post(
"/api/mcp/servers",
json={
**config,
"environment": {
**config["environment"],
"MINIMAX_API_KEY": "synthetic-only",
},
},
)
assert invalid.status_code == 422
assert invalid.json()["error"]["code"] == "MCP_ENVIRONMENT_INVALID"
created = client.post("/api/mcp/servers", json=config)
assert created.status_code == 201
server_id = created.json()["server_id"]
saved = client.put(
f"/api/mcp/servers/{server_id}/secrets/MINIMAX_API_KEY",
json={"secret": "synthetic-only"},
)
assert saved.status_code == 200
current = client.get(f"/api/mcp/servers/{server_id}")
assert current.json()["secret_environment"] == {"MINIMAX_API_KEY": True}
assert "synthetic-only" not in current.text
assert "synthetic-only" not in service._path.read_text(encoding="utf-8")
_, credentials_path = service.credentials._paths()
assert "synthetic-only" not in credentials_path.read_text(encoding="utf-8")
assert not current.json()["enabled"] # Saving never starts a third-party process.
client.close()
@pytest.mark.parametrize("operation", ["create", "trust"])
def test_mcp_lifecycle_lock_contention_keeps_event_loop_responsive(
monkeypatch,
operation,
) -> None:
from app import routes
from app.agent.tools import ToolRegistry
from app.config import get_settings
from app.contracts import McpServerCreateRequest, McpServerTrustRequest
from app.extensions.mcp_registry import McpServerRegistry
from app.providers.credentials import EncryptedCredentialStore
service = McpServerRegistry(
ToolRegistry(),
EncryptedCredentialStore(),
get_settings().data_dir,
allow_process_launch=True,
)
request = McpServerCreateRequest(
name="Lock contention fixture", command="not-executed"
)
server = service.create(request)
monkeypatch.setattr(routes, "container", SimpleNamespace(mcp_servers=service))
entered = threading.Event()
locked = threading.Event()
release = threading.Event()
original = getattr(service, operation)
def observed(*args):
entered.set()
return original(*args)
def hold_lifecycle_lock():
with service._lifecycle_lock:
locked.set()
release.wait(timeout=5)
monkeypatch.setattr(service, operation, observed)
holder = threading.Thread(target=hold_lifecycle_lock, daemon=True)
holder.start()
# An independent watchdog lets the test fail rather than hang if a regression
# blocks the event loop itself (an asyncio timeout alone cannot catch that).
watchdog = threading.Timer(5, release.set)
watchdog.start()
async def exercise():
pending = asyncio.create_task(
routes.create_mcp_server(request)
if operation == "create"
else routes.trust_mcp_server(
server.server_id,
McpServerTrustRequest(command_digest=server.command_digest),
)
)
try:
assert await asyncio.to_thread(entered.wait, 2)
assert not pending.done()
assert not release.is_set()
assert (await health()).status == "ok"
finally:
release.set()
await pending
try:
assert locked.wait(timeout=2)
asyncio.run(exercise())
finally:
release.set()
watchdog.cancel()
holder.join(timeout=2)
def test_service_status() -> None:
response = asyncio.run(service_status())
@@ -93,7 +259,9 @@ def test_core_collections_are_typed() -> None:
assert notes.items == []
assert notes.page.limit == 20
assert [skill.manifest.skill_id for skill in skills.items] == ["knowledge-assistant"]
assert [skill.manifest.skill_id for skill in skills.items] == [
"knowledge-assistant"
]
assert skills.items[0].status == "ready"
assert [plugin.manifest.plugin_id for plugin in plugins.items] == ["text-tools"]
assert plugins.items[0].status == "ready"
@@ -114,9 +282,15 @@ def test_provider_presets_include_openai_and_deepseek() -> None:
def test_provider_presets_static_route_precedes_provider_id_route() -> None:
from app.routes import router
get_paths = [route.path for route in router.routes if "GET" in getattr(route, "methods", set())]
get_paths = [
route.path
for route in router.routes
if "GET" in getattr(route, "methods", set())
]
assert get_paths.index("/api/providers/presets") < get_paths.index("/api/providers/{provider_id}")
assert get_paths.index("/api/providers/presets") < get_paths.index(
"/api/providers/{provider_id}"
)
def test_openapi_contains_documented_frontend_interfaces() -> None:
+212 -7
View File
@@ -1,4 +1,5 @@
import asyncio
import hashlib
import json
import sys
import threading
@@ -111,7 +112,9 @@ def test_update_disables_server_and_revokes_command_trust() -> None:
service.shutdown()
def test_update_remains_retryable_when_removed_secret_cleanup_fails(monkeypatch) -> None:
def test_update_remains_retryable_when_removed_secret_cleanup_fails(
monkeypatch,
) -> None:
service = registry()
created = service.create(request())
service.put_secret(created.server_id, "TEST_MCP_SECRET", "keep-until-retry")
@@ -166,7 +169,9 @@ def test_unavailable_server_removes_bridge_host(monkeypatch) -> None:
removed: list[str] = []
monkeypatch.setattr(service.bridge, "remove", removed.append)
service._unavailable(created.server_id, "connection lost")
generation = object()
service._generations[created.server_id] = generation
service._unavailable(created.server_id, generation, "connection lost")
current = service.get(created.server_id)
assert removed == [f"mcp.{created.server_id}"]
@@ -175,6 +180,180 @@ def test_unavailable_server_removes_bridge_host(monkeypatch) -> None:
service.shutdown()
def test_old_failure_callback_cannot_stop_replacement_host(monkeypatch) -> None:
service = registry()
callbacks = []
original_start = service.bridge.start
def capture_callback(*args, **kwargs):
callbacks.append(args[4])
return original_start(*args, **kwargs)
monkeypatch.setattr(service.bridge, "start", capture_callback)
created = service.create(request(secret_environment_keys=[]))
service.trust(created.server_id, created.command_digest)
callback_thread = None
try:
service.test(created.server_id)
service.enable(created.server_id)
old_callback = callbacks[-1]
callback_started = threading.Event()
callback_finished = threading.Event()
def delayed_failure():
callback_started.set()
old_callback(f"mcp.{created.server_id}", "delayed old failure")
callback_finished.set()
# Queue the old callback while a replacement owns the lifecycle lock.
with service._lifecycle_lock:
callback_thread = threading.Thread(target=delayed_failure, daemon=True)
callback_thread.start()
assert callback_started.wait(timeout=2)
service.disable(created.server_id)
service.enable(created.server_id)
assert callback_finished.wait(timeout=2)
assert service.get(created.server_id).enabled is True
assert service.get(created.server_id).status == "ready"
assert service.tools.definitions()
callbacks[-1](f"mcp.{created.server_id}", "current failure")
assert service.get(created.server_id).enabled is False
assert service.get(created.server_id).status == "unhealthy"
finally:
service.shutdown()
if callback_thread is not None:
callback_thread.join(timeout=2)
def test_header_case_only_rename_preserves_secret() -> None:
service = registry()
config = {
"name": "HTTP",
"transport": "streamable_http",
"url": "https://example.test/mcp",
"secret_header_keys": ["Authorization"],
}
created = service.create(McpServerCreateRequest(**config))
service.put_secret(created.server_id, "Authorization", "synthetic", kind="header")
config["secret_header_keys"] = ["authorization"]
updated = service.update(
created.server_id, McpServerUpdateRequest(**config, version=created.version)
)
assert updated.secret_headers == {"authorization": True}
assert (
service.credentials.resolve(
service._secret_id(created.server_id, "authorization", "header")
)
== "synthetic"
)
def test_environment_secrets_are_case_sensitive_and_delete_independently() -> None:
service = registry()
created = service.create(request(secret_environment_keys=["TOKEN", "token"]))
service.put_secret(created.server_id, "TOKEN", "upper")
service.put_secret(created.server_id, "token", "lower")
assert (
service.credentials.resolve(service._secret_id(created.server_id, "TOKEN"))
== "upper"
)
assert (
service.credentials.resolve(service._secret_id(created.server_id, "token"))
== "lower"
)
service.delete_secret(created.server_id, "TOKEN")
assert service.get(created.server_id).secret_environment == {
"TOKEN": False,
"token": True,
}
def test_legacy_environment_credential_migration_is_idempotent() -> None:
service = registry()
created = service.create(request(secret_environment_keys=["TOKEN"]))
suffix = hashlib.sha256(b"environment\0token").hexdigest()[:20]
legacy_id = f"mcp.{created.server_id}.{suffix}"
service.credentials.put(legacy_id, "legacy-value")
service._records[created.server_id]["secret_environment_version"] = 1
service._write()
migrated = registry()
assert migrated.get(created.server_id).secret_environment == {"TOKEN": True}
assert (
migrated.credentials.resolve(migrated._secret_id(created.server_id, "TOKEN"))
== "legacy-value"
)
assert not migrated.credentials.has(legacy_id)
migrated.put_secret(created.server_id, "TOKEN", "new-value")
assert (
registry().credentials.resolve(migrated._secret_id(created.server_id, "TOKEN"))
== "new-value"
)
def test_ambiguous_legacy_credentials_are_not_assigned_to_two_variables() -> None:
service = registry()
created = service.create(request(secret_environment_keys=["TOKEN", "token"]))
suffix = hashlib.sha256(b"environment\0token").hexdigest()[:20]
legacy_id = f"mcp.{created.server_id}.{suffix}"
service.credentials.put(legacy_id, "cannot-reconstruct-originals")
service._records[created.server_id]["secret_environment_version"] = 1
service._write()
migrated = registry()
current = migrated.get(created.server_id)
assert current.secret_environment == {"TOKEN": False, "token": False}
assert current.enabled is False
assert current.last_test_succeeded is None
assert migrated.credentials.has(
legacy_id
) # Keep the original ciphertext recoverable.
migrated.put_secret(created.server_id, "TOKEN", "upper")
migrated.put_secret(created.server_id, "token", "lower")
assert registry().get(created.server_id).secret_environment == {
"TOKEN": True,
"token": True,
}
migrated.delete(created.server_id)
assert not migrated.credentials.has(legacy_id)
def test_credential_id_migration_keeps_new_values_and_is_atomic(monkeypatch) -> None:
credentials = EncryptedCredentialStore()
credentials.put("mcp.old", "old-value")
credentials.put("mcp.new", "new-value")
original_write = credentials._write_tokens
def fail_write(_tokens):
raise CredentialStoreError("synthetic failure")
monkeypatch.setattr(credentials, "_write_tokens", fail_write)
with pytest.raises(CredentialStoreError):
credentials.move_many({"mcp.old": "mcp.new"})
assert credentials.resolve("mcp.old") == "old-value"
assert credentials.resolve("mcp.new") == "new-value"
monkeypatch.setattr(credentials, "_write_tokens", original_write)
credentials.move_many({"mcp.old": "mcp.new"})
assert credentials.resolve("mcp.old") is None
assert credentials.resolve("mcp.new") == "new-value"
def test_ambiguous_legacy_secret_is_not_resurrected_after_removing_a_key() -> None:
service = registry()
created = service.create(request(secret_environment_keys=["TOKEN", "token"]))
legacy_id = service._legacy_environment_secret_id(created.server_id, "TOKEN")
service.credentials.put(legacy_id, "ambiguous-old-value")
service._records[created.server_id]["secret_environment_version"] = 1
service._write()
migrated = registry()
migrated.update(
created.server_id,
McpServerUpdateRequest(
**request(secret_environment_keys=["token"]).model_dump(),
version=created.version,
),
)
assert registry().get(created.server_id).secret_environment == {"token": False}
def test_production_rejects_process_launch_even_after_approval() -> None:
service = registry(launch=False)
created = service.create(request(secret_environment_keys=[]))
@@ -184,6 +363,29 @@ def test_production_rejects_process_launch_even_after_approval() -> None:
assert error.value.code == "MCP_SANDBOX_REQUIRED"
@pytest.mark.parametrize("startup,tool", [(120, 300), (1.5, 2.5)])
def test_server_timeouts_survive_bridge_adaptation_and_reload(startup, tool) -> None:
service = registry()
created = service.create(
request(
secret_environment_keys=[],
startup_timeout_seconds=startup,
tool_timeout_seconds=tool,
)
)
service.trust(created.server_id, created.command_digest)
try:
tested = service.test(created.server_id)
assert tested.last_test_succeeded is True
assert tested.startup_timeout_seconds == startup
assert tested.tool_timeout_seconds == tool
restored = registry().get(created.server_id)
assert restored.startup_timeout_seconds == startup
assert restored.tool_timeout_seconds == tool
finally:
service.shutdown()
def test_enable_requires_successful_test_and_update_checks_version() -> None:
service = registry()
created = service.create(request(secret_environment_keys=[]))
@@ -307,15 +509,18 @@ def test_lifecycle_operations_are_serialized_and_tool_names_are_isolated() -> No
service.test(server.server_id)
with ThreadPoolExecutor(max_workers=4) as pool:
enabled = list(pool.map(lambda item: service.enable(item.server_id), servers * 2))
enabled = list(
pool.map(lambda item: service.enable(item.server_id), servers * 2)
)
assert all(item.enabled for item in enabled)
names = [
item.name
for item in service.tools.definitions()
if item.source == "mcp_server"
item.name for item in service.tools.definitions() if item.source == "mcp_server"
]
assert len(names) == len(set(names))
assert all(any(name.startswith(f"mcp.{item.server_id}.") for name in names) for item in servers)
assert all(
any(name.startswith(f"mcp.{item.server_id}.") for name in names)
for item in servers
)
with ThreadPoolExecutor(max_workers=4) as pool:
list(pool.map(lambda item: service.disable(item.server_id), servers * 2))
+83
View File
@@ -0,0 +1,83 @@
import json
from contextlib import closing
import httpx
import pytest
from app.extensions import mcp
class ChunkStream(httpx.SyncByteStream):
def __init__(self, chunks):
self.chunks = chunks
self.bytes_read = 0
def __iter__(self):
for chunk in self.chunks:
self.bytes_read += len(chunk)
yield chunk
def test_sse_rejects_unterminated_line_before_reading_entire_stream(monkeypatch):
monkeypatch.setattr(mcp, "MAX_MCP_MESSAGE_BYTES", 1024)
stream = ChunkStream([b"x" * 256] * 256)
with (
closing(httpx.Response(200, stream=stream)) as response,
pytest.raises(mcp.McpBridgeError, match="too large"),
):
list(mcp._iter_sse(response))
assert stream.bytes_read == 1280
def test_sse_limits_combined_event_before_partial_line_is_complete(monkeypatch):
monkeypatch.setattr(mcp, "MAX_MCP_MESSAGE_BYTES", 32)
stream = ChunkStream(
[b"data: 123456789\n", b"data: 123456789\n", b"x", b"not-read"]
)
with (
closing(httpx.Response(200, stream=stream)) as response,
pytest.raises(mcp.McpBridgeError, match="too large"),
):
list(mcp._iter_sse(response))
assert stream.bytes_read == 33
@pytest.mark.parametrize("separator", [b"\n", b"\r", b"\r\n"])
@pytest.mark.parametrize("chunk_size", [1, 2, 7, 1024])
def test_sse_preserves_utf8_and_line_endings_across_chunks(separator, chunk_size):
payload = json.dumps(
{"jsonrpc": "2.0", "id": 1, "result": {"text": "中文"}}, ensure_ascii=False
)
wire = b"\xef\xbb\xbf" + separator.join(
[
b": heartbeat",
b"event: message",
b"id: replay-1",
("data: " + payload).encode(),
b"",
b"",
]
)
stream = ChunkStream(
[wire[index : index + chunk_size] for index in range(0, len(wire), chunk_size)]
)
with closing(httpx.Response(200, stream=stream)) as response:
assert list(mcp._iter_sse(response)) == [("message", "replay-1", payload)]
def test_sse_event_limit_resets_between_events(monkeypatch):
monkeypatch.setattr(mcp, "MAX_MCP_MESSAGE_BYTES", 16)
with closing(
httpx.Response(200, stream=ChunkStream([b"data: one\n\ndata: two\r\r"]))
) as response:
assert list(mcp._iter_sse(response)) == [
("message", None, "one"),
("message", None, "two"),
]
def test_sse_preserves_multiline_data_and_final_unterminated_line():
with closing(
httpx.Response(200, stream=ChunkStream([b"data: first\ndata: last"]))
) as response:
assert list(mcp._iter_sse(response)) == [("message", None, "first\nlast")]