203 lines
6.5 KiB
Python
203 lines
6.5 KiB
Python
"""
|
||
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
|