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.
170 lines
4.7 KiB
170 lines
4.7 KiB
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_migu_music))
|
|
|
|
// 记忆中心(拾忆):只读工具。写操作留在 home 的 memory 模块,
|
|
// 在这里复制一份写路径的校验/幂等/提醒重算必然漂移。
|
|
this.RegisterComp(new(tool_get_memory_items))
|
|
this.RegisterComp(new(tool_get_memory_stats))
|
|
this.RegisterComp(new(tool_get_memory_report))
|
|
|
|
// 会议纪要(只读):回答「我上次开会定了啥」
|
|
this.RegisterComp(new(tool_search_meeting_notes))
|
|
|
|
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
|
|
}
|
|
|
|
// 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"))
|
|
}
|
|
|