ai_scheduler/internal/biz/advice_chat.go

246 lines
7.8 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 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 == "complete" {
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: "历史聊天记录:\n" + pkg.JsonStringIgonErr(chatList),
},
{
Role: openai.ChatMessageRoleUser,
Content: chat.Content,
},
}
return message, nil
}
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
cursor, err := a.mongo.Co(a.advicerChatHisMongo).Find(ctx, filter, options.Find().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
}
return
}
func (a *AdviceChatBiz) taskPrompt(session *dbmodel.AiAdviceSession) string {
return "[当前时间]" + time.Now().Format("2006-01-02 15:04:05")
}
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),
},
{
Role: openai.ChatMessageRoleSystem,
Content: a.assistantPrompt(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(constants.BasePrompt2)
return prompt.String()
}
func (a *AdviceChatBiz) assistantPrompt(chatData *entitys.ChatData) string {
return pkg.JsonStringIgonErr(chatData)
}
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
}