huoyan-enterprise/backend/core/mcp_config.py

203 lines
6.5 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""
MCP 多服务器配置解析。
支持三种写法(可组合,会去重 URL
1. **JSON推荐** — ``MCP_SERVERS``::
[{"name":"local-excel","url":"http://localhost:9090/mcp","transport":"streamable_http","api_key":"k1"},
{"name":"remote","url":"http://example.com/mcp","api_key":"k2"}]
2. **逗号分隔 URL** — ``MCP_URLS``(共用 ``MCP_TRANSPORT`` / ``MCP_API_KEY``::
local|http://localhost:9090/mcp,remote|http://example.com/mcp
或仅 URL: http://localhost:9090/mcp,http://example.com/mcp
3. **单地址(兼容旧版)** — ``MCP_URL`` + ``MCP_TRANSPORT`` + ``MCP_API_KEY``
"""
from __future__ import annotations
import json
import re
from dataclasses import dataclass
from typing import Any, Optional
from urllib.parse import urlparse
from logger.logging import get_logger
logger = get_logger(__name__)
_SERVER_NAME_RE = re.compile(r"^[a-zA-Z][a-zA-Z0-9_-]{0,62}$")
@dataclass(frozen=True)
class McpServerEntry:
name: str
url: str
transport: str = "streamable_http"
api_key: Optional[str] = None
def _sanitize_server_name(raw: str, fallback: str) -> str:
name = (raw or "").strip()
if name and _SERVER_NAME_RE.match(name):
return name
return fallback
def _parse_name_url_token(token: str, index: int) -> Optional[McpServerEntry]:
token = token.strip()
if not token:
return None
if "|" in token:
name_part, url_part = token.split("|", 1)
url = url_part.strip()
name = _sanitize_server_name(name_part, f"mcp-{index}")
else:
url = token
parsed = urlparse(url)
host = (parsed.hostname or "server").replace(".", "-")
name = _sanitize_server_name(host, f"mcp-{index}")
if not url.startswith(("http://", "https://")):
logger.warning("跳过无效 MCP URL需 http/https: {}", token[:80])
return None
return McpServerEntry(name=name, url=url)
def _parse_mcp_urls_list(
raw: str,
*,
default_transport: str,
default_api_key: Optional[str],
) -> list[McpServerEntry]:
entries: list[McpServerEntry] = []
for i, token in enumerate(raw.split(",")):
base = _parse_name_url_token(token, i)
if base:
entries.append(
McpServerEntry(
name=base.name,
url=base.url,
transport=default_transport,
api_key=default_api_key,
)
)
return entries
def _parse_mcp_servers_json(raw: str) -> list[McpServerEntry]:
try:
data = json.loads(raw)
except json.JSONDecodeError as e:
logger.error("MCP_SERVERS JSON 解析失败: {}", e)
return []
if not isinstance(data, list):
logger.error("MCP_SERVERS 必须是 JSON 数组")
return []
entries: list[McpServerEntry] = []
for i, item in enumerate(data):
if not isinstance(item, dict):
logger.warning("MCP_SERVERS[{}] 不是对象,已跳过", i)
continue
url = str(item.get("url") or "").strip()
if not url.startswith(("http://", "https://")):
logger.warning("MCP_SERVERS[{}] url 无效,已跳过", i)
continue
name = _sanitize_server_name(str(item.get("name") or ""), f"mcp-{i}")
transport = str(item.get("transport") or "streamable_http").strip() or "streamable_http"
api_key = item.get("api_key") or item.get("apiKey")
if api_key is not None:
api_key = str(api_key).strip() or None
entries.append(
McpServerEntry(name=name, url=url, transport=transport, api_key=api_key)
)
return entries
def parse_mcp_server_entries(settings: Any) -> list[McpServerEntry]:
"""从 Settings 解析全部 MCP 服务器条目(去重 URL后者覆盖同名"""
default_transport = (getattr(settings, "mcp_transport", None) or "streamable_http").strip()
default_api_key = getattr(settings, "mcp_api_key", None)
if default_api_key:
default_api_key = str(default_api_key).strip() or None
ordered: list[McpServerEntry] = []
# 1) 旧版单地址
legacy_url = (getattr(settings, "mcp_url", None) or "").strip()
if legacy_url:
ordered.append(
McpServerEntry(
name="default",
url=legacy_url,
transport=default_transport,
api_key=default_api_key,
)
)
# 2) 逗号分隔多地址
mcp_urls = (getattr(settings, "mcp_urls", None) or "").strip()
if mcp_urls:
ordered.extend(
_parse_mcp_urls_list(
mcp_urls,
default_transport=default_transport,
default_api_key=default_api_key,
)
)
# 3) JSON 完整配置(优先级最高,可覆盖 transport/key
mcp_servers_json = (getattr(settings, "mcp_servers", None) or "").strip()
if mcp_servers_json:
ordered.extend(_parse_mcp_servers_json(mcp_servers_json))
# 去重:按 url 保留最后一次出现的配置
by_url: dict[str, McpServerEntry] = {}
for entry in ordered:
by_url[entry.url.rstrip("/")] = entry
# 确保 name 唯一
used_names: set[str] = set()
result: list[McpServerEntry] = []
for entry in by_url.values():
name = entry.name
suffix = 1
while name in used_names:
name = f"{entry.name}-{suffix}"
suffix += 1
used_names.add(name)
result.append(
McpServerEntry(
name=name,
url=entry.url,
transport=entry.transport,
api_key=entry.api_key,
)
)
return result
def has_mcp_servers(settings: Any) -> bool:
return len(parse_mcp_server_entries(settings)) > 0
def build_mcp_connections(settings: Any) -> dict[str, dict[str, Any]]:
"""构建 ``MultiServerMCPClient`` 所需的 connections 字典。"""
connections: dict[str, dict[str, Any]] = {}
for entry in parse_mcp_server_entries(settings):
conn: dict[str, Any] = {
"transport": entry.transport or "streamable_http",
"url": entry.url,
}
if entry.api_key:
conn["headers"] = {"Authorization": f"Bearer {entry.api_key}"}
connections[entry.name] = conn
logger.info(
"MCP 服务器已注册: name={}, transport={}, url={}",
entry.name,
conn["transport"],
entry.url,
)
return connections