"""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}