diff --git a/backend/internal/ws/handler_test.go b/backend/internal/ws/handler_test.go new file mode 100644 index 0000000..97cf697 --- /dev/null +++ b/backend/internal/ws/handler_test.go @@ -0,0 +1,563 @@ +package ws + +import ( + "encoding/base64" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/gin-gonic/gin" + "github.com/gorilla/websocket" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "context" + + "github.com/hhs/camtalk/internal/logger" + "github.com/hhs/camtalk/internal/models" + "github.com/hhs/camtalk/internal/orchestrator" + "github.com/hhs/camtalk/internal/session" +) + +func init() { + logger.Init("debug", "console") + gin.SetMode(gin.TestMode) +} + +// --- Mock Orchestrator --- + +// MockOrchestrator 实现 orchestrator.Orchestrator 接口, +// 模拟完整的 STT → LLM → TTS 管道,通过 sender 推送消息。 +type MockOrchestrator struct { + // STTResult 模拟的语音识别结果 + STTResult string + // LLMDeltas 模拟的 LLM 流式输出 + LLMDeltas []string + // TTSAudios 模拟的 TTS 音频数据(每项一个 base64 编码的 MP3 片段) + TTSAudios []string + // Err 如果非 nil,ProcessQuery 直接返回此错误 + Err error + // Delay 每个消息之间的延迟(用于 interrupt 测试) + Delay time.Duration +} + +func (m *MockOrchestrator) ProcessQuery( + ctx context.Context, + sessionID string, + req models.WsQuery, + history []models.Message, + sender orchestrator.Sender, +) error { + if m.Err != nil { + sender.SendError(models.WsError{ + Type: "error", + RequestID: req.RequestID, + Code: "INTERNAL_ERROR", + Message: m.Err.Error(), + }) + return m.Err + } + + // Step 1: 发送 STT 结果 + if m.STTResult != "" { + _ = sender.SendSTTResult(models.WsSTTResult{ + Type: "stt_result", + RequestID: req.RequestID, + Text: m.STTResult, + IsFinal: true, + }) + } + if m.Delay > 0 { + select { + case <-ctx.Done(): + return nil // 中断视为正常完成 + case <-time.After(m.Delay): + } + } + + // Step 2: 发送 LLM chunks + var fullText strings.Builder + for _, delta := range m.LLMDeltas { + select { + case <-ctx.Done(): + return nil // 中断视为正常完成 + default: + } + fullText.WriteString(delta) + _ = sender.SendLLMChunk(models.WsLLMChunk{ + Type: "llm_chunk", + RequestID: req.RequestID, + Delta: delta, + Role: "assistant", + }) + if m.Delay > 0 { + select { + case <-ctx.Done(): + return nil // 中断视为正常完成 + case <-time.After(m.Delay): + } + } + } + + // Step 3: 发送 TTS 音频 + for i, audio := range m.TTSAudios { + select { + case <-ctx.Done(): + return nil // 中断视为正常完成 + default: + } + isLast := i == len(m.TTSAudios)-1 + _ = sender.SendTTSAudio(models.WsTTSAudio{ + Type: "tts_audio", + RequestID: req.RequestID, + Audio: audio, + MimeType: "audio/mp3", + IsLast: isLast, + }) + } + + // Step 4: 发送 llm_done + _ = sender.SendLLMDone(models.WsLLMDone{ + Type: "llm_done", + RequestID: req.RequestID, + FullText: fullText.String(), + Model: "gpt-4o", + LatencyMs: 100, + }) + + return nil +} + +// --- 测试辅助函数 --- + +// setupTestServer 创建测试用 Gin 服务器和 WebSocket URL。 +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() }) + + r := gin.New() + r.GET("/ws", ServeWS(sessionMgr, orch)) + + srv := httptest.NewServer(r) + + // 构造 WebSocket URL + wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws" + + return srv, wsURL +} + +// connectWS 建立 WebSocket 连接并返回 conn。 +func connectWS(t *testing.T, wsURL string) *websocket.Conn { + t.Helper() + + conn, _, err := websocket.DefaultDialer.Dial(wsURL, nil) + require.NoError(t, err, "WebSocket 连接失败") + t.Cleanup(func() { conn.Close() }) + return conn +} + +// readJSON 从 WebSocket 读取一条 JSON 消息。 +func readJSON(t *testing.T, conn *websocket.Conn) map[string]any { + t.Helper() + + conn.SetReadDeadline(time.Now().Add(5 * time.Second)) + var msg map[string]any + err := conn.ReadJSON(&msg) + require.NoError(t, err, "读取 WebSocket 消息失败") + return msg +} + +// --- 测试用例 --- + +// TestWS_Connected 验证连接建立后收到 connected 消息。 +func TestWS_Connected(t *testing.T) { + srv, wsURL := setupTestServer(t, &MockOrchestrator{}) + defer srv.Close() + + conn := connectWS(t, wsURL) + + msg := readJSON(t, conn) + assert.Equal(t, "connected", msg["type"]) + assert.NotEmpty(t, msg["session_id"]) + assert.Equal(t, "0.1.0", msg["server_version"]) +} + +// TestWS_PingPong 验证 ping/pong 心跳。 +func TestWS_PingPong(t *testing.T) { + srv, wsURL := setupTestServer(t, &MockOrchestrator{}) + defer srv.Close() + + conn := connectWS(t, wsURL) + + // 读取 connected 消息 + _ = readJSON(t, conn) + + // 发送 ping + err := conn.WriteJSON(map[string]string{"type": "ping"}) + require.NoError(t, err) + + // 读取 pong + msg := readJSON(t, conn) + assert.Equal(t, "pong", msg["type"]) +} + +// TestWS_QueryFullFlow 验证完整的 query → stt_result → llm_chunk → tts_audio → llm_done 流程。 +func TestWS_QueryFullFlow(t *testing.T) { + audioB64 := base64.StdEncoding.EncodeToString([]byte("fake-audio-data")) + imageB64 := base64.StdEncoding.EncodeToString([]byte("fake-image-data")) + + mock := &MockOrchestrator{ + STTResult: "你好,世界", + LLMDeltas: []string{"你好", ",世界!"}, + TTSAudios: []string{base64.StdEncoding.EncodeToString([]byte("mp3-data-1")), base64.StdEncoding.EncodeToString([]byte("mp3-data-2"))}, + } + + srv, wsURL := setupTestServer(t, mock) + defer srv.Close() + + conn := connectWS(t, wsURL) + + // 1. 读取 connected + connected := readJSON(t, conn) + assert.Equal(t, "connected", connected["type"]) + sessionID := connected["session_id"].(string) + assert.NotEmpty(t, sessionID) + + // 2. 发送 query + queryMsg := models.WsQuery{ + Type: "query", + RequestID: "req-test-001", + Image: imageB64, + Audio: audioB64, + MimeType: "audio/pcm", + } + err := conn.WriteJSON(queryMsg) + require.NoError(t, err) + + // 3. 读取 stt_result + sttResult := readJSON(t, conn) + assert.Equal(t, "stt_result", sttResult["type"]) + assert.Equal(t, "req-test-001", sttResult["request_id"]) + assert.Equal(t, "你好,世界", sttResult["text"]) + assert.Equal(t, true, sttResult["is_final"]) + + // 4. 读取 llm_chunk 消息 + chunk1 := readJSON(t, conn) + assert.Equal(t, "llm_chunk", chunk1["type"]) + assert.Equal(t, "req-test-001", chunk1["request_id"]) + assert.Equal(t, "你好", chunk1["delta"]) + assert.Equal(t, "assistant", chunk1["role"]) + + chunk2 := readJSON(t, conn) + assert.Equal(t, "llm_chunk", chunk2["type"]) + assert.Equal(t, ",世界!", chunk2["delta"]) + + // 5. 读取 tts_audio 消息 + tts1 := readJSON(t, conn) + assert.Equal(t, "tts_audio", tts1["type"]) + assert.Equal(t, "req-test-001", tts1["request_id"]) + assert.NotEmpty(t, tts1["audio"]) + assert.Equal(t, "audio/mp3", tts1["mime_type"]) + assert.Equal(t, false, tts1["is_last"]) + + tts2 := readJSON(t, conn) + assert.Equal(t, "tts_audio", tts2["type"]) + assert.Equal(t, true, tts2["is_last"]) + + // 6. 读取 llm_done + llmDone := readJSON(t, conn) + assert.Equal(t, "llm_done", llmDone["type"]) + assert.Equal(t, "req-test-001", llmDone["request_id"]) + assert.Equal(t, "你好,世界!", llmDone["full_text"]) + assert.Equal(t, "gpt-4o", llmDone["model"]) + assert.NotNil(t, llmDone["latency_ms"]) +} + +// TestWS_QuerySTTOnly 验证只有 STT 结果、无 LLM 输出的场景。 +func TestWS_QuerySTTOnly(t *testing.T) { + audioB64 := base64.StdEncoding.EncodeToString([]byte("audio")) + + mock := &MockOrchestrator{ + STTResult: "测试语音", + // LLMDeltas 为空 → 不发送 llm_chunk + // TTSAudios 为空 → 不发送 tts_audio + } + + srv, wsURL := setupTestServer(t, mock) + defer srv.Close() + + conn := connectWS(t, wsURL) + _ = readJSON(t, conn) // connected + + err := conn.WriteJSON(models.WsQuery{ + Type: "query", + RequestID: "req-stt-only", + Audio: audioB64, + }) + require.NoError(t, err) + + // 应收到 stt_result + stt := readJSON(t, conn) + assert.Equal(t, "stt_result", stt["type"]) + assert.Equal(t, "测试语音", stt["text"]) + + // 应收到 llm_done(即使没有 chunk) + done := readJSON(t, conn) + assert.Equal(t, "llm_done", done["type"]) + assert.Equal(t, "", done["full_text"]) +} + +// TestWS_UnknownMessageType 验证未知消息类型返回 error。 +func TestWS_UnknownMessageType(t *testing.T) { + srv, wsURL := setupTestServer(t, &MockOrchestrator{}) + defer srv.Close() + + conn := connectWS(t, wsURL) + _ = readJSON(t, conn) // connected + + err := conn.WriteJSON(map[string]string{"type": "unknown_type"}) + require.NoError(t, err) + + errMsg := readJSON(t, conn) + assert.Equal(t, "error", errMsg["type"]) + assert.Equal(t, "INVALID_MESSAGE", errMsg["code"]) + assert.Contains(t, errMsg["message"], "unknown message type") +} + +// TestWS_InvalidJSON 验证无效 JSON 返回 error。 +func TestWS_InvalidJSON(t *testing.T) { + srv, wsURL := setupTestServer(t, &MockOrchestrator{}) + defer srv.Close() + + conn := connectWS(t, wsURL) + _ = readJSON(t, conn) // connected + + err := conn.WriteMessage(websocket.TextMessage, []byte("not-json")) + require.NoError(t, err) + + errMsg := readJSON(t, conn) + assert.Equal(t, "error", errMsg["type"]) + assert.Equal(t, "INVALID_MESSAGE", errMsg["code"]) +} + +// TestWS_MultipleQueries 验证同一连接上可以发送多次 query。 +func TestWS_MultipleQueries(t *testing.T) { + mock := &MockOrchestrator{ + STTResult: "识别结果", + LLMDeltas: []string{"回复"}, + } + + srv, wsURL := setupTestServer(t, mock) + defer srv.Close() + + conn := connectWS(t, wsURL) + _ = readJSON(t, conn) // connected + + audioB64 := base64.StdEncoding.EncodeToString([]byte("audio")) + + for i := 0; i < 3; i++ { + reqID := "req-multi-" + string(rune('0'+i)) + err := conn.WriteJSON(models.WsQuery{ + Type: "query", + RequestID: reqID, + Audio: audioB64, + }) + require.NoError(t, err) + + // 每次应收到完整的响应序列 + stt := readJSON(t, conn) + assert.Equal(t, "stt_result", stt["type"], "第 %d 次 query", i+1) + + chunk := readJSON(t, conn) + assert.Equal(t, "llm_chunk", chunk["type"], "第 %d 次 query", i+1) + + done := readJSON(t, conn) + assert.Equal(t, "llm_done", done["type"], "第 %d 次 query", i+1) + } +} + +// TestWS_Interrupt 验证 interrupt 取消正在进行的请求。 +func TestWS_Interrupt(t *testing.T) { + // 使用较长延迟模拟慢请求 + mock := &MockOrchestrator{ + STTResult: "识别文本", + LLMDeltas: []string{"第一句", "第二句", "第三句", "第四句", "第五句"}, + Delay: 200 * time.Millisecond, + } + + srv, wsURL := setupTestServer(t, mock) + defer srv.Close() + + conn := connectWS(t, wsURL) + _ = readJSON(t, conn) // connected + + audioB64 := base64.StdEncoding.EncodeToString([]byte("audio")) + + // 发送 query + err := conn.WriteJSON(models.WsQuery{ + Type: "query", + RequestID: "req-interrupt", + Audio: audioB64, + }) + require.NoError(t, err) + + // 收到 stt_result + stt := readJSON(t, conn) + assert.Equal(t, "stt_result", stt["type"]) + + // 收到第一个 llm_chunk + chunk1 := readJSON(t, conn) + assert.Equal(t, "llm_chunk", chunk1["type"]) + + // 发送 interrupt + err = conn.WriteJSON(map[string]string{"type": "interrupt"}) + require.NoError(t, err) + + // 等待 interrupt 生效 + time.Sleep(500 * time.Millisecond) + + // 验证连接仍然存活(可以发 ping 收 pong) + require.NoError(t, conn.WriteJSON(map[string]string{"type": "ping"})) + + var pong map[string]any + conn.SetReadDeadline(time.Now().Add(5 * time.Second)) + require.NoError(t, conn.ReadJSON(&pong), "interrupt 后连接应仍存活") + assert.Equal(t, "pong", pong["type"]) +} + +// TestWS_DisconnectCleanup 验证断开连接时清理资源。 +func TestWS_DisconnectCleanup(t *testing.T) { + // 使用较长延迟模拟慢请求 + mock := &MockOrchestrator{ + STTResult: "识别文本", + LLMDeltas: []string{"长回复第一部分", "长回复第二部分"}, + Delay: 500 * time.Millisecond, + } + + srv, wsURL := setupTestServer(t, mock) + defer srv.Close() + + conn := connectWS(t, wsURL) + _ = readJSON(t, conn) // connected + + audioB64 := base64.StdEncoding.EncodeToString([]byte("audio")) + + // 发送 query + err := conn.WriteJSON(models.WsQuery{ + Type: "query", + RequestID: "req-disconnect", + Audio: audioB64, + }) + require.NoError(t, err) + + // 收到 stt_result + stt := readJSON(t, conn) + assert.Equal(t, "stt_result", stt["type"]) + + // 关闭连接(模拟客户端断开) + conn.Close() + + // 等待一小段时间让服务器处理断开 + time.Sleep(300 * time.Millisecond) + + // 如果没有 panic 或 goroutine 泄漏,测试通过 + // (Go test 的 -race 检测器会捕获数据竞争) +} + +// TestWS_SessionCreated 验证每次连接都创建新会话。 +func TestWS_SessionCreated(t *testing.T) { + srv, wsURL := setupTestServer(t, &MockOrchestrator{}) + defer srv.Close() + + // 第一次连接 + conn1 := connectWS(t, wsURL) + msg1 := readJSON(t, conn1) + sid1 := msg1["session_id"].(string) + conn1.Close() + + time.Sleep(100 * time.Millisecond) + + // 第二次连接 + conn2 := connectWS(t, wsURL) + sid2 := readJSON(t, conn2)["session_id"].(string) + + assert.NotEmpty(t, sid1) + assert.NotEmpty(t, sid2) + assert.NotEqual(t, sid1, sid2, "两次连接应创建不同的会话") +} + +// TestWS_QueryWithoutImage 验证不带图片的 query。 +func TestWS_QueryWithoutImage(t *testing.T) { + mock := &MockOrchestrator{ + STTResult: "纯语音输入", + LLMDeltas: []string{"收到"}, + } + + srv, wsURL := setupTestServer(t, mock) + defer srv.Close() + + conn := connectWS(t, wsURL) + _ = readJSON(t, conn) // connected + + audioB64 := base64.StdEncoding.EncodeToString([]byte("audio")) + + err := conn.WriteJSON(models.WsQuery{ + Type: "query", + RequestID: "req-no-image", + Audio: audioB64, + // Image 为空 + }) + require.NoError(t, err) + + stt := readJSON(t, conn) + assert.Equal(t, "stt_result", stt["type"]) + assert.Equal(t, "纯语音输入", stt["text"]) + + chunk := readJSON(t, conn) + assert.Equal(t, "llm_chunk", chunk["type"]) + + done := readJSON(t, conn) + assert.Equal(t, "llm_done", done["type"]) +} + +// TestWS_QueryWithTTSDisabled 验证 TTS 未启用时不应收到 tts_audio。 +func TestWS_QueryWithTTSDisabled(t *testing.T) { + // MockOrchestrator 的 TTSAudios 为空 → 不发送 tts_audio + mock := &MockOrchestrator{ + STTResult: "语音", + LLMDeltas: []string{"回复"}, + // TTSAudios 留空 + } + + srv, wsURL := setupTestServer(t, mock) + defer srv.Close() + + conn := connectWS(t, wsURL) + _ = readJSON(t, conn) // connected + + audioB64 := base64.StdEncoding.EncodeToString([]byte("audio")) + + err := conn.WriteJSON(models.WsQuery{ + Type: "query", + RequestID: "req-no-tts", + Audio: audioB64, + }) + require.NoError(t, err) + + stt := readJSON(t, conn) + assert.Equal(t, "stt_result", stt["type"]) + + chunk := readJSON(t, conn) + assert.Equal(t, "llm_chunk", chunk["type"]) + + done := readJSON(t, conn) + assert.Equal(t, "llm_done", done["type"]) + + // 不应有 tts_audio 消息;设置短超时验证 + conn.SetReadDeadline(time.Now().Add(200 * time.Millisecond)) + var extra map[string]any + err = conn.ReadJSON(&extra) + assert.Error(t, err, "不应有额外消息") +}