增加供应商智能添加功能
This commit is contained in:
@@ -0,0 +1,506 @@
|
||||
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})
|
||||
}
|
||||
Reference in New Issue
Block a user