305 lines
9.3 KiB
Python
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}
|