go-common中增加mcp抽离成通用方法(编排由业务方自行处理)
This commit is contained in:
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]) + "…"
|
||||
}
|
||||
Reference in New Issue
Block a user