300 lines
9.2 KiB
Go
300 lines
9.2 KiB
Go
// 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
|
||
}
|