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:
@@ -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)
|
||||||
}
|
}
|
||||||
|
|
||||||
// 初始化 UserRepository(内存模式用于无 DB 场景)
|
userRepo = store.NewPgUserRepository(pool)
|
||||||
var userRepo store.UserRepository
|
msgRepo = store.NewPgMessageRepository(pool)
|
||||||
|
logger.Log.Infow("postgres storage initialized", "driver", cfg.Storage.Driver)
|
||||||
|
} else {
|
||||||
userRepo = store.NewMemUserRepository()
|
userRepo = store.NewMemUserRepository()
|
||||||
|
logger.Log.Info("using in-memory storage")
|
||||||
|
}
|
||||||
|
|
||||||
// 初始化 Session Manager(MVP 默认内存实现)
|
// 初始化 Session Manager
|
||||||
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
|
||||||
|
|||||||
@@ -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 条
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
80
backend/internal/store/migrate.go
Normal file
80
backend/internal/store/migrate.go
Normal 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
|
||||||
|
}
|
||||||
9
backend/migrations/embed.go
Normal file
9
backend/migrations/embed.go
Normal file
@@ -0,0 +1,9 @@
|
|||||||
|
// Package migrations 提供数据库迁移 SQL 文件的嵌入式访问。
|
||||||
|
package migrations
|
||||||
|
|
||||||
|
import "embed"
|
||||||
|
|
||||||
|
// FS 包含所有迁移 SQL 文件。
|
||||||
|
//
|
||||||
|
//go:embed *.sql
|
||||||
|
var FS embed.FS
|
||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user