UNPKG

@huggingface/mcp-client

Version:
218 lines (197 loc) 6.66 kB
import { Client } from "@modelcontextprotocol/sdk/client/index.js"; import type { StdioServerParameters } from "@modelcontextprotocol/sdk/client/stdio.js"; import { StdioClientTransport } from "@modelcontextprotocol/sdk/client/stdio.js"; import { InferenceClient } from "@huggingface/inference"; import type { InferenceProviderOrPolicy } from "@huggingface/inference"; import type { ChatCompletionInputMessage, ChatCompletionInputTool, ChatCompletionStreamOutput, ChatCompletionStreamOutputDeltaToolCall, } from "@huggingface/tasks/src/tasks/chat-completion/inference"; import { version as packageVersion } from "../package.json"; import { debug } from "./utils"; import type { ServerConfig } from "./types"; import type { Transport } from "@modelcontextprotocol/sdk/shared/transport"; import { SSEClientTransport } from "@modelcontextprotocol/sdk/client/sse.js"; import { StreamableHTTPClientTransport } from "@modelcontextprotocol/sdk/client/streamableHttp.js"; import { ResultFormatter } from "./ResultFormatter.js"; type ToolName = string; export interface ChatCompletionInputMessageTool extends ChatCompletionInputMessage { role: "tool"; tool_call_id: string; content: string; name?: string; } export class McpClient { protected client: InferenceClient; protected provider: InferenceProviderOrPolicy | undefined; protected model: string; private clients: Map<ToolName, Client> = new Map(); public readonly availableTools: ChatCompletionInputTool[] = []; constructor({ provider, endpointUrl, model, apiKey, }: ( | { provider: InferenceProviderOrPolicy; endpointUrl?: undefined; } | { endpointUrl: string; provider?: undefined; } ) & { model: string; apiKey: string; }) { this.client = endpointUrl ? new InferenceClient(apiKey, { endpointUrl: endpointUrl }) : new InferenceClient(apiKey); this.provider = provider; this.model = model; } async addMcpServers(servers: (ServerConfig | StdioServerParameters)[]): Promise<void> { await Promise.all(servers.map((s) => this.addMcpServer(s))); } async addMcpServer(server: ServerConfig | StdioServerParameters): Promise<void> { let transport: Transport; const asUrl = (url: string | URL): URL => { return typeof url === "string" ? new URL(url) : url; }; if (!("type" in server)) { transport = new StdioClientTransport({ ...server, env: { ...server.env, PATH: process.env.PATH ?? "" }, }); } else { switch (server.type) { case "stdio": transport = new StdioClientTransport({ ...server.config, env: { ...server.config.env, PATH: process.env.PATH ?? "" }, }); break; case "sse": transport = new SSEClientTransport(asUrl(server.config.url), server.config.options); break; case "http": transport = new StreamableHTTPClientTransport(asUrl(server.config.url), server.config.options); break; } } const mcp = new Client({ name: "@huggingface/mcp-client", version: packageVersion }); await mcp.connect(transport); const toolsResult = await mcp.listTools(); debug( "Connected to server with tools:", toolsResult.tools.map(({ name }) => name) ); for (const tool of toolsResult.tools) { this.clients.set(tool.name, mcp); } this.availableTools.push( ...toolsResult.tools.map((tool) => { return { type: "function", function: { name: tool.name, description: tool.description, parameters: tool.inputSchema, }, } satisfies ChatCompletionInputTool; }) ); } async *processSingleTurnWithTools( messages: ChatCompletionInputMessage[], opts: { exitLoopTools?: ChatCompletionInputTool[]; exitIfFirstChunkNoTool?: boolean; abortSignal?: AbortSignal; } = {} ): AsyncGenerator<ChatCompletionStreamOutput | ChatCompletionInputMessageTool> { debug("start of single turn"); const stream = this.client.chatCompletionStream({ provider: this.provider, model: this.model, messages, tools: opts.exitLoopTools ? [...opts.exitLoopTools, ...this.availableTools] : this.availableTools, tool_choice: "auto", signal: opts.abortSignal, }); const message = { role: "unknown", content: "", } satisfies ChatCompletionInputMessage; const finalToolCalls: Record<number, ChatCompletionStreamOutputDeltaToolCall> = {}; let numOfChunks = 0; for await (const chunk of stream) { if (opts.abortSignal?.aborted) { throw new Error("AbortError"); } yield chunk; debug(chunk.choices[0]); numOfChunks++; const delta = chunk.choices[0]?.delta; if (!delta) { continue; } if (delta.role) { message.role = delta.role; } if (delta.content) { message.content += delta.content; } for (const toolCall of delta.tool_calls ?? []) { // aggregating chunks into an encoded arguments JSON object if (!finalToolCalls[toolCall.index]) { finalToolCalls[toolCall.index] = toolCall; } if (finalToolCalls[toolCall.index].function.arguments === undefined) { finalToolCalls[toolCall.index].function.arguments = ""; } if (toolCall.function.arguments) { finalToolCalls[toolCall.index].function.arguments += toolCall.function.arguments; } } if (opts.exitIfFirstChunkNoTool && numOfChunks <= 2 && Object.keys(finalToolCalls).length === 0) { /// If no tool is present in chunk number 1 or 2, exit. return; } } messages.push(message); for (const toolCall of Object.values(finalToolCalls)) { const toolName = toolCall.function.name ?? "unknown"; /// TODO(Fix upstream type so this is always a string)^ const toolArgs = toolCall.function.arguments === "" ? {} : JSON.parse(toolCall.function.arguments); const toolMessage: ChatCompletionInputMessageTool = { role: "tool", tool_call_id: toolCall.id, content: "", name: toolName, }; if (opts.exitLoopTools?.map((t) => t.function.name).includes(toolName)) { messages.push(toolMessage); return yield toolMessage; } /// Get the appropriate session for this tool const client = this.clients.get(toolName); if (client) { const result = await client.callTool({ name: toolName, arguments: toolArgs, signal: opts.abortSignal }); toolMessage.content = ResultFormatter.format(result); } else { toolMessage.content = `Error: No session found for tool: ${toolName}`; } messages.push(toolMessage); yield toolMessage; } } async cleanup(): Promise<void> { const clients = new Set(this.clients.values()); await Promise.all([...clients].map((client) => client.close())); } async [Symbol.dispose](): Promise<void> { return this.cleanup(); } }