You can not select more than 25 topics
Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
123 lines
4.0 KiB
123 lines
4.0 KiB
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)
|
|
})
|
|
}
|
|
|