package controllers import ( "encoding/json" "fmt" "io" "strconv" "strings" "time" "server/models" "server/pkg/jwtutil" "server/services" beego "github.com/beego/beego/v2/server/web" ) // BackendMcpServerController MCP服务器配置控制器 type BackendMcpServerController struct { beego.Controller } func (c *BackendMcpServerController) mcpSrvClaims() (*jwtutil.Claims, error) { auth := c.Ctx.Request.Header.Get("Authorization") if auth == "" { return nil, fmt.Errorf("未登录") } parts := strings.SplitN(auth, " ", 2) if len(parts) != 2 || parts[0] != "Bearer" { return nil, fmt.Errorf("认证信息格式错误") } claims, err := jwtutil.ParseToken(parts[1]) if err != nil { return nil, fmt.Errorf("无效的token") } if claims.UserType != "backend" { return nil, fmt.Errorf("无权访问") } return claims, nil } func (c *BackendMcpServerController) mcpSrvErr(httpStatus, bizCode int, msg string) { c.Ctx.Output.SetStatus(httpStatus) c.Data["json"] = map[string]interface{}{"code": bizCode, "msg": msg} _ = c.ServeJSON() } func (c *BackendMcpServerController) mcpSrvOk(data interface{}) { c.Data["json"] = map[string]interface{}{"code": 200, "msg": "success", "data": data} _ = c.ServeJSON() } type mcpServerPayload struct { Name string `json:"name"` Transport string `json:"transport"` Command string `json:"command"` Args string `json:"args"` Env string `json:"env"` URL string `json:"url"` Headers string `json:"headers"` Description string `json:"description"` Provider string `json:"provider"` FromMarket string `json:"from_market"` Enabled int8 `json:"enabled"` Remark string `json:"remark"` } func (p *mcpServerPayload) validate() error { if strings.TrimSpace(p.Name) == "" { return fmt.Errorf("服务名称不能为空") } p.Transport = strings.ToLower(strings.TrimSpace(p.Transport)) if p.Transport == "" { p.Transport = "http" } switch p.Transport { case "stdio": if strings.TrimSpace(p.Command) == "" { return fmt.Errorf("stdio 传输必须填写启动命令") } case "http", "sse": if strings.TrimSpace(p.URL) == "" { return fmt.Errorf("请填写服务地址 URL") } default: return fmt.Errorf("不支持的传输类型: %s(可选 stdio/http/sse)", p.Transport) } return nil } func (p *mcpServerPayload) toModel(claims *jwtutil.Claims) models.BackendMcpServer { return models.BackendMcpServer{ TenantID: fmt.Sprintf("%d", claims.TenantId), UserID: uint64(claims.UserID), Name: strings.TrimSpace(p.Name), Transport: p.Transport, Command: strings.TrimSpace(p.Command), Args: strings.TrimSpace(p.Args), Env: strings.TrimSpace(p.Env), URL: strings.TrimSpace(p.URL), Headers: strings.TrimSpace(p.Headers), Description: strings.TrimSpace(p.Description), Provider: strings.TrimSpace(p.Provider), FromMarket: strings.TrimSpace(p.FromMarket), Enabled: p.Enabled, Remark: strings.TrimSpace(p.Remark), CreateTime: time.Now(), UpdateTime: time.Now(), } } // List GET /backend/mcp/server/list func (c *BackendMcpServerController) List() { claims, err := c.mcpSrvClaims() if err != nil { c.mcpSrvErr(401, 401, err.Error()) return } var list []models.BackendMcpServer _, err = models.Orm.QueryTable(new(models.BackendMcpServer)). Filter("tenant_id", fmt.Sprintf("%d", claims.TenantId)). Filter("user_id", uint64(claims.UserID)). Filter("delete_time__isnull", true). OrderBy("-id"). All(&list) if err != nil { c.mcpSrvErr(500, 500, "查询失败: "+err.Error()) return } c.mcpSrvOk(map[string]interface{}{"list": list}) } // Create POST /backend/mcp/server func (c *BackendMcpServerController) Create() { claims, err := c.mcpSrvClaims() if err != nil { c.mcpSrvErr(401, 401, err.Error()) return } body, err := io.ReadAll(c.Ctx.Request.Body) if err != nil { c.mcpSrvErr(400, 400, "读取请求体失败") return } var p mcpServerPayload if err := json.Unmarshal(body, &p); err != nil { c.mcpSrvErr(400, 400, "参数格式错误") return } if err := p.validate(); err != nil { c.mcpSrvErr(400, 400, err.Error()) return } m := p.toModel(claims) id, err := models.Orm.Insert(&m) if err != nil { c.mcpSrvErr(500, 500, "创建失败: "+err.Error()) return } c.mcpSrvOk(map[string]interface{}{"id": id}) } // Update PUT /backend/mcp/server/:id func (c *BackendMcpServerController) Update() { claims, err := c.mcpSrvClaims() if err != nil { c.mcpSrvErr(401, 401, err.Error()) return } idStr := c.Ctx.Input.Param(":id") id, err := strconv.ParseUint(idStr, 10, 64) if err != nil { c.mcpSrvErr(400, 400, "ID格式错误") return } srv := models.BackendMcpServer{ID: id} if err := models.Orm.Read(&srv); err != nil { c.mcpSrvErr(404, 404, "服务不存在") return } if srv.TenantID != fmt.Sprintf("%d", claims.TenantId) || srv.UserID != uint64(claims.UserID) { c.mcpSrvErr(403, 403, "无权操作") return } body, err := io.ReadAll(c.Ctx.Request.Body) if err != nil { c.mcpSrvErr(400, 400, "读取请求体失败") return } var p mcpServerPayload if err := json.Unmarshal(body, &p); err != nil { c.mcpSrvErr(400, 400, "参数格式错误") return } if err := p.validate(); err != nil { c.mcpSrvErr(400, 400, err.Error()) return } srv.Name = strings.TrimSpace(p.Name) srv.Transport = p.Transport srv.Command = strings.TrimSpace(p.Command) srv.Args = strings.TrimSpace(p.Args) srv.Env = strings.TrimSpace(p.Env) srv.URL = strings.TrimSpace(p.URL) srv.Headers = strings.TrimSpace(p.Headers) srv.Description = strings.TrimSpace(p.Description) srv.Provider = strings.TrimSpace(p.Provider) srv.FromMarket = strings.TrimSpace(p.FromMarket) srv.Enabled = p.Enabled srv.Remark = strings.TrimSpace(p.Remark) srv.UpdateTime = time.Now() _, err = models.Orm.Update(&srv) if err != nil { c.mcpSrvErr(500, 500, "更新失败: "+err.Error()) return } // 配置变更后断开旧连接,下次使用时按新配置重连 services.McpClientManager.Disconnect(srv.ID) c.mcpSrvOk(nil) } // Delete DELETE /backend/mcp/server/:id func (c *BackendMcpServerController) Delete() { claims, err := c.mcpSrvClaims() if err != nil { c.mcpSrvErr(401, 401, err.Error()) return } idStr := c.Ctx.Input.Param(":id") id, err := strconv.ParseUint(idStr, 10, 64) if err != nil { c.mcpSrvErr(400, 400, "ID格式错误") return } srv := models.BackendMcpServer{ID: id} if err := models.Orm.Read(&srv); err != nil { c.mcpSrvErr(404, 404, "服务不存在") return } if srv.TenantID != fmt.Sprintf("%d", claims.TenantId) || srv.UserID != uint64(claims.UserID) { c.mcpSrvErr(403, 403, "无权操作") return } now := time.Now() srv.DeleteTime = &now if _, err := models.Orm.Update(&srv, "delete_time"); err != nil { c.mcpSrvErr(500, 500, "删除失败: "+err.Error()) return } services.McpClientManager.Disconnect(srv.ID) c.mcpSrvOk(nil) } // Toggle PUT /backend/mcp/server/:id/toggle func (c *BackendMcpServerController) Toggle() { claims, err := c.mcpSrvClaims() if err != nil { c.mcpSrvErr(401, 401, err.Error()) return } idStr := c.Ctx.Input.Param(":id") id, err := strconv.ParseUint(idStr, 10, 64) if err != nil { c.mcpSrvErr(400, 400, "ID格式错误") return } srv := models.BackendMcpServer{ID: id} if err := models.Orm.Read(&srv); err != nil { c.mcpSrvErr(404, 404, "服务不存在") return } if srv.TenantID != fmt.Sprintf("%d", claims.TenantId) || srv.UserID != uint64(claims.UserID) { c.mcpSrvErr(403, 403, "无权操作") return } body, err := io.ReadAll(c.Ctx.Request.Body) if err != nil { c.mcpSrvErr(400, 400, "读取请求体失败") return } var req struct { Enabled *int8 `json:"enabled"` } if len(body) > 0 { _ = json.Unmarshal(body, &req) } if req.Enabled != nil { srv.Enabled = *req.Enabled } else if srv.Enabled == 1 { srv.Enabled = 0 } else { srv.Enabled = 1 } srv.UpdateTime = time.Now() if _, err := models.Orm.Update(&srv, "enabled", "update_time"); err != nil { c.mcpSrvErr(500, 500, "更新失败: "+err.Error()) return } c.mcpSrvOk(map[string]interface{}{"enabled": srv.Enabled}) } // Test POST /backend/mcp/server/test // 使用服务器配置新建连接并列出工具(不进缓存),同时回写连接状态 func (c *BackendMcpServerController) Test() { claims, err := c.mcpSrvClaims() if err != nil { c.mcpSrvErr(401, 401, err.Error()) return } body, err := io.ReadAll(c.Ctx.Request.Body) if err != nil { c.mcpSrvErr(400, 400, "读取请求体失败") return } var req struct { ID uint64 `json:"id"` } _ = json.Unmarshal(body, &req) var srv models.BackendMcpServer if req.ID > 0 { srv = models.BackendMcpServer{ID: req.ID} if err := models.Orm.Read(&srv); err != nil { c.mcpSrvErr(404, 404, "服务不存在") return } if srv.TenantID != fmt.Sprintf("%d", claims.TenantId) || srv.UserID != uint64(claims.UserID) { c.mcpSrvErr(403, 403, "无权操作") return } } else { // 未保存的配置直接测试 var p mcpServerPayload if err := json.Unmarshal(body, &p); err != nil { c.mcpSrvErr(400, 400, "参数格式错误") return } if err := p.validate(); err != nil { c.mcpSrvErr(400, 400, err.Error()) return } srv = p.toModel(claims) } tools, err := services.McpClientManager.TestConnection(&srv) if err != nil { // 回写失败状态 if req.ID > 0 { srv.Status = 2 srv.LastError = err.Error() srv.ToolCount = 0 srv.UpdateTime = time.Now() _, _ = models.Orm.Update(&srv, "status", "last_error", "tool_count", "update_time") } c.mcpSrvErr(400, 400, "连接失败: "+err.Error()) return } // 回写成功状态 if req.ID > 0 { srv.Status = 1 srv.LastError = "" srv.ToolCount = len(tools) srv.UpdateTime = time.Now() _, _ = models.Orm.Update(&srv, "status", "last_error", "tool_count", "update_time") } c.mcpSrvOk(map[string]interface{}{ "success": true, "tools": tools, "count": len(tools), }) } // Tools POST /backend/mcp/server/:id/tools // 使用缓存连接列出工具(供会话注入与前端查看) func (c *BackendMcpServerController) Tools() { claims, err := c.mcpSrvClaims() if err != nil { c.mcpSrvErr(401, 401, err.Error()) return } idStr := c.Ctx.Input.Param(":id") id, err := strconv.ParseUint(idStr, 10, 64) if err != nil { c.mcpSrvErr(400, 400, "ID格式错误") return } srv := models.BackendMcpServer{ID: id} if err := models.Orm.Read(&srv); err != nil { c.mcpSrvErr(404, 404, "服务不存在") return } if srv.TenantID != fmt.Sprintf("%d", claims.TenantId) || srv.UserID != uint64(claims.UserID) { c.mcpSrvErr(403, 403, "无权操作") return } tools, err := services.McpClientManager.EnsureConnected(&srv) if err != nil { srv.Status = 2 srv.LastError = err.Error() srv.UpdateTime = time.Now() _, _ = models.Orm.Update(&srv, "status", "last_error", "update_time") c.mcpSrvErr(400, 400, "连接失败: "+err.Error()) return } if srv.Status != 1 || srv.ToolCount != len(tools) { srv.Status = 1 srv.LastError = "" srv.ToolCount = len(tools) srv.UpdateTime = time.Now() _, _ = models.Orm.Update(&srv, "status", "last_error", "tool_count", "update_time") } c.mcpSrvOk(map[string]interface{}{"tools": tools, "count": len(tools)}) } // Market GET /backend/mcp/server/market func (c *BackendMcpServerController) Market() { if _, err := c.mcpSrvClaims(); err != nil { c.mcpSrvErr(401, 401, err.Error()) return } c.mcpSrvOk(map[string]interface{}{"list": services.GetMcpMarket()}) } // AddFromMarket POST /backend/mcp/server/from-market // 从市场一键添加服务 func (c *BackendMcpServerController) AddFromMarket() { claims, err := c.mcpSrvClaims() if err != nil { c.mcpSrvErr(401, 401, err.Error()) return } body, err := io.ReadAll(c.Ctx.Request.Body) if err != nil { c.mcpSrvErr(400, 400, "读取请求体失败") return } var req struct { Key string `json:"key"` Name string `json:"name"` Enabled *int8 `json:"enabled"` } if err := json.Unmarshal(body, &req); err != nil { c.mcpSrvErr(400, 400, "参数格式错误") return } item, ok := services.FindMarketItem(req.Key) if !ok { c.mcpSrvErr(404, 404, "市场不存在该服务: "+req.Key) return } name := strings.TrimSpace(req.Name) if name == "" { name = item.Name } // 防止重复添加同一市场服务 count, _ := models.Orm.QueryTable(new(models.BackendMcpServer)). Filter("tenant_id", fmt.Sprintf("%d", claims.TenantId)). Filter("user_id", uint64(claims.UserID)). Filter("from_market", item.Key). Filter("delete_time__isnull", true). Count() if count > 0 { c.mcpSrvErr(400, 400, "该市场服务已添加,请勿重复添加") return } enabled := int8(0) if req.Enabled != nil { enabled = *req.Enabled } argsJSON, _ := json.Marshal(item.Args) envJSON, _ := json.Marshal(item.Env) srv := models.BackendMcpServer{ TenantID: fmt.Sprintf("%d", claims.TenantId), UserID: uint64(claims.UserID), Name: name, Transport: item.Transport, Command: item.Command, Args: string(argsJSON), Env: string(envJSON), URL: item.URL, Description: item.Description, Provider: item.Provider, FromMarket: item.Key, Enabled: enabled, CreateTime: time.Now(), UpdateTime: time.Now(), } id, err := models.Orm.Insert(&srv) if err != nil { c.mcpSrvErr(500, 500, "添加失败: "+err.Error()) return } c.mcpSrvOk(map[string]interface{}{"id": id}) }