Files
CamTalk/backend/internal/session/memory.go

586 lines
15 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 session
import (
"context"
"encoding/json"
"sort"
"sync"
"time"
"github.com/google/uuid"
"github.com/hhs/camtalk/internal/logger"
"github.com/hhs/camtalk/internal/models"
"github.com/hhs/camtalk/internal/store"
)
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{}
msgRepo store.MessageRepository // 可选消息持久化Write-Through
sessRepo store.SessionRepository // 可选会话持久化Write-Through
}
// Option MemoryManager 的函数式选项。
type Option func(*MemoryManager)
// WithMessageRepository 注入消息持久化仓库,启用 Write-Through 模式。
func WithMessageRepository(repo store.MessageRepository) Option {
return func(m *MemoryManager) {
m.msgRepo = repo
}
}
// WithSessionRepository 注入会话持久化仓库,启用会话元数据 Write-Through 模式。
func WithSessionRepository(repo store.SessionRepository) Option {
return func(m *MemoryManager) {
m.sessRepo = repo
}
}
// NewMemoryManager 创建内存版 SessionManager。
// ttl 为会话过期时间maxHistory 为对话历史上限0 表示使用默认值 20
// opts 为可选配置,如 WithMessageRepository 启用消息持久化。
func NewMemoryManager(ttl time.Duration, maxHistory int, opts ...Option) *MemoryManager {
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{}),
}
for _, opt := range opts {
opt(m)
}
// 启动后台清理 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
}
// Create 创建新会话。userID 为空表示匿名会话。
func (m *MemoryManager) Create(ctx context.Context, userID string, config models.SessionConfig) (string, error) {
m.mu.Lock()
id := uuid.New().String()
now := time.Now()
m.sessions[id] = &sessionEntry{
session: models.Session{
ID: id,
UserID: userID,
Title: models.DefaultSessionTitle,
CreatedAt: now,
UpdatedAt: now,
Config: config,
},
history: make([]models.Message, 0),
lastActive: now,
}
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)
}
}()
}
logger.Log.Debugw("session created", "session", id, "user_id", userID)
return id, nil
}
// Get 获取会话。内存中不存在时,尝试从 PG 加载(透明恢复)。
func (m *MemoryManager) Get(ctx context.Context, sessionID string) (*models.Session, error) {
m.mu.RLock()
entry, ok := m.sessions[sessionID]
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
}
return nil, ErrSessionNotFound
}
// UpdateConfig 更新会话配置。
func (m *MemoryManager) UpdateConfig(ctx context.Context, sessionID string, patch models.SessionConfigPatch) error {
m.mu.Lock()
entry, ok := m.sessions[sessionID]
if !ok || m.isExpired(entry) {
m.mu.Unlock()
return ErrSessionNotFound
}
patch.Apply(&entry.session.Config)
entry.lastActive = time.Now()
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)
}
}()
}
logger.Log.Debugw("session config updated", "session", sessionID)
return nil
}
// UpdateTitle 更新会话标题。
func (m *MemoryManager) UpdateTitle(ctx context.Context, sessionID string, title string) error {
m.mu.Lock()
entry, ok := m.sessions[sessionID]
if !ok || m.isExpired(entry) {
m.mu.Unlock()
return ErrSessionNotFound
}
entry.session.Title = title
entry.session.UpdatedAt = time.Now()
entry.lastActive = time.Now()
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)
}
}()
}
logger.Log.Debugw("session title updated", "session", sessionID, "title", title)
return nil
}
// ListByUser 获取用户的对话列表(分页,按 UpdatedAt 降序)。
// 若配置了 SessionRepository从 PG 查询(包含内存中已过期的会话)。
// 若配置了 MessageRepository消息统计从 PostgreSQL 聚合查询(更准确)。
func (m *MemoryManager) ListByUser(ctx context.Context, userID string, page, size int) ([]ConversationSummary, int, error) {
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) {
m.mu.RLock()
var list []ConversationSummary
var sessionIDs []string
for _, entry := range m.sessions {
if entry.session.UserID != userID || m.isExpired(entry) {
continue
}
summary := ConversationSummary{
ID: entry.session.ID,
Title: entry.session.Title,
UpdatedAt: entry.lastActive,
}
summary.MessageCount = len(entry.history)
if len(entry.history) > 0 {
summary.LastMessage = entry.history[len(entry.history)-1].Content
}
list = append(list, summary)
sessionIDs = append(sessionIDs, entry.session.ID)
}
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
}
}
}
}
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
}
// 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。
// 若配置了 MessageRepository消息会异步写入 PostgreSQLWrite-Through
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) {
m.mu.Unlock()
return ErrSessionNotFound
}
entry.history = append(entry.history, msg)
// 自动更新标题:首条 user 消息时,如果标题为默认值,自动更新为消息前 20 字符
if msg.Role == "user" && entry.session.Title == models.DefaultSessionTitle {
entry.session.Title = generateTitle(msg.Content)
}
// 超过上限时裁剪,保留最新的 maxHistory 条
if len(entry.history) > m.maxHistory {
entry.history = entry.history[len(entry.history)-m.maxHistory:]
}
now := time.Now()
entry.lastActive = now
entry.session.UpdatedAt = now
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)
}
}()
}
return nil
}
// generateTitle 从首条消息生成对话标题(取前 20 个字符)。
func generateTitle(firstMessage string) string {
runes := []rune(firstMessage)
if len(runes) > 20 {
return string(runes[:20]) + "…"
}
return firstMessage
}
// 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)
}
// 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 显式销毁会话。
func (m *MemoryManager) Destroy(ctx context.Context, sessionID string) error {
m.mu.Lock()
if _, ok := m.sessions[sessionID]; !ok {
m.mu.Unlock()
return ErrSessionNotFound
}
delete(m.sessions, sessionID)
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)
}
}()
}
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
}
// 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,
}
}