Files
yunzerwebsiteallinone/go/controllers/backend_mcp_server.go
T

507 lines
13 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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)")
}
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})
}