feat: add MCP tool integration
This commit is contained in:
304
backend/routes/mcp_servers.py
Normal file
304
backend/routes/mcp_servers.py
Normal 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}
|
||||
@@ -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 {}),
|
||||
|
||||
Reference in New Issue
Block a user