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

362 lines
9.2 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"
"sync/atomic"
"time"
"github.com/hhs/camtalk/internal/logger"
"github.com/hhs/camtalk/internal/models"
"github.com/hhs/camtalk/internal/store"
)
// TieredManager 三级存储 SessionManager 实现。
//
// L1内存→ L2Redis→ L3PostgreSQL
//
// 读L1 miss → L2 miss → L3回填到 L1+L2
// 写L1 → L2同步→ L3异步
// 降级Redis 不可用时,回退到 L1+L3 模式
type TieredManager struct {
l1 *MemoryManager // L1: 内存缓存
l2 *RedisManager // L2: Redis可选
sessRepo store.SessionRepository // L3: PostgreSQL 会话持久化(可选)
msgRepo store.MessageRepository // L3: PostgreSQL 消息持久化(可选)
redisOK atomic.Bool // Redis 健康状态
stopCh chan struct{} // 停止信号
}
// TieredOption TieredManager 的函数式选项。
type TieredOption func(*TieredManager)
// WithTieredSessionRepository 注入 L3 会话持久化仓库。
func WithTieredSessionRepository(repo store.SessionRepository) TieredOption {
return func(m *TieredManager) {
m.sessRepo = repo
}
}
// WithTieredMessageRepository 注入 L3 消息持久化仓库。
func WithTieredMessageRepository(repo store.MessageRepository) TieredOption {
return func(m *TieredManager) {
m.msgRepo = repo
}
}
// NewTieredManager 创建三级存储 SessionManager。
// l2 为 nil 时降级为 L1+L3 模式。
func NewTieredManager(
ttl time.Duration,
maxHistory int,
l2 *RedisManager,
opts ...TieredOption,
) *TieredManager {
m := &TieredManager{
l2: l2,
stopCh: make(chan struct{}),
}
for _, opt := range opts {
opt(m)
}
// 初始化 L1内存注入 L3 仓库实现 Write-Through
var l1Opts []Option
if m.sessRepo != nil {
l1Opts = append(l1Opts, WithSessionRepository(m.sessRepo))
}
if m.msgRepo != nil {
l1Opts = append(l1Opts, WithMessageRepository(m.msgRepo))
}
m.l1 = NewMemoryManager(ttl, maxHistory, l1Opts...)
// 初始化 Redis 健康状态
if l2 != nil {
m.redisOK.Store(true)
go m.healthCheck()
}
return m
}
// healthCheck 定期检查 Redis 健康状态。
func (m *TieredManager) healthCheck() {
ticker := time.NewTicker(30 * time.Second)
defer ticker.Stop()
for {
select {
case <-ticker.C:
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
err := m.l2.Ping(ctx)
cancel()
wasOK := m.redisOK.Load()
isOK := err == nil
m.redisOK.Store(isOK)
if wasOK && !isOK {
logger.Log.Warn("Redis connection lost, degrading to L1+L3 mode")
} else if !wasOK && isOK {
logger.Log.Info("Redis connection restored, resuming L1+L2+L3 mode")
}
case <-m.stopCh:
return
}
}
}
// isRedisOK 检查 Redis 是否可用。
func (m *TieredManager) isRedisOK() bool {
return m.l2 != nil && m.redisOK.Load()
}
// Stop 停止 TieredManager清理后台 goroutine
func (m *TieredManager) Stop() {
close(m.stopCh)
m.l1.Stop()
}
// Create 创建新会话。
// 写入L1 → L2同步→ L3异步
func (m *TieredManager) Create(ctx context.Context, userID string, config models.SessionConfig) (string, error) {
// L1: 内存
id, err := m.l1.Create(ctx, userID, config)
if err != nil {
return "", err
}
// L2: Redis同步使用 L1 生成的 ID 保证一致性
if m.isRedisOK() {
if _, err := m.l2.CreateWithID(ctx, id, userID, config); err != nil {
logger.Log.Warnw("Redis Create failed, continuing without L2",
"session", id, "error", err)
}
}
// L3: PostgreSQL异步由 L1 的 Write-Through 处理)
return id, nil
}
// Get 获取会话。
// 读取L1 → L2回填 L1→ L3回填 L1+L2
func (m *TieredManager) Get(ctx context.Context, sessionID string) (*models.Session, error) {
// L1: 内存
sess, err := m.l1.Get(ctx, sessionID)
if err == nil {
return sess, nil
}
if err != ErrSessionNotFound {
return nil, err
}
// L2: Redis
if m.isRedisOK() {
sess, err = m.l2.Get(ctx, sessionID)
if err == nil {
// 回填 L1
history, _ := m.l2.GetHistory(ctx, sessionID, 0)
m.l1.LoadSession(sess, history)
return sess, nil
}
if err != ErrSessionNotFound {
logger.Log.Warnw("Redis Get failed",
"session", sessionID, "error", err)
}
}
// L3: PostgreSQL由 L1 的 Cache-Aside 处理)
// L1.Get 已经实现了从 PostgreSQL 恢复的逻辑
return nil, ErrSessionNotFound
}
// UpdateConfig 更新会话配置。
// 写入L1 → L2同步→ L3异步
func (m *TieredManager) UpdateConfig(ctx context.Context, sessionID string, patch models.SessionConfigPatch) error {
// L1: 内存
if err := m.l1.UpdateConfig(ctx, sessionID, patch); err != nil {
return err
}
// L2: Redis同步
if m.isRedisOK() {
if err := m.l2.UpdateConfig(ctx, sessionID, patch); err != nil {
logger.Log.Warnw("Redis UpdateConfig failed",
"session", sessionID, "error", err)
}
}
// L3: PostgreSQL异步由 L1 的 Write-Through 处理)
return nil
}
// UpdateTitle 更新会话标题。
// 写入L1 → L2同步→ L3异步
func (m *TieredManager) UpdateTitle(ctx context.Context, sessionID string, title string) error {
// L1: 内存
if err := m.l1.UpdateTitle(ctx, sessionID, title); err != nil {
return err
}
// L2: Redis同步
if m.isRedisOK() {
if err := m.l2.UpdateTitle(ctx, sessionID, title); err != nil {
logger.Log.Warnw("Redis UpdateTitle failed",
"session", sessionID, "error", err)
}
}
// L3: PostgreSQL异步由 L1 的 Write-Through 处理)
return nil
}
// ListByUser 查询用户的会话列表。
// 读取L1 + L2 + L3 合并去重
func (m *TieredManager) ListByUser(ctx context.Context, userID string, page, size int) ([]ConversationSummary, int, error) {
// 优先使用 L1已集成 L3 回退逻辑)
return m.l1.ListByUser(ctx, userID, page, size)
}
// GetHistory 获取对话历史。
// 读取L1 → L2回填 L1→ L3回填 L1+L2
func (m *TieredManager) GetHistory(ctx context.Context, sessionID string, limit int) ([]models.Message, error) {
// L1: 内存
msgs, err := m.l1.GetHistory(ctx, sessionID, limit)
if err == nil && len(msgs) > 0 {
return msgs, nil
}
// L2: Redis
if m.isRedisOK() {
msgs, err = m.l2.GetHistory(ctx, sessionID, limit)
if err == nil && len(msgs) > 0 {
// 回填 L1通过 Get 触发)
m.l1.Get(ctx, sessionID)
return msgs, nil
}
}
// L3: PostgreSQL由 L1 的 Cache-Aside 处理)
return nil, ErrSessionNotFound
}
// AppendMessage 追加消息。
// 写入L1 → L2同步→ L3异步
func (m *TieredManager) AppendMessage(ctx context.Context, sessionID string, msg models.Message) error {
// L1: 内存Write-Through 到 L3
if err := m.l1.AppendMessage(ctx, sessionID, msg); err != nil {
return err
}
// L2: Redis同步
if m.isRedisOK() {
if err := m.l2.AppendMessage(ctx, sessionID, msg); err != nil {
logger.Log.Warnw("Redis AppendMessage failed",
"session", sessionID, "error", err)
}
}
return nil
}
// SetActiveRequest 设置当前活跃请求。
func (m *TieredManager) SetActiveRequest(ctx context.Context, sessionID string, requestID string) error {
// L1: 内存
if err := m.l1.SetActiveRequest(ctx, sessionID, requestID); err != nil {
return err
}
// L2: Redis同步
if m.isRedisOK() {
if err := m.l2.SetActiveRequest(ctx, sessionID, requestID); err != nil {
logger.Log.Warnw("Redis SetActiveRequest failed",
"session", sessionID, "error", err)
}
}
return nil
}
// GetActiveRequestID 获取当前活跃请求 ID。
func (m *TieredManager) GetActiveRequestID(ctx context.Context, sessionID string) (string, error) {
// L1: 内存
id, err := m.l1.GetActiveRequestID(ctx, sessionID)
if err == nil && id != "" {
return id, nil
}
// L2: Redis
if m.isRedisOK() {
id, err = m.l2.GetActiveRequestID(ctx, sessionID)
if err == nil && id != "" {
return id, nil
}
}
return "", nil
}
// ClearActiveRequest 清除当前活跃请求。
func (m *TieredManager) ClearActiveRequest(ctx context.Context, sessionID string) error {
// L1: 内存
if err := m.l1.ClearActiveRequest(ctx, sessionID); err != nil {
return err
}
// L2: Redis同步
if m.isRedisOK() {
if err := m.l2.ClearActiveRequest(ctx, sessionID); err != nil {
logger.Log.Warnw("Redis ClearActiveRequest failed",
"session", sessionID, "error", err)
}
}
return nil
}
// Touch 刷新会话活跃时间。
func (m *TieredManager) Touch(ctx context.Context, sessionID string) error {
// L1: 内存
if err := m.l1.Touch(ctx, sessionID); err != nil {
return err
}
// L2: Redis同步
if m.isRedisOK() {
if err := m.l2.Touch(ctx, sessionID); err != nil {
logger.Log.Warnw("Redis Touch failed",
"session", sessionID, "error", err)
}
}
return nil
}
// Destroy 销毁会话。
// 写入L1 → L2 → L3
func (m *TieredManager) Destroy(ctx context.Context, sessionID string) error {
// L1: 内存Write-Through 到 L3
if err := m.l1.Destroy(ctx, sessionID); err != nil {
return err
}
// L2: Redis同步
if m.isRedisOK() {
if err := m.l2.Destroy(ctx, sessionID); err != nil {
logger.Log.Warnw("Redis Destroy failed",
"session", sessionID, "error", err)
}
}
return nil
}
// ActiveCount 返回活跃会话数量。
func (m *TieredManager) ActiveCount() int {
return m.l1.ActiveCount()
}