// Package mcp 提供 MCP(Model Context Protocol)**Client** 能力:连接外部 MCP // server(容器/服务),发现工具、调用工具。 // // 定位(见 .cursor/skills/go-common/SKILL.md):本包只做「连接 + 调用」基础设施, // 不包含任何路由 / 参数抽取 / 编排逻辑 —— 那属于业务侧的 Prompt/Agent 编排, // 由消费方基于 Manager 暴露的 ListTools/CallTool 自行实现(可参考各业务项目里的 // ToolOrchestrator 写法)。 // // 传输基于官方 SDK(github.com/modelcontextprotocol/go-sdk),支持 // streamable-http / sse,绝不手写 JSON-RPC。 package mcp import ( "context" "encoding/json" "fmt" "net/http" "strings" "sync" "time" "git.toowon.com/jimmy/go-common/config" mcpsdk "github.com/modelcontextprotocol/go-sdk/mcp" ) const ( // clientName / clientVersion 在 initialize 握手时上报给 server,便于对端日志识别调用方。 clientName = "go-common-mcp-client" clientVersion = "1.0.0" ) // ToolDescriptor 从 MCP server 拉取的工具元信息。 type ToolDescriptor struct { Name string // 原始工具名(不含命名空间前缀) Description string // 工具说明 InputSchema json.RawMessage // JSON Schema 原文,供上层抽参/校验 } // ToolCallResult 单次工具调用结果。 type ToolCallResult struct { Text string // 合并后的文本内容(MCP text 内容块拼接) IsError bool // 工具自身报告的错误(非协议层错误) Raw json.RawMessage // 原始 CallToolResult JSON(含结构化内容/文件信息等) } // Client 单个 MCP server 连接的最小接口;可替换实现(参考 storage.Storage 的做法), // 便于测试时注入 fake 实现。 type Client interface { // Name 返回 server 逻辑名。 Name() string // ListTools 列出该 server 暴露的全部工具(带 inputSchema)。 ListTools(ctx context.Context) ([]ToolDescriptor, error) // CallTool 调用指定工具(name 为原始工具名,不含命名空间前缀)。 CallTool(ctx context.Context, name string, args map[string]any) (ToolCallResult, error) // Close 关闭会话并释放资源。 Close() error } // sdkClient 是 Client 的默认实现:封装与单个 MCP server 的会话。 // 惰性连接(首次调用才连)+ 失败后自动重连;调用超时不影响整个会话生命周期。 type sdkClient struct { cfg config.MCPServerConfig callTimeout time.Duration sdk *mcpsdk.Client mu sync.Mutex session *mcpsdk.ClientSession cancel context.CancelFunc } // NewClient 创建单个 MCP server 客户端(此时不建立连接,首次调用时惰性连接)。 func NewClient(cfg config.MCPServerConfig, callTimeout time.Duration) Client { if callTimeout <= 0 { callTimeout = 60 * time.Second } return &sdkClient{ cfg: cfg, callTimeout: callTimeout, sdk: mcpsdk.NewClient(&mcpsdk.Implementation{Name: clientName, Version: clientVersion}, nil), } } // Name 返回 server 逻辑名。 func (c *sdkClient) Name() string { return c.cfg.Name } // ListTools 列出 server 暴露的全部工具(带 inputSchema)。失败时重置会话以便下次重连。 func (c *sdkClient) ListTools(ctx context.Context) ([]ToolDescriptor, error) { session, err := c.ensureSession(ctx) if err != nil { return nil, err } opCtx, cancel := context.WithTimeout(ctx, c.callTimeout) defer cancel() res, err := session.ListTools(opCtx, nil) if err != nil { c.reset() return nil, fmt.Errorf("list tools from mcp server %q: %w", c.cfg.Name, err) } out := make([]ToolDescriptor, 0, len(res.Tools)) for _, t := range res.Tools { if t == nil { continue } schema, mErr := json.Marshal(t.InputSchema) if mErr != nil { schema = nil } out = append(out, ToolDescriptor{ Name: t.Name, Description: t.Description, InputSchema: schema, }) } return out, nil } // CallTool 调用指定工具(name 为原始工具名,不含命名空间前缀)。 func (c *sdkClient) CallTool(ctx context.Context, name string, args map[string]any) (ToolCallResult, error) { session, err := c.ensureSession(ctx) if err != nil { return ToolCallResult{}, err } opCtx, cancel := context.WithTimeout(ctx, c.callTimeout) defer cancel() res, err := session.CallTool(opCtx, &mcpsdk.CallToolParams{Name: name, Arguments: args}) if err != nil { c.reset() return ToolCallResult{}, fmt.Errorf("call tool %q on mcp server %q: %w", name, c.cfg.Name, err) } raw, _ := json.Marshal(res) return ToolCallResult{ Text: joinTextContent(res.Content), IsError: res.IsError, Raw: raw, }, nil } // Close 关闭会话并释放资源。 func (c *sdkClient) Close() error { c.mu.Lock() defer c.mu.Unlock() return c.closeLocked() } // ensureSession 确保会话已连接(SDK 在 Connect 内自动完成 initialize 握手),返回可用会话。 // // 注意:Connect 传入的 context 会被用于会话的整个生命周期(含后台监听),因此必须使用 // 后台 context 而非调用超时 context;否则调用结束取消 context 会断开会话。这里通过 // goroutine + 定时器为「连接握手」单独施加超时,超时则取消后台 context 放弃本次连接。 func (c *sdkClient) ensureSession(ctx context.Context) (*mcpsdk.ClientSession, error) { c.mu.Lock() defer c.mu.Unlock() if c.session != nil { return c.session, nil } transport, err := c.newTransport() if err != nil { return nil, err } baseCtx, cancel := context.WithCancel(context.Background()) type result struct { session *mcpsdk.ClientSession err error } ch := make(chan result, 1) go func() { s, e := c.sdk.Connect(baseCtx, transport, nil) ch <- result{session: s, err: e} }() select { case <-ctx.Done(): cancel() return nil, fmt.Errorf("connect mcp server %q canceled: %w", c.cfg.Name, ctx.Err()) case <-time.After(c.callTimeout): cancel() return nil, fmt.Errorf("connect mcp server %q timeout after %s", c.cfg.Name, c.callTimeout) case r := <-ch: if r.err != nil { cancel() return nil, fmt.Errorf("connect mcp server %q: %w", c.cfg.Name, r.err) } c.session = r.session c.cancel = cancel return r.session, nil } } // newTransport 按传输类型构造官方 SDK 传输;带鉴权时注入自定义 HTTP 客户端添加 Authorization 头。 func (c *sdkClient) newTransport() (mcpsdk.Transport, error) { httpClient := c.authHTTPClient() switch normalizeTransport(c.cfg.Transport) { case "sse": return &mcpsdk.SSEClientTransport{Endpoint: c.cfg.URL, HTTPClient: httpClient}, nil case "streamable-http": // DisableStandaloneSSE:仅做请求-响应式调用,不维持服务端主动推送的持久 SSE 流, // 兼容性更好、超时语义更清晰。 return &mcpsdk.StreamableClientTransport{ Endpoint: c.cfg.URL, HTTPClient: httpClient, DisableStandaloneSSE: true, }, nil default: return nil, fmt.Errorf("unsupported mcp transport %q for server %q", c.cfg.Transport, c.cfg.Name) } } // authHTTPClient 在配置了凭证时返回注入 Authorization 头的 HTTP 客户端;否则返回 nil(SDK 用默认客户端)。 func (c *sdkClient) authHTTPClient() *http.Client { value := authHeaderValue(c.cfg.AuthToken) if value == "" { return nil } return &http.Client{Transport: &authRoundTripper{base: http.DefaultTransport, value: value}} } // reset 重置会话,使下次调用重新连接(处理连接中断/对端重启)。 func (c *sdkClient) reset() { c.mu.Lock() defer c.mu.Unlock() _ = c.closeLocked() } // closeLocked 在持有 mu 的前提下关闭会话。 func (c *sdkClient) closeLocked() error { var err error if c.session != nil { err = c.session.Close() c.session = nil } if c.cancel != nil { c.cancel() c.cancel = nil } return err } // authRoundTripper 为每个出站请求注入固定的 Authorization 头。 type authRoundTripper struct { base http.RoundTripper value string } func (a *authRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) { r := req.Clone(req.Context()) r.Header.Set("Authorization", a.value) base := a.base if base == nil { base = http.DefaultTransport } return base.RoundTrip(r) } // joinTextContent 拼接 MCP 返回内容中的所有 text 块(忽略图片/音频等其他类型)。 func joinTextContent(contents []mcpsdk.Content) string { var b strings.Builder for _, ct := range contents { if tc, ok := ct.(*mcpsdk.TextContent); ok { if b.Len() > 0 { b.WriteString("\n") } b.WriteString(tc.Text) } } return b.String() } // normalizeTransport 归一化传输类型,空值/http 视为 streamable-http。 func normalizeTransport(t string) string { switch strings.TrimSpace(strings.ToLower(t)) { case "", "http", "streamable-http", "streamablehttp": return "streamable-http" case "sse": return "sse" default: return strings.TrimSpace(strings.ToLower(t)) } } // authHeaderValue 归一化鉴权头:已带 Bearer/Basic 前缀则原样使用,否则补 "Bearer "。 func authHeaderValue(token string) string { token = strings.TrimSpace(token) if token == "" { return "" } lower := strings.ToLower(token) if strings.HasPrefix(lower, "bearer ") || strings.HasPrefix(lower, "basic ") { return token } return "Bearer " + token }