- TestWS_AuthMissingToken: 无 token 返回 401 - TestWS_AuthInvalidToken: 无效 token 返回 401 - TestWS_AuthExpiredToken: 过期 token 返回 401 - TestWS_AuthValidToken: 有效 token 成功连接 - TestWS_AuthConversationIDResume: conversation_id 恢复已有对话 - TestWS_AuthConversationIDNotFound: 不存在的 conversation_id 返回 401 - TestWS_AuthConversationIDOwnership: 非 owner 访问返回 401
730 lines
20 KiB
Go
730 lines
20 KiB
Go
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 如果非 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。
|
||
// 返回的 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))
|
||
|
||
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))
|
||
|
||
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")
|
||
}
|