Files
ai-video-fullstack/backend/routes/mcp_servers.py
2026-07-18 00:00:06 +08:00

305 lines
9.3 KiB
Python

"""Workspace MCP Server CRUD and explicit remote tool synchronization."""
from __future__ import annotations
import re
import uuid
from datetime import datetime, timezone
from db.models import AssistantToolBinding, McpServer, Tool
from db.session import get_session
from fastapi import APIRouter, Depends, HTTPException
from models import RuntimeMcpServer
from schemas import McpServerOut, McpServerUpsert, McpSyncResult, ToolOut
from services.auth import require_admin
from services.masking import mask_secrets, merge_secrets
from services.tool_resources import tool_to_out
from services.tools import McpClientError, McpToolClient
from sqlalchemy import func, select
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession
router = APIRouter(
prefix="/api/mcp-servers",
tags=["mcp-servers"],
dependencies=[Depends(require_admin)],
)
def _runtime_server(server: McpServer) -> RuntimeMcpServer:
config = server.config or {}
secrets = server.secrets or {}
return RuntimeMcpServer(
id=server.id,
name=server.name,
transport=server.transport,
url=server.url,
timeout_seconds=int(config.get("timeout_seconds") or 30),
headers={
str(key): str(value)
for key, value in (config.get("headers") or {}).items()
},
secret_headers={
str(key): str(value)
for key, value in (secrets.get("headers") or {}).items()
},
)
async def _to_out(session: AsyncSession, server: McpServer) -> McpServerOut:
tool_count = int(
(
await session.execute(
select(func.count(Tool.id)).where(Tool.mcp_server_id == server.id)
)
).scalar_one()
)
config = server.config or {}
secrets = server.secrets or {}
return McpServerOut(
id=server.id,
name=server.name,
description=server.description,
transport=server.transport, # type: ignore[arg-type]
url=server.url,
timeout_seconds=int(config.get("timeout_seconds") or 30),
headers=dict(config.get("headers") or {}),
secret_headers=mask_secrets(secrets.get("headers") or {}),
status=server.status, # type: ignore[arg-type]
tool_count=tool_count,
last_synced_at=(
server.last_synced_at.isoformat() if server.last_synced_at else None
),
updated_at=server.updated_at.isoformat() if server.updated_at else None,
)
def _apply_body(
server: McpServer,
body: McpServerUpsert,
) -> None:
server.name = body.name.strip()
server.description = body.description.strip()
server.transport = body.transport
server.url = body.url.strip()
server.config = {
"timeout_seconds": body.timeout_seconds,
"headers": body.headers,
}
server.secrets = merge_secrets(
{"headers": body.secret_headers},
server.secrets or {},
)
server.status = body.status
def _function_slug(value: str) -> str:
slug = re.sub(r"[^a-z0-9]+", "_", value.lower()).strip("_")
if not slug or not slug[0].isalpha():
slug = f"tool_{slug}"
return slug[:64]
async def _available_function_name(
session: AsyncSession,
server: McpServer,
remote_name: str,
) -> str:
base = _function_slug(f"mcp_{server.name}_{remote_name}")
candidate = base
suffix = 2
while (
await session.execute(
select(Tool.id).where(Tool.function_name == candidate).limit(1)
)
).scalar_one_or_none():
tail = f"_{suffix}"
candidate = f"{base[: 64 - len(tail)]}{tail}"
suffix += 1
return candidate
@router.get("", response_model=list[McpServerOut])
async def list_mcp_servers(session: AsyncSession = Depends(get_session)):
rows = (
await session.execute(
select(McpServer).order_by(McpServer.updated_at.desc())
)
).scalars().all()
return [await _to_out(session, server) for server in rows]
@router.post("", response_model=McpServerOut)
async def create_mcp_server(
body: McpServerUpsert,
session: AsyncSession = Depends(get_session),
):
server = McpServer(
id=f"mcp_{uuid.uuid4().hex[:12]}",
name=body.name.strip(),
url=body.url.strip(),
)
_apply_body(server, body)
session.add(server)
await session.commit()
await session.refresh(server)
return await _to_out(session, server)
@router.get("/{server_id}", response_model=McpServerOut)
async def get_mcp_server(
server_id: str,
session: AsyncSession = Depends(get_session),
):
server = await session.get(McpServer, server_id)
if not server:
raise HTTPException(404, "MCP Server 不存在")
return await _to_out(session, server)
@router.put("/{server_id}", response_model=McpServerOut)
async def update_mcp_server(
server_id: str,
body: McpServerUpsert,
session: AsyncSession = Depends(get_session),
):
server = await session.get(McpServer, server_id)
if not server:
raise HTTPException(404, "MCP Server 不存在")
_apply_body(server, body)
await session.commit()
await session.refresh(server)
return await _to_out(session, server)
@router.post("/{server_id}/sync", response_model=McpSyncResult)
async def sync_mcp_tools(
server_id: str,
session: AsyncSession = Depends(get_session),
):
server = await session.get(McpServer, server_id)
if not server:
raise HTTPException(404, "MCP Server 不存在")
if server.status != "active":
raise HTTPException(400, "请先启用 MCP Server")
try:
discovered = await McpToolClient().list_tools(_runtime_server(server))
except McpClientError as exc:
raise HTTPException(400, str(exc)) from exc
existing_rows = (
await session.execute(
select(Tool).where(Tool.mcp_server_id == server.id)
)
).scalars().all()
existing = {
str(tool.remote_tool_name): tool
for tool in existing_rows
if tool.remote_tool_name
}
created = 0
updated = 0
synchronized: list[Tool] = []
for remote in discovered:
remote_name = str(remote["name"])
tool = existing.get(remote_name)
if tool is None:
tool = Tool(
id=f"tool_{uuid.uuid4().hex[:12]}",
name=f"{server.name} · {remote_name}",
function_name=await _available_function_name(
session,
server,
remote_name,
),
type="mcp",
mcp_server_id=server.id,
remote_tool_name=remote_name,
description=str(remote.get("description") or ""),
definition={},
secrets={},
status="active",
)
session.add(tool)
created += 1
else:
updated += 1
previous_config = (tool.definition or {}).get("config") or {}
tool.description = str(remote.get("description") or tool.description)
tool.definition = {
"schema_version": 1,
"type": "mcp",
"config": {
"remote_tool_name": remote_name,
"input_schema": remote.get("input_schema") or {},
"schema_hash": remote.get("schema_hash") or "",
"dynamic_variable_assignments": previous_config.get(
"dynamic_variable_assignments"
)
or {},
},
}
synchronized.append(tool)
server.last_synced_at = datetime.now(timezone.utc)
try:
await session.commit()
except IntegrityError as exc:
await session.rollback()
raise HTTPException(409, "同步工具时发生名称冲突,请重试") from exc
for tool in synchronized:
await session.refresh(tool)
await session.refresh(server)
server_out = await _to_out(session, server)
tool_outputs: list[ToolOut] = [tool_to_out(tool) for tool in synchronized]
return McpSyncResult(
server=server_out,
created=created,
updated=updated,
tools=tool_outputs,
)
@router.get("/{server_id}/tools", response_model=list[ToolOut])
async def list_mcp_tools(
server_id: str,
session: AsyncSession = Depends(get_session),
):
if not await session.get(McpServer, server_id):
raise HTTPException(404, "MCP Server 不存在")
rows = (
await session.execute(
select(Tool)
.where(Tool.mcp_server_id == server_id)
.order_by(Tool.name)
)
).scalars().all()
return [tool_to_out(tool) for tool in rows]
@router.delete("/{server_id}")
async def delete_mcp_server(
server_id: str,
session: AsyncSession = Depends(get_session),
):
server = await session.get(McpServer, server_id)
if not server:
raise HTTPException(404, "MCP Server 不存在")
in_use = (
await session.execute(
select(AssistantToolBinding.tool_id)
.join(Tool, Tool.id == AssistantToolBinding.tool_id)
.where(Tool.mcp_server_id == server_id)
.limit(1)
)
).scalar_one_or_none()
if in_use:
raise HTTPException(409, "MCP 工具正被助手引用,请先解绑")
await session.delete(server)
await session.commit()
return {"ok": True}