feat: 添加集成测试并同步文档 #42
563
backend/internal/ws/handler_test.go
Normal file
563
backend/internal/ws/handler_test.go
Normal file
@@ -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, "不应有额外消息")
|
||||||
|
}
|
||||||
165
docs/03-接口文档.md
165
docs/03-接口文档.md
@@ -256,7 +256,6 @@ POST /api/sessions
|
|||||||
Content-Type: application/json
|
Content-Type: application/json
|
||||||
|
|
||||||
{
|
{
|
||||||
"user_id": "optional-user-id",
|
|
||||||
"config": {
|
"config": {
|
||||||
"tts_enabled": true,
|
"tts_enabled": true,
|
||||||
"detail_level": "low",
|
"detail_level": "low",
|
||||||
@@ -301,28 +300,18 @@ Go 网关内部与外部 AI 服务(Deepgram STT、GPT-4o、OpenAI TTS)的调
|
|||||||
语音识别:接收前端采集的音频,返回识别文本。
|
语音识别:接收前端采集的音频,返回识别文本。
|
||||||
|
|
||||||
```go
|
```go
|
||||||
// STTService 语音识别服务契约。
|
// Service 语音识别服务契约。
|
||||||
type STTService interface {
|
type Service interface {
|
||||||
// Recognize 识别一段完整音频,返回最终文本。
|
// Recognize 识别一段完整音频,返回最终文本。
|
||||||
Recognize(ctx context.Context, audio []byte, opts STTOptions) (string, error)
|
Recognize(ctx context.Context, audio []byte, opts Options) (string, error)
|
||||||
|
|
||||||
// RecognizeStream 流式识别(边说边识别,可选实现)。
|
|
||||||
// audioStream 持续接收音频片段,返回的 channel 持续输出中间结果。
|
|
||||||
RecognizeStream(ctx context.Context, audioStream <-chan []byte, opts STTOptions) (<-chan STTPartial, error)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// STTOptions 语音识别参数。
|
// Options 语音识别参数。
|
||||||
type STTOptions struct {
|
type Options struct {
|
||||||
Encoding string // "pcm_s16le" — 前端 VAD 输出格式
|
Encoding string // "pcm_s16le" — 前端 VAD 输出格式
|
||||||
SampleRate int // 16000 — 前端麦克风采样率
|
SampleRate int // 16000 — 前端麦克风采样率
|
||||||
Language string // "zh-CN"
|
Language string // "zh-CN"
|
||||||
}
|
}
|
||||||
|
|
||||||
// STTPartial 流式识别的中间/最终结果。
|
|
||||||
type STTPartial struct {
|
|
||||||
Text string
|
|
||||||
IsFinal bool
|
|
||||||
}
|
|
||||||
```
|
```
|
||||||
|
|
||||||
**Deepgram 接入约定**:
|
**Deepgram 接入约定**:
|
||||||
@@ -336,23 +325,23 @@ type STTPartial struct {
|
|||||||
多模态推理:接收图像 + 文本 + 对话历史,流式返回回复。
|
多模态推理:接收图像 + 文本 + 对话历史,流式返回回复。
|
||||||
|
|
||||||
```go
|
```go
|
||||||
// LLMService 多模态大模型服务契约。
|
// Service 多模态大模型服务契约。
|
||||||
type LLMService interface {
|
type Service interface {
|
||||||
// ChatStream 流式推理,返回增量文本的 channel。
|
// ChatStream 流式推理,返回增量文本的 channel。
|
||||||
// 调用方必须消费 channel 直到 Done=true,否则需 cancel ctx 以释放连接。
|
// 调用方必须消费 channel 直到 Done=true,否则需 cancel ctx 以释放连接。
|
||||||
ChatStream(ctx context.Context, req LLMRequest) (<-chan LLMChunk, error)
|
ChatStream(ctx context.Context, req Request) (<-chan Chunk, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
// LLMRequest 推理请求。
|
// Request 推理请求。
|
||||||
type LLMRequest struct {
|
type Request struct {
|
||||||
Image []byte // JPEG 图片(已从 Base64 解码)
|
Image []byte // JPEG 图片(已从 Base64 解码)
|
||||||
Text string // 用户语音识别后的文本
|
Text string // 用户语音识别后的文本
|
||||||
History []Message // 最近 N 轮对话历史
|
History []models.Message // 最近 N 轮对话历史
|
||||||
Language string // "zh-CN"
|
Language string // "zh-CN"
|
||||||
}
|
}
|
||||||
|
|
||||||
// LLMChunk 流式推理的一个增量片段。
|
// Chunk 流式推理的一个增量片段。
|
||||||
type LLMChunk struct {
|
type Chunk struct {
|
||||||
Delta string // 增量文本
|
Delta string // 增量文本
|
||||||
Done bool // 是否结束
|
Done bool // 是否结束
|
||||||
TokensUsed *TokenUsage // 仅 Done=true 时有值
|
TokensUsed *TokenUsage // 仅 Done=true 时有值
|
||||||
@@ -388,24 +377,24 @@ user: [图片 + 用户语音文本]
|
|||||||
语音合成:接收文本流,输出音频 chunk 流。
|
语音合成:接收文本流,输出音频 chunk 流。
|
||||||
|
|
||||||
```go
|
```go
|
||||||
// TTSService 语音合成服务契约。
|
// Service 语音合成服务契约。
|
||||||
type TTSService interface {
|
type Service interface {
|
||||||
// SynthesizeStream 流式合成。
|
// SynthesizeStream 流式合成。
|
||||||
// textStream 接收句子级文本(由 Orchestrator 的句子切分器产出),
|
// textStream 接收句子级文本(由 Orchestrator 的句子切分器产出),
|
||||||
// 返回的 channel 持续输出 MP3 音频 chunk。
|
// 返回的 channel 持续输出 MP3 音频 chunk。
|
||||||
SynthesizeStream(ctx context.Context, textStream <-chan string, opts TTSOptions) (<-chan TTSChunk, error)
|
SynthesizeStream(ctx context.Context, textStream <-chan string, opts Options) (<-chan Chunk, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
// TTSOptions 合成参数。
|
// Options 合成参数。
|
||||||
type TTSOptions struct {
|
type Options struct {
|
||||||
Voice string // "alloy" | "nova" | "shimmer" | ...
|
Voice string // "alloy" | "nova" | "shimmer" | ...
|
||||||
Speed float64 // 1.0 为正常语速
|
Speed float64 // 1.0 为正常语速
|
||||||
OutputFmt string // "mp3" — 固定使用 MP3,浏览器原生支持
|
OutputFmt string // "mp3" — 固定使用 MP3,浏览器原生支持
|
||||||
SampleRate int // 24000
|
SampleRate int // 24000
|
||||||
}
|
}
|
||||||
|
|
||||||
// TTSChunk 一个音频片段。
|
// Chunk 一个音频片段。
|
||||||
type TTSChunk struct {
|
type Chunk struct {
|
||||||
Audio []byte // MP3 音频数据(未 Base64 编码,由发送层编码)
|
Audio []byte // MP3 音频数据(未 Base64 编码,由发送层编码)
|
||||||
IsLast bool // 是否为最后一片
|
IsLast bool // 是否为最后一片
|
||||||
}
|
}
|
||||||
@@ -446,73 +435,32 @@ LLM 流式输出: "这" "是一" "朵红色" "的花。" "它看起" "来很美
|
|||||||
### Orchestrator 接口
|
### Orchestrator 接口
|
||||||
|
|
||||||
```go
|
```go
|
||||||
// Orchestrator AI 编排器,协调 STT → LLM → TTS 全链路。
|
// Orchestrator AI 编排器接口。
|
||||||
type Orchestrator struct {
|
type Orchestrator interface {
|
||||||
stt STTService
|
// ProcessQuery 处理一次完整的视觉对话请求。
|
||||||
llm LLMService
|
// 通过 sender 向前端实时推送 stt_result、llm_chunk、llm_done、tts_audio 消息。
|
||||||
tts TTSService
|
ProcessQuery(ctx context.Context, sessionID string, req models.WsQuery,
|
||||||
|
history []models.Message, sender Sender) error
|
||||||
}
|
}
|
||||||
|
|
||||||
// ProcessQuery 处理一次完整的视觉对话请求。
|
// Sender 抽象 WebSocket 消息推送能力,便于测试时 mock。
|
||||||
// 通过 client 向前端实时推送 stt_result、llm_chunk、llm_done、tts_audio 消息。
|
type Sender interface {
|
||||||
func (o *Orchestrator) ProcessQuery(ctx context.Context, client MessageSender, req *QueryRequest) {
|
SendSTTResult(result models.WsSTTResult) error
|
||||||
ctx, cancel := context.WithTimeout(ctx, 10*time.Second)
|
SendLLMChunk(chunk models.WsLLMChunk) error
|
||||||
defer cancel()
|
SendLLMDone(done models.WsLLMDone) error
|
||||||
|
SendTTSAudio(audio models.WsTTSAudio) error
|
||||||
// Step 1: STT — 识别用户语音
|
SendError(err models.WsError) error
|
||||||
text, err := o.stt.Recognize(ctx, req.Audio, STTOptions{
|
|
||||||
Encoding: "pcm_s16le", SampleRate: 16000, Language: "zh-CN",
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
client.SendError(req.RequestID, "STT_ERROR", err.Error())
|
|
||||||
return
|
|
||||||
}
|
|
||||||
client.SendSTTResult(req.RequestID, text, true)
|
|
||||||
|
|
||||||
// Step 2: LLM 流式输出 + 句子切分
|
|
||||||
llmStream, _ := o.llm.ChatStream(ctx, LLMRequest{
|
|
||||||
Image: req.Image, Text: text, Language: "zh-CN",
|
|
||||||
})
|
|
||||||
|
|
||||||
sentenceCh := make(chan string, 4)
|
|
||||||
go func() {
|
|
||||||
defer close(sentenceCh)
|
|
||||||
var buf strings.Builder
|
|
||||||
var fullText strings.Builder
|
|
||||||
for chunk := range llmStream {
|
|
||||||
// 即时推送文字给客户端(逐 token 显示)
|
|
||||||
client.SendLLMChunk(req.RequestID, chunk.Delta)
|
|
||||||
fullText.WriteString(chunk.Delta)
|
|
||||||
buf.WriteString(chunk.Delta)
|
|
||||||
// 遇到句子边界就吐出
|
|
||||||
if isSentenceEnd(chunk.Delta) {
|
|
||||||
sentenceCh <- buf.String()
|
|
||||||
buf.Reset()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
// 最后一段不足一句的也吐出
|
|
||||||
if buf.Len() > 0 {
|
|
||||||
sentenceCh <- buf.String()
|
|
||||||
}
|
|
||||||
// 推送 llm_done
|
|
||||||
client.SendLLMDone(req.RequestID, fullText.String(), chunk.TokensUsed, chunk.Model)
|
|
||||||
}()
|
|
||||||
|
|
||||||
// Step 3: TTS 并行消费句子流
|
|
||||||
ttsStream, _ := o.tts.SynthesizeStream(ctx, sentenceCh, TTSOptions{
|
|
||||||
Voice: "alloy", OutputFmt: "mp3", SampleRate: 24000,
|
|
||||||
})
|
|
||||||
for chunk := range ttsStream {
|
|
||||||
client.SendTTSAudio(req.RequestID, chunk.Audio, chunk.IsLast)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// isSentenceEnd 判断 delta 中是否包含句子结束标志。
|
|
||||||
func isSentenceEnd(delta string) bool {
|
|
||||||
return strings.ContainsAny(delta, "。!?\n.!?\n")
|
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
**Pipeline 实现**(`internal/orchestrator/pipeline.go`):
|
||||||
|
1. Base64 解码音频/图片
|
||||||
|
2. 调用 `stt.Recognize()` → 发送 `stt_result`
|
||||||
|
3. 调用 `llm.ChatStream()` 获取流式输出,goroutine 消费 token → 发送 `llm_chunk` + 句子切分
|
||||||
|
4. 另一 goroutine 从句子 channel 读取 → 调用 `tts.SynthesizeStream()` → 发送 `tts_audio`
|
||||||
|
5. 流结束 → 发送 `llm_done`
|
||||||
|
6. TTS 失败静默跳过,STT/LLM 失败发送对应 error 消息
|
||||||
|
|
||||||
### 并发控制
|
### 并发控制
|
||||||
|
|
||||||
- 每个 `ProcessQuery` 调用在独立 goroutine 中运行
|
- 每个 `ProcessQuery` 调用在独立 goroutine 中运行
|
||||||
@@ -584,9 +532,9 @@ session:{id}:history → List (对话历史)
|
|||||||
### 接口定义
|
### 接口定义
|
||||||
|
|
||||||
```go
|
```go
|
||||||
// SessionManager 会话管理器。
|
// Manager 会话管理器接口。
|
||||||
// WebSocket Handler 通过此接口操作会话,不直接接触 Redis。
|
// WebSocket Handler 通过此接口操作会话,不直接接触存储层。
|
||||||
type SessionManager interface {
|
type Manager interface {
|
||||||
// Create 创建新会话,返回 session ID。
|
// Create 创建新会话,返回 session ID。
|
||||||
Create(ctx context.Context, config models.SessionConfig) (string, error)
|
Create(ctx context.Context, config models.SessionConfig) (string, error)
|
||||||
|
|
||||||
@@ -605,14 +553,20 @@ type SessionManager interface {
|
|||||||
// SetActiveRequest 标记当前正在处理的请求 ID(interrupt 用)。
|
// SetActiveRequest 标记当前正在处理的请求 ID(interrupt 用)。
|
||||||
SetActiveRequest(ctx context.Context, sessionID string, requestID string) error
|
SetActiveRequest(ctx context.Context, sessionID string, requestID string) error
|
||||||
|
|
||||||
|
// GetActiveRequestID 获取当前活跃请求 ID。
|
||||||
|
GetActiveRequestID(ctx context.Context, sessionID string) (string, error)
|
||||||
|
|
||||||
// ClearActiveRequest 清除活跃请求标记(请求完成或中断后)。
|
// ClearActiveRequest 清除活跃请求标记(请求完成或中断后)。
|
||||||
ClearActiveRequest(ctx context.Context, sessionID string) error
|
ClearActiveRequest(ctx context.Context, sessionID string) error
|
||||||
|
|
||||||
// Touch 刷新 TTL(心跳时调用)。
|
// Touch 刷新 TTL(心跳时调用)。
|
||||||
Touch(ctx context.Context, sessionID string) error
|
Touch(ctx context.Context, sessionID string) error
|
||||||
|
|
||||||
// Destroy 显式销毁会话(REST API DELETE 或连接断开清理)。
|
// Destroy 显式销毁会话(REST API DELETE)。
|
||||||
Destroy(ctx context.Context, sessionID string) error
|
Destroy(ctx context.Context, sessionID string) error
|
||||||
|
|
||||||
|
// ActiveCount 返回当前活跃会话数(健康检查用)。
|
||||||
|
ActiveCount() int
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
@@ -629,7 +583,8 @@ case "query":
|
|||||||
|
|
||||||
history, _ := sessionMgr.GetHistory(ctx, sessionID, 20) // 获取对话上下文
|
history, _ := sessionMgr.GetHistory(ctx, sessionID, 20) // 获取对话上下文
|
||||||
|
|
||||||
go orchestrator.ProcessQuery(ctx, client, &msg, history) // 异步编排
|
sender := &WSClient{client: client, requestID: msg.RequestID}
|
||||||
|
go orch.ProcessQuery(ctx, sessionID, msg, history, sender) // 异步编排
|
||||||
|
|
||||||
// interrupt 分支
|
// interrupt 分支
|
||||||
case "interrupt":
|
case "interrupt":
|
||||||
@@ -648,26 +603,30 @@ case "interrupt":
|
|||||||
联调阶段无 Redis 时,用同一接口的内存实现:
|
联调阶段无 Redis 时,用同一接口的内存实现:
|
||||||
|
|
||||||
```go
|
```go
|
||||||
type InMemorySessionManager struct {
|
type MemoryManager struct {
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
sessions map[string]*sessionEntry
|
sessions map[string]*sessionEntry
|
||||||
|
ttl time.Duration
|
||||||
|
maxHistory int
|
||||||
|
stopCleaner chan struct{}
|
||||||
}
|
}
|
||||||
|
|
||||||
type sessionEntry struct {
|
type sessionEntry struct {
|
||||||
session models.Session
|
session models.Session
|
||||||
history []models.Message
|
history []models.Message
|
||||||
activeReqID string
|
activeReqID string
|
||||||
|
lastActive time.Time
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
注入时根据配置切换:
|
注入时根据配置切换:
|
||||||
|
|
||||||
```go
|
```go
|
||||||
var sessionMgr SessionManager
|
var sessionMgr session.Manager
|
||||||
if cfg.Redis.Addr != "" {
|
if cfg.Redis.Addr != "" {
|
||||||
sessionMgr = NewRedisSessionManager(redisClient, 30*time.Minute, 20)
|
sessionMgr = session.NewRedisManager(redisClient, 30*time.Minute, 20)
|
||||||
} else {
|
} else {
|
||||||
sessionMgr = NewInMemorySessionManager()
|
sessionMgr = session.NewMemoryManager(30*time.Minute, 20)
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
@@ -933,8 +892,6 @@ type QueryRequest struct {
|
|||||||
type Message struct {
|
type Message struct {
|
||||||
Role string `json:"role"` // "user" | "assistant"
|
Role string `json:"role"` // "user" | "assistant"
|
||||||
Content string `json:"content"`
|
Content string `json:"content"`
|
||||||
ImageURL string `json:"image_url,omitempty"`
|
|
||||||
TokensUsed int `json:"tokens_used,omitempty"`
|
|
||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user