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

Merged
huanghaosheng merged 55 commits from develop into main 2026-06-14 19:25:37 +08:00
6 changed files with 164 additions and 18 deletions
Showing only changes of commit 542720695d - Show all commits

View File

@@ -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)
}
// 初始化 UserRepository内存模式用于无 DB 场景)
var userRepo store.UserRepository
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")
}
// 初始化 Session ManagerMVP 默认内存实现)
// 初始化 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

View File

@@ -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 条

View File

@@ -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
}

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
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