package mcp import ( "context" "fmt" "net/http" "runtime/debug" "time" "yunyan/lego/sys/log" "github.com/mark3labs/mcp-go/mcp" "github.com/mark3labs/mcp-go/server" ) const ( ToolGroup_GLOBAL = "GLOBAL" // 全局工具组 ToolGroup_OVERSEAS = "OVERSEAS" // 海外工具组 ToolGroup_CHINA = "CHINA" // 国内工具组 ) type ( ITool interface { Tool() (tool mcp.Tool) Handl(ctx context.Context, request mcp.CallToolRequest) (*mcp.CallToolResult, error) } ) func stringInArray(target string, array []string) bool { for _, item := range array { if item == target { return true } } return false } // safeLogErrorf 记录错误日志,日志系统未初始化等异常不向外扩散 func safeLogErrorf(format string, args ...interface{}) { defer func() { _ = recover() }() log.Errorf(format, args...) } // safeToolHandler 包装工具处理函数,保证任何情况下都给大模型返回一个可读的结果: // 1. 捕获 handler 内部 panic(参数异常导致的空指针、数组越界等),不让单个工具打崩整条请求; // 2. 兜底 handler 返回 (nil, nil) 的情况——mcp-go 的 request_handler 会直接解引用结果(*result), // result 为 nil 时必然 panic; // 3. 把 error 转成工具错误结果(IsError),让模型能看到失败原因并自行重试/换参数, // 而不是整条 JSON-RPC 请求失败。 func safeToolHandler(name string, handler server.ToolHandlerFunc) server.ToolHandlerFunc { return func(ctx context.Context, request mcp.CallToolRequest) (result *mcp.CallToolResult, err error) { defer func() { if r := recover(); r != nil { //先给出结果,再记录日志,保证日志系统本身出问题也不影响返回 result = mcp.NewToolResultError(fmt.Sprintf("工具 %s 执行异常,请确认调用参数是否正确后重试", name)) err = nil safeLogErrorf("tool %s panic recovered: %v\nstack: %s", name, r, string(debug.Stack())) } }() if result, err = handler(ctx, request); err != nil { safeLogErrorf("tool %s exec err: %v", name, err) result = mcp.NewToolResultError(fmt.Sprintf("工具 %s 执行失败: %s", name, err.Error())) err = nil return } if result == nil { safeLogErrorf("tool %s return nil result", name) result = mcp.NewToolResultError(fmt.Sprintf("工具 %s 未返回结果,请确认调用参数是否正确后重试", name)) } return } } func isValidDate(dateStr string) bool { _, err := time.Parse("2006-01-02", dateStr) return err == nil } // withRecover 添加panic恢复中间件,防止单个请求的panic导致整个服务崩溃 func withRecover(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { defer func() { if err := recover(); err != nil { // 记录panic信息 log.Errorf("Panic recovered in %s %s: %v", r.Method, r.URL.Path, err) // 打印堆栈信息用于调试 if stack := debug.Stack(); len(stack) > 0 { log.Errorf("Stack trace: %s", string(stack)) } // 返回500错误给客户端 w.Header().Set("Content-Type", "application/json") w.WriteHeader(http.StatusInternalServerError) w.Write([]byte(`{"error":"Internal server error","message":"服务器内部错误,请稍后重试"}`)) } }() next.ServeHTTP(w, r) }) } func withCORS(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { origin := r.Header.Get("Origin") if origin != "" { w.Header().Set("Access-Control-Allow-Origin", origin) } else { w.Header().Set("Access-Control-Allow-Origin", "*") } w.Header().Set("Access-Control-Allow-Methods", "GET, POST, OPTIONS") w.Header().Set("Access-Control-Allow-Headers", "Content-Type, Authorization, X-Session-Id, Mcp-Session-Id") w.Header().Set("Access-Control-Expose-Headers", "Mcp-Session-Id") w.Header().Set("Access-Control-Allow-Credentials", "true") // 处理预检请求 if r.Method == http.MethodOptions { w.WriteHeader(http.StatusOK) return } next.ServeHTTP(w, r) }) }