diff --git a/backend/internal/ws/handler.go b/backend/internal/ws/handler.go index 4e17bc8..5aa2acd 100644 --- a/backend/internal/ws/handler.go +++ b/backend/internal/ws/handler.go @@ -10,6 +10,7 @@ import ( "github.com/gin-gonic/gin" "github.com/gorilla/websocket" + "github.com/hhs/camtalk/internal/config" "github.com/hhs/camtalk/internal/errors" "github.com/hhs/camtalk/internal/logger" "github.com/hhs/camtalk/internal/models" @@ -17,8 +18,23 @@ import ( "github.com/hhs/camtalk/internal/session" ) -var upgrader = websocket.Upgrader{ - CheckOrigin: func(r *http.Request) bool { return true }, // 开发阶段允许所有来源 +// newUpgrader 根据配置创建 WebSocket upgrader。 +func newUpgrader(cfg *config.Config) websocket.Upgrader { + allowedOrigins := cfg.Server.AllowedOrigins + return websocket.Upgrader{ + CheckOrigin: func(r *http.Request) bool { + if len(allowedOrigins) == 0 { + return true // 未配置则允许所有来源(开发模式) + } + origin := r.Header.Get("Origin") + for _, o := range allowedOrigins { + if o == origin || o == "*" { + return true + } + } + return false + }, + } } // Client 代表一个 WebSocket 客户端连接。 @@ -75,13 +91,21 @@ func (w *WSClient) SendError(err models.WsError) error { } // ServeWS 处理 WebSocket 升级请求。 -func ServeWS(sessionMgr session.Manager, orch orchestrator.Orchestrator) gin.HandlerFunc { +func ServeWS(sessionMgr session.Manager, orch orchestrator.Orchestrator, cfg *config.Config) gin.HandlerFunc { + upgrader := newUpgrader(cfg) + heartbeatInterval := time.Duration(cfg.Server.HeartbeatInterval) * time.Second + heartbeatTimeout := time.Duration(cfg.Server.HeartbeatTimeout) * time.Second + version := cfg.App.Version + + maxHistory := cfg.Session.MaxHistory + return func(c *gin.Context) { - serveWS(c, sessionMgr, orch) + serveWS(c, sessionMgr, orch, upgrader, heartbeatInterval, heartbeatTimeout, version, maxHistory) } } -func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orchestrator) { +func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orchestrator, + upgrader websocket.Upgrader, heartbeatInterval, heartbeatTimeout time.Duration, version string, maxHistory int) { conn, err := upgrader.Upgrade(c.Writer, c.Request, nil) if err != nil { logger.Log.Errorw("websocket upgrade failed", "error", err) @@ -108,7 +132,7 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche _ = client.SendJSON(models.WsConnected{ Type: "connected", SessionID: sessionID, - ServerVersion: "0.1.0", + ServerVersion: version, }) logger.Log.Infow("client connected", "session", sessionID) @@ -122,12 +146,12 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche // 启动心跳检查 goroutine done := make(chan struct{}) go func() { - ticker := time.NewTicker(30 * time.Second) + ticker := time.NewTicker(heartbeatInterval) defer ticker.Stop() for { select { case <-ticker.C: - if time.Since(lastPong) > 60*time.Second { + if time.Since(lastPong) > heartbeatTimeout { logger.Log.Warnw("heartbeat timeout", "session", sessionID) conn.Close() return @@ -181,7 +205,7 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche } // 获取对话历史 - history, _ := client.sessionMgr.GetHistory(context.Background(), sessionID, 20) + history, _ := client.sessionMgr.GetHistory(context.Background(), sessionID, maxHistory) // 创建可取消的 context ctx, cancel := context.WithCancel(context.Background())