Merge pull request 'feat: 对接 PostgreSQL 存储层' (#98) from build/backend into develop

Reviewed-on: http://8.161.227.145:3000/XEngineers/CamTalk/pulls/98
This commit was merged in pull request #98.
This commit is contained in:
2026-06-14 18:46:05 +08:00
6 changed files with 164 additions and 18 deletions

View File

@@ -22,6 +22,7 @@ import (
"github.com/hhs/camtalk/internal/session" "github.com/hhs/camtalk/internal/session"
"github.com/hhs/camtalk/internal/store" "github.com/hhs/camtalk/internal/store"
"github.com/hhs/camtalk/internal/ws" "github.com/hhs/camtalk/internal/ws"
migrations "github.com/hhs/camtalk/migrations"
) )
// Version 通过构建时 -ldflags 注入,如: // Version 通过构建时 -ldflags 注入,如:
@@ -50,27 +51,39 @@ func main() {
ctx, cancel := context.WithCancel(context.Background()) ctx, cancel := context.WithCancel(context.Background())
defer cancel() defer cancel()
var userRepo store.UserRepository
var msgRepo store.MessageRepository
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)
if err != nil { if err != nil {
logger.Log.Fatalw("failed to connect to postgres", "error", err) logger.Log.Fatalw("failed to connect to postgres", "error", err)
} }
defer pool.Close() defer pool.Close()
logger.Log.Infow("postgres connected", "driver", cfg.Storage.Driver)
// TODO: Phase 2 - 初始化 UserRepository 和 MessageRepository // 执行数据库迁移
_ = pool if err := store.RunMigrations(ctx, pool, migrations.FS); err != nil {
logger.Log.Fatalw("failed to run migrations", "error", err)
}
userRepo = store.NewPgUserRepository(pool)
msgRepo = store.NewPgMessageRepository(pool)
logger.Log.Infow("postgres storage initialized", "driver", cfg.Storage.Driver)
} else {
userRepo = store.NewMemUserRepository()
logger.Log.Info("using in-memory storage")
} }
// 初始化 UserRepository内存模式用于无 DB 场景) // 初始化 Session Manager
var userRepo store.UserRepository
userRepo = store.NewMemUserRepository()
// 初始化 Session ManagerMVP 默认内存实现)
var sessionMgr session.Manager var sessionMgr session.Manager
// TODO: 当 Redis 配置非空时切换为 RedisManager var sessionOpts []session.Option
if msgRepo != nil {
sessionOpts = append(sessionOpts, session.WithMessageRepository(msgRepo))
}
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,
sessionOpts...,
) )
defer sessionMgr.(*session.MemoryManager).Stop() defer sessionMgr.(*session.MemoryManager).Stop()
@@ -141,7 +154,7 @@ func main() {
authHandler.RegisterRoutes(apiGroup) authHandler.RegisterRoutes(apiGroup)
// Conversation REST 端点 // Conversation REST 端点
convHandler := api.NewConversationHandler(sessionMgr, tokenMgr) convHandler := api.NewConversationHandler(sessionMgr, tokenMgr, msgRepo)
convHandler.RegisterRoutes(apiGroup) convHandler.RegisterRoutes(apiGroup)
// WebSocket // WebSocket

View File

@@ -12,19 +12,23 @@ import (
apperr "github.com/hhs/camtalk/internal/errors" apperr "github.com/hhs/camtalk/internal/errors"
"github.com/hhs/camtalk/internal/models" "github.com/hhs/camtalk/internal/models"
"github.com/hhs/camtalk/internal/session" "github.com/hhs/camtalk/internal/session"
"github.com/hhs/camtalk/internal/store"
) )
// ConversationHandler 提供对话相关的 REST 端点。 // ConversationHandler 提供对话相关的 REST 端点。
type ConversationHandler struct { type ConversationHandler struct {
sessionMgr session.Manager sessionMgr session.Manager
tokenMgr *auth.TokenManager tokenMgr *auth.TokenManager
msgRepo store.MessageRepository // 可选,为 nil 时 fallback 到内存查询
} }
// NewConversationHandler 创建 ConversationHandler。 // NewConversationHandler 创建 ConversationHandler。
func NewConversationHandler(sessionMgr session.Manager, tokenMgr *auth.TokenManager) *ConversationHandler { // msgRepo 可选,为 nil 时消息查询走内存。
func NewConversationHandler(sessionMgr session.Manager, tokenMgr *auth.TokenManager, msgRepo store.MessageRepository) *ConversationHandler {
return &ConversationHandler{ return &ConversationHandler{
sessionMgr: sessionMgr, sessionMgr: sessionMgr,
tokenMgr: tokenMgr, tokenMgr: tokenMgr,
msgRepo: msgRepo,
} }
} }
@@ -215,7 +219,7 @@ func (h *ConversationHandler) Delete(c *gin.Context) {
// //
// 查询参数: // 查询参数:
// - limit: 返回消息数量上限,默认 50 // - limit: 返回消息数量上限,默认 50
// - before: 消息偏移量(用于分页),返回此偏移量之前的消息 // - before: 消息 ID 游标(用于分页),返回此 ID 之前的消息
func (h *ConversationHandler) GetMessages(c *gin.Context) { func (h *ConversationHandler) GetMessages(c *gin.Context) {
sessionID := c.Param("id") sessionID := c.Param("id")
@@ -229,9 +233,27 @@ func (h *ConversationHandler) GetMessages(c *gin.Context) {
limit = 50 limit = 50
} }
before, _ := strconv.Atoi(c.DefaultQuery("before", "0")) beforeID, _ := strconv.ParseInt(c.DefaultQuery("before", "0"), 10, 64)
// 获取全量历史(内存实现中 history 是全量存储的 // 优先从 PostgreSQL 查询(支持持久化后的全量历史
if h.msgRepo != nil {
messages, err := h.msgRepo.GetMessages(c.Request.Context(), sessionID, limit, beforeID)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{
"code": apperr.CodeInternalError,
"message": "failed to get messages",
})
return
}
count, _ := h.msgRepo.GetMessageCount(c.Request.Context(), sessionID)
c.JSON(http.StatusOK, gin.H{
"messages": messages,
"total": count,
})
return
}
// fallback从内存查询
allMessages, err := h.sessionMgr.GetHistory(c.Request.Context(), sessionID, 0) allMessages, err := h.sessionMgr.GetHistory(c.Request.Context(), sessionID, 0)
if err != nil { if err != nil {
if errors.Is(err, session.ErrSessionNotFound) { if errors.Is(err, session.ErrSessionNotFound) {
@@ -250,9 +272,9 @@ func (h *ConversationHandler) GetMessages(c *gin.Context) {
total := len(allMessages) total := len(allMessages)
// before > 0 表示取 before 之前的消息(不含 before 位置 // beforeID > 0 表示偏移量(兼容旧接口语义
if before > 0 && before <= total { if beforeID > 0 && int(beforeID) <= total {
allMessages = allMessages[:before] allMessages = allMessages[:beforeID]
} }
// 取最后 limit 条 // 取最后 limit 条

View File

@@ -93,7 +93,7 @@ func newConvTestRouter(mgr session.Manager) (*gin.Engine, *auth.TokenManager) {
gin.SetMode(gin.TestMode) gin.SetMode(gin.TestMode)
r := gin.New() r := gin.New()
tm := auth.NewTokenManager("test-secret", 15*time.Minute, 7*24*time.Hour) tm := auth.NewTokenManager("test-secret", 15*time.Minute, 7*24*time.Hour)
h := api.NewConversationHandler(mgr, tm) h := api.NewConversationHandler(mgr, tm, nil)
h.RegisterRoutes(r.Group("/api")) h.RegisterRoutes(r.Group("/api"))
return r, tm return r, tm
} }

View File

@@ -0,0 +1,80 @@
package store
import (
"context"
"fmt"
"io/fs"
"sort"
"strings"
"github.com/jackc/pgx/v5/pgxpool"
"github.com/hhs/camtalk/internal/logger"
)
// RunMigrations 从给定的 fs.FS 中读取 *.up.sql 文件并按版本号顺序执行。
// 已执行过的版本会跳过(通过 schema_migrations 表记录)。
func RunMigrations(ctx context.Context, pool *pgxpool.Pool, fsys fs.FS) error {
// 确保 schema_migrations 表存在
if _, err := pool.Exec(ctx, `CREATE TABLE IF NOT EXISTS schema_migrations (
version INTEGER PRIMARY KEY,
applied_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
)`); err != nil {
return fmt.Errorf("create schema_migrations table: %w", err)
}
// 收集所有 *.up.sql 文件
entries, err := fs.ReadDir(fsys, ".")
if err != nil {
return fmt.Errorf("read migrations dir: %w", err)
}
var files []string
for _, e := range entries {
if !e.IsDir() && strings.HasSuffix(e.Name(), ".up.sql") {
files = append(files, e.Name())
}
}
sort.Strings(files)
for _, name := range files {
// 从文件名提取版本号,如 "001_users.up.sql" → 1
var version int
if _, err := fmt.Sscanf(name, "%d_", &version); err != nil {
return fmt.Errorf("parse version from %s: %w", name, err)
}
// 检查是否已执行
var exists bool
if err := pool.QueryRow(ctx,
`SELECT EXISTS(SELECT 1 FROM schema_migrations WHERE version = $1)`, version,
).Scan(&exists); err != nil {
return fmt.Errorf("check migration version %d: %w", version, err)
}
if exists {
logger.Log.Debugw("migration already applied", "version", version, "file", name)
continue
}
// 读取并执行
content, err := fs.ReadFile(fsys, name)
if err != nil {
return fmt.Errorf("read migration %s: %w", name, err)
}
if _, err := pool.Exec(ctx, string(content)); err != nil {
return fmt.Errorf("execute migration %s: %w", name, err)
}
// 记录已执行
if _, err := pool.Exec(ctx,
`INSERT INTO schema_migrations (version) VALUES ($1)`, version,
); err != nil {
return fmt.Errorf("record migration %d: %w", version, err)
}
logger.Log.Infow("migration applied", "version", version, "file", name)
}
return nil
}

View File

@@ -0,0 +1,9 @@
// Package migrations 提供数据库迁移 SQL 文件的嵌入式访问。
package migrations
import "embed"
// FS 包含所有迁移 SQL 文件。
//
//go:embed *.sql
var FS embed.FS

View File

@@ -19,10 +19,32 @@ services:
container_name: camtalk-backend container_name: camtalk-backend
environment: environment:
- APP_ENV=production - APP_ENV=production
- CAMTALK_STORAGE_DRIVER=postgres
- CAMTALK_STORAGE_DSN=postgres://camtalk:camtalk123@postgres:5432/camtalk?sslmode=disable
- CAMTALK_AUTH_JWT_SECRET=78uWBBAF8XEQEotKDlrnlnd4y8i4WN3E4zXmNmC8BYQ=
depends_on:
- postgres
networks: networks:
- camtalk-net - camtalk-net
restart: unless-stopped restart: unless-stopped
postgres:
# 阿里云镜像,避免从 Docker Hub 拉取超时
image: registry.cn-hangzhou.aliyuncs.com/library/postgres:15-alpine
container_name: camtalk-postgres
environment:
POSTGRES_USER: camtalk
POSTGRES_PASSWORD: camtalk123
POSTGRES_DB: camtalk
volumes:
- pgdata:/var/lib/postgresql/data
networks:
- camtalk-net
restart: unless-stopped
volumes:
pgdata:
networks: networks:
camtalk-net: camtalk-net:
driver: bridge driver: bridge