Files
CamTalk/backend/internal/ws/handler_test.go
hhs 0b2bbb827d feat: Phase 8.1 - WS 集成测试 (handler_test.go)
使用 Gin test server + gorilla/websocket client 验证完整 WebSocket 流程:
- connected 连接建立、ping/pong 心跳
- 完整 query→stt_result→llm_chunk→tts_audio→llm_done 流程
- interrupt 中断、disconnect 断开清理
- 无效 JSON、未知消息类型错误处理
- 多次查询、无图片查询、TTS 未启用等场景
- 每次连接创建独立会话
2026-06-13 16:37:51 +08:00

564 lines
15 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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 如果非 nilProcessQuery 直接返回此错误
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, "不应有额外消息")
}