diff --git a/backend/internal/ws/handler_test.go b/backend/internal/ws/handler_test.go index 585acba..69f1207 100644 --- a/backend/internal/ws/handler_test.go +++ b/backend/internal/ws/handler_test.go @@ -2,6 +2,7 @@ package ws import ( "encoding/base64" + "net/http" "net/http/httptest" "strings" "testing" @@ -573,3 +574,156 @@ func TestWS_QueryWithTTSDisabled(t *testing.T) { err = conn.ReadJSON(&extra) assert.Error(t, err, "不应有额外消息") } + +// --- 认证测试辅助 --- + +// setupTestServerEx 创建测试服务器,返回 tokenMgr 和 sessionMgr 以便测试控制。 +func setupTestServerEx(t *testing.T, orch orchestrator.Orchestrator) (*httptest.Server, *auth.TokenManager, *session.MemoryManager) { + 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}, + } + r.GET("/ws", ServeWS(sessionMgr, orch, cfg, tokenMgr)) + + srv := httptest.NewServer(r) + return srv, tokenMgr, sessionMgr +} + +// httpGet 发送 HTTP GET 并返回状态码。 +func httpGet(t *testing.T, url string) int { + t.Helper() + resp, err := http.Get(url) + require.NoError(t, err) + resp.Body.Close() + return resp.StatusCode +} + +// --- 认证测试用例 --- + +// TestWS_AuthMissingToken 验证无 token 时返回 401。 +func TestWS_AuthMissingToken(t *testing.T) { + srv, _, _ := setupTestServerEx(t, &MockOrchestrator{}) + defer srv.Close() + + httpURL := srv.URL + "/ws" + status := httpGet(t, httpURL) + assert.Equal(t, http.StatusUnauthorized, status) +} + +// TestWS_AuthInvalidToken 验证无效 token 时返回 401。 +func TestWS_AuthInvalidToken(t *testing.T) { + srv, _, _ := setupTestServerEx(t, &MockOrchestrator{}) + defer srv.Close() + + httpURL := srv.URL + "/ws?token=invalid-token" + status := httpGet(t, httpURL) + assert.Equal(t, http.StatusUnauthorized, status) +} + +// TestWS_AuthExpiredToken 验证过期 token 时返回 401。 +func TestWS_AuthExpiredToken(t *testing.T) { + // 创建一个 access TTL 极短的 tokenMgr + sessionMgr := session.NewMemoryManager(5*time.Minute, 20) + defer sessionMgr.Stop() + + tokenMgr := auth.NewTokenManager("test-secret", -1*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}, + } + r.GET("/ws", ServeWS(sessionMgr, &MockOrchestrator{}, cfg, tokenMgr)) + srv := httptest.NewServer(r) + defer srv.Close() + + token, _, err := tokenMgr.GeneratePair("test-user", "testuser") + require.NoError(t, err) + + httpURL := srv.URL + "/ws?token=" + token + status := httpGet(t, httpURL) + assert.Equal(t, http.StatusUnauthorized, status) +} + +// TestWS_AuthValidToken 验证有效 token 能成功建立 WS 连接。 +func TestWS_AuthValidToken(t *testing.T) { + srv, tokenMgr, _ := setupTestServerEx(t, &MockOrchestrator{}) + defer srv.Close() + + token, _, err := tokenMgr.GeneratePair("user-1", "alice") + require.NoError(t, err) + + wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws?token=" + token + conn := connectWS(t, wsURL) + + msg := readJSON(t, conn) + assert.Equal(t, "connected", msg["type"]) + assert.NotEmpty(t, msg["session_id"]) +} + +// TestWS_AuthConversationIDResume 验证通过 conversation_id 恢复已有对话。 +func TestWS_AuthConversationIDResume(t *testing.T) { + srv, tokenMgr, sessionMgr := setupTestServerEx(t, &MockOrchestrator{}) + defer srv.Close() + + userID := "user-1" + + // 先创建一个属于该用户的 session + ctx := context.Background() + sessionID, err := sessionMgr.Create(ctx, userID, models.DefaultConfig()) + require.NoError(t, err) + + token, _, err := tokenMgr.GeneratePair(userID, "alice") + require.NoError(t, err) + + // 带 conversation_id 连接 + wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + + "/ws?token=" + token + "&conversation_id=" + sessionID + conn := connectWS(t, wsURL) + + msg := readJSON(t, conn) + assert.Equal(t, "connected", msg["type"]) + assert.Equal(t, sessionID, msg["session_id"], "应复用已有 session") +} + +// TestWS_AuthConversationIDNotFound 验证 conversation_id 不存在时返回 401。 +func TestWS_AuthConversationIDNotFound(t *testing.T) { + srv, tokenMgr, _ := setupTestServerEx(t, &MockOrchestrator{}) + defer srv.Close() + + token, _, err := tokenMgr.GeneratePair("user-1", "alice") + require.NoError(t, err) + + httpURL := srv.URL + "/ws?token=" + token + "&conversation_id=nonexistent-id" + status := httpGet(t, httpURL) + assert.Equal(t, http.StatusUnauthorized, status) +} + +// TestWS_AuthConversationIDOwnership 验证 conversation_id 不属于当前用户时返回 401。 +func TestWS_AuthConversationIDOwnership(t *testing.T) { + srv, tokenMgr, sessionMgr := setupTestServerEx(t, &MockOrchestrator{}) + defer srv.Close() + + ctx := context.Background() + // user-A 创建 session + sessionID, err := sessionMgr.Create(ctx, "user-A", models.DefaultConfig()) + require.NoError(t, err) + + // user-B 尝试连接该 session + token, _, err := tokenMgr.GeneratePair("user-B", "bob") + require.NoError(t, err) + + httpURL := srv.URL + "/ws?token=" + token + "&conversation_id=" + sessionID + status := httpGet(t, httpURL) + assert.Equal(t, http.StatusUnauthorized, status, "非 owner 访问应返回 401") +}