增加相关功能
This commit is contained in:
@@ -0,0 +1,184 @@
|
||||
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"),
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user