Merge pull request 'feat: 添加 PostgreSQL 服务并挂载数据库迁移脚本' #101

Merged
huanghaosheng merged 55 commits from develop into main 2026-06-14 19:25:37 +08:00
7 changed files with 377 additions and 35 deletions
Showing only changes of commit 16105f5ac6 - Show all commits

View File

@@ -53,6 +53,7 @@ func main() {
var userRepo store.UserRepository var userRepo store.UserRepository
var msgRepo store.MessageRepository var msgRepo store.MessageRepository
var sessRepo store.SessionRepository
if cfg.Storage.Driver == "postgres" { if cfg.Storage.Driver == "postgres" {
pool, err := store.NewPostgresPool(ctx, cfg.Storage.DSN) pool, err := store.NewPostgresPool(ctx, cfg.Storage.DSN)
@@ -68,6 +69,7 @@ func main() {
userRepo = store.NewPgUserRepository(pool) userRepo = store.NewPgUserRepository(pool)
msgRepo = store.NewPgMessageRepository(pool) msgRepo = store.NewPgMessageRepository(pool)
sessRepo = store.NewPgSessionRepository(pool)
logger.Log.Infow("postgres storage initialized", "driver", cfg.Storage.Driver) logger.Log.Infow("postgres storage initialized", "driver", cfg.Storage.Driver)
} else { } else {
userRepo = store.NewMemUserRepository() userRepo = store.NewMemUserRepository()
@@ -80,6 +82,9 @@ func main() {
if msgRepo != nil { if msgRepo != nil {
sessionOpts = append(sessionOpts, session.WithMessageRepository(msgRepo)) sessionOpts = append(sessionOpts, session.WithMessageRepository(msgRepo))
} }
if sessRepo != nil {
sessionOpts = append(sessionOpts, session.WithSessionRepository(sessRepo))
}
sessionMgr = session.NewMemoryManager( sessionMgr = session.NewMemoryManager(
time.Duration(cfg.Session.TTL)*time.Minute, time.Duration(cfg.Session.TTL)*time.Minute,
cfg.Session.MaxHistory, cfg.Session.MaxHistory,

View File

@@ -2,6 +2,7 @@ package session
import ( import (
"context" "context"
"encoding/json"
"sort" "sort"
"sync" "sync"
"time" "time"
@@ -34,7 +35,8 @@ type MemoryManager struct {
ttl time.Duration ttl time.Duration
maxHistory int maxHistory int
stopCleaner chan struct{} stopCleaner chan struct{}
msgRepo store.MessageRepository // 可选消息持久化Write-Through msgRepo store.MessageRepository // 可选消息持久化Write-Through
sessRepo store.SessionRepository // 可选会话持久化Write-Through
} }
// Option MemoryManager 的函数式选项。 // Option MemoryManager 的函数式选项。
@@ -47,6 +49,13 @@ func WithMessageRepository(repo store.MessageRepository) Option {
} }
} }
// WithSessionRepository 注入会话持久化仓库,启用会话元数据 Write-Through 模式。
func WithSessionRepository(repo store.SessionRepository) Option {
return func(m *MemoryManager) {
m.sessRepo = repo
}
}
// NewMemoryManager 创建内存版 SessionManager。 // NewMemoryManager 创建内存版 SessionManager。
// ttl 为会话过期时间maxHistory 为对话历史上限0 表示使用默认值 20 // ttl 为会话过期时间maxHistory 为对话历史上限0 表示使用默认值 20
// opts 为可选配置,如 WithMessageRepository 启用消息持久化。 // opts 为可选配置,如 WithMessageRepository 启用消息持久化。
@@ -114,9 +123,8 @@ func (m *MemoryManager) isExpired(entry *sessionEntry) bool {
} }
// Create 创建新会话。userID 为空表示匿名会话。 // Create 创建新会话。userID 为空表示匿名会话。
func (m *MemoryManager) Create(_ context.Context, userID string, config models.SessionConfig) (string, error) { func (m *MemoryManager) Create(ctx context.Context, userID string, config models.SessionConfig) (string, error) {
m.mu.Lock() m.mu.Lock()
defer m.mu.Unlock()
id := uuid.New().String() id := uuid.New().String()
now := time.Now() now := time.Now()
@@ -132,73 +140,181 @@ func (m *MemoryManager) Create(_ context.Context, userID string, config models.S
history: make([]models.Message, 0), history: make([]models.Message, 0),
lastActive: now, 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) logger.Log.Debugw("session created", "session", id, "user_id", userID)
return id, nil return id, nil
} }
// Get 获取会话。 // Get 获取会话。内存中不存在时,尝试从 PG 加载(透明恢复)。
func (m *MemoryManager) Get(_ context.Context, sessionID string) (*models.Session, error) { func (m *MemoryManager) Get(ctx context.Context, sessionID string) (*models.Session, error) {
m.mu.RLock() m.mu.RLock()
defer m.mu.RUnlock()
entry, ok := m.sessions[sessionID] entry, ok := m.sessions[sessionID]
if !ok || m.isExpired(entry) { if ok && !m.isExpired(entry) {
return nil, ErrSessionNotFound 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
} }
sess := entry.session // 复制一份返回 return nil, ErrSessionNotFound
return &sess, nil
} }
// UpdateConfig 更新会话配置。 // UpdateConfig 更新会话配置。
func (m *MemoryManager) UpdateConfig(_ context.Context, sessionID string, patch models.SessionConfigPatch) error { func (m *MemoryManager) UpdateConfig(ctx context.Context, sessionID string, patch models.SessionConfigPatch) error {
m.mu.Lock() m.mu.Lock()
defer m.mu.Unlock()
entry, ok := m.sessions[sessionID] entry, ok := m.sessions[sessionID]
if !ok || m.isExpired(entry) { if !ok || m.isExpired(entry) {
m.mu.Unlock()
return ErrSessionNotFound return ErrSessionNotFound
} }
patch.Apply(&entry.session.Config) patch.Apply(&entry.session.Config)
entry.lastActive = time.Now() 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) logger.Log.Debugw("session config updated", "session", sessionID)
return nil return nil
} }
// UpdateTitle 更新会话标题。 // UpdateTitle 更新会话标题。
func (m *MemoryManager) UpdateTitle(_ context.Context, sessionID string, title string) error { func (m *MemoryManager) UpdateTitle(ctx context.Context, sessionID string, title string) error {
m.mu.Lock() m.mu.Lock()
defer m.mu.Unlock()
entry, ok := m.sessions[sessionID] entry, ok := m.sessions[sessionID]
if !ok || m.isExpired(entry) { if !ok || m.isExpired(entry) {
m.mu.Unlock()
return ErrSessionNotFound return ErrSessionNotFound
} }
entry.session.Title = title entry.session.Title = title
entry.session.UpdatedAt = time.Now() entry.session.UpdatedAt = time.Now()
entry.lastActive = 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) logger.Log.Debugw("session title updated", "session", sessionID, "title", title)
return nil return nil
} }
// ListByUser 获取用户的对话列表(分页,按 UpdatedAt 降序)。 // ListByUser 获取用户的对话列表(分页,按 UpdatedAt 降序)。
// 若配置了 SessionRepository从 PG 查询(包含内存中已过期的会话)。
// 若配置了 MessageRepository消息统计从 PostgreSQL 聚合查询(更准确)。 // 若配置了 MessageRepository消息统计从 PostgreSQL 聚合查询(更准确)。
func (m *MemoryManager) ListByUser(ctx context.Context, userID string, page, size int) ([]ConversationSummary, int, error) { 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() m.mu.RLock()
// 收集该用户的所有 session
var list []ConversationSummary var list []ConversationSummary
var sessionIDs []string var sessionIDs []string
for _, entry := range m.sessions { for _, entry := range m.sessions {
if entry.session.UserID != userID { if entry.session.UserID != userID || m.isExpired(entry) {
continue
}
if m.isExpired(entry) {
continue continue
} }
summary := ConversationSummary{ summary := ConversationSummary{
@@ -206,7 +322,6 @@ func (m *MemoryManager) ListByUser(ctx context.Context, userID string, page, siz
Title: entry.session.Title, Title: entry.session.Title,
UpdatedAt: entry.lastActive, UpdatedAt: entry.lastActive,
} }
// 先用内存值填充,后续可能被 PG 统计覆盖
summary.MessageCount = len(entry.history) summary.MessageCount = len(entry.history)
if len(entry.history) > 0 { if len(entry.history) > 0 {
summary.LastMessage = entry.history[len(entry.history)-1].Content summary.LastMessage = entry.history[len(entry.history)-1].Content
@@ -216,7 +331,7 @@ func (m *MemoryManager) ListByUser(ctx context.Context, userID string, page, siz
} }
m.mu.RUnlock() m.mu.RUnlock()
// 若配置了 msgRepo从 PostgreSQL 获取更准确的消息统计 // 从 PG 获取更准确的消息统计
if m.msgRepo != nil && len(sessionIDs) > 0 { if m.msgRepo != nil && len(sessionIDs) > 0 {
if stats, err := m.msgRepo.GetSessionMessageStats(ctx, sessionIDs); err == nil { if stats, err := m.msgRepo.GetSessionMessageStats(ctx, sessionIDs); err == nil {
for i := range list { for i := range list {
@@ -225,25 +340,14 @@ func (m *MemoryManager) ListByUser(ctx context.Context, userID string, page, siz
list[i].MessageCount = s.MessageCount list[i].MessageCount = s.MessageCount
} }
} }
} else {
logger.Log.Warnw("get session message stats failed, falling back to in-memory", "error", err)
} }
} }
// 按 UpdatedAt 降序排序
sort.Slice(list, func(i, j int) bool { sort.Slice(list, func(i, j int) bool {
return list[i].UpdatedAt.After(list[j].UpdatedAt) return list[i].UpdatedAt.After(list[j].UpdatedAt)
}) })
total := len(list) total := len(list)
// 分页
if page <= 0 {
page = 1
}
if size <= 0 {
size = 20
}
start := (page - 1) * size start := (page - 1) * size
if start >= total { if start >= total {
return []ConversationSummary{}, total, nil return []ConversationSummary{}, total, nil
@@ -425,15 +529,26 @@ func (m *MemoryManager) Touch(_ context.Context, sessionID string) error {
} }
// Destroy 显式销毁会话。 // Destroy 显式销毁会话。
func (m *MemoryManager) Destroy(_ context.Context, sessionID string) error { func (m *MemoryManager) Destroy(ctx context.Context, sessionID string) error {
m.mu.Lock() m.mu.Lock()
defer m.mu.Unlock()
if _, ok := m.sessions[sessionID]; !ok { if _, ok := m.sessions[sessionID]; !ok {
m.mu.Unlock()
return ErrSessionNotFound return ErrSessionNotFound
} }
delete(m.sessions, sessionID) 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) logger.Log.Debugw("session destroyed", "session", sessionID)
return nil return nil
} }
@@ -452,3 +567,19 @@ func (m *MemoryManager) ActiveCount() int {
} }
return 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,
}
}

View File

@@ -0,0 +1,47 @@
package store
import (
"context"
"errors"
"time"
)
var (
// ErrSessionNotFound 会话不存在。
ErrSessionNotFound = errors.New("session not found")
)
// SessionRepository 会话持久化接口。
type SessionRepository interface {
// Save 创建或更新会话UPSERT
Save(ctx context.Context, s SessionRecord) error
// FindByID 根据 ID 查询会话。
FindByID(ctx context.Context, id string) (*SessionRecord, error)
// FindByUser 查询用户的会话列表(分页,按 updated_at 降序)。
// 返回 (列表, 总数, error)。
FindByUser(ctx context.Context, userID string, page, size int) ([]SessionRecord, int, error)
// UpdateTitle 更新会话标题。
UpdateTitle(ctx context.Context, id string, title string) error
// UpdateConfig 更新会话配置。
UpdateConfig(ctx context.Context, id string, configJSON []byte) error
// Touch 刷新 updated_at。
Touch(ctx context.Context, id string) error
// Delete 删除会话。
Delete(ctx context.Context, id string) error
}
// SessionRecord 持久化会话模型store 层)。
type SessionRecord struct {
ID string
UserID string
Title string
Config []byte // JSON 编码的 SessionConfig
CreatedAt time.Time
UpdatedAt time.Time
}

View File

@@ -0,0 +1,146 @@
package store
import (
"context"
"errors"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
)
// PgSessionRepository 基于 PostgreSQL 的 SessionRepository 实现。
type PgSessionRepository struct {
pool *pgxpool.Pool
}
// NewPgSessionRepository 创建 PgSessionRepository。
func NewPgSessionRepository(pool *pgxpool.Pool) *PgSessionRepository {
return &PgSessionRepository{pool: pool}
}
func (r *PgSessionRepository) Save(ctx context.Context, s SessionRecord) error {
_, err := r.pool.Exec(ctx,
`INSERT INTO sessions (id, user_id, title, config, created_at, updated_at)
VALUES ($1, $2, $3, $4, $5, $6)
ON CONFLICT (id) DO UPDATE SET
title = EXCLUDED.title,
config = EXCLUDED.config,
updated_at = EXCLUDED.updated_at`,
s.ID, s.UserID, s.Title, s.Config, s.CreatedAt, s.UpdatedAt,
)
return err
}
func (r *PgSessionRepository) FindByID(ctx context.Context, id string) (*SessionRecord, error) {
var s SessionRecord
err := r.pool.QueryRow(ctx,
`SELECT id, user_id, title, config, created_at, updated_at
FROM sessions WHERE id = $1`, id,
).Scan(&s.ID, &s.UserID, &s.Title, &s.Config, &s.CreatedAt, &s.UpdatedAt)
if errors.Is(err, pgx.ErrNoRows) {
return nil, ErrSessionNotFound
}
if err != nil {
return nil, err
}
return &s, nil
}
func (r *PgSessionRepository) FindByUser(ctx context.Context, userID string, page, size int) ([]SessionRecord, int, error) {
if page <= 0 {
page = 1
}
if size <= 0 {
size = 20
}
offset := (page - 1) * size
// 查询总数
var total int
if err := r.pool.QueryRow(ctx,
`SELECT COUNT(*) FROM sessions WHERE user_id = $1`, userID,
).Scan(&total); err != nil {
return nil, 0, err
}
// 查询列表
rows, err := r.pool.Query(ctx,
`SELECT id, user_id, title, config, created_at, updated_at
FROM sessions
WHERE user_id = $1
ORDER BY updated_at DESC
LIMIT $2 OFFSET $3`,
userID, size, offset,
)
if err != nil {
return nil, 0, err
}
defer rows.Close()
var list []SessionRecord
for rows.Next() {
var s SessionRecord
if err := rows.Scan(&s.ID, &s.UserID, &s.Title, &s.Config, &s.CreatedAt, &s.UpdatedAt); err != nil {
return nil, 0, err
}
list = append(list, s)
}
if err := rows.Err(); err != nil {
return nil, 0, err
}
return list, total, nil
}
func (r *PgSessionRepository) UpdateTitle(ctx context.Context, id string, title string) error {
tag, err := r.pool.Exec(ctx,
`UPDATE sessions SET title = $2, updated_at = NOW() WHERE id = $1`,
id, title,
)
if err != nil {
return err
}
if tag.RowsAffected() == 0 {
return ErrSessionNotFound
}
return nil
}
func (r *PgSessionRepository) UpdateConfig(ctx context.Context, id string, configJSON []byte) error {
tag, err := r.pool.Exec(ctx,
`UPDATE sessions SET config = $2, updated_at = NOW() WHERE id = $1`,
id, configJSON,
)
if err != nil {
return err
}
if tag.RowsAffected() == 0 {
return ErrSessionNotFound
}
return nil
}
func (r *PgSessionRepository) Touch(ctx context.Context, id string) error {
tag, err := r.pool.Exec(ctx,
`UPDATE sessions SET updated_at = NOW() WHERE id = $1`, id,
)
if err != nil {
return err
}
if tag.RowsAffected() == 0 {
return ErrSessionNotFound
}
return nil
}
func (r *PgSessionRepository) Delete(ctx context.Context, id string) error {
tag, err := r.pool.Exec(ctx,
`DELETE FROM sessions WHERE id = $1`, id,
)
if err != nil {
return err
}
if tag.RowsAffected() == 0 {
return ErrSessionNotFound
}
return nil
}

View File

@@ -0,0 +1 @@
DROP TABLE IF EXISTS sessions;

View File

@@ -0,0 +1,11 @@
CREATE TABLE IF NOT EXISTS sessions (
id UUID PRIMARY KEY,
user_id UUID NOT NULL,
title VARCHAR(256) NOT NULL DEFAULT '新对话',
config JSONB NOT NULL DEFAULT '{}',
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
);
CREATE INDEX IF NOT EXISTS idx_sessions_user_id ON sessions (user_id);
CREATE INDEX IF NOT EXISTS idx_sessions_user_updated ON sessions (user_id, updated_at DESC);

View File

@@ -38,6 +38,7 @@ services:
POSTGRES_DB: camtalk POSTGRES_DB: camtalk
volumes: volumes:
- pgdata:/var/lib/postgresql/data - pgdata:/var/lib/postgresql/data
- ./backend/migrations:/docker-entrypoint-initdb.d
networks: networks:
- camtalk-net - camtalk-net
restart: unless-stopped restart: unless-stopped