之前前端 TTSPlayer 攒齐所有音频片段后才播放,导致文字全部显示后才开始语音。 改为后端每句 TTS 发送 is_last: true,前端收到每句即加入播放队列,第一句到达即开始播放。 - 后端 Chunk 结构体新增 Final 字段,区分句子结束和整轮结束 - 前端 TTSPlayer 重写为队列式播放,onended 回调自动衔接下一句 - 同步更新接口文档和测试用例
379 lines
9.3 KiB
Go
379 lines
9.3 KiB
Go
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)
|
||
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)
|
||
sender.SendError(models.WsError{
|
||
Type: "error",
|
||
RequestID: req.RequestID,
|
||
Code: "STT_ERROR",
|
||
Message: "语音识别失败",
|
||
})
|
||
return err
|
||
}
|
||
userText = sttResult
|
||
|
||
// 发送 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)
|
||
llmReq := llm.Request{
|
||
Image: image,
|
||
Text: userText,
|
||
History: history,
|
||
Language: 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
|
||
}
|