2026-06-13 15:25:23 +08:00
|
|
|
|
package session
|
|
|
|
|
|
|
|
|
|
|
|
import (
|
|
|
|
|
|
"context"
|
2026-06-14 18:54:18 +08:00
|
|
|
|
"encoding/json"
|
2026-06-14 17:35:20 +08:00
|
|
|
|
"sort"
|
2026-06-13 15:25:23 +08:00
|
|
|
|
"sync"
|
|
|
|
|
|
"time"
|
|
|
|
|
|
|
|
|
|
|
|
"github.com/google/uuid"
|
|
|
|
|
|
|
|
|
|
|
|
"github.com/hhs/camtalk/internal/logger"
|
|
|
|
|
|
"github.com/hhs/camtalk/internal/models"
|
2026-06-14 17:54:33 +08:00
|
|
|
|
"github.com/hhs/camtalk/internal/store"
|
2026-06-13 15:25:23 +08:00
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
const (
|
|
|
|
|
|
defaultTTL = 30 * time.Minute
|
|
|
|
|
|
defaultHistorySize = 20
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
// sessionEntry 内部会话条目。
|
|
|
|
|
|
type sessionEntry struct {
|
|
|
|
|
|
session models.Session
|
|
|
|
|
|
history []models.Message
|
|
|
|
|
|
activeReqID string
|
|
|
|
|
|
lastActive time.Time
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// MemoryManager 基于内存的 SessionManager 实现。
|
|
|
|
|
|
// 适用于 MVP 和无 Redis 的开发环境。
|
|
|
|
|
|
type MemoryManager struct {
|
|
|
|
|
|
mu sync.RWMutex
|
|
|
|
|
|
sessions map[string]*sessionEntry
|
|
|
|
|
|
ttl time.Duration
|
|
|
|
|
|
maxHistory int
|
|
|
|
|
|
stopCleaner chan struct{}
|
2026-06-14 18:54:18 +08:00
|
|
|
|
msgRepo store.MessageRepository // 可选,消息持久化(Write-Through)
|
|
|
|
|
|
sessRepo store.SessionRepository // 可选,会话持久化(Write-Through)
|
2026-06-14 17:54:33 +08:00
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// Option MemoryManager 的函数式选项。
|
|
|
|
|
|
type Option func(*MemoryManager)
|
|
|
|
|
|
|
|
|
|
|
|
// WithMessageRepository 注入消息持久化仓库,启用 Write-Through 模式。
|
|
|
|
|
|
func WithMessageRepository(repo store.MessageRepository) Option {
|
|
|
|
|
|
return func(m *MemoryManager) {
|
|
|
|
|
|
m.msgRepo = repo
|
|
|
|
|
|
}
|
2026-06-13 15:25:23 +08:00
|
|
|
|
}
|
|
|
|
|
|
|
2026-06-14 18:54:18 +08:00
|
|
|
|
// WithSessionRepository 注入会话持久化仓库,启用会话元数据 Write-Through 模式。
|
|
|
|
|
|
func WithSessionRepository(repo store.SessionRepository) Option {
|
|
|
|
|
|
return func(m *MemoryManager) {
|
|
|
|
|
|
m.sessRepo = repo
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
|
2026-06-13 15:25:23 +08:00
|
|
|
|
// NewMemoryManager 创建内存版 SessionManager。
|
|
|
|
|
|
// ttl 为会话过期时间,maxHistory 为对话历史上限(0 表示使用默认值 20)。
|
2026-06-14 17:54:33 +08:00
|
|
|
|
// opts 为可选配置,如 WithMessageRepository 启用消息持久化。
|
|
|
|
|
|
func NewMemoryManager(ttl time.Duration, maxHistory int, opts ...Option) *MemoryManager {
|
2026-06-13 15:25:23 +08:00
|
|
|
|
if ttl <= 0 {
|
|
|
|
|
|
ttl = defaultTTL
|
|
|
|
|
|
}
|
|
|
|
|
|
if maxHistory <= 0 {
|
|
|
|
|
|
maxHistory = defaultHistorySize
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
m := &MemoryManager{
|
|
|
|
|
|
sessions: make(map[string]*sessionEntry),
|
|
|
|
|
|
ttl: ttl,
|
|
|
|
|
|
maxHistory: maxHistory,
|
|
|
|
|
|
stopCleaner: make(chan struct{}),
|
|
|
|
|
|
}
|
|
|
|
|
|
|
2026-06-14 17:54:33 +08:00
|
|
|
|
for _, opt := range opts {
|
|
|
|
|
|
opt(m)
|
|
|
|
|
|
}
|
|
|
|
|
|
|
2026-06-13 15:25:23 +08:00
|
|
|
|
// 启动后台清理 goroutine,每分钟清除过期会话。
|
|
|
|
|
|
go m.cleanLoop()
|
|
|
|
|
|
|
|
|
|
|
|
return m
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// cleanLoop 后台定期清理过期会话。
|
|
|
|
|
|
func (m *MemoryManager) cleanLoop() {
|
|
|
|
|
|
ticker := time.NewTicker(1 * time.Minute)
|
|
|
|
|
|
defer ticker.Stop()
|
|
|
|
|
|
for {
|
|
|
|
|
|
select {
|
|
|
|
|
|
case <-ticker.C:
|
|
|
|
|
|
m.cleanExpired()
|
|
|
|
|
|
case <-m.stopCleaner:
|
|
|
|
|
|
return
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// cleanExpired 清除所有过期会话。
|
|
|
|
|
|
func (m *MemoryManager) cleanExpired() {
|
|
|
|
|
|
m.mu.Lock()
|
|
|
|
|
|
defer m.mu.Unlock()
|
|
|
|
|
|
|
|
|
|
|
|
now := time.Now()
|
|
|
|
|
|
for id, entry := range m.sessions {
|
|
|
|
|
|
if now.Sub(entry.lastActive) > m.ttl {
|
|
|
|
|
|
delete(m.sessions, id)
|
|
|
|
|
|
logger.Log.Debugw("session expired (cleaner)", "session", id)
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// Stop 停止后台清理 goroutine。应用退出前调用。
|
|
|
|
|
|
func (m *MemoryManager) Stop() {
|
|
|
|
|
|
close(m.stopCleaner)
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// isExpired 检查会话是否过期(调用方需持锁或在已知 entry 存在时调用)。
|
|
|
|
|
|
func (m *MemoryManager) isExpired(entry *sessionEntry) bool {
|
|
|
|
|
|
return time.Since(entry.lastActive) > m.ttl
|
|
|
|
|
|
}
|
|
|
|
|
|
|
2026-06-14 17:35:20 +08:00
|
|
|
|
// Create 创建新会话。userID 为空表示匿名会话。
|
2026-06-14 18:54:18 +08:00
|
|
|
|
func (m *MemoryManager) Create(ctx context.Context, userID string, config models.SessionConfig) (string, error) {
|
2026-06-13 15:25:23 +08:00
|
|
|
|
m.mu.Lock()
|
|
|
|
|
|
|
|
|
|
|
|
id := uuid.New().String()
|
|
|
|
|
|
now := time.Now()
|
|
|
|
|
|
m.sessions[id] = &sessionEntry{
|
|
|
|
|
|
session: models.Session{
|
|
|
|
|
|
ID: id,
|
2026-06-14 17:35:20 +08:00
|
|
|
|
UserID: userID,
|
|
|
|
|
|
Title: models.DefaultSessionTitle,
|
2026-06-13 15:25:23 +08:00
|
|
|
|
CreatedAt: now,
|
2026-06-14 17:35:20 +08:00
|
|
|
|
UpdatedAt: now,
|
2026-06-13 15:25:23 +08:00
|
|
|
|
Config: config,
|
|
|
|
|
|
},
|
|
|
|
|
|
history: make([]models.Message, 0),
|
|
|
|
|
|
lastActive: now,
|
|
|
|
|
|
}
|
2026-06-14 18:54:18 +08:00
|
|
|
|
m.mu.Unlock()
|
|
|
|
|
|
|
|
|
|
|
|
// Write-Through:异步写 PG
|
|
|
|
|
|
if m.sessRepo != nil {
|
|
|
|
|
|
go func() {
|
|
|
|
|
|
cfgJSON, _ := json.Marshal(config)
|
|
|
|
|
|
if err := m.sessRepo.Save(ctx, store.SessionRecord{
|
|
|
|
|
|
ID: id, UserID: userID, Title: models.DefaultSessionTitle,
|
|
|
|
|
|
Config: cfgJSON, CreatedAt: now, UpdatedAt: now,
|
|
|
|
|
|
}); err != nil {
|
|
|
|
|
|
logger.Log.Warnw("persist session failed", "session", id, "error", err)
|
|
|
|
|
|
}
|
|
|
|
|
|
}()
|
|
|
|
|
|
}
|
2026-06-13 15:25:23 +08:00
|
|
|
|
|
2026-06-14 17:35:20 +08:00
|
|
|
|
logger.Log.Debugw("session created", "session", id, "user_id", userID)
|
2026-06-13 15:25:23 +08:00
|
|
|
|
return id, nil
|
|
|
|
|
|
}
|
|
|
|
|
|
|
2026-06-14 18:54:18 +08:00
|
|
|
|
// Get 获取会话。内存中不存在时,尝试从 PG 加载(透明恢复)。
|
|
|
|
|
|
func (m *MemoryManager) Get(ctx context.Context, sessionID string) (*models.Session, error) {
|
2026-06-13 15:25:23 +08:00
|
|
|
|
m.mu.RLock()
|
|
|
|
|
|
entry, ok := m.sessions[sessionID]
|
2026-06-14 18:54:18 +08:00
|
|
|
|
if ok && !m.isExpired(entry) {
|
|
|
|
|
|
sess := entry.session
|
|
|
|
|
|
m.mu.RUnlock()
|
|
|
|
|
|
return &sess, nil
|
|
|
|
|
|
}
|
|
|
|
|
|
m.mu.RUnlock()
|
|
|
|
|
|
|
|
|
|
|
|
// 内存未命中,尝试从 PG 加载
|
|
|
|
|
|
if m.sessRepo != nil {
|
|
|
|
|
|
rec, err := m.sessRepo.FindByID(ctx, sessionID)
|
|
|
|
|
|
if err != nil {
|
|
|
|
|
|
return nil, ErrSessionNotFound
|
|
|
|
|
|
}
|
|
|
|
|
|
sess := m.recordToSession(rec)
|
|
|
|
|
|
// 加载到内存(含消息历史)
|
|
|
|
|
|
if m.msgRepo != nil {
|
|
|
|
|
|
_ = m.LoadSessionFromRepo(ctx, sess)
|
|
|
|
|
|
} else {
|
|
|
|
|
|
_ = m.LoadSession(sess, nil)
|
|
|
|
|
|
}
|
|
|
|
|
|
return sess, nil
|
2026-06-13 15:25:23 +08:00
|
|
|
|
}
|
|
|
|
|
|
|
2026-06-14 18:54:18 +08:00
|
|
|
|
return nil, ErrSessionNotFound
|
2026-06-13 15:25:23 +08:00
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// UpdateConfig 更新会话配置。
|
2026-06-14 18:54:18 +08:00
|
|
|
|
func (m *MemoryManager) UpdateConfig(ctx context.Context, sessionID string, patch models.SessionConfigPatch) error {
|
2026-06-13 15:25:23 +08:00
|
|
|
|
m.mu.Lock()
|
|
|
|
|
|
|
|
|
|
|
|
entry, ok := m.sessions[sessionID]
|
|
|
|
|
|
if !ok || m.isExpired(entry) {
|
2026-06-14 18:54:18 +08:00
|
|
|
|
m.mu.Unlock()
|
2026-06-13 15:25:23 +08:00
|
|
|
|
return ErrSessionNotFound
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
patch.Apply(&entry.session.Config)
|
|
|
|
|
|
entry.lastActive = time.Now()
|
2026-06-14 18:54:18 +08:00
|
|
|
|
cfg := entry.session.Config
|
|
|
|
|
|
m.mu.Unlock()
|
|
|
|
|
|
|
|
|
|
|
|
// Write-Through:异步更新 PG
|
|
|
|
|
|
if m.sessRepo != nil {
|
|
|
|
|
|
go func() {
|
|
|
|
|
|
cfgJSON, _ := json.Marshal(cfg)
|
|
|
|
|
|
if err := m.sessRepo.UpdateConfig(ctx, sessionID, cfgJSON); err != nil {
|
|
|
|
|
|
logger.Log.Warnw("update session config in DB failed", "session", sessionID, "error", err)
|
|
|
|
|
|
}
|
|
|
|
|
|
}()
|
|
|
|
|
|
}
|
2026-06-13 15:25:23 +08:00
|
|
|
|
|
|
|
|
|
|
logger.Log.Debugw("session config updated", "session", sessionID)
|
|
|
|
|
|
return nil
|
|
|
|
|
|
}
|
|
|
|
|
|
|
2026-06-14 17:35:20 +08:00
|
|
|
|
// UpdateTitle 更新会话标题。
|
2026-06-14 18:54:18 +08:00
|
|
|
|
func (m *MemoryManager) UpdateTitle(ctx context.Context, sessionID string, title string) error {
|
2026-06-14 17:35:20 +08:00
|
|
|
|
m.mu.Lock()
|
|
|
|
|
|
|
|
|
|
|
|
entry, ok := m.sessions[sessionID]
|
|
|
|
|
|
if !ok || m.isExpired(entry) {
|
2026-06-14 18:54:18 +08:00
|
|
|
|
m.mu.Unlock()
|
2026-06-14 17:35:20 +08:00
|
|
|
|
return ErrSessionNotFound
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
entry.session.Title = title
|
|
|
|
|
|
entry.session.UpdatedAt = time.Now()
|
|
|
|
|
|
entry.lastActive = time.Now()
|
2026-06-14 18:54:18 +08:00
|
|
|
|
m.mu.Unlock()
|
|
|
|
|
|
|
|
|
|
|
|
// Write-Through:异步更新 PG
|
|
|
|
|
|
if m.sessRepo != nil {
|
|
|
|
|
|
go func() {
|
|
|
|
|
|
if err := m.sessRepo.UpdateTitle(ctx, sessionID, title); err != nil {
|
|
|
|
|
|
logger.Log.Warnw("update session title in DB failed", "session", sessionID, "error", err)
|
|
|
|
|
|
}
|
|
|
|
|
|
}()
|
|
|
|
|
|
}
|
2026-06-14 17:35:20 +08:00
|
|
|
|
|
|
|
|
|
|
logger.Log.Debugw("session title updated", "session", sessionID, "title", title)
|
|
|
|
|
|
return nil
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// ListByUser 获取用户的对话列表(分页,按 UpdatedAt 降序)。
|
2026-06-14 18:54:18 +08:00
|
|
|
|
// 若配置了 SessionRepository,从 PG 查询(包含内存中已过期的会话)。
|
2026-06-14 17:58:42 +08:00
|
|
|
|
// 若配置了 MessageRepository,消息统计从 PostgreSQL 聚合查询(更准确)。
|
|
|
|
|
|
func (m *MemoryManager) ListByUser(ctx context.Context, userID string, page, size int) ([]ConversationSummary, int, error) {
|
2026-06-14 18:54:18 +08:00
|
|
|
|
if page <= 0 {
|
|
|
|
|
|
page = 1
|
|
|
|
|
|
}
|
|
|
|
|
|
if size <= 0 {
|
|
|
|
|
|
size = 20
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// 优先从 PG 查询会话列表(包含已过期的会话)
|
|
|
|
|
|
if m.sessRepo != nil {
|
|
|
|
|
|
recs, total, err := m.sessRepo.FindByUser(ctx, userID, page, size)
|
|
|
|
|
|
if err != nil {
|
|
|
|
|
|
logger.Log.Warnw("list sessions from DB failed, falling back to in-memory", "error", err)
|
|
|
|
|
|
return m.listByUserFromMemory(ctx, userID, page, size)
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
list := make([]ConversationSummary, 0, len(recs))
|
|
|
|
|
|
var sessionIDs []string
|
|
|
|
|
|
for _, rec := range recs {
|
|
|
|
|
|
list = append(list, ConversationSummary{
|
|
|
|
|
|
ID: rec.ID,
|
|
|
|
|
|
Title: rec.Title,
|
|
|
|
|
|
UpdatedAt: rec.UpdatedAt,
|
|
|
|
|
|
})
|
|
|
|
|
|
sessionIDs = append(sessionIDs, rec.ID)
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// 用内存中的消息数填充
|
|
|
|
|
|
m.mu.RLock()
|
|
|
|
|
|
for i := range list {
|
|
|
|
|
|
if entry, ok := m.sessions[list[i].ID]; ok {
|
|
|
|
|
|
list[i].MessageCount = len(entry.history)
|
|
|
|
|
|
if len(entry.history) > 0 {
|
|
|
|
|
|
list[i].LastMessage = entry.history[len(entry.history)-1].Content
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
m.mu.RUnlock()
|
|
|
|
|
|
|
|
|
|
|
|
// 从 PG 获取更准确的消息统计
|
|
|
|
|
|
if m.msgRepo != nil && len(sessionIDs) > 0 {
|
|
|
|
|
|
if stats, err := m.msgRepo.GetSessionMessageStats(ctx, sessionIDs); err == nil {
|
|
|
|
|
|
for i := range list {
|
|
|
|
|
|
if s, ok := stats[list[i].ID]; ok {
|
|
|
|
|
|
list[i].LastMessage = s.LastMessage
|
|
|
|
|
|
list[i].MessageCount = s.MessageCount
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
return list, total, nil
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// fallback:纯内存查询
|
|
|
|
|
|
return m.listByUserFromMemory(ctx, userID, page, size)
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// listByUserFromMemory 从内存中获取用户的对话列表(无 PG 时的 fallback)。
|
|
|
|
|
|
func (m *MemoryManager) listByUserFromMemory(ctx context.Context, userID string, page, size int) ([]ConversationSummary, int, error) {
|
2026-06-14 17:35:20 +08:00
|
|
|
|
m.mu.RLock()
|
|
|
|
|
|
|
|
|
|
|
|
var list []ConversationSummary
|
2026-06-14 17:58:42 +08:00
|
|
|
|
var sessionIDs []string
|
2026-06-14 17:35:20 +08:00
|
|
|
|
for _, entry := range m.sessions {
|
2026-06-14 18:54:18 +08:00
|
|
|
|
if entry.session.UserID != userID || m.isExpired(entry) {
|
2026-06-14 17:35:20 +08:00
|
|
|
|
continue
|
|
|
|
|
|
}
|
|
|
|
|
|
summary := ConversationSummary{
|
2026-06-14 17:58:42 +08:00
|
|
|
|
ID: entry.session.ID,
|
|
|
|
|
|
Title: entry.session.Title,
|
|
|
|
|
|
UpdatedAt: entry.lastActive,
|
2026-06-14 17:35:20 +08:00
|
|
|
|
}
|
2026-06-14 17:58:42 +08:00
|
|
|
|
summary.MessageCount = len(entry.history)
|
2026-06-14 17:35:20 +08:00
|
|
|
|
if len(entry.history) > 0 {
|
|
|
|
|
|
summary.LastMessage = entry.history[len(entry.history)-1].Content
|
|
|
|
|
|
}
|
|
|
|
|
|
list = append(list, summary)
|
2026-06-14 17:58:42 +08:00
|
|
|
|
sessionIDs = append(sessionIDs, entry.session.ID)
|
|
|
|
|
|
}
|
|
|
|
|
|
m.mu.RUnlock()
|
|
|
|
|
|
|
2026-06-14 18:54:18 +08:00
|
|
|
|
// 从 PG 获取更准确的消息统计
|
2026-06-14 17:58:42 +08:00
|
|
|
|
if m.msgRepo != nil && len(sessionIDs) > 0 {
|
|
|
|
|
|
if stats, err := m.msgRepo.GetSessionMessageStats(ctx, sessionIDs); err == nil {
|
|
|
|
|
|
for i := range list {
|
|
|
|
|
|
if s, ok := stats[list[i].ID]; ok {
|
|
|
|
|
|
list[i].LastMessage = s.LastMessage
|
|
|
|
|
|
list[i].MessageCount = s.MessageCount
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
2026-06-14 17:35:20 +08:00
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
sort.Slice(list, func(i, j int) bool {
|
|
|
|
|
|
return list[i].UpdatedAt.After(list[j].UpdatedAt)
|
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
|
|
total := len(list)
|
|
|
|
|
|
start := (page - 1) * size
|
|
|
|
|
|
if start >= total {
|
|
|
|
|
|
return []ConversationSummary{}, total, nil
|
|
|
|
|
|
}
|
|
|
|
|
|
end := start + size
|
|
|
|
|
|
if end > total {
|
|
|
|
|
|
end = total
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
return list[start:end], total, nil
|
|
|
|
|
|
}
|
|
|
|
|
|
|
2026-06-13 15:25:23 +08:00
|
|
|
|
// GetHistory 获取最近 N 轮对话历史。
|
|
|
|
|
|
func (m *MemoryManager) GetHistory(_ context.Context, sessionID string, limit int) ([]models.Message, error) {
|
|
|
|
|
|
m.mu.RLock()
|
|
|
|
|
|
defer m.mu.RUnlock()
|
|
|
|
|
|
|
|
|
|
|
|
entry, ok := m.sessions[sessionID]
|
|
|
|
|
|
if !ok || m.isExpired(entry) {
|
|
|
|
|
|
return nil, ErrSessionNotFound
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
if limit <= 0 || limit > len(entry.history) {
|
|
|
|
|
|
limit = len(entry.history)
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// 返回最近 limit 条的副本
|
|
|
|
|
|
result := make([]models.Message, limit)
|
|
|
|
|
|
copy(result, entry.history[len(entry.history)-limit:])
|
|
|
|
|
|
return result, nil
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// AppendMessage 追加一条对话消息,同时刷新 TTL。
|
2026-06-14 17:54:33 +08:00
|
|
|
|
// 若配置了 MessageRepository,消息会异步写入 PostgreSQL(Write-Through)。
|
2026-06-13 15:25:23 +08:00
|
|
|
|
func (m *MemoryManager) AppendMessage(_ context.Context, sessionID string, msg models.Message) error {
|
|
|
|
|
|
m.mu.Lock()
|
|
|
|
|
|
|
|
|
|
|
|
entry, ok := m.sessions[sessionID]
|
|
|
|
|
|
if !ok || m.isExpired(entry) {
|
2026-06-14 17:54:33 +08:00
|
|
|
|
m.mu.Unlock()
|
2026-06-13 15:25:23 +08:00
|
|
|
|
return ErrSessionNotFound
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
entry.history = append(entry.history, msg)
|
|
|
|
|
|
|
2026-06-14 17:35:20 +08:00
|
|
|
|
// 自动更新标题:首条 user 消息时,如果标题为默认值,自动更新为消息前 20 字符
|
|
|
|
|
|
if msg.Role == "user" && entry.session.Title == models.DefaultSessionTitle {
|
|
|
|
|
|
entry.session.Title = generateTitle(msg.Content)
|
|
|
|
|
|
}
|
|
|
|
|
|
|
2026-06-13 15:25:23 +08:00
|
|
|
|
// 超过上限时裁剪,保留最新的 maxHistory 条
|
|
|
|
|
|
if len(entry.history) > m.maxHistory {
|
|
|
|
|
|
entry.history = entry.history[len(entry.history)-m.maxHistory:]
|
|
|
|
|
|
}
|
|
|
|
|
|
|
2026-06-14 17:35:20 +08:00
|
|
|
|
now := time.Now()
|
|
|
|
|
|
entry.lastActive = now
|
|
|
|
|
|
entry.session.UpdatedAt = now
|
2026-06-14 17:54:33 +08:00
|
|
|
|
m.mu.Unlock()
|
|
|
|
|
|
|
|
|
|
|
|
// Write-Through:异步写冷存储,不阻塞调用方
|
|
|
|
|
|
if m.msgRepo != 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)
|
|
|
|
|
|
}
|
|
|
|
|
|
}()
|
|
|
|
|
|
}
|
|
|
|
|
|
|
2026-06-13 15:25:23 +08:00
|
|
|
|
return nil
|
|
|
|
|
|
}
|
|
|
|
|
|
|
2026-06-14 17:35:20 +08:00
|
|
|
|
// generateTitle 从首条消息生成对话标题(取前 20 个字符)。
|
|
|
|
|
|
func generateTitle(firstMessage string) string {
|
|
|
|
|
|
runes := []rune(firstMessage)
|
|
|
|
|
|
if len(runes) > 20 {
|
|
|
|
|
|
return string(runes[:20]) + "…"
|
|
|
|
|
|
}
|
|
|
|
|
|
return firstMessage
|
|
|
|
|
|
}
|
|
|
|
|
|
|
2026-06-14 17:56:33 +08:00
|
|
|
|
// LoadSession 从外部存储加载会话到内存热存储。
|
|
|
|
|
|
// 用于 conversation_id 恢复场景:WS 连接时会话不在内存中,从 PostgreSQL 加载。
|
|
|
|
|
|
// 若会话已在内存中,返回 nil(幂等)。
|
|
|
|
|
|
func (m *MemoryManager) LoadSession(sess *models.Session, messages []models.Message) error {
|
|
|
|
|
|
m.mu.Lock()
|
|
|
|
|
|
defer m.mu.Unlock()
|
|
|
|
|
|
|
|
|
|
|
|
if _, ok := m.sessions[sess.ID]; ok {
|
|
|
|
|
|
return nil // 已在内存中,无需重复加载
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
m.sessions[sess.ID] = &sessionEntry{
|
|
|
|
|
|
session: *sess,
|
|
|
|
|
|
history: messages,
|
|
|
|
|
|
lastActive: time.Now(),
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
logger.Log.Debugw("session loaded from DB", "session", sess.ID, "messages", len(messages))
|
|
|
|
|
|
return nil
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// LoadSessionFromRepo 从 MessageRepository 加载会话消息并注册到内存。
|
|
|
|
|
|
// 适用于已注入 MessageRepository 的场景,调用方只需传入 session 元数据。
|
|
|
|
|
|
func (m *MemoryManager) LoadSessionFromRepo(ctx context.Context, sess *models.Session) error {
|
|
|
|
|
|
if m.msgRepo == nil {
|
|
|
|
|
|
return m.LoadSession(sess, nil)
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// 从冷存储加载全部消息(limit=0 表示全量)
|
|
|
|
|
|
stored, err := m.msgRepo.GetMessages(ctx, sess.ID, 0, 0)
|
|
|
|
|
|
if err != nil {
|
|
|
|
|
|
return err
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
messages := make([]models.Message, len(stored))
|
|
|
|
|
|
for i, s := range stored {
|
|
|
|
|
|
messages[i] = models.Message{Role: s.Role, Content: s.Content}
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
return m.LoadSession(sess, messages)
|
|
|
|
|
|
}
|
|
|
|
|
|
|
2026-06-13 15:25:23 +08:00
|
|
|
|
// SetActiveRequest 标记当前正在处理的请求 ID。
|
|
|
|
|
|
func (m *MemoryManager) SetActiveRequest(_ context.Context, sessionID string, requestID string) error {
|
|
|
|
|
|
m.mu.Lock()
|
|
|
|
|
|
defer m.mu.Unlock()
|
|
|
|
|
|
|
|
|
|
|
|
entry, ok := m.sessions[sessionID]
|
|
|
|
|
|
if !ok || m.isExpired(entry) {
|
|
|
|
|
|
return ErrSessionNotFound
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
entry.activeReqID = requestID
|
|
|
|
|
|
entry.lastActive = time.Now()
|
|
|
|
|
|
return nil
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// GetActiveRequestID 获取当前活跃请求 ID。
|
|
|
|
|
|
func (m *MemoryManager) GetActiveRequestID(_ context.Context, sessionID string) (string, error) {
|
|
|
|
|
|
m.mu.RLock()
|
|
|
|
|
|
defer m.mu.RUnlock()
|
|
|
|
|
|
|
|
|
|
|
|
entry, ok := m.sessions[sessionID]
|
|
|
|
|
|
if !ok || m.isExpired(entry) {
|
|
|
|
|
|
return "", ErrSessionNotFound
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
return entry.activeReqID, nil
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// ClearActiveRequest 清除活跃请求标记。
|
|
|
|
|
|
func (m *MemoryManager) ClearActiveRequest(_ context.Context, sessionID string) error {
|
|
|
|
|
|
m.mu.Lock()
|
|
|
|
|
|
defer m.mu.Unlock()
|
|
|
|
|
|
|
|
|
|
|
|
entry, ok := m.sessions[sessionID]
|
|
|
|
|
|
if !ok || m.isExpired(entry) {
|
|
|
|
|
|
return ErrSessionNotFound
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
entry.activeReqID = ""
|
|
|
|
|
|
entry.lastActive = time.Now()
|
|
|
|
|
|
return nil
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// Touch 刷新 TTL。
|
|
|
|
|
|
func (m *MemoryManager) Touch(_ context.Context, sessionID string) error {
|
|
|
|
|
|
m.mu.Lock()
|
|
|
|
|
|
defer m.mu.Unlock()
|
|
|
|
|
|
|
|
|
|
|
|
entry, ok := m.sessions[sessionID]
|
|
|
|
|
|
if !ok || m.isExpired(entry) {
|
|
|
|
|
|
return ErrSessionNotFound
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
entry.lastActive = time.Now()
|
|
|
|
|
|
return nil
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// Destroy 显式销毁会话。
|
2026-06-14 18:54:18 +08:00
|
|
|
|
func (m *MemoryManager) Destroy(ctx context.Context, sessionID string) error {
|
2026-06-13 15:25:23 +08:00
|
|
|
|
m.mu.Lock()
|
|
|
|
|
|
|
|
|
|
|
|
if _, ok := m.sessions[sessionID]; !ok {
|
2026-06-14 18:54:18 +08:00
|
|
|
|
m.mu.Unlock()
|
2026-06-13 15:25:23 +08:00
|
|
|
|
return ErrSessionNotFound
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
delete(m.sessions, sessionID)
|
2026-06-14 18:54:18 +08:00
|
|
|
|
m.mu.Unlock()
|
|
|
|
|
|
|
|
|
|
|
|
// Write-Through:异步删除 PG
|
|
|
|
|
|
if m.sessRepo != nil {
|
|
|
|
|
|
go func() {
|
|
|
|
|
|
if err := m.sessRepo.Delete(ctx, sessionID); err != nil {
|
|
|
|
|
|
logger.Log.Warnw("delete session from DB failed", "session", sessionID, "error", err)
|
|
|
|
|
|
}
|
|
|
|
|
|
}()
|
|
|
|
|
|
}
|
|
|
|
|
|
|
2026-06-13 15:25:23 +08:00
|
|
|
|
logger.Log.Debugw("session destroyed", "session", sessionID)
|
|
|
|
|
|
return nil
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
// ActiveCount 返回当前活跃会话数。
|
|
|
|
|
|
func (m *MemoryManager) ActiveCount() int {
|
|
|
|
|
|
m.mu.RLock()
|
|
|
|
|
|
defer m.mu.RUnlock()
|
|
|
|
|
|
|
|
|
|
|
|
now := time.Now()
|
|
|
|
|
|
count := 0
|
|
|
|
|
|
for _, entry := range m.sessions {
|
|
|
|
|
|
if now.Sub(entry.lastActive) <= m.ttl {
|
|
|
|
|
|
count++
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
return count
|
|
|
|
|
|
}
|
2026-06-14 18:54:18 +08:00
|
|
|
|
|
|
|
|
|
|
// recordToSession 将 store.SessionRecord 转换为 models.Session。
|
|
|
|
|
|
func (m *MemoryManager) recordToSession(rec *store.SessionRecord) *models.Session {
|
|
|
|
|
|
cfg := models.DefaultConfig()
|
|
|
|
|
|
if len(rec.Config) > 0 {
|
|
|
|
|
|
_ = json.Unmarshal(rec.Config, &cfg)
|
|
|
|
|
|
}
|
|
|
|
|
|
return &models.Session{
|
|
|
|
|
|
ID: rec.ID,
|
|
|
|
|
|
UserID: rec.UserID,
|
|
|
|
|
|
Title: rec.Title,
|
|
|
|
|
|
CreatedAt: rec.CreatedAt,
|
|
|
|
|
|
UpdatedAt: rec.UpdatedAt,
|
|
|
|
|
|
Config: cfg,
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|