增加供应商智能添加功能

This commit is contained in:
2026-09-04 00:09:32 +08:00
parent d8936d9730
commit 20be7f4817
29 changed files with 4317 additions and 1476 deletions
File diff suppressed because it is too large Load Diff
+4 -4
View File
@@ -56,6 +56,7 @@ type aiProviderPayload struct {
ApiBase string `json:"api_base"`
ApiKey string `json:"api_key"`
Models []string `json:"models"`
DefaultModel string `json:"default_model"` // 默认使用的模型
IsDefault int8 `json:"is_default"`
Status int8 `json:"status"`
Remark string `json:"remark"`
@@ -108,11 +109,8 @@ func (c *BackendAiProviderController) List() {
return
}
// 脱敏api_key + 解析模型列表
// 解析模型列表(不脱敏api_key,前端需要完整api_key进行测试)
for i := range list {
if list[i].ApiKey != "" && len(list[i].ApiKey) > 8 {
list[i].ApiKey = list[i].ApiKey[:4] + "****" + list[i].ApiKey[len(list[i].ApiKey)-4:]
}
list[i].ModelsList = parseProviderModels(list[i].Models)
}
@@ -171,6 +169,7 @@ func (c *BackendAiProviderController) Create() {
ApiBase: strings.TrimSpace(p.ApiBase),
ApiKey: strings.TrimSpace(p.ApiKey),
Models: string(modelsJSON),
DefaultModel: strings.TrimSpace(p.DefaultModel),
IsDefault: p.IsDefault,
Status: p.Status,
Remark: strings.TrimSpace(p.Remark),
@@ -240,6 +239,7 @@ func (c *BackendAiProviderController) Update() {
modelsJSON, _ := json.Marshal(p.Models)
provider.Models = string(modelsJSON)
}
provider.DefaultModel = strings.TrimSpace(p.DefaultModel)
provider.IsDefault = p.IsDefault
provider.Status = p.Status
provider.Remark = strings.TrimSpace(p.Remark)
+411 -29
View File
@@ -4,6 +4,8 @@ import (
"encoding/json"
"fmt"
"io"
"math"
"strconv"
"strings"
"time"
@@ -53,6 +55,75 @@ type smartGeneratePayload struct {
Type string `json:"type"` // supplier / customer
}
// extractJSONFromText 从文本中提取JSON对象(查找 { 和 } 之间的内容)
func extractJSONFromText(text string) string {
text = strings.TrimSpace(text)
// 找到第一个 { 和最后一个 }
start := strings.Index(text, "{")
end := strings.LastIndex(text, "}")
if start == -1 || end == -1 || end <= start {
return ""
}
// 提取可能的JSON字符串
jsonStr := text[start : end+1]
// 验证是否是有效的JSON(允许额外内容)
var test map[string]interface{}
if err := json.Unmarshal([]byte(jsonStr), &test); err == nil {
return jsonStr
}
// 尝试更宽松的提取:从第一个 { 开始找完整的JSON结构
return findCompleteJSON(text, start)
}
// findCompleteJSON 从指定位置开始查找完整的JSON对象
func findCompleteJSON(text string, startPos int) string {
if startPos == -1 || startPos >= len(text) {
return ""
}
stack := 0
start := -1
for i := startPos; i < len(text); i++ {
ch := text[i]
switch ch {
case '{':
if stack == 0 {
start = i
}
stack++
case '}':
stack--
if stack == 0 && start != -1 {
// 找到完整的JSON对象
jsonStr := text[start : i+1]
// 验证是否是有效JSON
var test map[string]interface{}
if err := json.Unmarshal([]byte(jsonStr), &test); err == nil {
return jsonStr
}
}
case '"':
// 跳过字符串中的括号
i++
for i < len(text) && text[i] != '"' {
if text[i] == '\\' {
i++ // 跳过转义字符
}
i++
}
}
}
return ""
}
// GenerateCompany POST /backend/ai/smart-generate/company
// 根据公司名称智能生成详细信息
func (c *BackendAiSmartGenerateController) GenerateCompany() {
@@ -97,26 +168,53 @@ func (c *BackendAiSmartGenerateController) GenerateCompany() {
}
provider := providers[0]
// 确定使用的模型(第一个)
// 确定使用的模型
providerModels := parseProviderModels(provider.Models)
if len(providerModels) == 0 {
c.sgJsonErr(400, 400, "默认接入配置没有可用模型")
return
}
// 优先使用默认模型,如果没有设置则使用第一个模型
useModel := providerModels[0]
if provider.DefaultModel != "" {
// 检查default_model是否在models列表中
for _, m := range providerModels {
if m == provider.DefaultModel {
useModel = m
break
}
}
}
// 调试:打印使用的模型(使用fmt替代beego.Info)
fmt.Printf("SmartGenerate: Using model=%s, default_model=%s, tenant=%d\n", useModel, provider.DefaultModel, claims.TenantId)
// 构造prompt
typeLabel := "供应商"
nameField := "supplier_name"
typeField := "supplier_type"
typeAllowed := "原材料供应商/设备供应商/服务供应商/其他"
if p.Type == "customer" {
typeLabel = "客户"
nameField = "customer_name"
typeField = "customer_type"
typeAllowed = "企业/政府机构/国企/教育机构/个人"
}
prompt := fmt.Sprintf(`你是一个企业信息助手。请根据%s名称"%s",生成该公司的详细信息,以JSON格式返回,包含以下字段:
prompt := fmt.Sprintf(`你是一个企业信息助手。请根据%s名称"%s",先用可用的MCP工具(如天眼查 search_companies / get_company_basic_profile 等)查询该企业的真实信息,然后输出以下两部分内容:
【第一部分:企业画像(默认展示内容)】
用清晰的结构化markdown输出该企业画像,包含:基础工商信息(法定代表人、注册资本、成立日期、登记状态、企业类型、登记机关、注册地址、所属行业)、规模与人员(人员规模、参保人数)、企业标签(如高新技术企业/专精特新等)、经营范围(节选)、经营概况(对外投资/招投标/商标/专利数量)、风险提示(司法案件/合作风险等)。有真实数据的填真实数据,查询不到的字段注明"暂无",不要编造数据。
【第二部分:结构化JSON(供系统自动填充)】
在回复最后输出一个【合法的JSON对象】,必须以 { 开头、以 } 结尾,键和字符串值都要用英文双引号,键值之间用英文冒号,键值对之间用英文逗号。例如:
{"customer_name":"连云港如年实业有限公司","contact_person":"张昊","registered_capital":"6000万元"}
包含以下字段(每个字段一个键值,查不到就留空字符串):
- %s: 公司全称
- %s: %s类型(必须从以下取值中选一个:%s)
- contact_person: 联系人
- contact_phone: 联系电话
- contact_phone: 联系电话(务必从查询结果中提取真实电话,缺失时与注册电话互相补齐,不要编造号码)
- contact_email: 邮箱
- address: 公司地址
- industry: 所属行业
@@ -132,53 +230,337 @@ func (c *BackendAiSmartGenerateController) GenerateCompany() {
- bank_name: 开户银行
- bank_account: 银行账号
- registered_address: 注册地址
- registered_phone: 注册电话
- registered_phone: 注册电话(与联系电话互相补齐,缺失时用联系电话,不要编造号码)
- remark: 备注
请只返回JSON对象,不要返回其他任何文字、解释或markdown代码块标记。如果某些信息不确定,请留空字符串。`, typeLabel, p.CompanyName, nameField)
注意:JSON对象必须放在回复最后且用 { } 包裹,不要用markdown代码块包裹,不要输出"customer_name: xxx"这类非JSON格式。`, typeLabel, p.CompanyName, nameField, typeField, typeLabel, typeAllowed)
messages := []openaiMessage{
{Role: "user", Content: prompt},
}
// 调用AI(非流式,超时120s)
reply, err := callAI(provider, useModel, "", messages)
// 调用AI(非流式,含已启用 MCP 工具:AI 会先调用天眼查等工具获取真实数据再生成)
reply, err := runToolLoop(claims, provider, useModel, "", messages)
if err != nil {
c.sgJsonErr(500, 500, "AI生成失败: "+err.Error())
return
}
// 解析AI返回的JSON
reply = strings.TrimSpace(reply)
// 去掉可能的markdown代码块标记
reply = strings.TrimPrefix(reply, "```json")
reply = strings.TrimPrefix(reply, "```")
reply = strings.TrimSuffix(reply, "```")
// 解析AI返回的JSON(前面的markdown企业画像作为raw返回给前端展示)
reply = strings.TrimSpace(reply)
var result map[string]interface{}
if err := json.Unmarshal([]byte(reply), &result); err != nil {
// 解析失败,返回原始文本让用户自己处理
c.sgOk(map[string]interface{}{
"raw": reply,
"parsed": false,
"name": p.CompanyName,
"type": p.Type,
"generate_time": time.Now().Format("2006-01-02 15:04:05"),
})
return
var jsonStr string
// 方法1: 尝试直接解析整个响应
if err := json.Unmarshal([]byte(reply), &result); err == nil {
jsonStr = reply
} else {
// 方法2: 从回复中提取最后一个完整JSON对象(结构化字段在回复末尾)
jsonStr = extractLastJSON(reply)
if jsonStr == "" {
// 方法3: 兜底——模型可能输出 "字段名: 值" 的非JSON格式文本,逐行摘取到对应字段
result = parseSmartGenKV(reply, nameField)
if len(result) == 0 {
// 全部失败,返回原始文本让用户自己处理
c.sgOk(map[string]interface{}{
"raw": reply,
"parsed": false,
"error": "AI返回格式异常,无法解析JSON",
"name": p.CompanyName,
"type": p.Type,
"generate_time": time.Now().Format("2006-01-02 15:04:05"),
})
return
}
// KV兜底成功,继续走后面的结果组装
c.sgOk(smartGenResult(reply, result, p.CompanyName, p.Type, nameField))
return
}
if err := json.Unmarshal([]byte(jsonStr), &result); err != nil {
c.sgOk(map[string]interface{}{
"raw": reply,
"parsed": false,
"error": fmt.Sprintf("解析JSON失败: %v", err),
"name": p.CompanyName,
"type": p.Type,
"generate_time": time.Now().Format("2006-01-02 15:04:05"),
})
return
}
}
c.sgOk(smartGenResult(reply, result, p.CompanyName, p.Type, nameField))
}
// smartGenResult 组装智能生成成功响应(含完整AI回复 raw),并对表单受限字段做规整
func smartGenResult(reply string, result map[string]interface{}, companyName, genType, nameField string) map[string]interface{} {
// 确保公司名称字段存在
if _, ok := result[nameField]; !ok || result[nameField] == "" {
result[nameField] = p.CompanyName
result[nameField] = companyName
}
c.sgOk(map[string]interface{}{
// 规整受限字段,使其符合新增客户/供应商表单的下拉选项
normalizeStrField(result, "enterprise_type", normalizeEnterpriseType)
normalizeStrField(result, "industry", normalizeIndustry)
normalizeStrField(result, "taxpayer_qualification", normalizeTaxpayer)
// 注册资本/实缴资本规整为纯数值(单位万元),与表单"万元"后缀一致
normalizeStrField(result, "registered_capital", normalizeCapital)
normalizeStrField(result, "paid_capital", normalizeCapital)
// 客户/供应商类型规整为表单下拉选项代码
if genType == "supplier" {
normalizeStrField(result, "supplier_type", normalizeSupplierType)
} else {
normalizeStrField(result, "customer_type", normalizeCustomerType)
}
// 联系电话与注册电话互相补齐,尽量保证有电话可填(不编造)
crossFillPhone(result, "contact_phone", "registered_phone")
return map[string]interface{}{
"parsed": true,
"data": result,
"name": p.CompanyName,
"type": p.Type,
"raw": reply, // AI完整回复(含企业画像markdown,供前端"AI响应数据"展示)
"name": companyName,
"type": genType,
"generate_time": time.Now().Format("2006-01-02 15:04:05"),
})
}
}
// normalizeStrField 对 result 中的字符串字段应用规整函数
func normalizeStrField(result map[string]interface{}, key string, fn func(string) string) {
if v, ok := result[key]; ok {
if s, ok2 := v.(string); ok2 {
result[key] = fn(s)
}
}
}
// crossFillPhone 两个字段有任一非空时互相补齐(如 联系电话 缺则用 注册电话,反之亦然)
func crossFillPhone(result map[string]interface{}, a, b string) {
av, aOK := result[a].(string)
bv, bOK := result[b].(string)
if aOK && strings.TrimSpace(av) == "" && bOK && strings.TrimSpace(bv) != "" {
result[a] = bv
}
if bOK && strings.TrimSpace(bv) == "" && aOK && strings.TrimSpace(av) != "" {
result[b] = av
}
}
// normalizeEnterpriseType 将AI返回的企业类型规整为新增表单下拉选项之一
func normalizeEnterpriseType(raw string) string {
s := strings.TrimSpace(raw)
// 去掉括号及其内容,如 "有限责任公司(自然人投资或控股)" -> "有限责任公司"
if i := strings.IndexAny(s, "(("); i > 0 {
s = strings.TrimSpace(s[:i])
}
switch {
case strings.Contains(s, "有限责任"):
return "有限责任公司"
case strings.Contains(s, "股份"):
return "股份有限公司"
case strings.Contains(s, "合伙"):
return "合伙企业"
case strings.Contains(s, "个人独资"):
return "个人独资企业"
case strings.Contains(s, "国有"):
return "国有企业"
case strings.Contains(s, "集体"):
return "集体企业"
case strings.Contains(s, "外商"):
return "外商投资企业"
case s == "其他":
return "其他"
case s == "":
return ""
default:
return "其他"
}
}
// normalizeIndustry 规整行业:去掉"大行业>细分行业"路径,取最后一段
func normalizeIndustry(raw string) string {
s := strings.TrimSpace(raw)
if i := strings.LastIndex(s, ">"); i >= 0 && i < len(s)-1 {
return strings.TrimSpace(s[i+1:])
}
return s
}
// normalizeTaxpayer 将AI返回的纳税人资质规整为下拉选项之一
func normalizeTaxpayer(raw string) string {
s := strings.TrimSpace(raw)
switch {
case strings.Contains(s, "一般"):
return "一般纳税人"
case strings.Contains(s, "小规模"):
return "小规模纳税人"
case s == "其他":
return "其他"
case s == "":
return ""
default:
return "其他"
}
}
// normalizeCustomerType 将AI返回的客户类型规整为下拉选项代码(1企业/2政府机构/3国企/4教育机构/5个人)
func normalizeCustomerType(raw string) string {
s := strings.TrimSpace(raw)
switch {
case s == "1" || s == "2" || s == "3" || s == "4" || s == "5":
return s
case strings.Contains(s, "政府"):
return "2"
case strings.Contains(s, "教育") || strings.Contains(s, "学校"):
return "4"
case strings.Contains(s, "个人") || strings.Contains(s, "个体"):
return "5"
case strings.Contains(s, "国企") || strings.Contains(s, "国有"):
return "3"
case s == "":
return ""
default:
return "1" // 企业
}
}
// normalizeSupplierType 将AI返回的供应商类型规整为下拉选项代码(1原材料/2设备/3服务/4其他)
func normalizeSupplierType(raw string) string {
s := strings.TrimSpace(raw)
switch {
case s == "1" || s == "2" || s == "3" || s == "4":
return s
case strings.Contains(s, "原材料"):
return "1"
case strings.Contains(s, "设备"):
return "2"
case strings.Contains(s, "服务"):
return "3"
case s == "":
return ""
default:
return "4" // 其他
}
}
// normalizeCapital 将AI返回的注册资本/实缴资本规整为纯数值(单位万元),与表单"万元"后缀一致。
// 例: "6000万人民币"->"6000", "6000.00万"->"6000", "1亿"->"10000", "5000万元"->"5000"
func normalizeCapital(raw string) string {
s := strings.TrimSpace(raw)
if s == "" {
return ""
}
// 去掉千分位逗号
s = strings.ReplaceAll(s, ",", "")
s = strings.ReplaceAll(s, ",", "")
// 提取数字部分(含小数与可能的负号)
i := 0
if len(s) > 0 && s[0] == '-' {
i = 1
}
for i < len(s) && ((s[i] >= '0' && s[i] <= '9') || s[i] == '.') {
i++
}
numStr := s[:i]
if numStr == "" || numStr == "." || numStr == "-" {
return ""
}
num, err := strconv.ParseFloat(numStr, 64)
if err != nil {
return ""
}
// 判断数字后的单位,统一换算为万元
rest := strings.TrimSpace(s[i:])
if strings.HasPrefix(rest, "亿") {
num *= 10000
} else if strings.HasPrefix(rest, "千") {
num *= 0.1
}
// 格式化:整数去小数位,非整数保留有效小数
if num == math.Trunc(num) {
return strconv.FormatFloat(num, 'f', 0, 64)
}
return strconv.FormatFloat(num, 'f', -1, 64)
}
// smartGenAllowedFields 智能生成JSON允许的字段名
var smartGenAllowedFields = map[string]bool{
"customer_name": true, "supplier_name": true, "contact_person": true,
"contact_phone": true, "contact_email": true, "address": true,
"industry": true, "registered_capital": true, "paid_capital": true,
"establish_date": true, "administrative_division": true, "enterprise_type": true,
"taxpayer_qualification": true, "business_scope": true, "invoice_title": true,
"tax_number": true, "bank_name": true, "bank_account": true,
"registered_address": true, "registered_phone": true, "remark": true,
"customer_type": true, "supplier_type": true,
}
// parseSmartGenKV 兜底解析:当模型输出 "字段名: 值" 的非JSON文本时,逐行摘取到对应字段
func parseSmartGenKV(text, nameField string) map[string]interface{} {
result := make(map[string]interface{})
lines := strings.Split(text, "\n")
for _, line := range lines {
line = strings.TrimSpace(line)
// 去掉可能的列表/引用前缀
line = strings.TrimLeft(line, "-*|#> `")
line = strings.TrimSpace(line)
idx := strings.Index(line, ":")
if idx <= 0 {
continue
}
key := strings.Trim(strings.TrimSpace(line[:idx]), "\"'`")
if !smartGenAllowedFields[key] && key != nameField {
continue
}
val := strings.TrimSpace(line[idx+1:])
val = strings.Trim(val, "\"'`")
val = strings.TrimSpace(val)
if val == "" || val == "null" || val == "NULL" || val == "undefined" || val == "暂无" {
val = ""
}
result[key] = val
}
return result
}
// extractLastJSON 从文本中提取最后一个完整的JSON对象(忽略markdown等前置内容)
func extractLastJSON(text string) string {
var lastValid string
depth := 0
start := -1
inStr := false
escaped := false
for i := 0; i < len(text); i++ {
ch := text[i]
if inStr {
if escaped {
escaped = false
} else if ch == '\\' {
escaped = true
} else if ch == '"' {
inStr = false
}
continue
}
switch ch {
case '"':
inStr = true
case '{':
if depth == 0 {
start = i
}
depth++
case '}':
depth--
if depth == 0 && start != -1 {
cand := text[start : i+1]
var test map[string]interface{}
if json.Unmarshal([]byte(cand), &test) == nil {
lastValid = cand
}
start = -1
}
}
}
return lastValid
}
+506
View File
@@ -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})
}
+265
View File
@@ -0,0 +1,265 @@
package controllers
import (
"encoding/json"
"fmt"
"io"
"strconv"
"strings"
"time"
"server/models"
"server/pkg/jwtutil"
beego "github.com/beego/beego/v2/server/web"
)
// BackendMcpToolController MCP工具控制器
type BackendMcpToolController struct {
beego.Controller
}
func (c *BackendMcpToolController) mcpClaims() (*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 *BackendMcpToolController) mcpJsonErr(httpStatus, bizCode int, msg string) {
c.Ctx.Output.SetStatus(httpStatus)
c.Data["json"] = map[string]interface{}{"code": bizCode, "msg": msg}
_ = c.ServeJSON()
}
func (c *BackendMcpToolController) mcpOk(data interface{}) {
c.Data["json"] = map[string]interface{}{"code": 200, "msg": "success", "data": data}
_ = c.ServeJSON()
}
type mcpToolPayload struct {
Name string `json:"name"`
Type string `json:"type"`
URL string `json:"url"`
Method string `json:"method"`
Headers string `json:"headers"`
Params string `json:"params"`
Body string `json:"body"`
Enabled int8 `json:"enabled"`
Remark string `json:"remark"`
}
// List GET /backend/mcp/tool/list
func (c *BackendMcpToolController) List() {
claims, err := c.mcpClaims()
if err != nil {
c.mcpJsonErr(401, 401, err.Error())
return
}
var list []models.BackendMcpTool
_, err = models.Orm.QueryTable(new(models.BackendMcpTool)).
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.mcpJsonErr(500, 500, "查询失败: "+err.Error())
return
}
c.mcpOk(map[string]interface{}{"list": list})
}
// Create POST /backend/mcp/tool
func (c *BackendMcpToolController) Create() {
claims, err := c.mcpClaims()
if err != nil {
c.mcpJsonErr(401, 401, err.Error())
return
}
body, err := io.ReadAll(c.Ctx.Request.Body)
if err != nil {
c.mcpJsonErr(400, 400, "读取请求体失败")
return
}
var p mcpToolPayload
if err := json.Unmarshal(body, &p); err != nil {
c.mcpJsonErr(400, 400, "参数格式错误")
return
}
if strings.TrimSpace(p.Name) == "" {
c.mcpJsonErr(400, 400, "工具名称不能为空")
return
}
if strings.TrimSpace(p.URL) == "" {
c.mcpJsonErr(400, 400, "API地址不能为空")
return
}
tool := models.BackendMcpTool{
TenantID: fmt.Sprintf("%d", claims.TenantId),
UserID: uint64(claims.UserID),
Name: strings.TrimSpace(p.Name),
Type: strings.TrimSpace(p.Type),
URL: strings.TrimSpace(p.URL),
Method: strings.TrimSpace(p.Method),
Headers: strings.TrimSpace(p.Headers),
Params: strings.TrimSpace(p.Params),
Body: strings.TrimSpace(p.Body),
Enabled: p.Enabled,
Remark: strings.TrimSpace(p.Remark),
CreateTime: time.Now(),
UpdateTime: time.Now(),
}
id, err := models.Orm.Insert(&tool)
if err != nil {
c.mcpJsonErr(500, 500, "创建失败: "+err.Error())
return
}
c.mcpOk(map[string]interface{}{"id": id})
}
// Update PUT /backend/mcp/tool/:id
func (c *BackendMcpToolController) Update() {
claims, err := c.mcpClaims()
if err != nil {
c.mcpJsonErr(401, 401, err.Error())
return
}
idStr := c.Ctx.Input.Param(":id")
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
c.mcpJsonErr(400, 400, "ID格式错误")
return
}
tool := models.BackendMcpTool{ID: id}
if err := models.Orm.Read(&tool); err != nil {
c.mcpJsonErr(404, 404, "工具不存在")
return
}
if tool.TenantID != fmt.Sprintf("%d", claims.TenantId) {
c.mcpJsonErr(403, 403, "无权操作")
return
}
body, err := io.ReadAll(c.Ctx.Request.Body)
if err != nil {
c.mcpJsonErr(400, 400, "读取请求体失败")
return
}
var p mcpToolPayload
if err := json.Unmarshal(body, &p); err != nil {
c.mcpJsonErr(400, 400, "参数格式错误")
return
}
tool.Name = strings.TrimSpace(p.Name)
tool.Type = strings.TrimSpace(p.Type)
tool.URL = strings.TrimSpace(p.URL)
tool.Method = strings.TrimSpace(p.Method)
tool.Headers = strings.TrimSpace(p.Headers)
tool.Params = strings.TrimSpace(p.Params)
tool.Body = strings.TrimSpace(p.Body)
tool.Enabled = p.Enabled
tool.Remark = strings.TrimSpace(p.Remark)
tool.UpdateTime = time.Now()
_, err = models.Orm.Update(&tool)
if err != nil {
c.mcpJsonErr(500, 500, "更新失败: "+err.Error())
return
}
c.mcpOk(nil)
}
// Delete DELETE /backend/mcp/tool/:id
func (c *BackendMcpToolController) Delete() {
claims, err := c.mcpClaims()
if err != nil {
c.mcpJsonErr(401, 401, err.Error())
return
}
idStr := c.Ctx.Input.Param(":id")
id, err := strconv.ParseUint(idStr, 10, 64)
if err != nil {
c.mcpJsonErr(400, 400, "ID格式错误")
return
}
tool := models.BackendMcpTool{ID: id}
if err := models.Orm.Read(&tool); err != nil {
c.mcpJsonErr(404, 404, "工具不存在")
return
}
if tool.TenantID != fmt.Sprintf("%d", claims.TenantId) {
c.mcpJsonErr(403, 403, "无权操作")
return
}
now := time.Now()
tool.DeleteTime = &now
_, err = models.Orm.Update(&tool, "delete_time")
if err != nil {
c.mcpJsonErr(500, 500, "删除失败: "+err.Error())
return
}
c.mcpOk(nil)
}
// Test POST /backend/mcp/tool/test
func (c *BackendMcpToolController) Test() {
// 验证用户登录状态
if _, err := c.mcpClaims(); err != nil {
c.mcpJsonErr(401, 401, err.Error())
return
}
body, err := io.ReadAll(c.Ctx.Request.Body)
if err != nil {
c.mcpJsonErr(400, 400, "读取请求体失败")
return
}
var p mcpToolPayload
if err := json.Unmarshal(body, &p); err != nil {
c.mcpJsonErr(400, 400, "参数格式错误")
return
}
if strings.TrimSpace(p.URL) == "" {
c.mcpJsonErr(400, 400, "API地址不能为空")
return
}
// TODO: 实现MCP工具测试逻辑
// 这里只是一个示例实现,实际测试逻辑需要根据工具类型和配置来实现
c.mcpOk(map[string]interface{}{
"success": true,
"message": "测试连接成功",
"method": p.Method,
"url": p.URL,
})
}