Files
voice_test/tts.go
T

103 lines
2.6 KiB
Go

package main
import (
"context"
"encoding/json"
"fmt"
"sync"
"time"
"github.com/gorilla/websocket"
)
// ttsSpeak: 百炼 CosyVoice 流式合成,音频块实时回调 onAudio。
func ttsSpeak(ctx context.Context, cfg agentCfg, text string, onAudio func([]byte)) error {
cctx, cancel := context.WithTimeout(ctx, 60*time.Second)
defer cancel()
conn, _, err := websocket.DefaultDialer.DialContext(cctx, cfg.BailianWssBaseURL,
map[string][]string{"Authorization": {"Bearer " + cfg.BailianKey}})
if err != nil {
return fmt.Errorf("连接 TTS: %w", err)
}
defer conn.Close()
taskID := uuid()
err = conn.WriteJSON(map[string]any{
"header": map[string]string{"action": "run-task", "task_id": taskID, "streaming": "duplex"},
"payload": map[string]any{
"task_group": "audio", "task": "tts", "function": "SpeechSynthesizer",
"model": cfg.TTSModel,
"parameters": map[string]any{
"text_type": "PlainText", "voice": cfg.TTSVoice,
"format": "pcm", "sample_rate": 16000,
},
"input": map[string]any{},
},
})
if err != nil {
return fmt.Errorf("TTS run-task: %w", err)
}
var once sync.Once
done := make(chan error, 1)
var buf []json.RawMessage // task-started / task-finished 事件
go func() {
for {
mt, data, err := conn.ReadMessage()
if err != nil {
select {
case done <- err:
default:
}
return
}
if mt == websocket.BinaryMessage {
onAudio(data)
continue
}
var ev struct {
Header struct {
Event string `json:"event"`
ErrorMessage string `json:"error_message"`
} `json:"header"`
}
if json.Unmarshal(data, &ev) != nil {
continue
}
switch ev.Header.Event {
case "task-started":
buf = append(buf, json.RawMessage("{}"))
once.Do(func() {
// task-started 到达后才允许发文本
if err := conn.WriteJSON(map[string]any{
"header": map[string]string{"action": "continue-task", "task_id": taskID, "streaming": "duplex"},
"payload": map[string]any{"input": map[string]string{"text": text}},
}); err != nil {
select {
case done <- err:
default:
}
}
conn.WriteJSON(map[string]any{
"header": map[string]string{"action": "finish-task", "task_id": taskID, "streaming": "duplex"},
"payload": map[string]any{"input": map[string]any{}},
})
})
case "task-finished":
select {
case done <- nil:
default:
}
return
case "task-failed":
select {
case done <- fmt.Errorf("TTS 失败: %s", ev.Header.ErrorMessage):
default:
}
return
}
}
}()
return <-done
}