From 0b2bbb827df48e5dabfb97d96c754505c6c8598b Mon Sep 17 00:00:00 2001 From: hhs <386998068@qq.com> Date: Sat, 13 Jun 2026 16:37:51 +0800 Subject: [PATCH 1/2] =?UTF-8?q?feat:=20Phase=208.1=20-=20WS=20=E9=9B=86?= =?UTF-8?q?=E6=88=90=E6=B5=8B=E8=AF=95=20(handler=5Ftest.go)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 使用 Gin test server + gorilla/websocket client 验证完整 WebSocket 流程: - connected 连接建立、ping/pong 心跳 - 完整 query→stt_result→llm_chunk→tts_audio→llm_done 流程 - interrupt 中断、disconnect 断开清理 - 无效 JSON、未知消息类型错误处理 - 多次查询、无图片查询、TTS 未启用等场景 - 每次连接创建独立会话 --- backend/internal/ws/handler_test.go | 563 ++++++++++++++++++++++++++++ 1 file changed, 563 insertions(+) create mode 100644 backend/internal/ws/handler_test.go diff --git a/backend/internal/ws/handler_test.go b/backend/internal/ws/handler_test.go new file mode 100644 index 0000000..97cf697 --- /dev/null +++ b/backend/internal/ws/handler_test.go @@ -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, "不应有额外消息") +} -- 2.49.1 From 9ce9d8c9f662f4356eda48a640e5c730686179f6 Mon Sep 17 00:00:00 2001 From: hhs <386998068@qq.com> Date: Sat, 13 Jun 2026 16:40:29 +0800 Subject: [PATCH 2/2] =?UTF-8?q?docs:=20Phase=208.2=20-=20=E5=90=8C?= =?UTF-8?q?=E6=AD=A5=E6=8E=A5=E5=8F=A3=E6=96=87=E6=A1=A3=E4=B8=8E=E4=BB=A3?= =?UTF-8?q?=E7=A0=81=E5=AE=9E=E7=8E=B0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 修正文档与代码的偏差: - SessionManager → Manager,补充 GetActiveRequestID/ActiveCount 方法 - Orchestrator 更新为接口 + Sender 抽象模式 - POST /api/sessions 移除未实现的 user_id 字段 - STT 移除未实现的 RecognizeStream 方法 - LLM/TTS 接口名更新为 Service/Request/Chunk/Options - Message 模型移除 ImageURL/TokensUsed(代码未使用) - 内存实现 MemoryManager 名称对齐 - WebSocket Handler 集成伪代码更新 ProcessQuery 签名 --- docs/03-接口文档.md | 179 +++++++++++++++++--------------------------- 1 file changed, 68 insertions(+), 111 deletions(-) diff --git a/docs/03-接口文档.md b/docs/03-接口文档.md index 948dba2..f1b7e1c 100644 --- a/docs/03-接口文档.md +++ b/docs/03-接口文档.md @@ -256,7 +256,6 @@ POST /api/sessions Content-Type: application/json { - "user_id": "optional-user-id", "config": { "tts_enabled": true, "detail_level": "low", @@ -301,28 +300,18 @@ Go 网关内部与外部 AI 服务(Deepgram STT、GPT-4o、OpenAI TTS)的调 语音识别:接收前端采集的音频,返回识别文本。 ```go -// STTService 语音识别服务契约。 -type STTService interface { +// Service 语音识别服务契约。 +type Service interface { // Recognize 识别一段完整音频,返回最终文本。 - Recognize(ctx context.Context, audio []byte, opts STTOptions) (string, error) - - // RecognizeStream 流式识别(边说边识别,可选实现)。 - // audioStream 持续接收音频片段,返回的 channel 持续输出中间结果。 - RecognizeStream(ctx context.Context, audioStream <-chan []byte, opts STTOptions) (<-chan STTPartial, error) + Recognize(ctx context.Context, audio []byte, opts Options) (string, error) } -// STTOptions 语音识别参数。 -type STTOptions struct { +// Options 语音识别参数。 +type Options struct { Encoding string // "pcm_s16le" — 前端 VAD 输出格式 SampleRate int // 16000 — 前端麦克风采样率 Language string // "zh-CN" } - -// STTPartial 流式识别的中间/最终结果。 -type STTPartial struct { - Text string - IsFinal bool -} ``` **Deepgram 接入约定**: @@ -336,23 +325,23 @@ type STTPartial struct { 多模态推理:接收图像 + 文本 + 对话历史,流式返回回复。 ```go -// LLMService 多模态大模型服务契约。 -type LLMService interface { +// Service 多模态大模型服务契约。 +type Service interface { // ChatStream 流式推理,返回增量文本的 channel。 // 调用方必须消费 channel 直到 Done=true,否则需 cancel ctx 以释放连接。 - ChatStream(ctx context.Context, req LLMRequest) (<-chan LLMChunk, error) + ChatStream(ctx context.Context, req Request) (<-chan Chunk, error) } -// LLMRequest 推理请求。 -type LLMRequest struct { - Image []byte // JPEG 图片(已从 Base64 解码) - Text string // 用户语音识别后的文本 - History []Message // 最近 N 轮对话历史 - Language string // "zh-CN" +// Request 推理请求。 +type Request struct { + Image []byte // JPEG 图片(已从 Base64 解码) + Text string // 用户语音识别后的文本 + History []models.Message // 最近 N 轮对话历史 + Language string // "zh-CN" } -// LLMChunk 流式推理的一个增量片段。 -type LLMChunk struct { +// Chunk 流式推理的一个增量片段。 +type Chunk struct { Delta string // 增量文本 Done bool // 是否结束 TokensUsed *TokenUsage // 仅 Done=true 时有值 @@ -388,24 +377,24 @@ user: [图片 + 用户语音文本] 语音合成:接收文本流,输出音频 chunk 流。 ```go -// TTSService 语音合成服务契约。 -type TTSService interface { +// Service 语音合成服务契约。 +type Service interface { // SynthesizeStream 流式合成。 // textStream 接收句子级文本(由 Orchestrator 的句子切分器产出), // 返回的 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 合成参数。 -type TTSOptions struct { +// Options 合成参数。 +type Options struct { Voice string // "alloy" | "nova" | "shimmer" | ... Speed float64 // 1.0 为正常语速 OutputFmt string // "mp3" — 固定使用 MP3,浏览器原生支持 SampleRate int // 24000 } -// TTSChunk 一个音频片段。 -type TTSChunk struct { +// Chunk 一个音频片段。 +type Chunk struct { Audio []byte // MP3 音频数据(未 Base64 编码,由发送层编码) IsLast bool // 是否为最后一片 } @@ -446,73 +435,32 @@ LLM 流式输出: "这" "是一" "朵红色" "的花。" "它看起" "来很美 ### Orchestrator 接口 ```go -// Orchestrator AI 编排器,协调 STT → LLM → TTS 全链路。 -type Orchestrator struct { - stt STTService - llm LLMService - tts TTSService +// Orchestrator AI 编排器接口。 +type Orchestrator interface { + // ProcessQuery 处理一次完整的视觉对话请求。 + // 通过 sender 向前端实时推送 stt_result、llm_chunk、llm_done、tts_audio 消息。 + ProcessQuery(ctx context.Context, sessionID string, req models.WsQuery, + history []models.Message, sender Sender) error } -// ProcessQuery 处理一次完整的视觉对话请求。 -// 通过 client 向前端实时推送 stt_result、llm_chunk、llm_done、tts_audio 消息。 -func (o *Orchestrator) ProcessQuery(ctx context.Context, client MessageSender, req *QueryRequest) { - ctx, cancel := context.WithTimeout(ctx, 10*time.Second) - defer cancel() - - // Step 1: STT — 识别用户语音 - 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") +// Sender 抽象 WebSocket 消息推送能力,便于测试时 mock。 +type Sender interface { + SendSTTResult(result models.WsSTTResult) error + SendLLMChunk(chunk models.WsLLMChunk) error + SendLLMDone(done models.WsLLMDone) error + SendTTSAudio(audio models.WsTTSAudio) error + SendError(err models.WsError) error } ``` +**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 中运行 @@ -584,9 +532,9 @@ session:{id}:history → List (对话历史) ### 接口定义 ```go -// SessionManager 会话管理器。 -// WebSocket Handler 通过此接口操作会话,不直接接触 Redis。 -type SessionManager interface { +// Manager 会话管理器接口。 +// WebSocket Handler 通过此接口操作会话,不直接接触存储层。 +type Manager interface { // Create 创建新会话,返回 session ID。 Create(ctx context.Context, config models.SessionConfig) (string, error) @@ -605,14 +553,20 @@ type SessionManager interface { // SetActiveRequest 标记当前正在处理的请求 ID(interrupt 用)。 SetActiveRequest(ctx context.Context, sessionID string, requestID string) error + // GetActiveRequestID 获取当前活跃请求 ID。 + GetActiveRequestID(ctx context.Context, sessionID string) (string, error) + // ClearActiveRequest 清除活跃请求标记(请求完成或中断后)。 ClearActiveRequest(ctx context.Context, sessionID string) error // Touch 刷新 TTL(心跳时调用)。 Touch(ctx context.Context, sessionID string) error - // Destroy 显式销毁会话(REST API DELETE 或连接断开清理)。 + // Destroy 显式销毁会话(REST API DELETE)。 Destroy(ctx context.Context, sessionID string) error + + // ActiveCount 返回当前活跃会话数(健康检查用)。 + ActiveCount() int } ``` @@ -629,7 +583,8 @@ case "query": 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 分支 case "interrupt": @@ -648,26 +603,30 @@ case "interrupt": 联调阶段无 Redis 时,用同一接口的内存实现: ```go -type InMemorySessionManager struct { - mu sync.RWMutex - sessions map[string]*sessionEntry +type MemoryManager struct { + mu sync.RWMutex + sessions map[string]*sessionEntry + ttl time.Duration + maxHistory int + stopCleaner chan struct{} } type sessionEntry struct { session models.Session history []models.Message activeReqID string + lastActive time.Time } ``` 注入时根据配置切换: ```go -var sessionMgr SessionManager +var sessionMgr session.Manager if cfg.Redis.Addr != "" { - sessionMgr = NewRedisSessionManager(redisClient, 30*time.Minute, 20) + sessionMgr = session.NewRedisManager(redisClient, 30*time.Minute, 20) } else { - sessionMgr = NewInMemorySessionManager() + sessionMgr = session.NewMemoryManager(30*time.Minute, 20) } ``` @@ -931,10 +890,8 @@ type QueryRequest struct { } type Message struct { - Role string `json:"role"` // "user" | "assistant" - Content string `json:"content"` - ImageURL string `json:"image_url,omitempty"` - TokensUsed int `json:"tokens_used,omitempty"` + Role string `json:"role"` // "user" | "assistant" + Content string `json:"content"` } ``` -- 2.49.1