package controllers import ( "encoding/json" "fmt" "io" "math" "strconv" "strings" "time" "server/models" "server/pkg/jwtutil" beego "github.com/beego/beego/v2/server/web" ) // BackendAiSmartGenerateController AI智能生成控制器 type BackendAiSmartGenerateController struct { beego.Controller } func (c *BackendAiSmartGenerateController) sgClaims() (*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 *BackendAiSmartGenerateController) sgJsonErr(httpStatus, bizCode int, msg string) { c.Ctx.Output.SetStatus(httpStatus) c.Data["json"] = map[string]interface{}{"code": bizCode, "msg": msg} _ = c.ServeJSON() } func (c *BackendAiSmartGenerateController) sgOk(data interface{}) { c.Data["json"] = map[string]interface{}{"code": 200, "msg": "success", "data": data} _ = c.ServeJSON() } type smartGeneratePayload struct { CompanyName string `json:"company_name"` 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() { claims, err := c.sgClaims() if err != nil { c.sgJsonErr(401, 401, err.Error()) return } body, err := io.ReadAll(c.Ctx.Request.Body) if err != nil { c.sgJsonErr(400, 400, "读取请求体失败") return } var p smartGeneratePayload if err := json.Unmarshal(body, &p); err != nil { c.sgJsonErr(400, 400, "参数格式错误") return } if strings.TrimSpace(p.CompanyName) == "" { c.sgJsonErr(400, 400, "公司名称不能为空") return } if p.Type != "supplier" && p.Type != "customer" { c.sgJsonErr(400, 400, "类型必须是 supplier 或 customer") return } // 查询租户默认AI配置 var providers []models.BackendAiProvider _, err = models.Orm.QueryTable(new(models.BackendAiProvider)). Filter("tenant_id", fmt.Sprintf("%d", claims.TenantId)). Filter("is_default", 1). Filter("status", 1). Filter("delete_time__isnull", true). Limit(1). All(&providers) if err != nil || len(providers) == 0 { c.sgJsonErr(400, 400, "当前企业没有配置智能体无法使用本功能,请联系管理员") return } 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",先用可用的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_email: 邮箱 - address: 公司地址 - industry: 所属行业 - registered_capital: 注册资本(如"100万元") - paid_capital: 实缴资本(如"50万元") - establish_date: 成立日期(YYYY-MM-DD格式,不确定则留空) - administrative_division: 行政区划(如"江苏省连云港市海州区") - enterprise_type: 企业类型(如"有限责任公司") - taxpayer_qualification: 纳税人资质("一般纳税人"或"小规模纳税人") - business_scope: 经营范围 - invoice_title: 发票抬头(公司全称) - tax_number: 统一社会信用代码/税号(18位,不确定则留空) - bank_name: 开户银行 - bank_account: 银行账号 - registered_address: 注册地址 - registered_phone: 注册电话(与联系电话互相补齐,缺失时用联系电话,不要编造号码) - remark: 备注 注意:JSON对象必须放在回复最后且用 { } 包裹,不要用markdown代码块包裹,不要输出"customer_name: xxx"这类非JSON格式。`, typeLabel, p.CompanyName, nameField, typeField, typeLabel, typeAllowed) messages := []openaiMessage{ {Role: "user", Content: prompt}, } // 调用AI(非流式,含已启用 MCP 工具:AI 会先调用天眼查等工具获取真实数据再生成) reply, usage, err := runToolLoop(claims, provider, useModel, "", messages) if err != nil { c.sgJsonErr(500, 500, "AI生成失败: "+err.Error()) return } // 解析AI返回的JSON(前面的markdown企业画像作为raw返回给前端展示) reply = strings.TrimSpace(reply) var result map[string]interface{} 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, "model": usage.Model, "tools": usage.Tools, "tool_count": usage.ToolCount, "rounds": usage.Rounds, "generate_time": time.Now().Format("2006-01-02 15:04:05"), }) return } // KV兜底成功,继续走后面的结果组装 c.sgOk(smartGenResult(reply, result, p.CompanyName, p.Type, nameField, usage)) 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, "model": usage.Model, "tools": usage.Tools, "tool_count": usage.ToolCount, "rounds": usage.Rounds, "generate_time": time.Now().Format("2006-01-02 15:04:05"), }) return } } c.sgOk(smartGenResult(reply, result, p.CompanyName, p.Type, nameField, usage)) } // aiGenNotFound 判断 AI 回复是否明确表示未查到企业信息(MCP 工具查询无结果) func aiGenNotFound(raw string) bool { for _, kw := range []string{"未找到", "未匹配", "查无", "未查询到", "未检索到", "no matching", "not found"} { if strings.Contains(raw, kw) { return true } } return false } // smartGenResult 组装智能生成成功响应(含完整AI回复 raw),并对表单受限字段做规整 func smartGenResult(reply string, result map[string]interface{}, companyName, genType, nameField string, usage *aiToolUsage) map[string]interface{} { // 确保公司名称字段存在 if _, ok := result[nameField]; !ok || result[nameField] == "" { result[nameField] = companyName } // 规整受限字段,使其符合新增客户/供应商表单的下拉选项 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, "found": !aiGenNotFound(reply), // MCP 是否查到企业信息 "data": result, "raw": reply, // AI完整回复(含企业画像markdown,供前端"AI响应数据"展示) "name": companyName, "type": genType, "model": usage.Model, // 实际使用的模型 "tools": usage.Tools, // 实际调用的 MCP 工具(server + tool) "tool_count": usage.ToolCount, "rounds": usage.Rounds, "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 }