go-common中增加mcp抽离成通用方法(编排由业务方自行处理)

This commit is contained in:
2026-07-25 11:27:34 +08:00
parent 121928733b
commit 97450f0739
14 changed files with 1112 additions and 9 deletions

299
mcp/client.go Normal file
View File

@@ -0,0 +1,299 @@
// Package mcp 提供 MCPModel Context Protocol**Client** 能力:连接外部 MCP
// server容器/服务),发现工具、调用工具。
//
// 定位(见 .cursor/skills/go-common/SKILL.md本包只做「连接 + 调用」基础设施,
// 不包含任何路由 / 参数抽取 / 编排逻辑 —— 那属于业务侧的 Prompt/Agent 编排,
// 由消费方基于 Manager 暴露的 ListTools/CallTool 自行实现(可参考各业务项目里的
// ToolOrchestrator 写法)。
//
// 传输基于官方 SDKgithub.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 客户端;否则返回 nilSDK 用默认客户端)。
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
}