diff --git a/backend/schemas.py b/backend/schemas.py index dfa0a2d..b01b118 100644 --- a/backend/schemas.py +++ b/backend/schemas.py @@ -22,7 +22,7 @@ KnowledgeRetrievalMode = Literal["automatic", "on_demand"] ToolType = Literal["end_call", "http", "mcp"] ToolStatus = Literal["active", "archived", "draft"] McpServerStatus = Literal["active", "archived", "draft"] -McpTransport = Literal["streamable_http"] +McpTransport = Literal["streamable_http", "sse"] ToolParameterType = Literal["string", "number", "integer", "boolean", "object", "array"] ToolParameterLocation = Literal["path", "query", "body", "header"] DynamicVariableType = Literal["string", "number", "boolean"] diff --git a/backend/services/tools/mcp_client.py b/backend/services/tools/mcp_client.py index bde76e0..d52f7b0 100644 --- a/backend/services/tools/mcp_client.py +++ b/backend/services/tools/mcp_client.py @@ -1,4 +1,4 @@ -"""One-shot Streamable HTTP MCP discovery and execution. +"""One-shot HTTP MCP discovery and execution. Connections intentionally live for one operation in the MVP. This keeps the async context ownership simple and correct; a session pool can be introduced @@ -15,6 +15,7 @@ from typing import Any, AsyncIterator import httpx from mcp import ClientSession +from mcp.client.sse import sse_client from mcp.client.streamable_http import streamable_http_client from models import RuntimeMcpServer @@ -41,16 +42,14 @@ class McpToolClient: """Minimal public MCP SDK wrapper shared by discovery and runtime calls.""" @asynccontextmanager - async def _session( + async def _transport_streams( self, server: RuntimeMcpServer, - ) -> AsyncIterator[ClientSession]: - if server.transport != "streamable_http": - raise McpClientError(f"暂不支持 MCP transport: {server.transport}") - - headers = {**server.headers, **server.secret_headers} - timeout = max(1, min(int(server.timeout_seconds), 120)) - try: + headers: dict[str, str], + timeout: int, + ) -> AsyncIterator[tuple[Any, Any]]: + """Open the configured MCP transport and expose its read/write streams.""" + if server.transport == "streamable_http": async with httpx.AsyncClient( headers=headers, timeout=httpx.Timeout(timeout), @@ -61,13 +60,38 @@ class McpToolClient: http_client=http_client, ) as streams: read_stream, write_stream, _ = streams - async with ClientSession( - read_stream, - write_stream, - read_timeout_seconds=timedelta(seconds=timeout), - ) as session: - await session.initialize() - yield session + yield read_stream, write_stream + return + + if server.transport == "sse": + async with sse_client( + server.url, + headers=headers, + timeout=timeout, + sse_read_timeout=timeout, + ) as streams: + yield streams + return + + raise McpClientError(f"暂不支持 MCP transport: {server.transport}") + + @asynccontextmanager + async def _session( + self, + server: RuntimeMcpServer, + ) -> AsyncIterator[ClientSession]: + headers = {**server.headers, **server.secret_headers} + timeout = max(1, min(int(server.timeout_seconds), 120)) + try: + async with self._transport_streams(server, headers, timeout) as streams: + read_stream, write_stream = streams + async with ClientSession( + read_stream, + write_stream, + read_timeout_seconds=timedelta(seconds=timeout), + ) as session: + await session.initialize() + yield session except McpClientError: raise except httpx.TimeoutException as exc: diff --git a/backend/tests/test_mcp_client.py b/backend/tests/test_mcp_client.py new file mode 100644 index 0000000..7488322 --- /dev/null +++ b/backend/tests/test_mcp_client.py @@ -0,0 +1,51 @@ +import unittest +from contextlib import asynccontextmanager +from unittest.mock import patch + +from models import RuntimeMcpServer +from schemas import McpServerUpsert +from services.tools.mcp_client import McpToolClient + + +class McpTransportTests(unittest.IsolatedAsyncioTestCase): + def test_server_schema_accepts_sse(self): + server = McpServerUpsert( + name="旧版 MCP 服务", + transport="sse", + url="https://mcp.example.com/sse", + ) + + self.assertEqual(server.transport, "sse") + + async def test_sse_transport_uses_sdk_sse_client(self): + captured: dict = {} + + @asynccontextmanager + async def fake_sse_client(url, **kwargs): + captured["url"] = url + captured.update(kwargs) + yield "read-stream", "write-stream" + + server = RuntimeMcpServer( + id="mcp_sse", + transport="sse", + url="https://mcp.example.com/sse", + ) + client = McpToolClient() + + with patch("services.tools.mcp_client.sse_client", fake_sse_client): + async with client._transport_streams( + server, + {"Authorization": "Bearer test"}, + 20, + ) as streams: + self.assertEqual(streams, ("read-stream", "write-stream")) + + self.assertEqual(captured["url"], server.url) + self.assertEqual(captured["headers"], {"Authorization": "Bearer test"}) + self.assertEqual(captured["timeout"], 20) + self.assertEqual(captured["sse_read_timeout"], 20) + + +if __name__ == "__main__": + unittest.main() diff --git a/frontend/src/components/pages/ComponentsToolsPage.tsx b/frontend/src/components/pages/ComponentsToolsPage.tsx index a2c827f..6db51bb 100644 --- a/frontend/src/components/pages/ComponentsToolsPage.tsx +++ b/frontend/src/components/pages/ComponentsToolsPage.tsx @@ -9,10 +9,17 @@ import { MoreHorizontal, Pencil, Plus, + ServerCog, Trash2, } from "lucide-react"; -import { McpServersSection } from "@/components/tools/McpServersSection"; +import { McpServerDialog } from "@/components/tools/McpServerDialog"; +import { + TOOL_DIALOG_CONTENT_CLASS, + ToolFormField as Field, + ToolFormSection as FieldSection, + ToolJsonField as JsonField, +} from "@/components/tools/tool-form-controls"; import { Badge } from "@/components/ui/badge"; import { Button } from "@/components/ui/button"; import { DataList, type DataListColumn } from "@/components/ui/data-list"; @@ -45,18 +52,23 @@ import { import { Switch } from "@/components/ui/switch"; import { Textarea } from "@/components/ui/textarea"; import { + mcpServersApi, toolsApi, type HttpToolDefinition, + type McpServer, type Tool, type ToolParameter, type ToolStatus, type ToolUpsert, } from "@/lib/api"; -type ToolKind = "end_call" | "http" | "mcp"; +type ToolKind = "end_call" | "http"; type HttpMethod = HttpToolDefinition["config"]["method"]; type ToolFilter = "全部" | "End Call" | "HTTP" | "MCP"; type SortOrder = "newest" | "oldest"; +type ToolResource = + | { kind: "tool"; id: string; tool: Tool } + | { kind: "mcp"; id: string; server: McpServer }; const toolFilters: readonly ToolFilter[] = ["全部", "End Call", "HTTP", "MCP"]; @@ -78,10 +90,6 @@ type ToolForm = { body: string; dynamicVariableAssignments: string; secretDynamicVariables: string; - mcpServerId: string; - remoteToolName: string; - inputSchema: string; - schemaHash: string; }; const EMPTY_OBJECT = "{}"; @@ -106,10 +114,6 @@ function blankForm(): ToolForm { body: EMPTY_OBJECT, dynamicVariableAssignments: EMPTY_OBJECT, secretDynamicVariables: EMPTY_OBJECT, - mcpServerId: "", - remoteToolName: "", - inputSchema: EMPTY_OBJECT, - schemaHash: "", }; } @@ -118,6 +122,9 @@ function pretty(value: unknown, fallback: string): string { } function formFromTool(tool: Tool): ToolForm { + if (tool.type === "mcp") { + throw new Error("MCP 服务需要使用 MCP 编辑窗口"); + } const base = { ...blankForm(), name: tool.name, functionName: tool.functionName }; base.type = tool.type; base.description = tool.description; @@ -128,17 +135,7 @@ function formFromTool(tool: Tool): ToolForm { base.captureReason = tool.definition.config.captureReason; return base; } - if (tool.definition.type === "mcp") { - base.mcpServerId = tool.mcpServerId ?? ""; - base.remoteToolName = tool.definition.config.remoteToolName; - base.inputSchema = pretty(tool.definition.config.inputSchema, EMPTY_OBJECT); - base.schemaHash = tool.definition.config.schemaHash; - base.dynamicVariableAssignments = pretty( - tool.definition.config.dynamicVariableAssignments ?? {}, - EMPTY_OBJECT, - ); - return base; - } + if (tool.definition.type === "mcp") return base; base.method = tool.definition.config.method; base.url = tool.definition.config.url; base.timeoutSeconds = String(tool.definition.config.timeoutSeconds); @@ -218,30 +215,6 @@ function payloadFromForm(form: ToolForm): ToolUpsert { }, }; } - if (form.type === "mcp") { - return { - name: form.name.trim(), - functionName: form.functionName, - description: form.description.trim(), - status: form.status, - mcpServerId: form.mcpServerId, - secrets: {}, - definition: { - schemaVersion: 1, - type: "mcp", - config: { - remoteToolName: form.remoteToolName, - inputSchema: parseObject(form.inputSchema, "MCP Input Schema"), - schemaHash: form.schemaHash, - dynamicVariableAssignments: parseObject( - form.dynamicVariableAssignments, - "变量赋值", - ) as Record, - }, - }, - }; - } - const timeoutSeconds = Number(form.timeoutSeconds); if (!Number.isInteger(timeoutSeconds) || timeoutSeconds < 1 || timeoutSeconds > 120) { throw new Error("超时时间必须是 1 到 120 秒之间的整数"); @@ -301,6 +274,7 @@ function updatedAtValue(value?: string | null): number { export function ComponentsToolsPage() { const [tools, setTools] = useState([]); + const [mcpServers, setMcpServers] = useState([]); const [loading, setLoading] = useState(true); const [error, setError] = useState(null); const [search, setSearch] = useState(""); @@ -309,6 +283,8 @@ export function ComponentsToolsPage() { const [currentPage, setCurrentPage] = useState(1); const [dialogOpen, setDialogOpen] = useState(false); const [editing, setEditing] = useState(null); + const [mcpDialogOpen, setMcpDialogOpen] = useState(false); + const [editingMcpServer, setEditingMcpServer] = useState(null); const [form, setForm] = useState(blankForm); const [saving, setSaving] = useState(false); const [formError, setFormError] = useState(null); @@ -316,13 +292,18 @@ export function ComponentsToolsPage() { const [deletingId, setDeletingId] = useState(null); const [duplicatingId, setDuplicatingId] = useState(null); - const loadTools = useCallback(async () => { + const loadResources = useCallback(async () => { setLoading(true); setError(null); try { - setTools(await toolsApi.list()); + const [nextTools, nextMcpServers] = await Promise.all([ + toolsApi.list(), + mcpServersApi.list(), + ]); + setTools(nextTools); + setMcpServers(nextMcpServers); } catch (loadError) { - setError(loadError instanceof Error ? loadError.message : "加载工具失败"); + setError(loadError instanceof Error ? loadError.message : "加载工具资源失败"); } finally { setLoading(false); } @@ -330,37 +311,60 @@ export function ComponentsToolsPage() { useEffect(() => { // eslint-disable-next-line react-hooks/set-state-in-effect - void loadTools(); - }, [loadTools]); + void loadResources(); + }, [loadResources]); - const filteredTools = useMemo(() => { - const query = search.trim().toLowerCase(); - return tools.filter((tool) => - (filter === "全部" || - (filter === "End Call" && tool.type === "end_call") || - (filter === "HTTP" && tool.type === "http") || - (filter === "MCP" && tool.type === "mcp")) && - (!query || - [tool.name, tool.functionName, tool.description].some((value) => - value.toLowerCase().includes(query), - )), + const resources = useMemo(() => { + const nativeTools = tools + .filter((tool) => tool.type !== "mcp") + .map((tool) => ({ kind: "tool", id: tool.id, tool }) as const); + const servers = mcpServers.map( + (server) => ({ kind: "mcp", id: server.id, server }) as const, ); - }, [filter, search, tools]); + return [...nativeTools, ...servers]; + }, [mcpServers, tools]); - const sortedTools = useMemo(() => { - return [...filteredTools].sort((a, b) => { - const diff = updatedAtValue(b.updatedAt) - updatedAtValue(a.updatedAt); + const filteredResources = useMemo(() => { + const query = search.trim().toLowerCase(); + return resources.filter((resource) => { + if (resource.kind === "mcp") { + return ( + (filter === "全部" || filter === "MCP") && + (!query || + [resource.server.name, resource.server.url, resource.server.description].some( + (value) => value.toLowerCase().includes(query), + )) + ); + } + const tool = resource.tool; + return ( + (filter === "全部" || + (filter === "End Call" && tool.type === "end_call") || + (filter === "HTTP" && tool.type === "http")) && + (!query || + [tool.name, tool.functionName, tool.description].some((value) => + value.toLowerCase().includes(query), + )) + ); + }); + }, [filter, resources, search]); + + const sortedResources = useMemo(() => { + return [...filteredResources].sort((a, b) => { + const aUpdatedAt = a.kind === "mcp" ? a.server.updatedAt : a.tool.updatedAt; + const bUpdatedAt = b.kind === "mcp" ? b.server.updatedAt : b.tool.updatedAt; + const diff = updatedAtValue(bUpdatedAt) - updatedAtValue(aUpdatedAt); if (diff !== 0) return sortOrder === "newest" ? diff : -diff; return a.id.localeCompare(b.id); }); - }, [filteredTools, sortOrder]); + }, [filteredResources, sortOrder]); const pageSize = 5; - const totalPages = Math.max(1, Math.ceil(sortedTools.length / pageSize)); + const totalPages = Math.max(1, Math.ceil(sortedResources.length / pageSize)); const safeCurrentPage = Math.min(currentPage, totalPages); const pageStart = (safeCurrentPage - 1) * pageSize; const pageEnd = pageStart + pageSize; - const paginatedTools = sortedTools.slice(pageStart, pageEnd); + const paginatedResources = sortedResources.slice(pageStart, pageEnd); function changeFilter(value: ToolFilter) { setFilter(value); @@ -380,6 +384,12 @@ export function ComponentsToolsPage() { setDialogOpen(true); } + function openCreateMcpServer() { + setDialogOpen(false); + setEditingMcpServer(null); + setMcpDialogOpen(true); + } + function openEdit(tool: Tool) { setEditing(tool); setForm(formFromTool(tool)); @@ -388,6 +398,11 @@ export function ComponentsToolsPage() { setDialogOpen(true); } + function openEditMcpServer(server: McpServer) { + setEditingMcpServer(server); + setMcpDialogOpen(true); + } + async function saveTool() { if (saving) return; setSaving(true); @@ -397,7 +412,7 @@ export function ComponentsToolsPage() { if (editing) await toolsApi.update(editing.id, payload); else await toolsApi.create(payload); setDialogOpen(false); - await loadTools(); + await loadResources(); } catch (saveError) { setFormError(saveError instanceof Error ? saveError.message : "保存失败"); } finally { @@ -410,7 +425,20 @@ export function ComponentsToolsPage() { setDeletingId(tool.id); try { await toolsApi.remove(tool.id); - await loadTools(); + await loadResources(); + } catch (removeError) { + setError(removeError instanceof Error ? removeError.message : "删除失败"); + } finally { + setDeletingId(null); + } + } + + async function removeMcpServer(server: McpServer) { + if (!window.confirm(`确认删除 MCP 工具资源“${server.name}”及其同步工具?`)) return; + setDeletingId(server.id); + try { + await mcpServersApi.remove(server.id); + await loadResources(); } catch (removeError) { setError(removeError instanceof Error ? removeError.message : "删除失败"); } finally { @@ -422,7 +450,7 @@ export function ComponentsToolsPage() { setDuplicatingId(id); try { await toolsApi.duplicate(id); - await loadTools(); + await loadResources(); } catch (duplicateError) { setError(duplicateError instanceof Error ? duplicateError.message : "复制失败"); } finally { @@ -430,18 +458,22 @@ export function ComponentsToolsPage() { } } - const columns: DataListColumn[] = [ + const columns: DataListColumn[] = [ { key: "name", header: "工具名称", width: "md:w-[360px]", - cell: (tool) => ( + cell: (resource) => ( <>
- {tool.name} + + {resource.kind === "mcp" ? resource.server.name : resource.tool.name} +
- {tool.functionName} + {resource.kind === "mcp" + ? `${resource.server.url} · ${resource.server.toolCount} 个工具` + : resource.tool.functionName}
), @@ -450,29 +482,19 @@ export function ComponentsToolsPage() { key: "type", header: "类型", width: "md:w-[128px]", - cell: (tool) => ( + cell: (resource) => ( - {tool.type === "end_call" - ? "End Call" - : tool.type === "mcp" - ? "MCP" + {resource.kind === "mcp" + ? "MCP" + : resource.tool.type === "end_call" + ? "End Call" : "HTTP"} ), }, - { - key: "status", - header: "状态", - width: "md:w-[156px]", - cell: (tool) => ( - - {tool.status === "active" ? "启用" : tool.status === "draft" ? "草稿" : "已归档"} - - ), - }, { key: "updated", width: "md:w-[176px]", @@ -491,74 +513,86 @@ export function ComponentsToolsPage() { ), cellClassName: "whitespace-nowrap tabular-nums text-muted-foreground", - cell: (tool) => formatTimestamp(tool.updatedAt), + cell: (resource) => + formatTimestamp( + resource.kind === "mcp" ? resource.server.updatedAt : resource.tool.updatedAt, + ), }, { key: "actions", header: "操作", align: "right", cellClassName: "flex justify-end gap-2", - cell: (tool) => ( - <> - - - - - - { + const name = resource.kind === "mcp" ? resource.server.name : resource.tool.name; + const resourceId = resource.id; + return ( + <> + + + + + + - {duplicatingId === tool.id ? ( - - ) : ( - - )} - 复制 - - { - event.preventDefault(); - void removeTool(tool); - }} - > - {deletingId === tool.id ? ( - - ) : ( - - )} - 删除 - - - - - ), + { + event.preventDefault(); + if (resource.kind === "tool") void duplicateTool(resource.tool.id); + }} + > + {duplicatingId === resourceId ? ( + + ) : ( + + )} + 复制 + + { + event.preventDefault(); + if (resource.kind === "mcp") void removeMcpServer(resource.server); + else void removeTool(resource.tool); + }} + > + {deletingId === resourceId ? ( + + ) : ( + + )} + 删除 + + + + + ); + }, }, ]; @@ -568,15 +602,31 @@ export function ComponentsToolsPage() { title="工具资源" description="管理可复用的助手工具,并将启用的工具绑定到提示词助手。" action={ - + + + + + + + + 添加普通工具 + + + + 添加 MCP 服务 + + + } /> - -
} /> - + columns={columns} - rows={paginatedTools} - rowKey={(tool) => tool.id} + rows={paginatedResources} + rowKey={(resource) => resource.id} loading={loading} loadingText="正在加载工具…" error={error} - onRetry={() => void loadTools()} + onRetry={() => void loadResources()} empty={{ - title: tools.length === 0 ? "暂无工具资源" : "未找到匹配的工具资源", + title: resources.length === 0 ? "暂无工具资源" : "未找到匹配的工具资源", description: - tools.length === 0 + resources.length === 0 ? "点击右上角「添加工具」开始。" : "请调整关键词或筛选条件后再试。", }} @@ -611,15 +661,15 @@ export function ComponentsToolsPage() { totalPages, onPageChange: setCurrentPage, summary: - filteredTools.length === 0 + filteredResources.length === 0 ? "没有数据" - : `显示 ${pageStart + 1}-${Math.min(pageEnd, filteredTools.length)} / 共 ${filteredTools.length} 个工具资源`, + : `显示 ${pageStart + 1}-${Math.min(pageEnd, filteredResources.length)} / 共 ${filteredResources.length} 个工具资源`, }} />
- + {editing ? "编辑工具资源" : "添加工具资源"} @@ -629,6 +679,21 @@ export function ComponentsToolsPage() {
+ + + - - - - - -