185 lines
5.3 KiB
Go
185 lines
5.3 KiB
Go
package controllers
|
||
|
||
import (
|
||
"encoding/json"
|
||
"fmt"
|
||
"io"
|
||
"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
|
||
}
|
||
|
||
// 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]
|
||
|
||
// 构造prompt
|
||
typeLabel := "供应商"
|
||
nameField := "supplier_name"
|
||
if p.Type == "customer" {
|
||
typeLabel = "客户"
|
||
nameField = "customer_name"
|
||
}
|
||
|
||
prompt := fmt.Sprintf(`你是一个企业信息助手。请根据%s名称"%s",生成该公司的详细信息,以JSON格式返回,包含以下字段:
|
||
- %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代码块标记。如果某些信息不确定,请留空字符串。`, typeLabel, p.CompanyName, nameField)
|
||
|
||
messages := []openaiMessage{
|
||
{Role: "user", Content: prompt},
|
||
}
|
||
|
||
// 调用AI(非流式,超时120s)
|
||
reply, err := callAI(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, "```")
|
||
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
|
||
}
|
||
|
||
// 确保公司名称字段存在
|
||
if _, ok := result[nameField]; !ok || result[nameField] == "" {
|
||
result[nameField] = p.CompanyName
|
||
}
|
||
|
||
c.sgOk(map[string]interface{}{
|
||
"parsed": true,
|
||
"data": result,
|
||
"name": p.CompanyName,
|
||
"type": p.Type,
|
||
"generate_time": time.Now().Format("2006-01-02 15:04:05"),
|
||
})
|
||
}
|