diff --git a/backend/internal/eino/graph.go b/backend/internal/eino/graph.go index f8a2449..8bf3a13 100644 --- a/backend/internal/eino/graph.go +++ b/backend/internal/eino/graph.go @@ -18,8 +18,9 @@ import ( const ( nodeSTT = "stt" nodeHistory = "history" - nodeLLM = "llm" - nodeSplitter = "splitter" + nodeLLM = "llm" + nodeMessageToString = "msg2str" + nodeSplitter = "splitter" nodeTTS = "tts" nodeDone = "done" ) @@ -69,6 +70,7 @@ func NewPipelineGraph( _ = g.AddLambdaNode(nodeSTT, NewSTTLambda(sttService)) _ = g.AddLambdaNode(nodeHistory, NewHistoryLambda(sessionMgr.GetHistory, maxHistory)) _ = g.AddChatModelNode(nodeLLM, chatModel) + _ = g.AddLambdaNode(nodeMessageToString, NewMessageToStringLambda()) _ = g.AddLambdaNode(nodeSplitter, NewSplitterLambda()) _ = g.AddLambdaNode(nodeTTS, NewTTSLambda( ttsService, @@ -83,7 +85,8 @@ func NewPipelineGraph( _ = g.AddEdge(compose.START, nodeSTT) _ = g.AddEdge(nodeSTT, nodeHistory) _ = g.AddEdge(nodeHistory, nodeLLM) - _ = g.AddEdge(nodeLLM, nodeSplitter) + _ = g.AddEdge(nodeLLM, nodeMessageToString) + _ = g.AddEdge(nodeMessageToString, nodeSplitter) _ = g.AddEdge(nodeSplitter, nodeTTS) _ = g.AddEdge(nodeTTS, nodeDone) _ = g.AddEdge(nodeDone, compose.END) diff --git a/backend/internal/eino/nodes_history.go b/backend/internal/eino/nodes_history.go index a558a6b..c48b488 100644 --- a/backend/internal/eino/nodes_history.go +++ b/backend/internal/eino/nodes_history.go @@ -40,30 +40,12 @@ func NewHistoryLambda(historyFetcher func(ctx context.Context, sessionID string, scenarioPrompt := llm.GetScenarioPrompt(scenario, language) systemPrompt := llm.BuildSystemPrompt(language, detailLevel, scenarioPrompt) - // 构建 system message(含图片) + // 构建 system message(仅文本,多模态内容只能放在 user 角色) systemMsg := &schema.Message{ Role: schema.System, Content: systemPrompt, } - // 如果有图片,添加到 system message 的多模态内容中 - if len(imageData) > 0 { - base64Str := base64.StdEncoding.EncodeToString(imageData) - mimeType := detectImageMimeType(imageData) - systemMsg.UserInputMultiContent = []schema.MessageInputPart{ - { - Type: schema.ChatMessagePartTypeImageURL, - Image: &schema.MessageInputImage{ - MessagePartCommon: schema.MessagePartCommon{ - Base64Data: &base64Str, - MIMEType: mimeType, - }, - Detail: schema.ImageURLDetailAuto, - }, - }, - } - } - messages := []*schema.Message{systemMsg} // 获取并追加历史消息 @@ -81,11 +63,37 @@ func NewHistoryLambda(historyFetcher func(ctx context.Context, sessionID string, } } - // 追加当前用户输入 - messages = append(messages, &schema.Message{ - Role: schema.User, - Content: sttOut.Text, - }) + // 追加当前用户输入(含图片,多模态内容只能放在 user 角色) + // 注意:不能同时设置 Content 和 UserInputMultiContent,需要统一放到 MultiContent 中 + if len(imageData) > 0 { + base64Str := base64.StdEncoding.EncodeToString(imageData) + mimeType := detectImageMimeType(imageData) + parts := []schema.MessageInputPart{ + { + Type: schema.ChatMessagePartTypeText, + Text: sttOut.Text, + }, + { + Type: schema.ChatMessagePartTypeImageURL, + Image: &schema.MessageInputImage{ + MessagePartCommon: schema.MessagePartCommon{ + Base64Data: &base64Str, + MIMEType: mimeType, + }, + Detail: schema.ImageURLDetailAuto, + }, + }, + } + messages = append(messages, &schema.Message{ + Role: schema.User, + UserInputMultiContent: parts, + }) + } else { + messages = append(messages, &schema.Message{ + Role: schema.User, + Content: sttOut.Text, + }) + } log.Infow("历史组装完成", "request_id", requestID, diff --git a/backend/internal/eino/nodes_splitter.go b/backend/internal/eino/nodes_splitter.go index 6f3de8b..e3cfafb 100644 --- a/backend/internal/eino/nodes_splitter.go +++ b/backend/internal/eino/nodes_splitter.go @@ -20,18 +20,49 @@ var sentenceDelimiters = map[rune]bool{ '?': true, } -// NewSplitterLambda 创建句子分割 Transform Lambda 节点。 -// 输入: StreamReader[string](LLM 完整文本的单帧流)→ 输出: StreamReader[[]string](句子数组流) +// NewMessageToStringLambda 创建 Message → String 转换 Lambda 节点。 +// 输入: *schema.Message → 输出: string // -// 在 Stream 模式下,框架自动将 ChatModel 的 StreamReader[*schema.Message] -// concat 为 string 后传入此节点。此节点将文本按句子边界切分, -// 每切出一个句子就输出一次,供 TTS 节点实时合成。 -func NewSplitterLambda() *compose.Lambda { - return compose.TransformableLambda(func(ctx context.Context, input *schema.StreamReader[string]) (*schema.StreamReader[[]string], error) { - sr, sw := schema.Pipe[[]string](8) +// 提取 Message.Content 文本,供 Splitter 节点消费。 +func NewMessageToStringLambda() *compose.Lambda { + return compose.TransformableLambda(func(ctx context.Context, input *schema.StreamReader[*schema.Message]) (*schema.StreamReader[string], error) { + sr, sw := schema.Pipe[string](8) go func() { defer sw.Close() + defer input.Close() + + for { + msg, err := input.Recv() + if err != nil { + if err == io.EOF { + return + } + sw.Send("", err) + return + } + if msg != nil && msg.Content != "" { + sw.Send(msg.Content, nil) + } + } + }() + + return sr, nil + }) +} + +// NewSplitterLambda 创建句子分割 Transform Lambda 节点。 +// 输入: StreamReader[string](LLM token 流)→ 输出: StreamReader[string](完整句子流) +// +// 逐字符累积,按句子分隔符切分。每切出一个完整句子就输出一次, +// 供下游 TTS 节点实时合成。 +func NewSplitterLambda() *compose.Lambda { + return compose.TransformableLambda(func(ctx context.Context, input *schema.StreamReader[string]) (*schema.StreamReader[string], error) { + sr, sw := schema.Pipe[string](8) + + go func() { + defer sw.Close() + defer input.Close() var buffer strings.Builder @@ -43,23 +74,22 @@ func NewSplitterLambda() *compose.Lambda { if buffer.Len() > 0 { text := strings.TrimSpace(buffer.String()) if text != "" { - sw.Send([]string{text}, nil) + sw.Send(text, nil) } } return } - sw.Send(nil, err) + sw.Send("", err) return } - // chunk 是 concat 后的完整文本(单帧流) // 逐字符累积,按句子分隔符切分 for _, r := range chunk { buffer.WriteRune(r) if sentenceDelimiters[r] { text := strings.TrimSpace(buffer.String()) if text != "" { - sw.Send([]string{text}, nil) + sw.Send(text, nil) } buffer.Reset() } diff --git a/backend/internal/eino/nodes_tts.go b/backend/internal/eino/nodes_tts.go index 5d4e7b7..09c86d1 100644 --- a/backend/internal/eino/nodes_tts.go +++ b/backend/internal/eino/nodes_tts.go @@ -3,90 +3,114 @@ package eino import ( "context" "encoding/base64" + "io" "github.com/cloudwego/eino/compose" + "github.com/cloudwego/eino/schema" "github.com/hhs/camtalk/internal/ai/tts" "github.com/hhs/camtalk/internal/logger" "github.com/hhs/camtalk/internal/models" ) -// NewTTSLambda 创建 TTS Lambda 节点。 -// 输入: []string(句子数组,框架自动从 StreamReader concat)→ 输出: struct{} +// NewTTSLambda 创建 TTS Transform Lambda 节点。 +// 输入: StreamReader[string](句子流)→ 输出: StreamReader[struct{}](结果流) // -// 将句子数组转为 channel,调用 ttsService.SynthesizeStream() 流式合成, +// 流式消费每个句子,调用 ttsService.SynthesizeStream() 合成, // 逐 chunk 推送 tts_audio 到客户端。TTS 失败静默跳过。 func NewTTSLambda(ttsService tts.Service, ttsVoice string, ttsSpeed float64, ttsOutputFmt string, ttsSampleRate int) *compose.Lambda { - return compose.InvokableLambda(func(ctx context.Context, sentences []string) (struct{}, error) { - log := logger.Log - sender := senderFromCtx(ctx) - requestID := requestIDFromCtx(ctx) - state := stateFromCtx(ctx) + return compose.TransformableLambda(func(ctx context.Context, input *schema.StreamReader[string]) (*schema.StreamReader[struct{}], error) { + sr, sw := schema.Pipe[struct{}](8) - // 检查 TTS 是否启用(从 State 或 context 获取) - // TTSEnabled 信息在 PipelineInput 中,通过 State 传递 - if state != nil { - state.mu.Lock() - ttsEnabled := true // 默认启用,由适配器通过 State 设置 - state.mu.Unlock() - if !ttsEnabled { - return struct{}{}, nil - } - } + go func() { + defer sw.Close() + defer input.Close() - if len(sentences) == 0 { - return struct{}{}, nil - } + log := logger.Log + sender := senderFromCtx(ctx) + requestID := requestIDFromCtx(ctx) - if sender == nil || requestID == "" { - return struct{}{}, nil - } - - log.Infow("开始 TTS 合成", "request_id", requestID, "sentence_count", len(sentences)) - - // 将句子数组转为 channel(ttsService.SynthesizeStream 需要 <-chan string) - sentenceCh := make(chan string, len(sentences)) - for _, s := range sentences { - sentenceCh <- s - } - close(sentenceCh) - - // 调用 TTS 服务 - ttsStream, err := ttsService.SynthesizeStream(ctx, sentenceCh, tts.Options{ - Voice: ttsVoice, - Speed: ttsSpeed, - OutputFmt: ttsOutputFmt, - SampleRate: ttsSampleRate, - }) - if err != nil { - log.Errorw("TTS 合成启动失败(已跳过)", "error", err, "request_id", requestID) - return struct{}{}, nil // TTS 失败不中断流程 - } - - // 消费 TTS 音频流,推送到客户端 - for chunk := range ttsStream { - select { - case <-ctx.Done(): - log.Infow("TTS 流被中断", "request_id", requestID) - return struct{}{}, ctx.Err() - default: + if sender == nil || requestID == "" { + // 消费并丢弃流 + for { + _, err := input.Recv() + if err != nil { + return + } + } } - 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) + // 收集句子,按批次合成 TTS + var sentences []string + for { + sentence, err := input.Recv() + if err != nil { + if err == io.EOF { + break + } + log.Errorw("TTS: stream recv error", "error", err, "request_id", requestID) + break + } + if sentence != "" { + sentences = append(sentences, sentence) + } } - } - log.Infow("TTS 合成完成", "request_id", requestID) - return struct{}{}, nil + if len(sentences) == 0 { + sw.Send(struct{}{}, nil) + return + } + + log.Infow("开始 TTS 合成", "request_id", requestID, "sentence_count", len(sentences)) + + // 将句子数组转为 channel + sentenceCh := make(chan string, len(sentences)) + for _, s := range sentences { + sentenceCh <- s + } + close(sentenceCh) + + // 调用 TTS 服务 + ttsStream, err := ttsService.SynthesizeStream(ctx, sentenceCh, tts.Options{ + Voice: ttsVoice, + Speed: ttsSpeed, + OutputFmt: ttsOutputFmt, + SampleRate: ttsSampleRate, + }) + if err != nil { + log.Errorw("TTS 合成启动失败(已跳过)", "error", err, "request_id", requestID) + sw.Send(struct{}{}, nil) + return + } + + // 消费 TTS 音频流,推送到客户端 + for chunk := range ttsStream { + select { + case <-ctx.Done(): + log.Infow("TTS 流被中断", "request_id", requestID) + sw.Send(struct{}{}, ctx.Err()) + return + default: + } + + 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) + } + } + + log.Infow("TTS 合成完成", "request_id", requestID) + sw.Send(struct{}{}, nil) + }() + + return sr, nil }) }