Files
CamTalk/backend/cmd/server/main.go
hhs 7adf81c6e5 feat: 集成限流器到服务
- main.go 初始化限流器(根据 Redis 可用性选择内存/Redis 实现)
- WebSocket handler 添加 query 消息限流(按 userID)
- Auth API 添加登录/注册限流(按 IP)
- refresh 和 logout 不限流(避免影响正常用户操作)
- 修复所有测试(传递 nil limiter 参数)
- 所有测试通过(包括 ws 和 api 集成测试)
2026-06-21 00:00:24 +08:00

291 lines
9.2 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/stt"
"github.com/hhs/camtalk/internal/ai/tts"
"github.com/hhs/camtalk/internal/config"
eino "github.com/hhs/camtalk/internal/eino"
"github.com/hhs/camtalk/internal/logger"
"github.com/hhs/camtalk/internal/ratelimit"
"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 rdb *redis.Client
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,
)
// 包装 userRepo 为带 Redis 缓存的版本refresh token 二级缓存)
userRepo = store.NewCachedUserRepository(userRepo, rdb, time.Duration(cfg.Auth.RefreshTTL)*time.Minute)
logger.Log.Infow("L2 Redis storage initialized",
"addr", cfg.Redis.Addr,
"db", cfg.Redis.DB,
"cached_user_repo", true)
}
// 初始化 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)
}
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)
}
// 初始化 Eino Graph + Orchestrator
pipelineGraph, err := eino.NewPipelineGraph(ctx, cfg, sttService, ttsService, sessionMgr)
if err != nil {
logger.Log.Fatalw("failed to create eino pipeline graph", "error", err)
}
orch := eino.NewEinoOrchestrator(pipelineGraph, sessionMgr, cfg.AI.LLM.Model)
// 初始化认证服务
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)
// 初始化限流器
var limiter ratelimit.Limiter
if cfg.RateLimit.Enabled {
if rdb != nil {
// 多实例:使用 Redis 令牌桶
limiter = ratelimit.NewRedisLimiter(rdb, cfg.RateLimit)
logger.Log.Info("rate limiter initialized with Redis backend")
} else {
// 单实例:使用内存令牌桶
limiter = ratelimit.NewMemoryLimiter(cfg.RateLimit)
logger.Log.Info("rate limiter initialized with in-memory backend")
}
defer limiter.Stop()
} else {
logger.Log.Info("rate limiter disabled")
}
// 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, limiter)
// Conversation REST 端点
convHandler := api.NewConversationHandler(sessionMgr, tokenMgr, msgRepo)
convHandler.RegisterRoutes(apiGroup)
// WebSocket
r.GET("/ws", ws.ServeWS(sessionMgr, orch, cfg, tokenMgr, limiter))
// 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(),
})
}
}