Files
CamTalk/backend/cmd/server/main.go
hhs 765cb34019 feat: 添加 Redis Docker 容器编排及三级存储架构
- docker-compose.yml 新增 Redis 服务及 camtalk-net 桥接网络
- 实现 TieredManager 三级存储(L1 内存 → L2 Redis → L3 PostgreSQL)
- config.go 新增 Redis/Persistence 配置类型及环境变量绑定
- 修复 CAMTALK_REDIS_ADDR 环境变量未被 Viper 绑定的问题
- .env.example 更新为三级存储配置并标注开发/部署地址差异
2026-06-19 22:53:06 +08:00

269 lines
8.6 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package main
import (
"context"
"errors"
"net/http"
"os/signal"
"strings"
"syscall"
"time"
"github.com/gin-gonic/gin"
"github.com/redis/go-redis/v9"
"github.com/hhs/camtalk/internal/api"
"github.com/hhs/camtalk/internal/auth"
"github.com/hhs/camtalk/internal/ai/llm"
"github.com/hhs/camtalk/internal/ai/stt"
"github.com/hhs/camtalk/internal/ai/tts"
"github.com/hhs/camtalk/internal/config"
"github.com/hhs/camtalk/internal/logger"
"github.com/hhs/camtalk/internal/orchestrator"
"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 注入,如:
// go build -ldflags "-X main.Version=v1.0.0" ./cmd/server
var Version string
var startTime = time.Now()
func main() {
// 加载配置(工作目录用于定位 .env 和 config.yaml
cfg, err := config.Load(".")
if err != nil {
panic("failed to load config: " + err.Error())
}
// 初始化日志
logger.Init(cfg.Log.Level, cfg.Log.Format)
defer logger.Sync()
logger.Log.Infow("config loaded",
"env", cfg.App.Env,
"addr", cfg.Server.Addr(),
)
// 初始化存储层三级存储架构L1 内存 → L2 Redis → L3 PostgreSQL
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
var userRepo store.UserRepository
var msgRepo store.MessageRepository
var sessRepo store.SessionRepository
// L3: PostgreSQL冷数据持久化层
dsn := cfg.Storage.Persistence.DSN
if dsn == "" {
dsn = cfg.Storage.DSN // 兼容旧配置
}
if cfg.Storage.Persistence.Enabled && cfg.Storage.Persistence.Driver == "postgres" {
if dsn == "" {
logger.Log.Fatalw("storage.persistence.dsn is required when persistence is enabled",
"hint", "set CAMTALK_STORAGE_DSN environment variable")
}
pool, err := store.NewPostgresPool(ctx, dsn)
if err != nil {
logger.Log.Fatalw("failed to connect to postgres", "error", err)
}
defer pool.Close()
// 执行数据库迁移
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)
sessRepo = store.NewPgSessionRepository(pool)
logger.Log.Infow("L3 PostgreSQL storage initialized", "driver", cfg.Storage.Persistence.Driver)
} else {
userRepo = store.NewMemUserRepository()
logger.Log.Info("using in-memory user storage")
}
// L2: Redis热数据分布式会话层
var redisMgr *session.RedisManager
if cfg.Storage.Redis.Enabled {
rdb := redis.NewClient(&redis.Options{
Addr: cfg.Redis.Addr,
Password: cfg.Redis.Password,
DB: cfg.Redis.DB,
})
// 验证 Redis 连接
if err := rdb.Ping(ctx).Err(); err != nil {
logger.Log.Fatalw("failed to connect to redis", "error", err)
}
redisMgr = session.NewRedisManager(
rdb,
time.Duration(cfg.Session.TTL)*time.Minute,
cfg.Session.MaxHistory,
)
logger.Log.Infow("L2 Redis storage initialized",
"addr", cfg.Redis.Addr,
"db", cfg.Redis.DB)
}
// 初始化 Session Manager三级存储
var sessionMgr session.Manager
if cfg.Storage.Redis.Enabled {
// L1 + L2 + L3 三级存储
var tieredOpts []session.TieredOption
if sessRepo != nil {
tieredOpts = append(tieredOpts, session.WithTieredSessionRepository(sessRepo))
}
if msgRepo != nil {
tieredOpts = append(tieredOpts, session.WithTieredMessageRepository(msgRepo))
}
tieredMgr := session.NewTieredManager(
time.Duration(cfg.Session.TTL)*time.Minute,
cfg.Session.MaxHistory,
redisMgr,
tieredOpts...,
)
sessionMgr = tieredMgr
defer tieredMgr.Stop()
logger.Log.Info("session manager initialized with L1+L2+L3 tiered storage")
} else {
// L1 + L3 两级存储(无 Redis
var sessionOpts []session.Option
if msgRepo != nil {
sessionOpts = append(sessionOpts, session.WithMessageRepository(msgRepo))
}
if sessRepo != nil {
sessionOpts = append(sessionOpts, session.WithSessionRepository(sessRepo))
}
memMgr := session.NewMemoryManager(
time.Duration(cfg.Session.TTL)*time.Minute,
cfg.Session.MaxHistory,
sessionOpts...,
)
sessionMgr = memMgr
defer memMgr.Stop()
logger.Log.Info("session manager initialized with L1+L3 storage (Redis disabled)")
}
// 初始化 AI 服务
logger.Log.Infow("initializing AI services",
"stt.provider", cfg.AI.STT.Provider,
"stt.model", cfg.AI.STT.Model,
"llm.provider", cfg.AI.LLM.Provider,
"llm.model", cfg.AI.LLM.Model,
"tts.provider", cfg.AI.TTS.Provider,
"tts.model", cfg.AI.TTS.Model,
"tts.voice", cfg.AI.TTS.Voice,
)
var sttService stt.Service
switch strings.ToLower(cfg.AI.STT.Provider) {
case "mimo", "xiaomi":
sttService = stt.NewMiMoService(cfg.AI.STT.APIKey, cfg.AI.STT.Model, cfg.AI.STT.Endpoint, cfg.AI.STT.Timeout, logger.Log)
logger.Log.Infow("STT service initialized", "provider", "mimo", "model", cfg.AI.STT.Model, "endpoint", cfg.AI.STT.Endpoint)
default:
sttService = stt.NewDeepgramService(cfg.AI.STT.APIKey, cfg.AI.STT.Model, cfg.AI.STT.Endpoint, cfg.AI.STT.Timeout, logger.Log)
logger.Log.Infow("STT service initialized", "provider", "deepgram", "model", cfg.AI.STT.Model)
}
llmService := llm.NewOpenAIService(cfg.AI.LLM.APIKey, cfg.AI.LLM.Model, cfg.AI.LLM.Endpoint, cfg.AI.LLM.Timeout, cfg.AI.LLM.HTTPClientTimeout, logger.Log)
logger.Log.Infow("LLM service initialized", "provider", cfg.AI.LLM.Provider, "model", cfg.AI.LLM.Model, "endpoint", cfg.AI.LLM.Endpoint, "timeout", cfg.AI.LLM.Timeout)
var ttsService tts.Service
switch strings.ToLower(cfg.AI.TTS.Provider) {
case "mimo", "xiaomi":
ttsService = tts.NewMiMoService(cfg.AI.TTS.APIKey, cfg.AI.TTS.Model, cfg.AI.TTS.Voice, cfg.AI.TTS.Endpoint, cfg.AI.TTS.Timeout, cfg.AI.TTS.HTTPClientTimeout, logger.Log)
logger.Log.Infow("TTS service initialized", "provider", "mimo", "model", cfg.AI.TTS.Model, "voice", cfg.AI.TTS.Voice, "endpoint", cfg.AI.TTS.Endpoint)
default:
ttsService = tts.NewOpenAIService(cfg.AI.TTS.APIKey, cfg.AI.TTS.Model, cfg.AI.TTS.Voice, cfg.AI.TTS.Endpoint, cfg.AI.TTS.Speed, cfg.AI.TTS.Timeout, cfg.AI.TTS.HTTPClientTimeout, logger.Log)
logger.Log.Infow("TTS service initialized", "provider", "openai", "model", cfg.AI.TTS.Model, "voice", cfg.AI.TTS.Voice, "speed", cfg.AI.TTS.Speed)
}
// 初始化 Orchestrator
orch := orchestrator.New(sttService, llmService, ttsService, sessionMgr, cfg)
// 初始化认证服务
tokenMgr := auth.NewTokenManager(
cfg.Auth.JWTSecret,
time.Duration(cfg.Auth.AccessTTL)*time.Minute,
time.Duration(cfg.Auth.RefreshTTL)*time.Minute,
)
authService := auth.NewAuthService(tokenMgr, userRepo)
// Gin 模式
if cfg.App.Env == "prod" {
gin.SetMode(gin.ReleaseMode)
}
r := gin.New()
r.Use(gin.Recovery())
// REST API
apiGroup := r.Group("/api")
{
apiGroup.GET("/health", healthHandler(sessionMgr, cfg))
}
// Session REST 端点
sessionHandler := api.NewSessionHandler(sessionMgr)
sessionHandler.RegisterRoutes(apiGroup)
// Auth REST 端点
authHandler := api.NewAuthHandler(authService, tokenMgr)
authHandler.RegisterRoutes(apiGroup)
// Conversation REST 端点
convHandler := api.NewConversationHandler(sessionMgr, tokenMgr, msgRepo)
convHandler.RegisterRoutes(apiGroup)
// WebSocket
r.GET("/ws", ws.ServeWS(sessionMgr, orch, cfg, tokenMgr))
// HTTP Server
srv := &http.Server{
Addr: cfg.Server.Addr(),
Handler: r,
ReadTimeout: time.Duration(cfg.Server.ReadTimeout) * time.Second,
WriteTimeout: time.Duration(cfg.Server.WriteTimeout) * time.Second,
}
// Graceful shutdown
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM)
defer stop()
go func() {
logger.Log.Infow("server starting", "addr", srv.Addr)
if err := srv.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) {
logger.Log.Fatalw("listen failed", "error", err)
}
}()
<-ctx.Done()
logger.Log.Info("shutting down...")
shutdownCtx, cancel := context.WithTimeout(context.Background(), time.Duration(cfg.Server.ShutdownTimeout)*time.Second)
defer cancel()
if err := srv.Shutdown(shutdownCtx); err != nil {
logger.Log.Errorw("shutdown error", "error", err)
}
logger.Log.Info("server stopped")
}
// healthHandler 健康检查。
func healthHandler(sessionMgr session.Manager, cfg *config.Config) gin.HandlerFunc {
return func(c *gin.Context) {
version := Version
if version == "" {
version = cfg.App.Version
}
c.JSON(200, gin.H{
"status": "ok",
"version": version,
"uptime_seconds": int(time.Since(startTime).Seconds()),
"active_sessions": sessionMgr.ActiveCount(),
})
}
}