package orchestrator import ( "context" "encoding/base64" "strings" "sync" "time" "unicode/utf8" "github.com/hhs/camtalk/internal/ai/llm" "github.com/hhs/camtalk/internal/ai/stt" "github.com/hhs/camtalk/internal/ai/tts" "github.com/hhs/camtalk/internal/config" "github.com/hhs/camtalk/internal/logger" "github.com/hhs/camtalk/internal/models" "github.com/hhs/camtalk/internal/session" ) // Pipeline 实现 Orchestrator 接口,管理 STT → LLM → TTS 流式管道。 type Pipeline struct { sttService stt.Service llmService llm.Service ttsService tts.Service sessionMgr session.Manager model string // LLM 模型名,用于 llm_done 上报 ttsVoice string // TTS 音色 ttsSpeed float64 // TTS 语速 ttsOutputFmt string // TTS 输出格式 ttsSampleRate int // TTS 输出采样率 } // New 创建 Pipeline 实例。 func New( sttService stt.Service, llmService llm.Service, ttsService tts.Service, sessionMgr session.Manager, cfg *config.Config, ) *Pipeline { return &Pipeline{ sttService: sttService, llmService: llmService, ttsService: ttsService, sessionMgr: sessionMgr, model: cfg.AI.LLM.Model, ttsVoice: cfg.AI.TTS.Voice, ttsSpeed: cfg.AI.TTS.Speed, ttsOutputFmt: cfg.AI.TTS.OutputFormat, ttsSampleRate: cfg.AI.TTS.SampleRate, } } // ProcessQuery 实现 Orchestrator 接口。 func (p *Pipeline) ProcessQuery( ctx context.Context, sessionID string, req models.WsQuery, history []models.Message, sender Sender, ) error { log := logger.Log startTime := time.Now() // 解码音频数据(文本输入模式可跳过) var audio []byte if req.Text == "" && req.Audio != "" { var err error audio, err = base64.StdEncoding.DecodeString(req.Audio) if err != nil { log.Errorw("音频解码失败", "error", err) sender.SendError(models.WsError{ Type: "error", RequestID: req.RequestID, Code: "INVALID_MESSAGE", Message: "音频数据解码失败", }) return err } } // 解码图片数据(可选) var image []byte if req.Image != "" { var err error image, err = base64.StdEncoding.DecodeString(req.Image) if err != nil { log.Errorw("图片解码失败", "error", err) sender.SendError(models.WsError{ Type: "error", RequestID: req.RequestID, Code: "INVALID_MESSAGE", Message: "图片数据解码失败", }) return err } } // 设置活跃请求 if err := p.sessionMgr.SetActiveRequest(ctx, sessionID, req.RequestID); err != nil { log.Errorw("设置活跃请求失败", "error", err) } defer p.sessionMgr.ClearActiveRequest(ctx, sessionID) // 获取会话配置 sess, err := p.sessionMgr.Get(ctx, sessionID) if err != nil { log.Errorw("获取会话失败", "error", err) sender.SendError(models.WsError{ Type: "error", RequestID: req.RequestID, Code: "SESSION_NOT_FOUND", Message: "会话不存在", }) return err } // Step 1: 获取用户文本(语音识别或直接使用输入文本) var userText string if req.Text != "" { // 文本输入模式:跳过 STT,直接使用用户输入的文本 log.Infow("使用文本输入", "request_id", req.RequestID, "text", req.Text) userText = req.Text // 发送 stt_result 以保持前端消息流一致性 if err := sender.SendSTTResult(models.WsSTTResult{ Type: "stt_result", RequestID: req.RequestID, Text: userText, IsFinal: true, }); err != nil { log.Errorw("发送 STT 结果失败", "error", err) } } else { // 语音模式:执行 STT 语音识别 log.Infow("开始语音识别", "request_id", req.RequestID, "audio_bytes", len(audio)) sttResult, err := p.sttService.Recognize(ctx, audio, stt.Options{ Encoding: "pcm_s16le", SampleRate: 16000, Language: sess.Config.Language, }) if err != nil { log.Errorw("语音识别失败", "error", err, "audio_bytes", len(audio)) sender.SendError(models.WsError{ Type: "error", RequestID: req.RequestID, Code: "STT_ERROR", Message: "语音识别失败: " + err.Error(), }) return err } userText = sttResult // STT 返回空文本:未识别到语音,发送结果后直接返回(不调 LLM) if strings.TrimSpace(userText) == "" { log.Infow("语音识别结果为空", "request_id", req.RequestID) userText = "(未识别到语音)" if err := sender.SendSTTResult(models.WsSTTResult{ Type: "stt_result", RequestID: req.RequestID, Text: userText, IsFinal: true, }); err != nil { log.Errorw("发送 STT 结果失败", "error", err) } // 发送空的 llm_done 以结束本轮处理 latency := time.Since(startTime).Milliseconds() _ = sender.SendLLMDone(models.WsLLMDone{ Type: "llm_done", RequestID: req.RequestID, FullText: "", Model: p.model, LatencyMs: latency, }) return nil } // 发送 STT 结果 if err := sender.SendSTTResult(models.WsSTTResult{ Type: "stt_result", RequestID: req.RequestID, Text: userText, IsFinal: true, }); err != nil { log.Errorw("发送 STT 结果失败", "error", err) } } // 追加用户消息到历史 p.sessionMgr.AppendMessage(ctx, sessionID, models.Message{ Role: "user", Content: userText, }) // Step 2+3: LLM 流式推理 + TTS 并行合成 log.Infow("开始 LLM 推理", "request_id", req.RequestID, "scenario", sess.Config.Scenario) llmReq := llm.Request{ Image: image, Text: userText, History: history, Language: sess.Config.Language, SystemPrompt: llm.GetScenarioPrompt(sess.Config.Scenario, sess.Config.Language), } llmStream, err := p.llmService.ChatStream(ctx, llmReq) if err != nil { log.Errorw("LLM 流式推理启动失败", "error", err) sender.SendError(models.WsError{ Type: "error", RequestID: req.RequestID, Code: "LLM_ERROR", Message: "LLM 推理失败", }) return err } // 创建句子切分器 sentenceCh := make(chan string, 4) splitter := NewSplitter(sentenceCh) // 并行:LLM 消费 + TTS 合成 var wg sync.WaitGroup var fullText string var ttsErr error // goroutine 1: 消费 LLM token + 句子切分 var tokenUsage *llm.TokenUsage wg.Add(1) go func() { defer wg.Done() defer close(sentenceCh) fullText, tokenUsage = p.consumeLLMStream(ctx, llmStream, req.RequestID, sender, splitter) }() // goroutine 2: TTS 合成(如果启用) if sess.Config.TTSEnabled { wg.Add(1) go func() { defer wg.Done() log.Infow("开始 TTS 合成", "request_id", req.RequestID) ttsErr = p.synthesizeTTS(ctx, sentenceCh, req.RequestID, sender) }() } else { // 如果 TTS 未启用,需要消费 sentenceCh 防止阻塞 go func() { for range sentenceCh { } }() } // 等待所有 goroutine 完成 wg.Wait() // TTS 失败静默跳过 if ttsErr != nil { log.Warnw("TTS 合成失败(已跳过)", "error", ttsErr) } // 追加助手消息到历史 p.sessionMgr.AppendMessage(ctx, sessionID, models.Message{ Role: "assistant", Content: fullText, }) // 发送 llm_done latency := time.Since(startTime).Milliseconds() done := models.WsLLMDone{ Type: "llm_done", RequestID: req.RequestID, FullText: fullText, Model: p.model, LatencyMs: latency, } if tokenUsage != nil { done.TokensUsed = struct { Prompt int `json:"prompt"` Completion int `json:"completion"` Total int `json:"total"` }{ Prompt: tokenUsage.Prompt, Completion: tokenUsage.Completion, Total: tokenUsage.Total, } } if err := sender.SendLLMDone(done); err != nil { log.Errorw("发送 llm_done 失败", "error", err) } log.Infow("查询处理完成", "request_id", req.RequestID, "latency_ms", latency, "text_length", utf8.RuneCountInString(fullText), ) return nil } // consumeLLMStream 消费 LLM 流式输出,发送 llm_chunk 并进行句子切分。 // 返回完整文本和 token 用量。 func (p *Pipeline) consumeLLMStream( ctx context.Context, stream <-chan llm.Chunk, requestID string, sender Sender, splitter *Splitter, ) (string, *llm.TokenUsage) { log := logger.Log var fullText strings.Builder var tokenUsage *llm.TokenUsage for chunk := range stream { // 检查上下文是否已取消 select { case <-ctx.Done(): log.Infow("LLM 流被中断", "request_id", requestID) return fullText.String(), tokenUsage default: } if chunk.Done { // 流结束,记录 token 用量 if chunk.TokensUsed != nil { tokenUsage = chunk.TokensUsed log.Infow("LLM 用量统计", "request_id", requestID, "prompt_tokens", tokenUsage.Prompt, "completion_tokens", tokenUsage.Completion, "total_tokens", tokenUsage.Total, ) } break } // 累积全文 fullText.WriteString(chunk.Delta) // 发送 llm_chunk if err := sender.SendLLMChunk(models.WsLLMChunk{ Type: "llm_chunk", RequestID: requestID, Delta: chunk.Delta, Role: "assistant", }); err != nil { log.Errorw("发送 llm_chunk 失败", "error", err) } // 句子切分 splitter.Feed(chunk.Delta) } // 刷新切分器中的剩余文本 splitter.Flush() return fullText.String(), tokenUsage } // synthesizeTTS 从句子 channel 读取文本,进行 TTS 合成并发送音频。 func (p *Pipeline) synthesizeTTS( ctx context.Context, sentenceCh <-chan string, requestID string, sender Sender, ) error { log := logger.Log ttsStream, err := p.ttsService.SynthesizeStream(ctx, sentenceCh, tts.Options{ Voice: p.ttsVoice, Speed: p.ttsSpeed, OutputFmt: p.ttsOutputFmt, SampleRate: p.ttsSampleRate, }) if err != nil { log.Errorw("TTS 合成启动失败", "error", err) return err } // 消费 TTS 音频流 for chunk := range ttsStream { // 检查上下文是否已取消 select { case <-ctx.Done(): log.Infow("TTS 流被中断", "request_id", requestID) return ctx.Err() default: } // Base64 编码音频数据 audioBase64 := base64.StdEncoding.EncodeToString(chunk.Audio) if err := sender.SendTTSAudio(models.WsTTSAudio{ Type: "tts_audio", RequestID: requestID, Audio: audioBase64, MimeType: "audio/mp3", IsLast: chunk.IsLast, Final: chunk.Final, }); err != nil { log.Errorw("发送 tts_audio 失败", "error", err) } } return nil }