375 lines
13 KiB
Go
375 lines
13 KiB
Go
package biz
|
||
|
||
import (
|
||
"ai_scheduler/internal/biz/llm_service/third_party"
|
||
"ai_scheduler/internal/data/constants"
|
||
"ai_scheduler/internal/data/mongo_model"
|
||
|
||
"ai_scheduler/internal/data/impl"
|
||
dbmodel "ai_scheduler/internal/data/model"
|
||
"ai_scheduler/internal/entitys"
|
||
"ai_scheduler/internal/pkg"
|
||
"ai_scheduler/utils"
|
||
"context"
|
||
"encoding/json"
|
||
"errors"
|
||
"fmt"
|
||
|
||
"github.com/google/uuid"
|
||
"github.com/sashabaranov/go-openai"
|
||
"go.mongodb.org/mongo-driver/bson"
|
||
"go.mongodb.org/mongo-driver/mongo/options"
|
||
"xorm.io/builder"
|
||
|
||
"strings"
|
||
"time"
|
||
)
|
||
|
||
type AdviceChatBiz struct {
|
||
openai *third_party.OpenAi
|
||
rdb *utils.Rdb
|
||
aiAdviceSessionImpl *impl.AiAdviceSessionImpl
|
||
aiAdviceModelSupImpl *impl.AiAdviceModelSupImpl
|
||
advicerChatHisMongo *mongo_model.AdvicerChatHisMongo
|
||
mongo *pkg.Mongo
|
||
}
|
||
|
||
func NewAdviceChatBiz(
|
||
openai *third_party.OpenAi,
|
||
rdb *utils.Rdb,
|
||
aiAdviceSessionImpl *impl.AiAdviceSessionImpl,
|
||
aiAdviceModelSupImpl *impl.AiAdviceModelSupImpl,
|
||
advicerChatHisMongo *mongo_model.AdvicerChatHisMongo,
|
||
mongo *pkg.Mongo,
|
||
) *AdviceChatBiz {
|
||
return &AdviceChatBiz{
|
||
openai: openai,
|
||
rdb: rdb,
|
||
aiAdviceSessionImpl: aiAdviceSessionImpl,
|
||
aiAdviceModelSupImpl: aiAdviceModelSupImpl,
|
||
advicerChatHisMongo: advicerChatHisMongo,
|
||
mongo: mongo,
|
||
}
|
||
}
|
||
|
||
func (a *AdviceChatBiz) contextCache(ctx context.Context, chatData *entitys.ChatData, req *entitys.AdvicerChatRegistReq, projectInfo *entitys.AdvicerProjectInfoRes) (promptJson string, contextCache string, err error) {
|
||
switch constants.Mode(projectInfo.ModelInfo.Mode) {
|
||
case constants.ModeResponse, constants.ModeContext:
|
||
//ModeContext依赖火山专有Context缓存API,统一降级为Responses存储+PreviousResponseID续写
|
||
prompt, err := a.buildBasePromptResponse(ctx, chatData, req)
|
||
if err != nil {
|
||
return "", "", err
|
||
}
|
||
cache, err := a.openai.CreateResponseMessages(ctx, projectInfo.ModelInfo.Key, projectInfo.ModelInfo.URL, projectInfo.ModelInfo.ChatModel, prompt, "")
|
||
if err != nil {
|
||
return "", "", err
|
||
}
|
||
contextCache = cache.ID
|
||
promptJson = pkg.JsonStringIgonErr(prompt)
|
||
default:
|
||
return "", "", fmt.Errorf("未知的mode类型:%d", projectInfo.ModelInfo.Mode)
|
||
}
|
||
return
|
||
}
|
||
|
||
func (a *AdviceChatBiz) Regis(ctx context.Context, chatData *entitys.ChatData, req *entitys.AdvicerChatRegistReq, projectInfo *entitys.AdvicerProjectInfoRes) (string, error) {
|
||
promptJson, contextCache, err := a.contextCache(ctx, chatData, req, projectInfo)
|
||
if err != nil {
|
||
return "", err
|
||
}
|
||
sessionId := uuid.New().String()
|
||
//创建会话
|
||
err = a.aiAdviceSessionImpl.Add(ctx, &dbmodel.AiAdviceSession{
|
||
SessionID: sessionId,
|
||
ProjectID: projectInfo.Base.ProjectID,
|
||
SupID: projectInfo.Base.ModelSupID,
|
||
AdvicerVersionID: req.AdvicerVersionId,
|
||
ClientID: req.ClientId,
|
||
TalkSkillID: req.TalkSkillId,
|
||
Mission: req.Mission,
|
||
ContextCache: contextCache,
|
||
CreateAt: time.Now(),
|
||
})
|
||
if err != nil {
|
||
return "", err
|
||
}
|
||
err = a.rdb.Rdb.SetEx(ctx, sessionId, promptJson, 3600*time.Second).Err()
|
||
|
||
return sessionId, err
|
||
}
|
||
|
||
func (a *AdviceChatBiz) Chat(ctx context.Context, chat *entitys.AdvicerChatReq) (assistant mongo_model.Assistant, err error) {
|
||
var session dbmodel.AiAdviceSession
|
||
cond := builder.NewCond()
|
||
cond = cond.And(builder.Eq{"session_id": chat.SessionId})
|
||
err = a.aiAdviceSessionImpl.GetOneBySearchToStrut(&cond, &session)
|
||
if err != nil {
|
||
return assistant, err
|
||
}
|
||
if session.SessionID == "" {
|
||
return assistant, errors.New("未找到会话信息")
|
||
}
|
||
if len(chat.Content) == 0 {
|
||
return assistant, nil
|
||
}
|
||
var modelInfo dbmodel.AiAdviceModelSup
|
||
cond = builder.NewCond()
|
||
cond = cond.And(builder.Eq{"sup_id": session.SupID})
|
||
err = a.aiAdviceModelSupImpl.GetOneBySearchToStrut(&cond, &modelInfo)
|
||
if err != nil {
|
||
return assistant, err
|
||
}
|
||
if modelInfo.SupID == 0 {
|
||
return assistant, errors.New("未找到模型信息")
|
||
}
|
||
chatHis, err := a.getChatHis(ctx, session.SessionID, 6)
|
||
if err != nil {
|
||
return assistant, err
|
||
}
|
||
prompt, err := a.buildChatPromptResponse(ctx, chat, &session, chatHis)
|
||
if err != nil {
|
||
return assistant, err
|
||
}
|
||
resContent, err := a.callLlmResponse(ctx, prompt, modelInfo.Key, modelInfo.URL, modelInfo.ChatModel, session.ContextCache)
|
||
if err != nil {
|
||
return assistant, err
|
||
}
|
||
|
||
result := resContent.GetOutputText()
|
||
if err = json.Unmarshal([]byte(result), &assistant); err != nil {
|
||
return assistant, err
|
||
}
|
||
var inToken, outToken int64
|
||
if resContent.Usage != nil {
|
||
inToken = int64(resContent.Usage.InputTokens)
|
||
outToken = int64(resContent.Usage.OutputTokens)
|
||
}
|
||
chatCtx, cancel := context.WithCancel(context.Background())
|
||
go func(session dbmodel.AiAdviceSession) {
|
||
defer cancel()
|
||
_, _ = a.mongo.Co(a.advicerChatHisMongo).InsertOne(chatCtx, &mongo_model.AdvicerChatHisMongo{
|
||
SessionId: chat.SessionId,
|
||
User: chat.Content,
|
||
Assistant: assistant,
|
||
InToken: inToken,
|
||
OutToken: outToken,
|
||
CreatAt: time.Now(),
|
||
})
|
||
if assistant.MissionStatus == "fail" || assistant.MissionStatus == "completed" {
|
||
cond = builder.NewCond()
|
||
cond = cond.And(builder.Eq{"session_id": chat.SessionId})
|
||
session.MissionStatus = assistant.MissionStatus
|
||
session.MissionCompleteDesc = assistant.MissionCompleteDesc
|
||
_ = a.aiAdviceSessionImpl.UpdateByCond(&cond, session)
|
||
}
|
||
}(session)
|
||
return
|
||
}
|
||
|
||
func (a *AdviceChatBiz) buildChatPromptResponse(ctx context.Context, chat *entitys.AdvicerChatReq, session *dbmodel.AiAdviceSession, chatList []mongo_model.AdvicerChatHisMongoEntity) ([]openai.ResponseInputMessage, error) {
|
||
message := []openai.ResponseInputMessage{
|
||
{
|
||
Role: openai.ChatMessageRoleSystem,
|
||
Content: a.taskPrompt(session),
|
||
},
|
||
{
|
||
Role: openai.ChatMessageRoleSystem,
|
||
Content: formatChatHis(chatList),
|
||
},
|
||
{
|
||
Role: openai.ChatMessageRoleUser,
|
||
Content: chat.Content,
|
||
},
|
||
}
|
||
return message, nil
|
||
}
|
||
|
||
// formatChatHis 将历史记录转写为「客户/我」的自然对话文本,更贴近真人聊天语境
|
||
func formatChatHis(chatList []mongo_model.AdvicerChatHisMongoEntity) string {
|
||
if len(chatList) == 0 {
|
||
return "(暂无历史聊天记录,这是你与客户的第一轮对话)"
|
||
}
|
||
var b strings.Builder
|
||
b.WriteString("以下是你(顾问)与客户最近的聊天记录(按时间先后排列):\n")
|
||
for _, h := range chatList {
|
||
b.WriteString("客户:" + h.User + "\n")
|
||
b.WriteString("我:" + h.Assistant + "\n")
|
||
}
|
||
return b.String()
|
||
}
|
||
|
||
func (a *AdviceChatBiz) getChatHis(ctx context.Context, sessionId string, limit int64) (chatList []mongo_model.AdvicerChatHisMongoEntity, err error) {
|
||
chatList = make([]mongo_model.AdvicerChatHisMongoEntity, 0)
|
||
filter := bson.M{}
|
||
filter["sessionId"] = sessionId
|
||
// 按时间倒序取最近 limit 条(保证获取的是最新对话)
|
||
cursor, err := a.mongo.Co(a.advicerChatHisMongo).Find(ctx, filter,
|
||
options.Find().SetSort(bson.D{{Key: "creatAt", Value: -1}}).SetLimit(limit))
|
||
if err != nil {
|
||
return chatList, err
|
||
}
|
||
for cursor.Next(ctx) {
|
||
var chatHIS mongo_model.AdvicerChatHisMongo
|
||
if err := cursor.Decode(&chatHIS); err != nil {
|
||
return nil, err
|
||
}
|
||
chatList = append(chatList, chatHIS.Entity())
|
||
}
|
||
|
||
if err := cursor.Err(); err != nil {
|
||
return nil, err
|
||
}
|
||
// Mongo 按时间倒序查询,反转回时间正序后再拼入 prompt
|
||
for i, j := 0, len(chatList)-1; i < j; i, j = i+1, j-1 {
|
||
chatList[i], chatList[j] = chatList[j], chatList[i]
|
||
}
|
||
return
|
||
}
|
||
|
||
func (a *AdviceChatBiz) taskPrompt(session *dbmodel.AiAdviceSession) string {
|
||
var b strings.Builder
|
||
b.WriteString("[当前时间]" + time.Now().Format("2006-01-02 15:04:05"))
|
||
status := session.MissionStatus
|
||
if status != "completed" && status != "fail" {
|
||
status = "in_progress"
|
||
}
|
||
b.WriteString("\n[当前任务状态]" + status)
|
||
if session.MissionCompleteDesc != "" {
|
||
b.WriteString("\n[任务最新进展]" + session.MissionCompleteDesc)
|
||
}
|
||
return b.String()
|
||
}
|
||
|
||
func (a *AdviceChatBiz) buildBasePromptResponse(ctx context.Context, chatData *entitys.ChatData, req *entitys.AdvicerChatRegistReq) ([]openai.ResponseInputMessage, error) {
|
||
message := []openai.ResponseInputMessage{
|
||
{
|
||
Role: openai.ChatMessageRoleSystem,
|
||
Content: a.sysPrompt(chatData, req),
|
||
},
|
||
}
|
||
// 销售人设风格指令(自然语言规则 + 真实对话范例 few-shot),无人设数据时跳过
|
||
if persona := a.personaPrompt(chatData); persona != "" {
|
||
message = append(message, openai.ResponseInputMessage{
|
||
Role: openai.ChatMessageRoleSystem,
|
||
Content: persona,
|
||
})
|
||
}
|
||
message = append(message, openai.ResponseInputMessage{
|
||
Role: openai.ChatMessageRoleSystem,
|
||
Content: "[结构化数据]\n" + pkg.JsonStringIgonErr(chatData),
|
||
})
|
||
return message, nil
|
||
}
|
||
|
||
func (a *AdviceChatBiz) sysPrompt(chatData *entitys.ChatData, req *entitys.AdvicerChatRegistReq) string {
|
||
var prompt strings.Builder
|
||
prompt.WriteString(constants.BasePrompt)
|
||
prompt.WriteString(req.Mission)
|
||
prompt.WriteString("\n")
|
||
// 项目风控红线(来自项目模板 rule_dimension),最高优先级约束
|
||
if rule := strings.TrimSpace(chatData.RuleDimension); rule != "" {
|
||
prompt.WriteString(constants.RulePromptPrefix)
|
||
prompt.WriteString(rule)
|
||
prompt.WriteString("\n")
|
||
}
|
||
prompt.WriteString(constants.BasePrompt2)
|
||
return prompt.String()
|
||
|
||
}
|
||
|
||
// personaPrompt 会话注册时注入的人设风格提示(内部委托 formatPersonaPrompt)
|
||
func (a *AdviceChatBiz) personaPrompt(chatData *entitys.ChatData) string {
|
||
return formatPersonaPrompt(chatData.AdvicerVersion)
|
||
}
|
||
|
||
// formatPersonaPrompt 将销售版本的结构化人设编译为自然语言风格指令 + 真实对话范例(few-shot),
|
||
// 让模型模仿「本人」的说话方式,而不是机械读取数据;无人设数据时返回空串。
|
||
// chat 会话与策略引擎(AI 回复/主动触达)共用本函数,保证各链路语气一致。
|
||
func formatPersonaPrompt(v *mongo_model.AdvicerVersionMongoEntity) string {
|
||
if v == nil {
|
||
return ""
|
||
}
|
||
var b strings.Builder
|
||
items := 0
|
||
b.WriteString("[真人风格指令]以下是你本人的说话风格,对话时必须严格模仿,让对方认不出你是AI:\n")
|
||
// 方言特征
|
||
if v.DialectFeatures.Region != "" || len(v.DialectFeatures.KeyWords) > 0 {
|
||
b.WriteString("- 语言习惯:")
|
||
if v.DialectFeatures.Region != "" {
|
||
b.WriteString("说话带" + v.DialectFeatures.Region + "口音")
|
||
if v.DialectFeatures.Intensity > 0 {
|
||
b.WriteString(fmt.Sprintf(",方言使用强度约%d%%(自然穿插即可,别每句话都用)", int(v.DialectFeatures.Intensity*100)))
|
||
}
|
||
}
|
||
if len(v.DialectFeatures.KeyWords) > 0 {
|
||
b.WriteString(";常用口语词:" + strings.Join(v.DialectFeatures.KeyWords, "、"))
|
||
}
|
||
b.WriteString("\n")
|
||
items++
|
||
}
|
||
// 句式习惯
|
||
sp := v.SentencePatterns
|
||
if len(sp.OpeningMode)+len(sp.ExplanationMode)+len(sp.ConfirmationMode)+len(sp.SummaryMode)+len(sp.TransitionMode) > 0 {
|
||
b.WriteString("- 句式习惯:")
|
||
if len(sp.OpeningMode) > 0 {
|
||
b.WriteString("开场爱用「" + strings.Join(sp.OpeningMode, "」「") + "」;")
|
||
}
|
||
if len(sp.ExplanationMode) > 0 {
|
||
b.WriteString("讲解时爱说「" + strings.Join(sp.ExplanationMode, "」「") + "」;")
|
||
}
|
||
if len(sp.ConfirmationMode) > 0 {
|
||
b.WriteString("跟客户确认时爱用「" + strings.Join(sp.ConfirmationMode, "」「") + "」;")
|
||
}
|
||
if len(sp.SummaryMode) > 0 {
|
||
b.WriteString("总结时爱说「" + strings.Join(sp.SummaryMode, "」「") + "」;")
|
||
}
|
||
if len(sp.TransitionMode) > 0 {
|
||
b.WriteString("换话题时爱用「" + strings.Join(sp.TransitionMode, "」「") + "」")
|
||
}
|
||
b.WriteString("\n")
|
||
items++
|
||
}
|
||
// 语气基调
|
||
t := v.ToneTags
|
||
if t.Enthusiasm+t.Patience+t.Confidence+t.Friendliness+t.Persuasion > 0 {
|
||
b.WriteString(fmt.Sprintf("- 语气基调:热情度%d%%、耐心度%d%%、自信度%d%%、亲和力%d%%、说服力%d%%\n",
|
||
int(t.Enthusiasm*100), int(t.Patience*100), int(t.Confidence*100), int(t.Friendliness*100), int(t.Persuasion*100)))
|
||
items++
|
||
}
|
||
// 性格特征
|
||
if len(v.PersonalityTags) > 0 {
|
||
b.WriteString("- 性格特征:" + strings.Join(v.PersonalityTags, "、") + "\n")
|
||
items++
|
||
}
|
||
// 真实对话范例(few-shot,最多 5 条)
|
||
if len(v.SignatureDialogues) > 0 {
|
||
b.WriteString("- 以下是你以前和客户的真实原话,只模仿说话方式(语气/用词/断句),不要照抄内容:\n")
|
||
for i, d := range v.SignatureDialogues {
|
||
if i >= 5 {
|
||
break
|
||
}
|
||
if d.Dialogue == "" {
|
||
continue
|
||
}
|
||
if d.Context != "" {
|
||
b.WriteString(fmt.Sprintf(" 客户场景「%s」时,你说过:「%s」\n", d.Context, d.Dialogue))
|
||
} else {
|
||
b.WriteString(fmt.Sprintf(" 你说过:「%s」\n", d.Dialogue))
|
||
}
|
||
}
|
||
items++
|
||
}
|
||
if items == 0 {
|
||
return ""
|
||
}
|
||
return b.String()
|
||
}
|
||
|
||
func (a *AdviceChatBiz) callLlmResponse(ctx context.Context, request []openai.ResponseInputMessage, key string, url string, modelName string, id string) (*openai.CreateResponseResponse, error) {
|
||
res, err := a.openai.CreateResponseMessages(ctx, key, url, modelName, request, id)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
return res, nil
|
||
}
|