365 lines
11 KiB
Go
365 lines
11 KiB
Go
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]) + "…"
|
||
}
|