Files
CamTalk/backend/internal/orchestrator/pipeline.go

404 lines
10 KiB
Go
Raw Normal View History

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 语音识别
2026-06-14 19:54:22 +08:00
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 {
2026-06-14 19:54:22 +08:00
log.Errorw("语音识别失败", "error", err, "audio_bytes", len(audio))
sender.SendError(models.WsError{
Type: "error",
RequestID: req.RequestID,
Code: "STT_ERROR",
2026-06-14 19:54:22 +08:00
Message: "语音识别失败: " + err.Error(),
})
return err
}
userText = sttResult
2026-06-14 19:54:22 +08:00
// 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 并行合成
2026-06-14 20:32:30 +08:00
log.Infow("开始 LLM 推理", "request_id", req.RequestID, "scenario", sess.Config.Scenario)
llmReq := llm.Request{
2026-06-14 20:32:30 +08:00
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,
2026-06-13 19:57:35 +08:00
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{
2026-06-14 11:12:49 +08:00
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
}