refactor(volcano adapter): 重构火山 TTS 适配器请求体结构与错误处理
重构请求体为类型安全的结构体实现,替换原有的map动态构造方式;统一错误处理为fmt.Errorf包装原始错误,优化注释与代码格式,提升代码可维护性与可读性。
This commit is contained in:
+68
-44
@@ -17,6 +17,12 @@ import (
|
||||
"github.com/volcano-tts/tts-api/dto"
|
||||
)
|
||||
|
||||
// 火山引擎 TTS v3 HTTP 单向流式 API 客户端。
|
||||
// 官方文档:https://www.volcengine.com/docs/6561/1598757
|
||||
// 官方 Go 示例:请求体仅含 req_params;本实现按 commit 4aed966 经验额外带上
|
||||
// user.uid 和 namespace="UnidirectionalTTS"(早期用其它 namespace 出现过兼容性
|
||||
// 问题,显式指定最稳)。复刻音色场景额外带 req_params.model。
|
||||
|
||||
type HTTPClient struct {
|
||||
client *http.Client
|
||||
}
|
||||
@@ -50,6 +56,31 @@ func (h *HTTPClient) PostStream(url string, headers map[string]string, body []by
|
||||
return h.client.Do(req)
|
||||
}
|
||||
|
||||
// --- 请求体结构(对应火山 v3 API 请求 JSON) ---
|
||||
|
||||
type ttsRequest struct {
|
||||
User ttsUser `json:"user"`
|
||||
Namespace string `json:"namespace"`
|
||||
ReqParams ttsReqParams `json:"req_params"`
|
||||
}
|
||||
|
||||
type ttsUser struct {
|
||||
UID string `json:"uid"`
|
||||
}
|
||||
|
||||
type ttsReqParams struct {
|
||||
Text string `json:"text"`
|
||||
Speaker string `json:"speaker"`
|
||||
Model string `json:"model"`
|
||||
AudioParams ttsAudioParams `json:"audio_params"`
|
||||
}
|
||||
|
||||
type ttsAudioParams struct {
|
||||
Format string `json:"format"`
|
||||
SampleRate int `json:"sample_rate"`
|
||||
SpeechRate int `json:"speech_rate"`
|
||||
}
|
||||
|
||||
// convertSpeedToSpeechRate 把 OpenAI 风格的 speed(倍率)转成火山 v3 的 speech_rate(百分比)。
|
||||
// 文档规定 speech_rate 范围 [-50, 100],对应 0.5x ~ 2.0x 倍速。
|
||||
// 输入超出范围会被截断到边界值。
|
||||
@@ -148,10 +179,10 @@ func Synthesis(config *dto.ByteDanceTTSConfig, httpClient *HTTPClient, text stri
|
||||
|
||||
model := config.Model
|
||||
if model == "" {
|
||||
model = "seed-tts-2.0-standard"
|
||||
model = "seed-tts-2.0-standard" // 文档默认值,复刻音色可设为 seed-tts-2.0-expressive
|
||||
}
|
||||
|
||||
// 决定实际输出格式:优先用请求中指定的格式,否则用配置中的格式,最后默认 mp3
|
||||
// 决定实际输出格式:优先用请求中指定的格式,否则用配置中的格式,最后默认 mp3
|
||||
outputFormat := config.Format
|
||||
if requestFormat != "" {
|
||||
outputFormat = requestFormat
|
||||
@@ -160,7 +191,7 @@ func Synthesis(config *dto.ByteDanceTTSConfig, httpClient *HTTPClient, text stri
|
||||
outputFormat = "mp3"
|
||||
}
|
||||
|
||||
// 根据输出格式确定 API 请求格式(wav → pcm + 封装 header)
|
||||
// 根据输出格式确定 API 请求格式(wav → pcm + 本端封装 header)
|
||||
apiFormat, needWavHeader := resolveAPIFormat(outputFormat)
|
||||
|
||||
sampleRate := config.SampleRate
|
||||
@@ -168,61 +199,56 @@ func Synthesis(config *dto.ByteDanceTTSConfig, httpClient *HTTPClient, text stri
|
||||
sampleRate = 24000
|
||||
}
|
||||
|
||||
// 按火山 v3 HTTP Chunked 单向流式 API 文档构造请求体
|
||||
// https://www.volcengine.com/docs/6561/1598757
|
||||
params := map[string]interface{}{
|
||||
"user": map[string]interface{}{
|
||||
"uid": reqID,
|
||||
},
|
||||
"namespace": "UnidirectionalTTS",
|
||||
"req_params": map[string]interface{}{
|
||||
"text": text,
|
||||
"speaker": speaker,
|
||||
"model": model,
|
||||
"audio_params": map[string]interface{}{
|
||||
"format": apiFormat,
|
||||
"sample_rate": sampleRate,
|
||||
"speech_rate": speechRate,
|
||||
// 构造请求体:严格按 v3 API 文档 JSON 结构
|
||||
req := ttsRequest{
|
||||
User: ttsUser{UID: reqID},
|
||||
Namespace: "UnidirectionalTTS",
|
||||
ReqParams: ttsReqParams{
|
||||
Text: text,
|
||||
Speaker: speaker,
|
||||
Model: model,
|
||||
AudioParams: ttsAudioParams{
|
||||
Format: apiFormat,
|
||||
SampleRate: sampleRate,
|
||||
SpeechRate: speechRate,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
// 鉴权 header 按文档 v3 新版控制台方式
|
||||
body, err := json.Marshal(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("marshal TTS request: %w", err)
|
||||
}
|
||||
|
||||
// 鉴权 header 按 v3 新版控制台方式(Connection 由 Go http 默认 keep-alive)
|
||||
headers := map[string]string{
|
||||
"Content-Type": "application/json",
|
||||
"Connection": "keep-alive",
|
||||
"X-Api-Resource-Id": config.ResourceId,
|
||||
"X-Api-Request-Id": reqID,
|
||||
"X-Api-Key": config.ApiKey,
|
||||
// 请求用量返回,使合成结束时携带 usage 字段
|
||||
// 请求用量返回,合成结束时响应中携带 usage 字段
|
||||
"X-Control-Require-Usage-Tokens-Return": "*",
|
||||
"X-Api-Resource-Id": config.ResourceId, // 模型路由(seed-tts-2.0 / seed-icl-2.0)
|
||||
"X-Api-Request-Id": reqID,
|
||||
"X-Api-Key": config.ApiKey, // v3 鉴权 key
|
||||
}
|
||||
|
||||
bodyStr, err := json.Marshal(params)
|
||||
resp, err := httpClient.PostStream(config.URL, headers, body, config.Timeout)
|
||||
if err != nil {
|
||||
log.Printf("JSON marshal fail: %v", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
resp, err := httpClient.PostStream(config.URL, headers, bodyStr, config.Timeout)
|
||||
if err != nil {
|
||||
log.Printf("http post fail: %v", err)
|
||||
return nil, err
|
||||
return nil, fmt.Errorf("send TTS request: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
log.Printf("Failed to read error response body: %v", err)
|
||||
respBody, readErr := io.ReadAll(resp.Body)
|
||||
if readErr != nil {
|
||||
log.Printf("TTS service error: status=%d, read body fail: %v", resp.StatusCode, readErr)
|
||||
} else {
|
||||
log.Printf("TTS service error: status=%d, body=%s", resp.StatusCode, string(body))
|
||||
log.Printf("TTS service error: status=%d, body=%s", resp.StatusCode, string(respBody))
|
||||
}
|
||||
return nil, fmt.Errorf("TTS service error: status %d", resp.StatusCode)
|
||||
}
|
||||
|
||||
var audioData []byte
|
||||
scanner := bufio.NewScanner(resp.Body)
|
||||
// 初始 1MB / 最大 8MB,与示例同量级,留足 TTS 长文本 room
|
||||
scanner.Buffer(make([]byte, 1024*1024), 8*1024*1024)
|
||||
|
||||
for scanner.Scan() {
|
||||
@@ -261,12 +287,11 @@ func Synthesis(config *dto.ByteDanceTTSConfig, httpClient *HTTPClient, text stri
|
||||
case "TTSSentenceEnd":
|
||||
log.Printf("Sentence end: sequence=%d", v3Resp.Sequence)
|
||||
default:
|
||||
// 音频数据 chunk:data 字段为 base64 编码的音频片段
|
||||
// 音频数据 chunk:data 字段为 base64 编码的音频片段
|
||||
if v3Resp.Data != "" {
|
||||
chunk, err := base64.StdEncoding.DecodeString(v3Resp.Data)
|
||||
if err != nil {
|
||||
log.Printf("base64 decode fail: %v", err)
|
||||
return nil, err
|
||||
return nil, fmt.Errorf("decode audio chunk: %w", err)
|
||||
}
|
||||
audioData = append(audioData, chunk...)
|
||||
}
|
||||
@@ -274,15 +299,14 @@ func Synthesis(config *dto.ByteDanceTTSConfig, httpClient *HTTPClient, text stri
|
||||
}
|
||||
|
||||
if err := scanner.Err(); err != nil {
|
||||
log.Printf("read stream fail: %v", err)
|
||||
return nil, err
|
||||
return nil, fmt.Errorf("read TTS stream: %w", err)
|
||||
}
|
||||
|
||||
if len(audioData) == 0 {
|
||||
return nil, fmt.Errorf("no audio data received")
|
||||
return nil, fmt.Errorf("no audio data received from TTS service")
|
||||
}
|
||||
|
||||
// 若输出格式为 wav,需要在 pcm 数据前拼装完整的 wav header
|
||||
// 若输出格式为 wav,需要在 pcm 数据前拼装完整的 wav header
|
||||
if needWavHeader {
|
||||
wavHeader := buildWavHeader(len(audioData), sampleRate)
|
||||
wavData := make([]byte, 0, len(wavHeader)+len(audioData))
|
||||
|
||||
Reference in New Issue
Block a user