From 7e1102902d6e5fc0398ec8c6a099809aed6322c4 Mon Sep 17 00:00:00 2001 From: sun <3371392206@qq.com> Date: Tue, 30 Jun 2026 23:34:45 +0800 Subject: [PATCH] =?UTF-8?q?refactor(volcano=20adapter):=20=E9=87=8D?= =?UTF-8?q?=E6=9E=84=E7=81=AB=E5=B1=B1=20TTS=20=E9=80=82=E9=85=8D=E5=99=A8?= =?UTF-8?q?=E8=AF=B7=E6=B1=82=E4=BD=93=E7=BB=93=E6=9E=84=E4=B8=8E=E9=94=99?= =?UTF-8?q?=E8=AF=AF=E5=A4=84=E7=90=86?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 重构请求体为类型安全的结构体实现,替换原有的map动态构造方式;统一错误处理为fmt.Errorf包装原始错误,优化注释与代码格式,提升代码可维护性与可读性。 --- adapter/volcano/volcano.go | 122 ++++++++++++++++++++++--------------- 1 file changed, 73 insertions(+), 49 deletions(-) diff --git a/adapter/volcano/volcano.go b/adapter/volcano/volcano.go index 5a1cfc1..d51ed30 100644 --- a/adapter/volcano/volcano.go +++ b/adapter/volcano/volcano.go @@ -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 倍速。 // 输入超出范围会被截断到边界值。 @@ -88,13 +119,13 @@ func buildWavHeader(dataLen int, sampleRate int) []byte { binary.LittleEndian.PutUint32(header[4:8], uint32(36+dataLen)) copy(header[8:12], "WAVE") copy(header[12:16], "fmt ") - binary.LittleEndian.PutUint32(header[16:20], 16) // SubChunk1Size - binary.LittleEndian.PutUint16(header[20:22], 1) // PCM format - binary.LittleEndian.PutUint16(header[22:24], 1) // NumChannels + binary.LittleEndian.PutUint32(header[16:20], 16) // SubChunk1Size + binary.LittleEndian.PutUint16(header[20:22], 1) // PCM format + binary.LittleEndian.PutUint16(header[22:24], 1) // NumChannels binary.LittleEndian.PutUint32(header[24:28], uint32(sampleRate)) binary.LittleEndian.PutUint32(header[28:32], uint32(byteRate)) 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") binary.LittleEndian.PutUint32(header[40:44], uint32(dataLen)) @@ -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 字段 + "Content-Type": "application/json", + // 请求用量返回,合成结束时响应中携带 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))