156 lines
4.5 KiB
Go
156 lines
4.5 KiB
Go
package session
|
|
|
|
import (
|
|
sessionDao "ai-agent/workflow/dao/session"
|
|
sessionDto "ai-agent/workflow/model/dto/session"
|
|
"ai-agent/workflow/model/entity"
|
|
"context"
|
|
"fmt"
|
|
"sort"
|
|
|
|
"gitea.redpowerfuture.com/red-future/common/utils"
|
|
"github.com/gogf/gf/v2/os/gtime"
|
|
)
|
|
|
|
var SessionService = &sessionService{}
|
|
|
|
type sessionService struct{}
|
|
|
|
func (s *sessionService) Create(ctx context.Context, req *sessionDto.CreateSessionReq) (res *sessionDto.CreateSessionRes, err error) {
|
|
id, err := sessionDao.SessionDao.Insert(ctx, req)
|
|
if err != nil {
|
|
return
|
|
}
|
|
return &sessionDto.CreateSessionRes{Id: id}, nil
|
|
}
|
|
|
|
func (s *sessionService) List(ctx context.Context, req *sessionDto.ListSessionReq) (res *sessionDto.ListSessionRes, err error) {
|
|
user, err := utils.GetUserInfo(ctx)
|
|
if err != nil {
|
|
return
|
|
}
|
|
list, total, err := sessionDao.SessionDao.List(ctx, user.UserName, req.Page)
|
|
if err != nil {
|
|
return
|
|
}
|
|
res = &sessionDto.ListSessionRes{Total: total}
|
|
for _, item := range list {
|
|
res.List = append(res.List, &sessionDto.VOSession{
|
|
Id: item.Id,
|
|
SessionName: item.SessionName,
|
|
CreatedAt: item.CreatedAt,
|
|
})
|
|
}
|
|
return
|
|
}
|
|
|
|
func (s *sessionService) Delete(ctx context.Context, req *sessionDto.DeleteSessionReq) (err error) {
|
|
return sessionDao.SessionDao.DeleteCascade(ctx, req.Id)
|
|
}
|
|
|
|
// ListSessionResults 会话内全部结果:工作流 + 普通对话混排,按创建时间倒序
|
|
func (s *sessionService) ListSessionResults(ctx context.Context, req *sessionDto.ListSessionResultsReq) (res *sessionDto.ListSessionResultsRes, err error) {
|
|
wfList, err := sessionDao.WorkflowSessionResultDao.ListBySession(ctx, req.SessionId)
|
|
if err != nil {
|
|
return
|
|
}
|
|
chatList, err := sessionDao.ChatSessionResultDao.ListBySession(ctx, req.SessionId)
|
|
if err != nil {
|
|
return
|
|
}
|
|
|
|
type mixed struct {
|
|
createdAt *gtime.Time
|
|
vo *sessionDto.VOSessionResult
|
|
}
|
|
var items []mixed
|
|
for _, w := range wfList {
|
|
items = append(items, mixed{createdAt: w.CreatedAt, vo: wfResultVO(w)})
|
|
}
|
|
for _, c := range chatList {
|
|
items = append(items, mixed{createdAt: c.CreatedAt, vo: chatResultVO(c)})
|
|
}
|
|
sort.Slice(items, func(i, j int) bool {
|
|
return items[i].createdAt.After(items[j].createdAt)
|
|
})
|
|
|
|
res = new(sessionDto.ListSessionResultsRes)
|
|
for _, it := range items {
|
|
res.List = append(res.List, it.vo)
|
|
}
|
|
return
|
|
}
|
|
|
|
func wfResultVO(w *entity.WorkflowSessionResult) *sessionDto.VOSessionResult {
|
|
return &sessionDto.VOSessionResult{
|
|
ResultId: w.Id,
|
|
Type: "workflow",
|
|
Status: w.Status,
|
|
FlowId: w.FlowId,
|
|
FlowName: w.FlowName,
|
|
RequestParams: w.RequestParams,
|
|
ResultParams: w.ResultParams,
|
|
TotalTokens: w.TotalTokens,
|
|
TotalFee: w.TotalFee,
|
|
ErrorMsg: w.ErrorMessage,
|
|
CreatedAt: w.CreatedAt,
|
|
}
|
|
}
|
|
|
|
func chatResultVO(c *entity.ChatSessionResult) *sessionDto.VOSessionResult {
|
|
status := sessionDto.ResultStatusSuccess
|
|
if c.ErrorMessage != "" {
|
|
status = sessionDto.ResultStatusFailed
|
|
}
|
|
return &sessionDto.VOSessionResult{
|
|
ResultId: c.Id,
|
|
Type: "chat",
|
|
Status: status,
|
|
Question: c.Question,
|
|
Answer: c.Answer,
|
|
TotalTokens: c.TotalTokens,
|
|
TotalFee: c.TotalFee,
|
|
ErrorMsg: c.ErrorMessage,
|
|
CreatedAt: c.CreatedAt,
|
|
}
|
|
}
|
|
|
|
func (s *sessionService) DeleteResult(ctx context.Context, req *sessionDto.DeleteSessionResultReq) (err error) {
|
|
switch req.Type {
|
|
case "workflow":
|
|
return sessionDao.WorkflowSessionResultDao.SoftDelete(ctx, req.Id)
|
|
case "chat":
|
|
return sessionDao.ChatSessionResultDao.SoftDelete(ctx, req.Id)
|
|
default:
|
|
return fmt.Errorf("未知结果类型: %s", req.Type)
|
|
}
|
|
}
|
|
|
|
// ListWorkflowResults 工作流维度结果平铺列表(含已删除会话下的结果)
|
|
func (s *sessionService) ListWorkflowResults(ctx context.Context, req *sessionDto.ListWorkflowResultsReq) (res *sessionDto.ListWorkflowResultsRes, err error) {
|
|
user, err := utils.GetUserInfo(ctx)
|
|
if err != nil {
|
|
return
|
|
}
|
|
list, total, err := sessionDao.WorkflowSessionResultDao.ListWorkflowResults(ctx, user.UserName, req.FlowId, req.Page)
|
|
if err != nil {
|
|
return
|
|
}
|
|
res = &sessionDto.ListWorkflowResultsRes{Total: total}
|
|
for _, w := range list {
|
|
res.List = append(res.List, &sessionDto.VOWorkflowResult{
|
|
ResultId: w.Id,
|
|
SessionId: w.SessionId,
|
|
FlowId: w.FlowId,
|
|
FlowName: w.FlowName,
|
|
RequestParams: w.RequestParams,
|
|
ResultParams: w.ResultParams,
|
|
Status: w.Status,
|
|
TotalTokens: w.TotalTokens,
|
|
TotalFee: w.TotalFee,
|
|
ErrorMsg: w.ErrorMessage,
|
|
CreatedAt: w.CreatedAt,
|
|
})
|
|
}
|
|
return
|
|
} |