Files
ai-gateway-go/internal/workbench/mcp_client.go
T
superidou 5759c1862e AI Gateway Go 0.10.0 源码快照 + 旗舰版需求规划报告
M0-M7 已完成:核心网关(身份/RBAC/TOTP/OIDC/SAML/Provider/配额/路由/内容策略/审计/定价)+ 资源市场(MCP/Skills/数字员工)。
含 22 个 PostgreSQL 迁移、管理端/门户端前端源码、OpenAPI 契约、部署 compose。

Co-Authored-By: Claude <noreply@anthropic.com>
2026-08-12 11:45:54 +08:00

370 lines
12 KiB
Go

package workbench
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"strings"
"sync"
"time"
)
// mcpProtocolVersion is the Model Context Protocol version this client speaks.
// Streamable HTTP (2025-06-18) is the current stable transport; the server may
// negotiate an older one and we accept whatever it replies with.
const mcpProtocolVersion = "2025-06-18"
// MCPTool is one tool discovered from an MCP server via tools/list.
type MCPTool struct {
Name string `json:"name"`
Description string `json:"description"`
InputSchema json.RawMessage `json:"inputSchema"`
}
// MCPToolResult is the outcome of a tools/call.
type MCPToolResult struct {
Content string
IsError bool
}
// mcpServerState caches the negotiated session and discovered tool list for one
// server, keyed by server ID. initialize is expensive (two round trips) and
// stateless HTTP servers expect it per session, so we keep the session warm for
// cacheTTL and only re-handshake once it lapses.
type mcpServerState struct {
initAt time.Time
sessionID string
tools []MCPTool
toolsAt time.Time
}
// MCPClient is a minimal Model Context Protocol client over the streamable HTTP
// (and legacy SSE) transports. It speaks just enough of the protocol for the
// resource marketplace: initialize + notifications/initialized, tools/list for
// discovery, and tools/call for execution.
type MCPClient struct {
client *http.Client
cacheTTL time.Duration
mu sync.Mutex
states map[string]*mcpServerState
}
func NewMCPClient(allowPrivate bool, cacheTTL time.Duration) *MCPClient {
if cacheTTL <= 0 {
cacheTTL = 60 * time.Second
}
return &MCPClient{
client: &http.Client{
Timeout: 15 * time.Second,
Transport: &http.Transport{
DialContext: safeToolDial(allowPrivate),
ForceAttemptHTTP2: true,
TLSHandshakeTimeout: 5 * time.Second,
ResponseHeaderTimeout: 10 * time.Second,
MaxIdleConns: 64,
MaxIdleConnsPerHost: 16,
IdleConnTimeout: 90 * time.Second,
},
CheckRedirect: func(*http.Request, []*http.Request) error { return errors.New("MCP 服务器不允许重定向") },
},
cacheTTL: cacheTTL,
states: make(map[string]*mcpServerState),
}
}
// DiscoverTools returns the tools an MCP server advertises, reusing a cached
// list for cacheTTL. headers are the already-decrypted request headers (e.g.
// Authorization) for this server.
func (c *MCPClient) DiscoverTools(ctx context.Context, server MCPServer, headers map[string]string) ([]MCPTool, error) {
state, err := c.ensureInitialized(ctx, server, headers)
if err != nil {
return nil, err
}
if state.tools != nil && time.Since(state.toolsAt) < c.cacheTTL {
return state.tools, nil
}
result, err := c.call(ctx, server, headers, "tools/list", map[string]any{})
if err != nil {
return nil, err
}
var list struct {
Tools []MCPTool `json:"tools"`
}
if err = json.Unmarshal(result, &list); err != nil {
return nil, fmt.Errorf("MCP tools/list 响应无效: %w", err)
}
for i := range list.Tools {
if list.Tools[i].Name == "" {
return nil, errors.New("MCP 服务器返回了没有名称的工具")
}
if len(list.Tools[i].InputSchema) == 0 {
list.Tools[i].InputSchema = json.RawMessage(`{}`)
}
}
c.mu.Lock()
state.tools = list.Tools
state.toolsAt = time.Now()
c.mu.Unlock()
return list.Tools, nil
}
// CallTool invokes one tool on an MCP server and returns the concatenated text
// content. An isError result is surfaced as an error so callers treat it as a
// failed tool round rather than a successful empty answer.
func (c *MCPClient) CallTool(ctx context.Context, server MCPServer, headers map[string]string, name string, args map[string]any) (MCPToolResult, error) {
state, err := c.ensureInitialized(ctx, server, headers)
if err != nil {
return MCPToolResult{}, err
}
// The runtime exposes tools under the collision-proof prefix
// (mcp__{serverCode}__{toolName}); strip it before the wire call since the
// remote server only knows the unprefixed tool name.
if _, resolved, ok := resolveMCPTool(name); ok {
name = resolved
}
params := map[string]any{"name": name, "arguments": args}
result, err := c.call(ctx, server, headers, "tools/call", params)
if err != nil {
return MCPToolResult{}, err
}
var called struct {
Content []struct {
Type string `json:"type"`
Text string `json:"text"`
} `json:"content"`
IsError bool `json:"isError"`
}
if err = json.Unmarshal(result, &called); err != nil {
return MCPToolResult{}, fmt.Errorf("MCP tools/call 响应无效: %w", err)
}
text := strings.Builder{}
for _, block := range called.Content {
if block.Type == "text" {
if text.Len() > 0 {
text.WriteString("\n")
}
text.WriteString(block.Text)
}
}
_ = state // keep state warm check semantics obvious
if called.IsError {
return MCPToolResult{}, errors.New("MCP 工具执行失败: " + text.String())
}
return MCPToolResult{Content: text.String()}, nil
}
// ensureInitialized performs the MCP initialize handshake for a server if its
// session has lapsed (or no cached tools exist yet), then acknowledges with
// notifications/initialized. The handshake is guarded by the per-server cache
// so a burst of calls does not re-initialize every request.
func (c *MCPClient) ensureInitialized(ctx context.Context, server MCPServer, headers map[string]string) (*mcpServerState, error) {
c.mu.Lock()
state, ok := c.states[server.ID]
if ok && time.Since(state.initAt) < c.cacheTTL {
c.mu.Unlock()
return state, nil
}
c.mu.Unlock()
result, headersOut, err := c.handshake(ctx, server, headers)
if err != nil {
return nil, err
}
sessionID := headersOut.Get("Mcp-Session-Id")
if server.Transport != "" && server.Transport != "streamable-http" && server.Transport != "sse" {
return nil, fmt.Errorf("不支持的 MCP 传输方式 %s", server.Transport)
}
_ = result // negotiated protocol version is accepted as-is
c.mu.Lock()
state = &mcpServerState{initAt: time.Now(), sessionID: sessionID}
c.states[server.ID] = state
c.mu.Unlock()
// Best-effort acknowledgment; servers that require it will reject later
// calls and we will surface that error naturally.
c.sendNotification(ctx, server, headers, sessionID)
return state, nil
}
func (c *MCPClient) handshake(ctx context.Context, server MCPServer, headers map[string]string) (json.RawMessage, http.Header, error) {
params := map[string]any{
"protocolVersion": mcpProtocolVersion,
"capabilities": map[string]any{},
"clientInfo": map[string]any{"name": "LLMGuardX语枢", "version": "0.10.0"},
}
payload, _ := json.Marshal(mcpRequest{JSONRPC: "2.0", ID: 1, Method: "initialize", Params: params})
request, err := http.NewRequestWithContext(ctx, http.MethodPost, server.EndpointURL, bytes.NewReader(payload))
if err != nil {
return nil, nil, err
}
c.prepare(request, headers, "")
response, err := c.client.Do(request)
if err != nil {
return nil, nil, fmt.Errorf("MCP 连接失败: %w", err)
}
defer response.Body.Close()
result, err := readMCPBody(response)
if err != nil {
return nil, nil, err
}
return result, response.Header, nil
}
func (c *MCPClient) sendNotification(ctx context.Context, server MCPServer, headers map[string]string, sessionID string) {
payload, _ := json.Marshal(mcpRequest{JSONRPC: "2.0", ID: nil, Method: "notifications/initialized"})
request, err := http.NewRequestWithContext(ctx, http.MethodPost, server.EndpointURL, bytes.NewReader(payload))
if err != nil {
return
}
c.prepare(request, headers, sessionID)
response, err := c.client.Do(request)
if err != nil {
return
}
_, _ = io.Copy(io.Discard, io.LimitReader(response.Body, 1<<10))
response.Body.Close()
}
func (c *MCPClient) call(ctx context.Context, server MCPServer, headers map[string]string, method string, params any) (json.RawMessage, error) {
c.mu.Lock()
sessionID := ""
if state, ok := c.states[server.ID]; ok {
sessionID = state.sessionID
}
c.mu.Unlock()
payload, _ := json.Marshal(mcpRequest{JSONRPC: "2.0", ID: 1, Method: method, Params: params})
request, err := http.NewRequestWithContext(ctx, http.MethodPost, server.EndpointURL, bytes.NewReader(payload))
if err != nil {
return nil, err
}
c.prepare(request, headers, sessionID)
response, err := c.client.Do(request)
if err != nil {
return nil, fmt.Errorf("MCP 调用失败: %w", err)
}
defer response.Body.Close()
return readMCPBody(response)
}
func (c *MCPClient) prepare(request *http.Request, headers map[string]string, sessionID string) {
request.Header.Set("Content-Type", "application/json")
request.Header.Set("Accept", "application/json, text/event-stream")
for key, value := range headers {
request.Header.Set(key, value)
}
if sessionID != "" {
request.Header.Set("Mcp-Session-Id", sessionID)
}
}
type mcpRequest struct {
JSONRPC string `json:"jsonrpc"`
ID any `json:"id"`
Method string `json:"method"`
Params any `json:"params,omitempty"`
}
type mcpError struct {
Code int `json:"code"`
Message string `json:"message"`
}
type mcpResponse struct {
JSONRPC string `json:"jsonrpc"`
ID any `json:"id"`
Result json.RawMessage `json:"result"`
Error *mcpError `json:"error"`
}
// readMCPBody reads and parses a single JSON-RPC response. Streamable HTTP
// servers return application/json; legacy SSE servers stream data frames, from
// which the first complete JSON object is extracted.
func readMCPBody(response *http.Response) (json.RawMessage, error) {
code := response.StatusCode
raw, err := io.ReadAll(io.LimitReader(response.Body, (1<<20)+1))
if err != nil {
return nil, err
}
if len(raw) > 1<<20 {
return nil, errors.New("MCP 响应超过 1 MiB")
}
if code < 200 || code >= 300 {
return nil, fmt.Errorf("MCP 服务器返回 HTTP %d", code)
}
var body []byte
if strings.HasPrefix(response.Header.Get("Content-Type"), "text/event-stream") {
body, err = extractSSEJSON(raw)
if err != nil {
return nil, err
}
} else {
body = raw
}
var resp mcpResponse
if err = json.Unmarshal(body, &resp); err != nil {
return nil, fmt.Errorf("MCP 响应不是有效 JSON-RPC: %w", err)
}
if resp.Error != nil {
return nil, fmt.Errorf("MCP 服务器错误 (%d): %s", resp.Error.Code, resp.Error.Message)
}
if len(resp.Result) == 0 {
return nil, errors.New("MCP 服务器返回空结果")
}
return resp.Result, nil
}
// extractSSEJSON concatenates the first data frame's payload into one JSON
// document. MCP streamable-HTTP servers emit a single frame per request, so
// only the first frame boundary is consumed and the rest is ignored.
func extractSSEJSON(raw []byte) ([]byte, error) {
var data strings.Builder
for _, line := range strings.Split(string(raw), "\n") {
trimmed := strings.TrimSpace(line)
switch {
case strings.HasPrefix(trimmed, "data:"):
value := strings.TrimSpace(strings.TrimPrefix(trimmed, "data:"))
if value == "" {
continue
}
if data.Len() > 0 {
data.WriteString("\n")
}
data.WriteString(value)
case trimmed == "" && data.Len() > 0:
if json.Valid([]byte(data.String())) {
return []byte(data.String()), nil
}
data.Reset()
}
}
if data.Len() > 0 && json.Valid([]byte(data.String())) {
return []byte(data.String()), nil
}
return nil, errors.New("MCP 服务器未返回有效的 SSE 数据帧")
}
// toolCallPrefix namespaces MCP tools inside a shared tool list so different
// servers cannot collide. Format: mcp__{serverCode}__{toolName}.
const toolCallPrefix = "mcp__"
func mcpToolName(serverCode, tool string) string { return toolCallPrefix + serverCode + "__" + tool }
// resolveMCPTool splits a prefixed tool name back into its MCP server code and
// tool name. Returns ok=false for names that are not MCP-prefixed.
func resolveMCPTool(prefixed string) (serverCode, tool string, ok bool) {
if !strings.HasPrefix(prefixed, toolCallPrefix) {
return "", "", false
}
rest := strings.TrimPrefix(prefixed, toolCallPrefix)
parts := strings.SplitN(rest, "__", 2)
if len(parts) != 2 || parts[0] == "" || parts[1] == "" {
return "", "", false
}
return parts[0], parts[1], true
}