Files
yunzerwebsiteallinone/go/controllers/backend_ai_smart_generate.go
T

567 lines
18 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"
"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, 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,
"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] = 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,
"data": result,
"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
}