package mcp import ( "context" "yunyan/comm" "yunyan/modules" "fmt" "log" "net/http" "time" "yunyan/lego/core" "github.com/mark3labs/mcp-go/mcp" "github.com/mark3labs/mcp-go/server" ) func NewModule() core.IModule { m := new(Mcp) return m } type Mcp struct { modules.ModuleBase service core.IService server *server.MCPServer model *modelComp options *Options // openai openai.ISys // juhe juhe.ISys } func (this *Mcp) GetType() core.M_Modules { return comm.ModuleMcp } func (this *Mcp) NewOptions() (options core.IModuleOptions) { return new(Options) } func (this *Mcp) Init(service core.IService, module core.IModule, options core.IModuleOptions) (err error) { this.service = service this.options = options.(*Options) if err = this.ModuleBase.Init(service, module, options); err != nil { return } this.server = server.NewMCPServer( "mcp-server", "1.0.0", server.WithResourceCapabilities(true, true), server.WithPromptCapabilities(true), server.WithToolCapabilities(true), ) return } func (this *Mcp) Start() (err error) { if err = this.ModuleBase.Start(); err != nil { return } go this.run() return } func (this *Mcp) OnInstallComp() { this.ModuleBase.OnInstallComp() this.model = this.RegisterComp(new(modelComp)).(*modelComp) // this.RegisterComp(new(tool_weather)) this.RegisterComp(new(tool_gaode_map_geocode)) this.RegisterComp(new(tool_gaode_map_regeocode)) this.RegisterComp(new(tool_gaode_map_route_navigation)) this.RegisterComp(new(tool_gaode_map_weather)) this.RegisterComp(new(tool_gaode_map_around_search)) this.RegisterComp(new(tool_gaode_map_searchplace)) this.RegisterComp(new(tool_google_map_geocode)) this.RegisterComp(new(tool_google_map_regeocode)) this.RegisterComp(new(tool_google_map_nearbysearch)) this.RegisterComp(new(tool_google_map_route_navigation)) this.RegisterComp(new(tool_bocha_search)) this.RegisterComp(new(tool_tavily_search)) this.RegisterComp(new(tool_hefeng_weather)) this.RegisterComp(new(tool_music)) this.RegisterComp(new(tool_migu_music)) // this.RegisterComp(new(tool_music_pause)) // this.RegisterComp(new(tool_music_resume)) // this.RegisterComp(new(tool_music_prev)) // this.RegisterComp(new(tool_music_next)) // this.RegisterComp(new(tool_music_close)) this.RegisterComp(new(tool_allhelp_task)) this.RegisterComp(new(tool_get_user_tasks)) this.RegisterComp(new(tool_cancel_user_task)) } // run 启动 MCP 的 HTTP 服务入口 // 参数: // - 无 // // 返回值: // - 无 // // 异常: // - 服务启动失败会直接 log.Fatalf 退出进程 func (this *Mcp) run() { sseServer := server.NewSSEServer(this.server, server.WithBaseURL(this.options.Addr), server.WithStaticBasePath("/sse"), server.WithSSEEndpoint("/"), server.WithSSEContextFunc(authFromRequest), server.WithUseFullURLForMessageEndpoint(true), server.WithKeepAlive(true), ) streamableHTTPServer := server.NewStreamableHTTPServer(this.server, server.WithEndpointPath("/mcp"), server.WithStateLess(true), server.WithHTTPContextFunc(authFromRequest), ) mux := http.NewServeMux() mux.Handle("/mcp", streamableHTTPServer) mux.Handle("/mcp/", streamableHTTPServer) mux.Handle("/sse", sseServer) mux.Handle("/sse/", sseServer) // 应用中间件:先恢复panic,再处理CORS handlerWithRecover := withRecover(mux) // 包裹 CORS 中间件 handlerWithCORS := withCORS(handlerWithRecover) httpServer := &http.Server{ Addr: fmt.Sprintf(":%d", this.options.Port), Handler: handlerWithCORS, ReadTimeout: time.Minute * 5, WriteTimeout: time.Minute * 5, IdleTimeout: time.Minute * 5, } var startErr error if this.options.TLSCertFile != "" && this.options.TLSKeyFile != "" { startErr = httpServer.ListenAndServeTLS(this.options.TLSCertFile, this.options.TLSKeyFile) } else { startErr = httpServer.ListenAndServe() } if startErr != nil { log.Fatalf("Server error: %v", startErr) } // log.Printf("SSE server listening on :%d", this.options.Port) // if err := sseServer.Start(fmt.Sprintf(":%d", this.options.Port)); err != nil { // log.Fatalf("Server error: %v", err) // } } func (this *Mcp) AddTool(group string, tool mcp.Tool, handler server.ToolHandlerFunc) bool { if tools, ok := this.options.Groups[group]; ok { //不存在 组注册 if stringInArray(tool.Name, tools) { //存在 菜注册 this.server.AddTool(tool, safeToolHandler(tool.Name, handler)) return true } } return false } // authKey is a custom context key for storing the auth token. type authKey struct{} // withAuthKey adds an auth key to the context. func withAuthKey(ctx context.Context, auth string) context.Context { return context.WithValue(ctx, authKey{}, auth) } // authFromRequest extracts the auth token from the request headers. func authFromRequest(ctx context.Context, r *http.Request) context.Context { return withAuthKey(ctx, r.Header.Get("Authorization")) }