refactor(volcano adapter): 重构火山 TTS 适配器请求体结构与错误处理

重构请求体为类型安全的结构体实现,替换原有的map动态构造方式;统一错误处理为fmt.Errorf包装原始错误,优化注释与代码格式,提升代码可维护性与可读性。
This commit is contained in:
sun
2026-06-30 23:34:45 +08:00
parent 746da76fa4
commit 7e1102902d
+73 -49
View File
@@ -17,6 +17,12 @@ import (
"github.com/volcano-tts/tts-api/dto" "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 { type HTTPClient struct {
client *http.Client client *http.Client
} }
@@ -50,6 +56,31 @@ func (h *HTTPClient) PostStream(url string, headers map[string]string, body []by
return h.client.Do(req) 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(百分比)。 // convertSpeedToSpeechRate 把 OpenAI 风格的 speed(倍率)转成火山 v3 的 speech_rate(百分比)。
// 文档规定 speech_rate 范围 [-50, 100],对应 0.5x ~ 2.0x 倍速。 // 文档规定 speech_rate 范围 [-50, 100],对应 0.5x ~ 2.0x 倍速。
// 输入超出范围会被截断到边界值。 // 输入超出范围会被截断到边界值。
@@ -88,13 +119,13 @@ func buildWavHeader(dataLen int, sampleRate int) []byte {
binary.LittleEndian.PutUint32(header[4:8], uint32(36+dataLen)) binary.LittleEndian.PutUint32(header[4:8], uint32(36+dataLen))
copy(header[8:12], "WAVE") copy(header[8:12], "WAVE")
copy(header[12:16], "fmt ") copy(header[12:16], "fmt ")
binary.LittleEndian.PutUint32(header[16:20], 16) // SubChunk1Size binary.LittleEndian.PutUint32(header[16:20], 16) // SubChunk1Size
binary.LittleEndian.PutUint16(header[20:22], 1) // PCM format binary.LittleEndian.PutUint16(header[20:22], 1) // PCM format
binary.LittleEndian.PutUint16(header[22:24], 1) // NumChannels binary.LittleEndian.PutUint16(header[22:24], 1) // NumChannels
binary.LittleEndian.PutUint32(header[24:28], uint32(sampleRate)) binary.LittleEndian.PutUint32(header[24:28], uint32(sampleRate))
binary.LittleEndian.PutUint32(header[28:32], uint32(byteRate)) binary.LittleEndian.PutUint32(header[28:32], uint32(byteRate))
binary.LittleEndian.PutUint16(header[32:34], uint16(blockAlign)) binary.LittleEndian.PutUint16(header[32:34], uint16(blockAlign))
binary.LittleEndian.PutUint16(header[34:36], 16) // BitsPerSample binary.LittleEndian.PutUint16(header[34:36], 16) // BitsPerSample
copy(header[36:40], "data") copy(header[36:40], "data")
binary.LittleEndian.PutUint32(header[40:44], uint32(dataLen)) binary.LittleEndian.PutUint32(header[40:44], uint32(dataLen))
@@ -148,10 +179,10 @@ func Synthesis(config *dto.ByteDanceTTSConfig, httpClient *HTTPClient, text stri
model := config.Model model := config.Model
if 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 outputFormat := config.Format
if requestFormat != "" { if requestFormat != "" {
outputFormat = requestFormat outputFormat = requestFormat
@@ -160,7 +191,7 @@ func Synthesis(config *dto.ByteDanceTTSConfig, httpClient *HTTPClient, text stri
outputFormat = "mp3" outputFormat = "mp3"
} }
// 根据输出格式确定 API 请求格式(wav → pcm + 封装 header) // 根据输出格式确定 API 请求格式(wav → pcm + 本端封装 header)
apiFormat, needWavHeader := resolveAPIFormat(outputFormat) apiFormat, needWavHeader := resolveAPIFormat(outputFormat)
sampleRate := config.SampleRate sampleRate := config.SampleRate
@@ -168,61 +199,56 @@ func Synthesis(config *dto.ByteDanceTTSConfig, httpClient *HTTPClient, text stri
sampleRate = 24000 sampleRate = 24000
} }
// 按火山 v3 HTTP Chunked 单向流式 API 文档构造请求体 // 构造请求体:严格按 v3 API 文档 JSON 结构
// https://www.volcengine.com/docs/6561/1598757 req := ttsRequest{
params := map[string]interface{}{ User: ttsUser{UID: reqID},
"user": map[string]interface{}{ Namespace: "UnidirectionalTTS",
"uid": reqID, ReqParams: ttsReqParams{
}, Text: text,
"namespace": "UnidirectionalTTS", Speaker: speaker,
"req_params": map[string]interface{}{ Model: model,
"text": text, AudioParams: ttsAudioParams{
"speaker": speaker, Format: apiFormat,
"model": model, SampleRate: sampleRate,
"audio_params": map[string]interface{}{ SpeechRate: speechRate,
"format": apiFormat,
"sample_rate": sampleRate,
"speech_rate": 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{ headers := map[string]string{
"Content-Type": "application/json", "Content-Type": "application/json",
"Connection": "keep-alive", // 请求用量返回,合成结束时响应中携带 usage 字段
"X-Api-Resource-Id": config.ResourceId,
"X-Api-Request-Id": reqID,
"X-Api-Key": config.ApiKey,
// 请求用量返回,使合成结束时携带 usage 字段
"X-Control-Require-Usage-Tokens-Return": "*", "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 { if err != nil {
log.Printf("JSON marshal fail: %v", err) return nil, fmt.Errorf("send TTS request: %w", 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
} }
defer resp.Body.Close() defer resp.Body.Close()
if resp.StatusCode != http.StatusOK { if resp.StatusCode != http.StatusOK {
body, err := io.ReadAll(resp.Body) respBody, readErr := io.ReadAll(resp.Body)
if err != nil { if readErr != nil {
log.Printf("Failed to read error response body: %v", err) log.Printf("TTS service error: status=%d, read body fail: %v", resp.StatusCode, readErr)
} else { } 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) return nil, fmt.Errorf("TTS service error: status %d", resp.StatusCode)
} }
var audioData []byte var audioData []byte
scanner := bufio.NewScanner(resp.Body) scanner := bufio.NewScanner(resp.Body)
// 初始 1MB / 最大 8MB,与示例同量级,留足 TTS 长文本 room
scanner.Buffer(make([]byte, 1024*1024), 8*1024*1024) scanner.Buffer(make([]byte, 1024*1024), 8*1024*1024)
for scanner.Scan() { for scanner.Scan() {
@@ -261,12 +287,11 @@ func Synthesis(config *dto.ByteDanceTTSConfig, httpClient *HTTPClient, text stri
case "TTSSentenceEnd": case "TTSSentenceEnd":
log.Printf("Sentence end: sequence=%d", v3Resp.Sequence) log.Printf("Sentence end: sequence=%d", v3Resp.Sequence)
default: default:
// 音频数据 chunk:data 字段为 base64 编码的音频片段 // 音频数据 chunk:data 字段为 base64 编码的音频片段
if v3Resp.Data != "" { if v3Resp.Data != "" {
chunk, err := base64.StdEncoding.DecodeString(v3Resp.Data) chunk, err := base64.StdEncoding.DecodeString(v3Resp.Data)
if err != nil { if err != nil {
log.Printf("base64 decode fail: %v", err) return nil, fmt.Errorf("decode audio chunk: %w", err)
return nil, err
} }
audioData = append(audioData, chunk...) audioData = append(audioData, chunk...)
} }
@@ -274,15 +299,14 @@ func Synthesis(config *dto.ByteDanceTTSConfig, httpClient *HTTPClient, text stri
} }
if err := scanner.Err(); err != nil { if err := scanner.Err(); err != nil {
log.Printf("read stream fail: %v", err) return nil, fmt.Errorf("read TTS stream: %w", err)
return nil, err
} }
if len(audioData) == 0 { 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 { if needWavHeader {
wavHeader := buildWavHeader(len(audioData), sampleRate) wavHeader := buildWavHeader(len(audioData), sampleRate)
wavData := make([]byte, 0, len(wavHeader)+len(audioData)) wavData := make([]byte, 0, len(wavHeader)+len(audioData))