feat: add MCP tool integration

This commit is contained in:
Xin Wang
2026-07-18 00:00:06 +08:00
parent bdf3d3dd9c
commit e39bb48ba8
21 changed files with 1641 additions and 62 deletions

View File

@@ -0,0 +1,304 @@
"""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}

View File

@@ -7,7 +7,8 @@ from db.session import get_session
from fastapi import APIRouter, Depends, HTTPException
from schemas import ToolOut, ToolUpsert
from services.auth import require_admin
from services.masking import mask_secrets, merge_secrets
from services.masking import merge_secrets
from services.tool_resources import tool_to_out
from sqlalchemy import select
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession
@@ -20,21 +21,10 @@ router = APIRouter(
)
def _to_out(tool: Tool) -> ToolOut:
return ToolOut(
id=tool.id,
name=tool.name,
function_name=tool.function_name,
type=tool.type,
description=tool.description,
definition=tool.definition,
secrets=mask_secrets(tool.secrets or {}),
status=tool.status,
updated_at=tool.updated_at.isoformat() if tool.updated_at else None,
)
def _payload(body: ToolUpsert, stored_secrets: dict | None = None) -> dict:
def _payload(
body: ToolUpsert,
stored_secrets: dict | None = None,
) -> dict:
definition = body.definition.model_dump()
secrets = (
merge_secrets(body.secrets, stored_secrets or {})
@@ -45,6 +35,12 @@ def _payload(body: ToolUpsert, stored_secrets: dict | None = None) -> dict:
"name": body.name.strip(),
"function_name": body.function_name,
"type": definition["type"],
"mcp_server_id": body.mcp_server_id if definition["type"] == "mcp" else None,
"remote_tool_name": (
definition["config"]["remote_tool_name"]
if definition["type"] == "mcp"
else None
),
"description": body.description.strip(),
"definition": definition,
"secrets": secrets,
@@ -66,7 +62,7 @@ async def _commit(session: AsyncSession, tool: Tool) -> ToolOut:
await session.rollback()
raise HTTPException(409, "工具函数名已存在") from exc
await session.refresh(tool)
return _to_out(tool)
return tool_to_out(tool)
@router.get("", response_model=list[ToolOut])
@@ -74,11 +70,13 @@ async def list_tools(session: AsyncSession = Depends(get_session)):
rows = (
await session.execute(select(Tool).order_by(Tool.updated_at.desc()))
).scalars().all()
return [_to_out(tool) for tool in rows]
return [tool_to_out(tool) for tool in rows]
@router.post("", response_model=ToolOut)
async def create_tool(body: ToolUpsert, session: AsyncSession = Depends(get_session)):
if body.definition.type == "mcp":
raise HTTPException(400, "MCP 工具请通过 MCP Server 同步创建")
tool = Tool(id=f"tool_{uuid.uuid4().hex[:12]}", **_payload(body))
session.add(tool)
return await _commit(session, tool)
@@ -89,7 +87,7 @@ async def get_tool(tool_id: str, session: AsyncSession = Depends(get_session)):
tool = await session.get(Tool, tool_id)
if not tool:
raise HTTPException(404, "工具不存在")
return _to_out(tool)
return tool_to_out(tool)
@router.put("/{tool_id}", response_model=ToolOut)
@@ -101,6 +99,14 @@ async def update_tool(
tool = await session.get(Tool, tool_id)
if not tool:
raise HTTPException(404, "工具不存在")
if tool.type == "mcp":
if body.definition.type != "mcp":
raise HTTPException(400, "同步生成的 MCP 工具不能修改类型")
if (
body.mcp_server_id != tool.mcp_server_id
or body.definition.config.remote_tool_name != tool.remote_tool_name
):
raise HTTPException(400, "MCP Server 和远端工具名称不能修改")
for key, value in _payload(body, tool.secrets or {}).items():
setattr(tool, key, value)
return await _commit(session, tool)
@@ -111,6 +117,8 @@ async def duplicate_tool(tool_id: str, session: AsyncSession = Depends(get_sessi
source = await session.get(Tool, tool_id)
if not source:
raise HTTPException(404, "工具不存在")
if source.type == "mcp":
raise HTTPException(400, "MCP 工具由 Server 同步维护,不能单独复制")
function_name = _duplicate_function_name(source.function_name)
existing = (
await session.execute(select(Tool).where(Tool.function_name == function_name))
@@ -124,6 +132,8 @@ async def duplicate_tool(tool_id: str, session: AsyncSession = Depends(get_sessi
name=f"{source.name} 副本",
function_name=function_name,
type=source.type,
mcp_server_id=source.mcp_server_id,
remote_tool_name=source.remote_tool_name,
description=source.description,
definition=dict(source.definition or {}),
secrets=dict(source.secrets or {}),