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