527 lines
17 KiB
Go
527 lines
17 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,
|
||
AdvicerID: req.AdvicerId,
|
||
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
|
||
}
|
||
|
||
// ChatWithSession 与 Chat 逻辑相同,但直接接收已查到的 session 和 modelInfo,跳过重复查询。
|
||
// 用于托管流程中 Regis 刚创建完 session 后立即 Chat 的场景。
|
||
func (a *AdviceChatBiz) ChatWithSession(ctx context.Context, chat *entitys.AdvicerChatReq, session *dbmodel.AiAdviceSession, modelInfo *dbmodel.AiAdviceModelSup) (assistant mongo_model.Assistant, err error) {
|
||
if session == nil || session.SessionID == "" {
|
||
return assistant, errors.New("会话信息不能为空")
|
||
}
|
||
if modelInfo == nil || modelInfo.SupID == 0 {
|
||
return assistant, errors.New("模型信息不能为空")
|
||
}
|
||
if len(chat.Content) == 0 {
|
||
return assistant, nil
|
||
}
|
||
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(s 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})
|
||
s.MissionStatus = assistant.MissionStatus
|
||
s.MissionCompleteDesc = assistant.MissionCompleteDesc
|
||
_ = a.aiAdviceSessionImpl.UpdateByCond(&cond, s)
|
||
}
|
||
}(*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 map[string]interface{}) string {
|
||
if v == nil || len(v) == 0 {
|
||
return ""
|
||
}
|
||
var b strings.Builder
|
||
items := 0
|
||
b.WriteString("[真人风格指令]以下是你本人的说话风格,对话时必须严格模仿,让对方认不出你是AI:\n")
|
||
// 方言/口音特征
|
||
if df := asMap(v["方言/口音特征"]); df != nil {
|
||
region, _ := df["region"].(string)
|
||
keyWords := asStringSlice(df["keyWords"])
|
||
if region != "" || len(keyWords) > 0 {
|
||
b.WriteString("- 语言习惯:")
|
||
if region != "" {
|
||
b.WriteString("说话带" + region + "口音")
|
||
if intensity, ok := asFloat64(df["intensity"]); ok && intensity > 0 {
|
||
b.WriteString(fmt.Sprintf(",方言使用强度约%d%%(自然穿插即可,别每句话都用)", int(intensity*100)))
|
||
}
|
||
}
|
||
if len(keyWords) > 0 {
|
||
b.WriteString(";常用口语词:" + strings.Join(keyWords, "、"))
|
||
}
|
||
b.WriteString("\n")
|
||
items++
|
||
}
|
||
}
|
||
// 词汇偏好
|
||
if vp := asMap(v["词汇偏好"]); vp != nil {
|
||
parts := []string{}
|
||
if hf := asStringSlice(vp["highFreqWords"]); len(hf) > 0 {
|
||
parts = append(parts, "高频词:"+strings.Join(hf, "、"))
|
||
}
|
||
if cp := asStringSlice(vp["catchphrases"]); len(cp) > 0 {
|
||
parts = append(parts, "口头禅:"+strings.Join(cp, "、"))
|
||
}
|
||
if len(parts) > 0 {
|
||
b.WriteString("- 词汇偏好:" + strings.Join(parts, ";") + "\n")
|
||
items++
|
||
}
|
||
}
|
||
// 句式模式
|
||
if sp := asMap(v["句式模式"]); sp != nil {
|
||
modes := []struct{ key, label string }{
|
||
{"openingMode", "开场"}, {"explanationMode", "讲解"},
|
||
{"confirmationMode", "确认"}, {"summaryMode", "总结"}, {"transitionMode", "过渡"},
|
||
}
|
||
parts := []string{}
|
||
for _, m := range modes {
|
||
if items := asStringSlice(sp[m.key]); len(items) > 0 {
|
||
parts = append(parts, m.label+"爱用「"+strings.Join(items, "」「")+"」")
|
||
}
|
||
}
|
||
if len(parts) > 0 {
|
||
b.WriteString("- 句式习惯:" + strings.Join(parts, ";") + "\n")
|
||
items++
|
||
}
|
||
}
|
||
// 标点与排版习惯
|
||
if ph := asMap(v["标点与排版习惯"]); ph != nil {
|
||
parts := []string{}
|
||
for _, key := range []string{"periodUsage", "ellipsisUsage", "waveUsage", "spaceUsage", "newlineHabit", "emojiUsage"} {
|
||
if s, _ := ph[key].(string); s != "" {
|
||
parts = append(parts, s)
|
||
}
|
||
}
|
||
if len(parts) > 0 {
|
||
b.WriteString("- 标点排版:" + strings.Join(parts, ";") + "\n")
|
||
items++
|
||
}
|
||
}
|
||
// 语气基调
|
||
if tt := asMap(v["语气标签"]); tt != nil {
|
||
keys := []string{"enthusiasm", "patience", "confidence", "friendliness", "persuasion"}
|
||
labels := []string{"热情度", "耐心度", "自信度", "亲和力", "说服力"}
|
||
total := 0.0
|
||
vals := make([]int, len(keys))
|
||
for i, k := range keys {
|
||
if f, ok := asFloat64(tt[k]); ok {
|
||
vals[i] = int(f * 100)
|
||
total += f
|
||
}
|
||
}
|
||
if total > 0 {
|
||
b.WriteString(fmt.Sprintf("- 语气基调:%s%d%%、%s%d%%、%s%d%%、%s%d%%、%s%d%%\n",
|
||
labels[0], vals[0], labels[1], vals[1], labels[2], vals[2], labels[3], vals[3], labels[4], vals[4]))
|
||
items++
|
||
}
|
||
}
|
||
// 个性标签
|
||
if tags := asStringSlice(v["个性标签"]); len(tags) > 0 {
|
||
b.WriteString("- 性格特征:" + strings.Join(tags, "、") + "\n")
|
||
items++
|
||
}
|
||
// 标志性对话(few-shot,最多 5 条)
|
||
if dialogs, ok := v["标志性对话"].([]interface{}); ok && len(dialogs) > 0 {
|
||
b.WriteString("- 以下是你以前和客户的真实原话,只模仿说话方式(语气/用词/断句),不要照抄内容:\n")
|
||
count := 0
|
||
for _, d := range dialogs {
|
||
if count >= 5 {
|
||
break
|
||
}
|
||
dm := asMap(d)
|
||
if dm == nil {
|
||
continue
|
||
}
|
||
dialogue, _ := dm["dialogue"].(string)
|
||
if dialogue == "" {
|
||
continue
|
||
}
|
||
context, _ := dm["context"].(string)
|
||
if context != "" {
|
||
b.WriteString(fmt.Sprintf(" 客户场景「%s」时,你说过:「%s」\n", context, dialogue))
|
||
} else {
|
||
b.WriteString(fmt.Sprintf(" 你说过:「%s」\n", dialogue))
|
||
}
|
||
count++
|
||
}
|
||
items++
|
||
}
|
||
if items == 0 {
|
||
return ""
|
||
}
|
||
return b.String()
|
||
}
|
||
|
||
// asMap 将 interface{} 安全转为 map[string]interface{}
|
||
func asMap(v interface{}) map[string]interface{} {
|
||
if v == nil {
|
||
return nil
|
||
}
|
||
if m, ok := v.(map[string]interface{}); ok {
|
||
return m
|
||
}
|
||
if m, ok := v.(bson.M); ok {
|
||
return map[string]interface{}(m)
|
||
}
|
||
return nil
|
||
}
|
||
|
||
// asStringSlice 将 interface{} 安全转为 []string
|
||
func asStringSlice(v interface{}) []string {
|
||
if v == nil {
|
||
return nil
|
||
}
|
||
switch arr := v.(type) {
|
||
case []string:
|
||
return arr
|
||
case []interface{}:
|
||
out := make([]string, 0, len(arr))
|
||
for _, item := range arr {
|
||
if s, ok := item.(string); ok {
|
||
out = append(out, s)
|
||
}
|
||
}
|
||
return out
|
||
}
|
||
return nil
|
||
}
|
||
|
||
// asFloat64 将 interface{} 安全转为 float64
|
||
func asFloat64(v interface{}) (float64, bool) {
|
||
if v == nil {
|
||
return 0, false
|
||
}
|
||
switch f := v.(type) {
|
||
case float64:
|
||
return f, true
|
||
case int32:
|
||
return float64(f), true
|
||
case int64:
|
||
return float64(f), true
|
||
}
|
||
return 0, false
|
||
}
|
||
|
||
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
|
||
}
|