Files
ai-video-fullstack/backend/services/tools/mcp_client.py
Xin Wang f027ed99b7 feat(mcp): enhance MCP transport support and add server management UI
- Extend McpTransport to support "sse" in schemas.
- Refactor McpToolClient to handle both "streamable_http" and "sse" transports.
- Introduce McpServerDialog for managing MCP server configurations, including transport settings and tool synchronization.
- Replace McpServersSection with the new dialog component for improved server management.
- Add tests for MCP transport handling and server dialog functionality.
2026-07-19 12:04:11 +08:00

177 lines
6.5 KiB
Python

"""One-shot HTTP MCP discovery and execution.
Connections intentionally live for one operation in the MVP. This keeps the
async context ownership simple and correct; a session pool can be introduced
later if measurements show MCP handshakes are a meaningful source of latency.
"""
from __future__ import annotations
import hashlib
import json
from contextlib import asynccontextmanager
from datetime import timedelta
from typing import Any, AsyncIterator
import httpx
from mcp import ClientSession
from mcp.client.sse import sse_client
from mcp.client.streamable_http import streamable_http_client
from models import RuntimeMcpServer
MAX_MCP_RESULT_BYTES = 1_000_000
class McpClientError(RuntimeError):
"""Safe error surfaced to tool callers and the administration API."""
def schema_hash(input_schema: dict[str, Any]) -> str:
serialized = json.dumps(
input_schema,
ensure_ascii=False,
sort_keys=True,
separators=(",", ":"),
)
return hashlib.sha256(serialized.encode("utf-8")).hexdigest()
class McpToolClient:
"""Minimal public MCP SDK wrapper shared by discovery and runtime calls."""
@asynccontextmanager
async def _transport_streams(
self,
server: RuntimeMcpServer,
headers: dict[str, str],
timeout: int,
) -> AsyncIterator[tuple[Any, Any]]:
"""Open the configured MCP transport and expose its read/write streams."""
if server.transport == "streamable_http":
async with httpx.AsyncClient(
headers=headers,
timeout=httpx.Timeout(timeout),
follow_redirects=False,
) as http_client:
async with streamable_http_client(
server.url,
http_client=http_client,
) as streams:
read_stream, write_stream, _ = streams
yield read_stream, write_stream
return
if server.transport == "sse":
async with sse_client(
server.url,
headers=headers,
timeout=timeout,
sse_read_timeout=timeout,
) as streams:
yield streams
return
raise McpClientError(f"暂不支持 MCP transport: {server.transport}")
@asynccontextmanager
async def _session(
self,
server: RuntimeMcpServer,
) -> AsyncIterator[ClientSession]:
headers = {**server.headers, **server.secret_headers}
timeout = max(1, min(int(server.timeout_seconds), 120))
try:
async with self._transport_streams(server, headers, timeout) as streams:
read_stream, write_stream = streams
async with ClientSession(
read_stream,
write_stream,
read_timeout_seconds=timedelta(seconds=timeout),
) as session:
await session.initialize()
yield session
except McpClientError:
raise
except httpx.TimeoutException as exc:
raise McpClientError("MCP Server 连接或调用超时") from exc
except Exception as exc: # MCP SDK exposes several transport exceptions.
raise McpClientError(f"MCP Server 调用失败: {exc}") from exc
async def list_tools(self, server: RuntimeMcpServer) -> list[dict[str, Any]]:
"""Discover all remote tools, following MCP pagination cursors."""
discovered: list[dict[str, Any]] = []
async with self._session(server) as session:
cursor: str | None = None
while True:
result = await session.list_tools(cursor=cursor)
for tool in result.tools:
input_schema = dict(tool.inputSchema or {})
discovered.append(
{
"name": tool.name,
"description": tool.description or "",
"input_schema": input_schema,
"schema_hash": schema_hash(input_schema),
}
)
cursor = result.nextCursor
if not cursor:
break
return discovered
async def call_tool(
self,
server: RuntimeMcpServer,
remote_tool_name: str,
arguments: dict[str, Any],
) -> dict[str, Any]:
"""Execute one remote tool and normalize the result for local callers."""
async with self._session(server) as session:
result = await session.call_tool(
remote_tool_name,
arguments=arguments,
read_timeout_seconds=timedelta(seconds=server.timeout_seconds),
)
data = self._normalize_result(result)
encoded = json.dumps(data, ensure_ascii=False, default=str).encode("utf-8")
if len(encoded) > MAX_MCP_RESULT_BYTES:
raise McpClientError("MCP 工具响应超过 1 MB 限制")
if bool(result.isError):
message = data.get("text") or "MCP 工具返回执行错误"
raise McpClientError(str(message)[:2048])
return data
@staticmethod
def _normalize_result(result: Any) -> dict[str, Any]:
"""Keep useful text/JSON while deliberately omitting binary payloads."""
content: list[dict[str, Any]] = []
text_parts: list[str] = []
for item in result.content or []:
item_type = str(getattr(item, "type", "unknown"))
if item_type == "text":
text = str(getattr(item, "text", ""))
content.append({"type": "text", "text": text})
if text:
text_parts.append(text)
continue
if item_type == "resource":
resource = getattr(item, "resource", None)
text = getattr(resource, "text", None)
if isinstance(text, str):
content.append({"type": "resource", "text": text})
text_parts.append(text)
else:
content.append({"type": "resource", "omitted": True})
continue
# Audio and image base64 data are not sent into the LLM context in MVP.
content.append({"type": item_type, "omitted": True})
return {
"text": "\n".join(text_parts),
"structuredContent": result.structuredContent,
"content": content,
}