go-common中增加mcp抽离成通用方法(编排由业务方自行处理)
This commit is contained in:
299
mcp/client.go
Normal file
299
mcp/client.go
Normal file
@@ -0,0 +1,299 @@
|
||||
// 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
|
||||
}
|
||||
Reference in New Issue
Block a user