feat: 优化对话历史功能

This commit is contained in:
2026-06-20 19:57:36 +08:00
parent 2a7d4c74d4
commit ab07e01adf
15 changed files with 653 additions and 174 deletions

View File

@@ -46,7 +46,6 @@ func (e *EinoOrchestrator) ProcessQuery(
ctx context.Context,
sessionID string,
req models.WsQuery,
history []models.Message,
sender orchestrator.Sender,
) error {
log := logger.Log
@@ -112,15 +111,7 @@ func (e *EinoOrchestrator) ProcessQuery(
ctx = WithStartTime(ctx, startTime)
ctx = WithPipelineState(ctx, genLocalState(ctx))
// 6. 追加用户消息到历史
if req.Text != "" {
_ = e.sessionMgr.AppendMessage(ctx, sessionID, models.Message{
Role: "user",
Content: req.Text,
})
}
// 7. 调用 GraphStream 模式 + 运行时 Callback
// 6. 调用 GraphStream 模式 + 运行时 Callback
streamReader, err := e.graph.Runnable.Stream(ctx, input, e.callbacks)
if err != nil {
log.Errorw("Graph Stream 启动失败", "error", err)
@@ -133,7 +124,7 @@ func (e *EinoOrchestrator) ProcessQuery(
return err
}
// 8. 消费 StreamReader触发整条链路执行side effects 推送消息到客户端)
// 7. 消费 StreamReader触发整条链路执行side effects 推送消息到客户端)
var output PipelineOutput
for {
o, err := streamReader.Recv()
@@ -147,12 +138,28 @@ func (e *EinoOrchestrator) ProcessQuery(
output = o
}
// 8. 追加用户消息到历史(使用 STT 结果,兼容文本输入和语音输入)
userText := output.TranscribedText
if userText == "" {
userText = req.Text // fallback 到原始文本输入
}
if userText != "" {
if err := e.sessionMgr.AppendMessage(ctx, sessionID, models.Message{
Role: "user",
Content: userText,
}); err != nil {
log.Errorw("追加用户消息到历史失败", "session", sessionID, "error", err)
}
}
// 9. 追加助手消息到历史
if output.FullResponse != "" {
_ = e.sessionMgr.AppendMessage(ctx, sessionID, models.Message{
if err := e.sessionMgr.AppendMessage(ctx, sessionID, models.Message{
Role: "assistant",
Content: output.FullResponse,
})
}); err != nil {
log.Errorw("追加助手消息到历史失败", "session", sessionID, "error", err)
}
}
latency := time.Since(startTime).Milliseconds()

View File

@@ -14,13 +14,11 @@ type Orchestrator interface {
// ctx 用于整体超时和中断控制。
// sessionID 用于会话管理和历史获取。
// req 包含图像和音频数据。
// history 是最近的对话历史。
// sender 用于向客户端推送消息。
ProcessQuery(
ctx context.Context,
sessionID string,
req models.WsQuery,
history []models.Message,
sender Sender,
) error
}

View File

@@ -18,6 +18,7 @@ type ConversationSummary struct {
Title string `json:"title"`
LastMessage string `json:"last_message"`
MessageCount int `json:"message_count"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}

View File

@@ -142,11 +142,11 @@ func (m *MemoryManager) Create(ctx context.Context, userID string, config models
}
m.mu.Unlock()
// Write-Through异步写 PG
// Write-Through异步写 PG(使用 Background context避免 HTTP 请求结束后 context 被取消)
if m.sessRepo != nil {
go func() {
cfgJSON, _ := json.Marshal(config)
if err := m.sessRepo.Save(ctx, store.SessionRecord{
if err := m.sessRepo.Save(context.Background(), store.SessionRecord{
ID: id, UserID: userID, Title: models.DefaultSessionTitle,
Config: cfgJSON, CreatedAt: now, UpdatedAt: now,
}); err != nil {
@@ -204,11 +204,11 @@ func (m *MemoryManager) UpdateConfig(ctx context.Context, sessionID string, patc
cfg := entry.session.Config
m.mu.Unlock()
// Write-Through异步更新 PG
// Write-Through异步更新 PG(使用 Background context
if m.sessRepo != nil {
go func() {
cfgJSON, _ := json.Marshal(cfg)
if err := m.sessRepo.UpdateConfig(ctx, sessionID, cfgJSON); err != nil {
if err := m.sessRepo.UpdateConfig(context.Background(), sessionID, cfgJSON); err != nil {
logger.Log.Warnw("update session config in DB failed", "session", sessionID, "error", err)
}
}()
@@ -233,10 +233,10 @@ func (m *MemoryManager) UpdateTitle(ctx context.Context, sessionID string, title
entry.lastActive = time.Now()
m.mu.Unlock()
// Write-Through异步更新 PG
// Write-Through异步更新 PG(使用 Background context
if m.sessRepo != nil {
go func() {
if err := m.sessRepo.UpdateTitle(ctx, sessionID, title); err != nil {
if err := m.sessRepo.UpdateTitle(context.Background(), sessionID, title); err != nil {
logger.Log.Warnw("update session title in DB failed", "session", sessionID, "error", err)
}
}()
@@ -271,6 +271,7 @@ func (m *MemoryManager) ListByUser(ctx context.Context, userID string, page, siz
list = append(list, ConversationSummary{
ID: rec.ID,
Title: rec.Title,
CreatedAt: rec.CreatedAt,
UpdatedAt: rec.UpdatedAt,
})
sessionIDs = append(sessionIDs, rec.ID)
@@ -320,6 +321,7 @@ func (m *MemoryManager) listByUserFromMemory(ctx context.Context, userID string,
summary := ConversationSummary{
ID: entry.session.ID,
Title: entry.session.Title,
CreatedAt: entry.session.CreatedAt,
UpdatedAt: entry.lastActive,
}
summary.MessageCount = len(entry.history)
@@ -394,8 +396,10 @@ func (m *MemoryManager) AppendMessage(_ context.Context, sessionID string, msg m
entry.history = append(entry.history, msg)
// 自动更新标题:首条 user 消息时,如果标题为默认值,自动更新为消息前 20 字符
titleUpdated := false
if msg.Role == "user" && entry.session.Title == models.DefaultSessionTitle {
entry.session.Title = generateTitle(msg.Content)
titleUpdated = true
}
// 超过上限时裁剪,保留最新的 maxHistory 条
@@ -406,13 +410,31 @@ func (m *MemoryManager) AppendMessage(_ context.Context, sessionID string, msg m
now := time.Now()
entry.lastActive = now
entry.session.UpdatedAt = now
// 复制标题(释放锁后安全使用)
persistTitle := entry.session.Title
m.mu.Unlock()
// Write-Through异步写冷存储,不阻塞调用方
// Write-Through消息同步写入 PostgreSQL保证调用顺序 = 插入顺序,
// 避免用户消息和 AI 消息的异步 goroutine 执行顺序不确定导致排序错乱)
if m.msgRepo != nil {
if err := m.msgRepo.SaveMessage(context.Background(), sessionID, msg, 0); err != nil {
logger.Log.Warnw("persist message failed", "session", sessionID, "error", err)
}
}
// Write-Through异步更新会话元数据标题 + updated_at到 PostgreSQL
if m.sessRepo != nil {
go func() {
if err := m.msgRepo.SaveMessage(context.Background(), sessionID, msg, 0); err != nil {
logger.Log.Warnw("persist message failed", "session", sessionID, "error", err)
if titleUpdated {
if err := m.sessRepo.UpdateTitle(context.Background(), sessionID, persistTitle); err != nil {
logger.Log.Warnw("persist session title failed", "session", sessionID, "error", err)
}
} else {
// 即使标题没变,也要刷新 updated_at保证列表排序正确
if err := m.sessRepo.Touch(context.Background(), sessionID); err != nil {
logger.Log.Warnw("touch session in DB failed", "session", sessionID, "error", err)
}
}
}()
}
@@ -540,10 +562,10 @@ func (m *MemoryManager) Destroy(ctx context.Context, sessionID string) error {
delete(m.sessions, sessionID)
m.mu.Unlock()
// Write-Through异步删除 PG
// Write-Through异步删除 PG(使用 Background context
if m.sessRepo != nil {
go func() {
if err := m.sessRepo.Delete(ctx, sessionID); err != nil {
if err := m.sessRepo.Delete(context.Background(), sessionID); err != nil {
logger.Log.Warnw("delete session from DB failed", "session", sessionID, "error", err)
}
}()

View File

@@ -98,15 +98,13 @@ func ServeWS(sessionMgr session.Manager, orch orchestrator.Orchestrator, cfg *co
heartbeatTimeout := time.Duration(cfg.Server.HeartbeatTimeout) * time.Second
version := cfg.App.Version
maxHistory := cfg.Session.MaxHistory
return func(c *gin.Context) {
serveWS(c, sessionMgr, orch, upgrader, heartbeatInterval, heartbeatTimeout, version, maxHistory, tokenMgr)
serveWS(c, sessionMgr, orch, upgrader, heartbeatInterval, heartbeatTimeout, version, tokenMgr)
}
}
func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orchestrator,
upgrader websocket.Upgrader, heartbeatInterval, heartbeatTimeout time.Duration, version string, maxHistory int, tokenMgr *auth.TokenManager) {
upgrader websocket.Upgrader, heartbeatInterval, heartbeatTimeout time.Duration, version string, tokenMgr *auth.TokenManager) {
// --- JWT 认证upgrade 前完成,失败直接返回 HTTP 错误) ---
token := c.Query("token")
@@ -236,9 +234,6 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
logger.Log.Warnw("set active request failed", "session", sessionID, "error", err)
}
// 获取对话历史
history, _ := client.sessionMgr.GetHistory(context.Background(), sessionID, maxHistory)
// 创建可取消的 context
ctx, cancel := context.WithCancel(context.Background())
client.mu.Lock()
@@ -260,7 +255,7 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
_ = client.sessionMgr.ClearActiveRequest(context.Background(), sessionID)
}()
if err := client.orchestrator.ProcessQuery(ctx, sessionID, msg, history, sender); err != nil {
if err := client.orchestrator.ProcessQuery(ctx, sessionID, msg, sender); err != nil {
logger.Log.Errorw("process query failed", "session", sessionID, "request", msg.RequestID, "error", err)
}
}()

View File

@@ -48,7 +48,6 @@ func (m *MockOrchestrator) ProcessQuery(
ctx context.Context,
sessionID string,
req models.WsQuery,
history []models.Message,
sender orchestrator.Sender,
) error {
if m.Err != nil {