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
|
||||
}
|
||||
364
mcp/manager.go
Normal file
364
mcp/manager.go
Normal file
@@ -0,0 +1,364 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"git.toowon.com/jimmy/go-common/config"
|
||||
"git.toowon.com/jimmy/go-common/logger"
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
// toolNameSeparator 命名空间分隔符:FQName = <prefix>__<rawName>,防多 server 工具重名。
|
||||
const toolNameSeparator = "__"
|
||||
|
||||
// toolCallLogMaxRunes 工具返回写入日志时的最大字符数,避免大返回值把日志灌爆。
|
||||
const toolCallLogMaxRunes = 2048
|
||||
|
||||
// Tool 暴露给上层(业务编排/接口)的工具元信息,带命名空间全名与所属 server。
|
||||
type Tool struct {
|
||||
FQName string `json:"fq_name"` // 命名空间全名,如 word__create_document
|
||||
Name string `json:"name"` // 原始工具名(不含前缀)
|
||||
ServerName string `json:"server_name"` // 所属 MCP server 逻辑名
|
||||
Description string `json:"description,omitempty"` // 工具说明
|
||||
InputSchema json.RawMessage `json:"input_schema,omitempty"` // JSON Schema 原文,供上层抽参/校验
|
||||
}
|
||||
|
||||
// ToolResult 工具调用结果。
|
||||
type ToolResult struct {
|
||||
Text string `json:"text"` // 合并后的文本内容
|
||||
IsError bool `json:"is_error"` // 工具自身报告的错误
|
||||
Raw json.RawMessage `json:"raw,omitempty"` // 原始结果 JSON(含文件名/路径等结构化信息)
|
||||
}
|
||||
|
||||
// Manager 是 MCP 工具子系统的标准模块对象(`app.MCP()` 取得)。
|
||||
//
|
||||
// 职责:
|
||||
// - 按配置管理多个 MCP server 客户端;
|
||||
// - 工具加命名空间前缀防冲突;
|
||||
// - CallTool 路由到对应 server;
|
||||
// - allowedTools 白名单收敛;
|
||||
// - 对外暴露工具的 InputSchema(供业务侧抽参)。
|
||||
//
|
||||
// 本接口是「基础底座」:仅负责工具的发现与调用,不包含路由/抽参/编排逻辑
|
||||
// (那属于业务侧的 Prompt/Agent 编排,由消费方基于本接口自行实现)。
|
||||
type Manager interface {
|
||||
// Enabled 子系统是否启用(关闭或无可用 server 时为 false)。
|
||||
Enabled() bool
|
||||
// ListTools 聚合所有 server 的可用工具(已按白名单过滤、带命名空间前缀)。
|
||||
ListTools(ctx context.Context) ([]Tool, error)
|
||||
// CallTool 按命名空间全名路由并调用工具。
|
||||
CallTool(ctx context.Context, fqName string, args map[string]any) (ToolResult, error)
|
||||
// Close 关闭全部底层 server 连接。
|
||||
Close() error
|
||||
}
|
||||
|
||||
// serverEntry 单个 server 的运行态:底层客户端 + 命名空间前缀 + 白名单。
|
||||
type serverEntry struct {
|
||||
name string
|
||||
prefix string
|
||||
allowed map[string]bool // 空表示放开全部工具
|
||||
client Client
|
||||
}
|
||||
|
||||
// isAllowed 判断原始工具名是否在白名单内(白名单为空=全部放开)。
|
||||
func (e *serverEntry) isAllowed(rawName string) bool {
|
||||
if len(e.allowed) == 0 {
|
||||
return true
|
||||
}
|
||||
return e.allowed[rawName]
|
||||
}
|
||||
|
||||
// route 命名空间全名到「server + 原始工具名」的路由。
|
||||
type route struct {
|
||||
server *serverEntry
|
||||
rawName string
|
||||
}
|
||||
|
||||
// manager 是 Manager 的默认实现。
|
||||
type manager struct {
|
||||
log *logger.Logger
|
||||
enabled bool
|
||||
servers []*serverEntry // 稳定顺序,便于路由前缀匹配
|
||||
|
||||
mu sync.RWMutex
|
||||
routes map[string]route // FQName -> 路由缓存(ListTools 时刷新,CallTool 时按需补全)
|
||||
}
|
||||
|
||||
// NewManager 根据配置构建 MCP 工具子系统。
|
||||
// - log:go-common 日志对象(可为 nil,仅跳过调用日志);
|
||||
// - cfg:nil 或 Enabled=false 时返回的 Manager.Enabled() 恒为 false,其余方法安全返回空结果。
|
||||
//
|
||||
// 此处仅创建客户端对象(惰性连接),不会在启动时连接外部 MCP server,避免外部服务未就绪
|
||||
// 导致本服务启动失败。
|
||||
func NewManager(log *logger.Logger, cfg *config.MCPConfig) Manager {
|
||||
m := &manager{log: log, routes: make(map[string]route)}
|
||||
if cfg == nil {
|
||||
return m
|
||||
}
|
||||
m.enabled = cfg.Enabled
|
||||
|
||||
timeout := time.Duration(cfg.CallTimeoutSeconds) * time.Second
|
||||
for _, sc := range cfg.Servers {
|
||||
if !sc.IsEnabled() {
|
||||
continue
|
||||
}
|
||||
m.servers = append(m.servers, newServerEntry(sc, NewClient(sc, timeout)))
|
||||
}
|
||||
return m
|
||||
}
|
||||
|
||||
// newServerEntry 组装 server 运行态(前缀归一化 + 白名单集合化)。
|
||||
func newServerEntry(sc config.MCPServerConfig, client Client) *serverEntry {
|
||||
prefix := strings.TrimSpace(sc.ToolNamePrefix)
|
||||
if prefix == "" {
|
||||
prefix = strings.TrimSpace(sc.Name)
|
||||
}
|
||||
allowed := make(map[string]bool, len(sc.AllowedTools))
|
||||
for _, t := range sc.AllowedTools {
|
||||
if t = strings.TrimSpace(t); t != "" {
|
||||
allowed[t] = true
|
||||
}
|
||||
}
|
||||
return &serverEntry{name: sc.Name, prefix: prefix, allowed: allowed, client: client}
|
||||
}
|
||||
|
||||
// Enabled 子系统启用且至少配置了一个 server 时才视为可用。
|
||||
func (m *manager) Enabled() bool {
|
||||
return m != nil && m.enabled && len(m.servers) > 0
|
||||
}
|
||||
|
||||
// ListTools 聚合全部 server 工具。单个 server 失败不影响其余 server(仅记录告警),保证整体健壮性。
|
||||
func (m *manager) ListTools(ctx context.Context) ([]Tool, error) {
|
||||
if !m.Enabled() {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
// 重建路由缓存,保证与最新工具列表一致。
|
||||
m.routes = make(map[string]route)
|
||||
|
||||
var out []Tool
|
||||
for _, srv := range m.servers {
|
||||
tools, err := srv.client.ListTools(ctx)
|
||||
if err != nil {
|
||||
m.logError("mcp list tools from server failed", err, map[string]any{"server": srv.name})
|
||||
continue
|
||||
}
|
||||
for _, t := range tools {
|
||||
if !srv.isAllowed(t.Name) {
|
||||
continue
|
||||
}
|
||||
fq := srv.prefix + toolNameSeparator + t.Name
|
||||
m.routes[fq] = route{server: srv, rawName: t.Name}
|
||||
out = append(out, Tool{
|
||||
FQName: fq,
|
||||
Name: t.Name,
|
||||
ServerName: srv.name,
|
||||
Description: t.Description,
|
||||
InputSchema: t.InputSchema,
|
||||
})
|
||||
}
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// CallTool 按命名空间全名路由到对应 server 并调用。
|
||||
func (m *manager) CallTool(ctx context.Context, fqName string, args map[string]any) (ToolResult, error) {
|
||||
if !m.Enabled() {
|
||||
return ToolResult{}, fmt.Errorf("mcp subsystem disabled")
|
||||
}
|
||||
r, ok := m.resolve(fqName)
|
||||
if !ok {
|
||||
return ToolResult{}, fmt.Errorf("unknown mcp tool %q", fqName)
|
||||
}
|
||||
// 白名单二次校验,防止绕过 ListTools 直接调用未开放工具。
|
||||
if !r.server.isAllowed(r.rawName) {
|
||||
return ToolResult{}, fmt.Errorf("mcp tool %q is not allowed", fqName)
|
||||
}
|
||||
|
||||
callID := uuid.NewString()
|
||||
requestID := RequestIDFromContext(ctx)
|
||||
m.logToolCallAsync(toolCallLogSnapshot{
|
||||
phase: "request", requestID: requestID, callID: callID,
|
||||
fqName: fqName, server: r.server.name, rawTool: r.rawName,
|
||||
args: cloneAnyMap(args),
|
||||
})
|
||||
|
||||
res, err := r.server.client.CallTool(ctx, r.rawName, args)
|
||||
if err != nil {
|
||||
m.logToolCallAsync(toolCallLogSnapshot{
|
||||
phase: "error", requestID: requestID, callID: callID,
|
||||
fqName: fqName, server: r.server.name, rawTool: r.rawName,
|
||||
args: cloneAnyMap(args), errMsg: err.Error(),
|
||||
})
|
||||
return ToolResult{}, err
|
||||
}
|
||||
out := ToolResult{Text: res.Text, IsError: res.IsError, Raw: res.Raw}
|
||||
m.logToolCallAsync(toolCallLogSnapshot{
|
||||
phase: "response", requestID: requestID, callID: callID,
|
||||
fqName: fqName, server: r.server.name, rawTool: r.rawName,
|
||||
args: cloneAnyMap(args), hasResult: true, result: out,
|
||||
})
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// resolve 解析命名空间全名到路由:优先读缓存,未命中则按 server 前缀推导(无需先 ListTools)。
|
||||
func (m *manager) resolve(fqName string) (route, bool) {
|
||||
m.mu.RLock()
|
||||
r, ok := m.routes[fqName]
|
||||
m.mu.RUnlock()
|
||||
if ok {
|
||||
return r, true
|
||||
}
|
||||
|
||||
for _, srv := range m.servers {
|
||||
pfx := srv.prefix + toolNameSeparator
|
||||
if strings.HasPrefix(fqName, pfx) {
|
||||
raw := strings.TrimPrefix(fqName, pfx)
|
||||
if raw == "" {
|
||||
continue
|
||||
}
|
||||
r := route{server: srv, rawName: raw}
|
||||
m.mu.Lock()
|
||||
m.routes[fqName] = r
|
||||
m.mu.Unlock()
|
||||
return r, true
|
||||
}
|
||||
}
|
||||
return route{}, false
|
||||
}
|
||||
|
||||
// Close 关闭全部底层 server 连接。
|
||||
func (m *manager) Close() error {
|
||||
var firstErr error
|
||||
for _, srv := range m.servers {
|
||||
if err := srv.client.Close(); err != nil && firstErr == nil {
|
||||
firstErr = err
|
||||
}
|
||||
}
|
||||
return firstErr
|
||||
}
|
||||
|
||||
// --- 可选的请求追踪 ID(不强依赖任何业务 requestctx 包) ---
|
||||
|
||||
type requestIDKey struct{}
|
||||
|
||||
// WithRequestID 可选:将请求追踪 ID 写入 context,CallTool 的调用日志会附带该字段。
|
||||
// 典型用法:HTTP 层已经有 middleware.RequestID 生成的 ID 时,
|
||||
// `ctx = mcp.WithRequestID(ctx, middleware.GetRequestID(r.Context()))`。
|
||||
// 不调用时日志里的 request_id 为空,不影响功能。
|
||||
func WithRequestID(ctx context.Context, id string) context.Context {
|
||||
return context.WithValue(ctx, requestIDKey{}, id)
|
||||
}
|
||||
|
||||
// RequestIDFromContext 读取通过 WithRequestID 写入的请求追踪 ID;未设置时返回空字符串。
|
||||
func RequestIDFromContext(ctx context.Context) string {
|
||||
if v, ok := ctx.Value(requestIDKey{}).(string); ok {
|
||||
return v
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// --- 异步调用日志 ---
|
||||
|
||||
// toolCallLogSnapshot 异步日志快照;主流程只组装结构体,重活(截断/序列化/写日志)在 goroutine 内完成。
|
||||
type toolCallLogSnapshot struct {
|
||||
phase string
|
||||
requestID string
|
||||
callID string
|
||||
fqName string
|
||||
server string
|
||||
rawTool string
|
||||
args map[string]any
|
||||
hasResult bool
|
||||
result ToolResult
|
||||
errMsg string
|
||||
}
|
||||
|
||||
func (m *manager) logToolCallAsync(snap toolCallLogSnapshot) {
|
||||
if m.log == nil {
|
||||
return
|
||||
}
|
||||
go writeToolCallLog(m.log, snap)
|
||||
}
|
||||
|
||||
func (m *manager) logError(message string, err error, fields map[string]any) {
|
||||
if m.log == nil {
|
||||
return
|
||||
}
|
||||
if fields == nil {
|
||||
fields = make(map[string]any)
|
||||
}
|
||||
fields["error"] = err.Error()
|
||||
m.log.Error(message, fields)
|
||||
}
|
||||
|
||||
func writeToolCallLog(l *logger.Logger, snap toolCallLogSnapshot) {
|
||||
payload := map[string]any{
|
||||
"type": "mcp_tool_call",
|
||||
"phase": snap.phase,
|
||||
"request_id": snap.requestID,
|
||||
"call_id": snap.callID,
|
||||
"tool": snap.fqName,
|
||||
"server": snap.server,
|
||||
"raw_tool": snap.rawTool,
|
||||
}
|
||||
if len(snap.args) > 0 {
|
||||
payload["arguments"] = snap.args
|
||||
}
|
||||
if snap.hasResult {
|
||||
text := snap.result.Text
|
||||
totalRunes := len([]rune(text))
|
||||
truncated := totalRunes > toolCallLogMaxRunes
|
||||
if truncated {
|
||||
text = truncateRunes(text, toolCallLogMaxRunes)
|
||||
}
|
||||
resp := map[string]any{
|
||||
"is_error": snap.result.IsError,
|
||||
"text": text,
|
||||
"text_total_runes": totalRunes,
|
||||
"text_truncated": truncated,
|
||||
}
|
||||
if len(snap.result.Raw) > 0 {
|
||||
rawStr := string(snap.result.Raw)
|
||||
rawRunes := len([]rune(rawStr))
|
||||
if rawRunes > toolCallLogMaxRunes {
|
||||
resp["raw"] = truncateRunes(rawStr, toolCallLogMaxRunes)
|
||||
resp["raw_total_runes"] = rawRunes
|
||||
resp["raw_truncated"] = true
|
||||
} else {
|
||||
resp["raw"] = json.RawMessage(snap.result.Raw)
|
||||
}
|
||||
}
|
||||
payload["response"] = resp
|
||||
}
|
||||
if snap.errMsg != "" {
|
||||
payload["error"] = snap.errMsg
|
||||
}
|
||||
l.Info("mcp tool call", payload)
|
||||
}
|
||||
|
||||
func cloneAnyMap(m map[string]any) map[string]any {
|
||||
if len(m) == 0 {
|
||||
return nil
|
||||
}
|
||||
out := make(map[string]any, len(m))
|
||||
for k, v := range m {
|
||||
out[k] = v
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func truncateRunes(s string, max int) string {
|
||||
r := []rune(s)
|
||||
if len(r) <= max {
|
||||
return s
|
||||
}
|
||||
return string(r[:max]) + "…"
|
||||
}
|
||||
200
mcp/manager_test.go
Normal file
200
mcp/manager_test.go
Normal file
@@ -0,0 +1,200 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"git.toowon.com/jimmy/go-common/config"
|
||||
)
|
||||
|
||||
// fakeClient 是测试用的 Client 假实现,不依赖真实网络连接。
|
||||
type fakeClient struct {
|
||||
name string
|
||||
tools []ToolDescriptor
|
||||
listErr error
|
||||
callResult ToolCallResult
|
||||
callErr error
|
||||
closed bool
|
||||
|
||||
lastCalledTool string
|
||||
lastCalledArgs map[string]any
|
||||
}
|
||||
|
||||
func (f *fakeClient) Name() string { return f.name }
|
||||
|
||||
func (f *fakeClient) ListTools(ctx context.Context) ([]ToolDescriptor, error) {
|
||||
if f.listErr != nil {
|
||||
return nil, f.listErr
|
||||
}
|
||||
return f.tools, nil
|
||||
}
|
||||
|
||||
func (f *fakeClient) CallTool(ctx context.Context, name string, args map[string]any) (ToolCallResult, error) {
|
||||
f.lastCalledTool = name
|
||||
f.lastCalledArgs = args
|
||||
if f.callErr != nil {
|
||||
return ToolCallResult{}, f.callErr
|
||||
}
|
||||
return f.callResult, nil
|
||||
}
|
||||
|
||||
func (f *fakeClient) Close() error {
|
||||
f.closed = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func newTestManager(entries ...*serverEntry) *manager {
|
||||
return &manager{enabled: true, servers: entries, routes: make(map[string]route)}
|
||||
}
|
||||
|
||||
func TestManagerDisabledWithoutConfig(t *testing.T) {
|
||||
m := NewManager(nil, nil)
|
||||
if m.Enabled() {
|
||||
t.Fatal("Manager should be disabled when cfg is nil")
|
||||
}
|
||||
if _, err := m.CallTool(context.Background(), "x__y", nil); err == nil {
|
||||
t.Fatal("CallTool should error when disabled")
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerDisabledWhenNoServers(t *testing.T) {
|
||||
m := NewManager(nil, &config.MCPConfig{Enabled: true})
|
||||
if m.Enabled() {
|
||||
t.Fatal("Manager should be disabled when no servers configured")
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerListToolsNamespacesAndFiltersAllowlist(t *testing.T) {
|
||||
fc := &fakeClient{name: "word", tools: []ToolDescriptor{
|
||||
{Name: "create_document", Description: "create a doc"},
|
||||
{Name: "delete_document", Description: "delete a doc"},
|
||||
}}
|
||||
entry := newServerEntry(config.MCPServerConfig{
|
||||
Name: "word",
|
||||
ToolNamePrefix: "word",
|
||||
AllowedTools: []string{"create_document"},
|
||||
}, fc)
|
||||
m := newTestManager(entry)
|
||||
|
||||
tools, err := m.ListTools(context.Background())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(tools) != 1 {
|
||||
t.Fatalf("expected 1 tool after allowlist filter, got %d", len(tools))
|
||||
}
|
||||
if tools[0].FQName != "word__create_document" {
|
||||
t.Fatalf("unexpected FQName: %s", tools[0].FQName)
|
||||
}
|
||||
if tools[0].ServerName != "word" {
|
||||
t.Fatalf("unexpected ServerName: %s", tools[0].ServerName)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerListToolsSkipsFailingServer(t *testing.T) {
|
||||
ok := &fakeClient{name: "ok-server", tools: []ToolDescriptor{{Name: "t1"}}}
|
||||
bad := &fakeClient{name: "bad-server", listErr: errors.New("boom")}
|
||||
m := newTestManager(
|
||||
newServerEntry(config.MCPServerConfig{Name: "ok-server"}, ok),
|
||||
newServerEntry(config.MCPServerConfig{Name: "bad-server"}, bad),
|
||||
)
|
||||
|
||||
tools, err := m.ListTools(context.Background())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(tools) != 1 {
|
||||
t.Fatalf("expected 1 tool from healthy server, got %d", len(tools))
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerCallToolRoutesByFQNameWithoutPriorListTools(t *testing.T) {
|
||||
fc := &fakeClient{name: "fetch-web", callResult: ToolCallResult{Text: "hello"}}
|
||||
entry := newServerEntry(config.MCPServerConfig{Name: "fetch-web", ToolNamePrefix: "fetch-web"}, fc)
|
||||
m := newTestManager(entry)
|
||||
|
||||
res, err := m.CallTool(context.Background(), "fetch-web__fetch_url", map[string]any{"url": "https://example.com"})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if res.Text != "hello" {
|
||||
t.Fatalf("unexpected result text: %s", res.Text)
|
||||
}
|
||||
if fc.lastCalledTool != "fetch_url" {
|
||||
t.Fatalf("expected raw tool name %q, got %q", "fetch_url", fc.lastCalledTool)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerCallToolRejectsDisallowedTool(t *testing.T) {
|
||||
entry := newServerEntry(config.MCPServerConfig{
|
||||
Name: "word",
|
||||
ToolNamePrefix: "word",
|
||||
AllowedTools: []string{"create_document"},
|
||||
}, &fakeClient{name: "word"})
|
||||
m := newTestManager(entry)
|
||||
|
||||
if _, err := m.CallTool(context.Background(), "word__delete_document", nil); err == nil {
|
||||
t.Fatal("expected error for tool not in allowlist")
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerCallToolUnknownServerPrefix(t *testing.T) {
|
||||
entry := newServerEntry(config.MCPServerConfig{Name: "word", ToolNamePrefix: "word"}, &fakeClient{name: "word"})
|
||||
m := newTestManager(entry)
|
||||
|
||||
if _, err := m.CallTool(context.Background(), "unknown__tool", nil); err == nil {
|
||||
t.Fatal("expected error for unresolvable fq name")
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerCloseClosesAllClients(t *testing.T) {
|
||||
fc1 := &fakeClient{name: "a"}
|
||||
fc2 := &fakeClient{name: "b"}
|
||||
m := newTestManager(
|
||||
newServerEntry(config.MCPServerConfig{Name: "a"}, fc1),
|
||||
newServerEntry(config.MCPServerConfig{Name: "b"}, fc2),
|
||||
)
|
||||
if err := m.Close(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !fc1.closed || !fc2.closed {
|
||||
t.Fatal("Close should close all underlying clients")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRequestIDContext(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
if got := RequestIDFromContext(ctx); got != "" {
|
||||
t.Fatalf("expected empty request id, got %q", got)
|
||||
}
|
||||
ctx = WithRequestID(ctx, "req-123")
|
||||
if got := RequestIDFromContext(ctx); got != "req-123" {
|
||||
t.Fatalf("expected req-123, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeTransportAndAuthHeader(t *testing.T) {
|
||||
cases := map[string]string{
|
||||
"": "streamable-http",
|
||||
"http": "streamable-http",
|
||||
"streamable-http": "streamable-http",
|
||||
"SSE": "sse",
|
||||
"weird": "weird",
|
||||
}
|
||||
for in, want := range cases {
|
||||
if got := normalizeTransport(in); got != want {
|
||||
t.Errorf("normalizeTransport(%q) = %q, want %q", in, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
if got := authHeaderValue(""); got != "" {
|
||||
t.Fatalf("expected empty auth header, got %q", got)
|
||||
}
|
||||
if got := authHeaderValue("abc"); got != "Bearer abc" {
|
||||
t.Fatalf("expected Bearer prefix, got %q", got)
|
||||
}
|
||||
if got := authHeaderValue("Bearer abc"); got != "Bearer abc" {
|
||||
t.Fatalf("expected unchanged, got %q", got)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user