diff --git a/backend/cmd/server/main.go b/backend/cmd/server/main.go index 2a6275b..8accba9 100644 --- a/backend/cmd/server/main.go +++ b/backend/cmd/server/main.go @@ -22,6 +22,7 @@ import ( "github.com/hhs/camtalk/internal/session" "github.com/hhs/camtalk/internal/store" "github.com/hhs/camtalk/internal/ws" + migrations "github.com/hhs/camtalk/migrations" ) // Version 通过构建时 -ldflags 注入,如: @@ -50,27 +51,39 @@ func main() { ctx, cancel := context.WithCancel(context.Background()) defer cancel() + var userRepo store.UserRepository + var msgRepo store.MessageRepository + if cfg.Storage.Driver == "postgres" { pool, err := store.NewPostgresPool(ctx, cfg.Storage.DSN) if err != nil { logger.Log.Fatalw("failed to connect to postgres", "error", err) } 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 场景) - var userRepo store.UserRepository - userRepo = store.NewMemUserRepository() - - // 初始化 Session Manager(MVP 默认内存实现) + // 初始化 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( time.Duration(cfg.Session.TTL)*time.Minute, cfg.Session.MaxHistory, + sessionOpts..., ) defer sessionMgr.(*session.MemoryManager).Stop() @@ -141,7 +154,7 @@ func main() { authHandler.RegisterRoutes(apiGroup) // Conversation REST 端点 - convHandler := api.NewConversationHandler(sessionMgr, tokenMgr) + convHandler := api.NewConversationHandler(sessionMgr, tokenMgr, msgRepo) convHandler.RegisterRoutes(apiGroup) // WebSocket diff --git a/backend/internal/api/conversation.go b/backend/internal/api/conversation.go index f95e2c1..728cc9a 100644 --- a/backend/internal/api/conversation.go +++ b/backend/internal/api/conversation.go @@ -12,19 +12,23 @@ import ( apperr "github.com/hhs/camtalk/internal/errors" "github.com/hhs/camtalk/internal/models" "github.com/hhs/camtalk/internal/session" + "github.com/hhs/camtalk/internal/store" ) // ConversationHandler 提供对话相关的 REST 端点。 type ConversationHandler struct { sessionMgr session.Manager tokenMgr *auth.TokenManager + msgRepo store.MessageRepository // 可选,为 nil 时 fallback 到内存查询 } // 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{ sessionMgr: sessionMgr, tokenMgr: tokenMgr, + msgRepo: msgRepo, } } @@ -215,7 +219,7 @@ func (h *ConversationHandler) Delete(c *gin.Context) { // // 查询参数: // - limit: 返回消息数量上限,默认 50 -// - before: 消息偏移量(用于分页),返回此偏移量之前的消息 +// - before: 消息 ID 游标(用于分页),返回此 ID 之前的消息 func (h *ConversationHandler) GetMessages(c *gin.Context) { sessionID := c.Param("id") @@ -229,9 +233,27 @@ func (h *ConversationHandler) GetMessages(c *gin.Context) { 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) if err != nil { if errors.Is(err, session.ErrSessionNotFound) { @@ -250,9 +272,9 @@ func (h *ConversationHandler) GetMessages(c *gin.Context) { total := len(allMessages) - // before > 0 表示取 before 之前的消息(不含 before 位置) - if before > 0 && before <= total { - allMessages = allMessages[:before] + // beforeID > 0 时表示偏移量(兼容旧接口语义) + if beforeID > 0 && int(beforeID) <= total { + allMessages = allMessages[:beforeID] } // 取最后 limit 条 diff --git a/backend/internal/api/conversation_test.go b/backend/internal/api/conversation_test.go index 6b873fc..468d172 100644 --- a/backend/internal/api/conversation_test.go +++ b/backend/internal/api/conversation_test.go @@ -93,7 +93,7 @@ func newConvTestRouter(mgr session.Manager) (*gin.Engine, *auth.TokenManager) { gin.SetMode(gin.TestMode) r := gin.New() 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")) return r, tm } diff --git a/backend/internal/store/migrate.go b/backend/internal/store/migrate.go new file mode 100644 index 0000000..b3ea4ba --- /dev/null +++ b/backend/internal/store/migrate.go @@ -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 +} diff --git a/backend/migrations/embed.go b/backend/migrations/embed.go new file mode 100644 index 0000000..9b6970b --- /dev/null +++ b/backend/migrations/embed.go @@ -0,0 +1,9 @@ +// Package migrations 提供数据库迁移 SQL 文件的嵌入式访问。 +package migrations + +import "embed" + +// FS 包含所有迁移 SQL 文件。 +// +//go:embed *.sql +var FS embed.FS diff --git a/docker-compose.yml b/docker-compose.yml index f5f15de..053f375 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -19,10 +19,32 @@ services: container_name: camtalk-backend environment: - 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: - camtalk-net 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: camtalk-net: driver: bridge