refactor(volcano adapter): 重构火山 TTS 适配器请求体结构与错误处理
重构请求体为类型安全的结构体实现,替换原有的map动态构造方式;统一错误处理为fmt.Errorf包装原始错误,优化注释与代码格式,提升代码可维护性与可读性。
This commit is contained in:
+73
-49
@@ -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))
|
||||||
|
|||||||
Reference in New Issue
Block a user