## 功能概述 - 用户可创建、编辑、删除自定义情景 - 支持自定义情景名称、图标、描述、Prompt、首句引导 - 完整的权限隔离,用户只能管理自己的情景 - 深度集成 Eino 框架,动态加载自建情景 Prompt ## 后端实现 ### 数据库 - 新增 user_scenarios 表 - 支持用户配额(最多 20 个) - 字段验证:description 可选,prompt 最小 10 字符 ### API - GET /api/scenarios - 获取用户情景列表 - POST /api/scenarios - 创建情景 - GET /api/scenarios/:id - 获取详情 - PATCH /api/scenarios/:id - 更新情景 - DELETE /api/scenarios/:id - 删除情景 ### Eino 集成 - PipelineState 添加 UserID 字段 - nodes_history 动态加载用户自建情景 - GetScenarioPrompt 支持自建情景优先级 ## 前端实现 ### 组件 - CreateScenarioModal - 创建情景对话框 - EditScenarioModal - 编辑情景对话框 - ConfigPanel 改造 - 分组显示系统预置和自建情景 ### Hook - useScenarios - 合并系统和自建情景,提供 CRUD 接口 ### 国际化 - 中文、英文、日文翻译支持 ## 问题修复 - 修复 CORS 问题:使用 Vite 代理 - 统一验证规则:description 可选,prompt 最小 10 字符 - 修复数据库约束:使用 NULLIF 处理空字符串 ## 文件变更 新增文件: 13 个 修改文件: 14 个 详见文档: docs/自建情景功能完整文档.md
295 lines
9.4 KiB
Go
295 lines
9.4 KiB
Go
package main
|
||
|
||
import (
|
||
"context"
|
||
"errors"
|
||
"net/http"
|
||
"os/signal"
|
||
"strings"
|
||
"syscall"
|
||
"time"
|
||
|
||
"github.com/gin-gonic/gin"
|
||
"github.com/jackc/pgx/v5/pgxpool"
|
||
"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/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
|
||
var pool *pgxpool.Pool // 数据库连接池
|
||
|
||
// 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")
|
||
}
|
||
var err error
|
||
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
|
||
var userScenarioRepo store.UserScenarioRepository
|
||
if pool != nil {
|
||
userScenarioRepo = store.NewPostgresUserScenarioRepo(pool)
|
||
}
|
||
pipelineGraph, err := eino.NewPipelineGraph(ctx, cfg, sttService, ttsService, sessionMgr, userScenarioRepo)
|
||
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)
|
||
|
||
// 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)
|
||
|
||
// UserScenario REST 端点
|
||
if pool != nil {
|
||
userScenarioRepo := store.NewPostgresUserScenarioRepo(pool)
|
||
userScenarioHandler := api.NewUserScenarioHandler(userScenarioRepo)
|
||
scenarioGroup := apiGroup.Group("/scenarios")
|
||
scenarioGroup.Use(auth.AuthMiddleware(tokenMgr))
|
||
{
|
||
scenarioGroup.GET("", userScenarioHandler.List)
|
||
scenarioGroup.POST("", userScenarioHandler.Create)
|
||
scenarioGroup.GET("/:id", userScenarioHandler.Get)
|
||
scenarioGroup.PATCH("/:id", userScenarioHandler.Update)
|
||
scenarioGroup.DELETE("/:id", userScenarioHandler.Delete)
|
||
}
|
||
}
|
||
|
||
// WebSocket
|
||
r.GET("/ws", ws.ServeWS(sessionMgr, orch, cfg, tokenMgr, userScenarioRepo))
|
||
|
||
// 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(),
|
||
})
|
||
}
|
||
}
|