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)
|
||||
|
||||
@@ -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
@@ -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
@@ -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:
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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")]
|
||||
Reference in New Issue
Block a user