Files
CamTalk/backend/internal/ws/handler_test.go

729 lines
20 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"
"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/auth"
"github.com/hhs/camtalk/internal/config"
"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,
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。
// 返回的 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},
}
r.GET("/ws", ServeWS(sessionMgr, orch, cfg, tokenMgr, nil, nil))
srv := httptest.NewServer(r)
// 生成有效 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
}
// 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, "test", 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, "不应有额外消息")
}
// --- 认证测试辅助 ---
// 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, nil, nil))
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, nil))
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")
}