diff --git a/backend/internal/ws/handler.go b/backend/internal/ws/handler.go index 7274b68..8b98e38 100644 --- a/backend/internal/ws/handler.go +++ b/backend/internal/ws/handler.go @@ -107,6 +107,21 @@ func ServeWS(sessionMgr session.Manager, orch orchestrator.Orchestrator, cfg *co func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orchestrator, upgrader websocket.Upgrader, heartbeatInterval, heartbeatTimeout time.Duration, version string, maxHistory int, tokenMgr *auth.TokenManager) { + + // --- JWT 认证(upgrade 前完成,失败直接返回 HTTP 错误) --- + token := c.Query("token") + if token == "" { + c.JSON(http.StatusUnauthorized, gin.H{"error": "missing token"}) + return + } + claims, err := tokenMgr.ValidateAccess(token) + if err != nil { + c.JSON(http.StatusUnauthorized, gin.H{"error": "invalid token"}) + return + } + userID := claims.UserID + username := claims.Username + conn, err := upgrader.Upgrade(c.Writer, c.Request, nil) if err != nil { logger.Log.Errorw("websocket upgrade failed", "error", err) @@ -115,7 +130,7 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche defer conn.Close() // 创建会话 - sessionID, err := sessionMgr.Create(context.Background(), "", models.DefaultConfig()) + sessionID, err := sessionMgr.Create(context.Background(), userID, models.DefaultConfig()) if err != nil { logger.Log.Errorw("create session failed", "error", err) return @@ -135,7 +150,7 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche SessionID: sessionID, ServerVersion: version, }) - logger.Log.Infow("client connected", "session", sessionID) + logger.Log.Infow("client connected", "session", sessionID, "user_id", userID, "username", username) // 心跳检测 lastPong := time.Now() diff --git a/backend/internal/ws/handler_test.go b/backend/internal/ws/handler_test.go index 7fc43cb..585acba 100644 --- a/backend/internal/ws/handler_test.go +++ b/backend/internal/ws/handler_test.go @@ -133,25 +133,29 @@ func (m *MockOrchestrator) ProcessQuery( // --- 测试辅助函数 --- // setupTestServer 创建测试用 Gin 服务器和 WebSocket URL。 +// 返回的 wsURL 已包含有效 token,可直接连接。 func setupTestServer(t *testing.T, orch orchestrator.Orchestrator) (*httptest.Server, string) { t.Helper() sessionMgr := session.NewMemoryManager(5*time.Minute, 20) t.Cleanup(func() { sessionMgr.Stop() }) + tokenMgr := auth.NewTokenManager("test-secret", 15*time.Minute, 7*24*time.Hour) + r := gin.New() cfg := &config.Config{ App: config.AppConfig{Version: "test"}, Server: config.ServerConfig{HeartbeatInterval: 30, HeartbeatTimeout: 60}, Session: config.SessionConfig{MaxHistory: 20}, } - tokenMgr := auth.NewTokenManager("test-secret", 15*time.Minute, 7*24*time.Hour) r.GET("/ws", ServeWS(sessionMgr, orch, cfg, tokenMgr)) srv := httptest.NewServer(r) - // 构造 WebSocket URL - wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws" + // 生成有效 token 并构造 WebSocket URL + token, _, err := tokenMgr.GeneratePair("test-user", "testuser") + require.NoError(t, err) + wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws?token=" + token return srv, wsURL }