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, "不应有额外消息") }