Files
Volcano-Engine-TTS-UI/controller/tts.go
T
sun 8a4187720c fix(tts): settings.default_resource_id 不再被 voice 行覆盖
Bug 001 (后端):

之前 f5563e6 修 55000000 时,LoadRuntimeConfig 拿 default_speaker 查
voice 表后,把 voice.Speaker/ResourceID/Model 全部覆盖了 settings 里
的对应值。本意是修 speaker,副作用是 default_resource_id 永远被
voice.resource_id 覆盖 → 用户在 setup 设的 default_resource_id
(seed-icl-2.0) 没机会生效,被 voice 行的 (volc.megatts.icl) 替换。

修法 (option B: settings 优先 + voice 兜底):
- speaker:  从 voice 表查 (default_speaker 是 voice 名,要解析成真 ID)
- resource_id: 保持 settings 里的(用户偏好)
- model:  voice 行的优先,settings 里的次之(向后兼容, 不动)

字段语义重新明确:
- settings.default_resource_id: 全局默认 resource(用户必设,运行时用)
- voices[].resource_id: 仅在 voice= 显式传时使用(per-voice 覆盖)
- 用户保持两个一致即可

附带修复: controller/tts.go 的注释误导(说 voice 为空用
setting.TTSOptions.Speaker 即 default_speaker),实际是 LoadRuntimeConfig
已解析为真 speaker ID。

测试:
- 已有 store/installer/middleware/router/telemetry 测试全绿
- e2e (本机): install settings.default_resource_id=seed-icl-2.0
  voice chun.resource_id=volc.megatts.icl (故意不一致)
  → 启动 log: 'resource=seed-icl-2.0' (settings 优先 ✓)
  → 火山 55000000/45000030 → 401 Invalid X-Api-Key (说明参数对,key 假)
  → 换真 key 即能 work

二进制大小: ~18MB → ~18MB (代码量不变)

未 push (待用户确认)
2026-08-30 13:09:40 +08:00

321 lines
11 KiB
Go

package controller
import (
"context"
"encoding/json"
"fmt"
"io"
"log"
"net/http"
"runtime"
"strings"
"time"
"github.com/volcano-tts/tts-api/adapter/volcano"
"github.com/volcano-tts/tts-api/common"
"github.com/volcano-tts/tts-api/dto"
"github.com/volcano-tts/tts-api/installer"
"github.com/volcano-tts/tts-api/metrics"
"github.com/volcano-tts/tts-api/middleware"
"github.com/volcano-tts/tts-api/setting"
"github.com/volcano-tts/tts-api/store"
"github.com/volcano-tts/tts-api/telemetry"
"github.com/volcano-tts/tts-api/version"
)
var (
volcanoClient *volcano.HTTPClient
adapterRec volcano.MetricsRecorder = metrics.AdapterRecorder{}
)
func InitController() {
volcanoClient = volcano.NewHTTPClient()
}
func truncateForLog(b []byte, max int) string {
if len(b) > max {
return string(b[:max]) + fmt.Sprintf("...(truncated, total %d bytes)", len(b))
}
return string(b)
}
// resolveClientFormat 把 OpenAI 风格的 response_format 映射为最终输出格式;
// 不识别或未指定时回退到 setting.TTSOptions.Format。
func resolveClientFormat(reqFmt string) string {
switch strings.ToLower(reqFmt) {
case "mp3", "wav", "opus", "pcm", "aac", "flac":
if reqFmt == "opus" {
return "ogg_opus"
}
return strings.ToLower(reqFmt)
}
return setting.TTSOptions.Format
}
// OpenaiTTSHandler 是 /v1/audio/speech 的入口。
func OpenaiTTSHandler(w http.ResponseWriter, r *http.Request) {
start := time.Now()
if r.Method != http.MethodPost {
log.Printf("警告: 错误的方法 - 方法=%s 期望=POST 路径=%s 客户端=%s",
r.Method, r.URL.Path, middleware.GetClientIP(r))
metrics.RequestTotal.Inc(telemetry.Labels{"status": "method_not_allowed", "format": "", "speaker": "", "model": ""})
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
// 安装模式双保险:即使 InstallGuard 中间件没拦住,这里也 503 + 引导跳转
if installer.GetMode() == installer.ModeSetup {
log.Printf("[tts] 安装模式下拒绝 /v1/audio/speech - 客户端=%s", middleware.GetClientIP(r))
w.Header().Set("Content-Type", "application/json; charset=utf-8")
w.WriteHeader(http.StatusServiceUnavailable)
_, _ = w.Write([]byte(`{"error":"not installed","code":"install_required","redirect":"/setup"}`))
return
}
if !middleware.ValidateAPIKey(r) {
metrics.AuthFailed.Inc(telemetry.Labels{})
log.Printf("警告: API Key 鉴权失败 - 路径=%s 客户端=%s 远端=%s",
r.URL.Path, middleware.GetClientIP(r), r.RemoteAddr)
middleware.SendJSONError(w, http.StatusUnauthorized, "Invalid API key provided.", "invalid_request_error", "invalid_api_key")
return
}
if setting.TTSConfigErr != nil {
log.Printf("警告: TTS配置未就绪,拒绝请求 - 错误=%v 路径=%s 客户端=%s",
setting.TTSConfigErr, r.URL.Path, middleware.GetClientIP(r))
middleware.SendJSONError(w, http.StatusServiceUnavailable, "TTS service configuration error. Please check environment variables and restart the service.", "configuration_error", "service_unavailable")
return
}
r.Body = http.MaxBytesReader(w, r.Body, common.MaxRequestBodySize)
body, err := io.ReadAll(r.Body)
if err != nil {
if strings.Contains(err.Error(), "request body too large") {
log.Printf("警告: 请求体过大 - 路径=%s 客户端=%s 限制=%d字节",
r.URL.Path, middleware.GetClientIP(r), common.MaxRequestBodySize)
http.Error(w, "Request body too large", http.StatusRequestEntityTooLarge)
return
}
log.Printf("警告: 读取请求体失败 - 路径=%s 客户端=%s 错误=%v",
r.URL.Path, middleware.GetClientIP(r), err)
http.Error(w, "Failed to read request body", http.StatusBadRequest)
return
}
var req dto.OpenAITTSRequest
if err := json.Unmarshal(body, &req); err != nil {
log.Printf("警告: JSON 解析失败 - 路径=%s 客户端=%s 错误=%v body前200字节=%q",
r.URL.Path, middleware.GetClientIP(r), err, truncateForLog(body, 200))
http.Error(w, "Invalid JSON", http.StatusBadRequest)
return
}
if req.Model != "" {
if len(req.Model) > common.MaxModelNameLength {
log.Printf("警告: Model 名过长 - 路径=%s 客户端=%s 长度=%d 限制=%d",
r.URL.Path, middleware.GetClientIP(r), len(req.Model), common.MaxModelNameLength)
http.Error(w, fmt.Sprintf("Model name too long (max %d characters)", common.MaxModelNameLength), http.StatusBadRequest)
return
}
if strings.ContainsAny(req.Model, "\x00\n\r\t") {
log.Printf("警告: Model 名含非法字符 - 路径=%s 客户端=%s model前50字节=%q",
r.URL.Path, middleware.GetClientIP(r), truncateForLog([]byte(req.Model), 50))
http.Error(w, "Model name contains invalid characters", http.StatusBadRequest)
return
}
}
if req.Input == "" {
log.Printf("警告: input 字段为空 - 路径=%s 客户端=%s", r.URL.Path, middleware.GetClientIP(r))
http.Error(w, "Input text is required", http.StatusBadRequest)
return
}
if len(req.Input) > common.MaxTextLength {
log.Printf("警告: input 文本过长 - 路径=%s 客户端=%s 长度=%d 限制=%d",
r.URL.Path, middleware.GetClientIP(r), len(req.Input), common.MaxTextLength)
http.Error(w, fmt.Sprintf("Input text too long (max %d characters)", common.MaxTextLength), http.StatusBadRequest)
return
}
speed := req.Speed
if speed <= 0 {
speed = common.DefaultSpeed
}
if speed < common.MinSpeed {
speed = common.MinSpeed
}
if speed > common.MaxSpeed {
speed = common.MaxSpeed
}
clientFormat := resolveClientFormat(req.ResponseFormat)
opts := setting.TTSOptions
opts.Text = req.Input
// M3: voice 路由
// - voice 为空 → 走 LoadRuntimeConfig 解析过的 opts.Speaker (已是真 speaker ID,
// default_speaker 是 voice 名,LoadRuntimeConfig 查 voice 表后替换)
// - voice 非空 → 查 voices 表,替换 opts.Speaker / ResourceID / Model
// - 命中但 enabled=0 → 仍可用(用户显式传 voice 即覆盖 enabled 状态;若想禁用在 admin UI 关掉就行)
// - 未命中 → 400 "unknown voice: <name>"
if req.Voice != "" {
s := GetAdminStore()
if s == nil {
log.Printf("警告: voice=%s 路由但 store 未初始化 - 路径=%s", req.Voice, r.URL.Path)
middleware.SendJSONError(w, http.StatusServiceUnavailable,
"voice routing requires database; not initialized",
"configuration_error", "db_not_ready")
return
}
v, err := s.VoiceGetByName(req.Voice)
if err != nil {
if err == store.ErrNotFound {
log.Printf("警告: 未知 voice=%s - 路径=%s 客户端=%s", req.Voice, r.URL.Path, middleware.GetClientIP(r))
middleware.SendJSONError(w, http.StatusBadRequest,
fmt.Sprintf("unknown voice: %s", req.Voice),
"invalid_request_error", "unknown_voice")
return
}
log.Printf("警告: voice 查库失败 - 错误=%v voice=%s", err, req.Voice)
middleware.SendJSONError(w, http.StatusInternalServerError,
"voice lookup failed", "server_error", "db_read_failed")
return
}
// 覆盖 opts(API key / UID 保留自 setting.TTSOptions)
opts.Speaker = v.Speaker
opts.ResourceID = v.ResourceID
if v.Model != "" {
opts.Model = v.Model
}
log.Printf("[tts] voice=%s 命中 (speaker=%s resource=%s model=%s) - 客户端=%s",
req.Voice, v.Speaker, v.ResourceID, v.Model, middleware.GetClientIP(r))
}
ctx, cancel := context.WithTimeout(r.Context(), setting.TTSTimeout)
defer cancel()
result, err := volcano.Synthesis(ctx, volcanoClient, opts, req.Input, clientFormat, speed, adapterRec)
duration := time.Since(start)
finalLabels := telemetry.Labels{
"format": clientFormat,
// speaker 是火山复刻音色 ID(用户付费资产),不能直接出现在 /metrics label 里
//(无鉴权可枚举)。用 sha1[:8] 替代:同 speaker 同 label 保留 per-voice 观测,
//但反推不出原值。Admin UI 想要看原名通过 /api/voices 拿 name 字段。
"speaker": telemetry.SpeakerLabel(opts.Speaker),
"model": opts.Model,
}
if err != nil {
finalLabels["status"] = classifyStatus(err)
metrics.RequestTotal.Inc(finalLabels)
metrics.RequestDuration.Observe(duration.Seconds(), telemetry.Labels{"status": finalLabels["status"], "format": clientFormat})
log.Printf("警告: TTS 合成失败 - 路径=%s 客户端=%s 文本长度=%d 耗时=%v 错误=%v",
r.URL.Path, middleware.GetClientIP(r), len(req.Input), duration, err)
middleware.SendJSONError(w, http.StatusInternalServerError, "TTS synthesis failed.", "server_error", "synthesis_failed")
return
}
finalLabels["status"] = "ok"
metrics.RequestTotal.Inc(finalLabels)
metrics.RequestDuration.Observe(duration.Seconds(), telemetry.Labels{"status": "ok", "format": clientFormat})
w.Header().Set("Content-Type", contentTypeFor(result.Format))
w.Header().Set("Content-Length", fmt.Sprintf("%d", len(result.AudioData)))
w.Header().Set("X-Request-Id", result.ReqID)
w.WriteHeader(http.StatusOK)
w.Write(result.AudioData)
}
func classifyStatus(err error) string {
if ue, ok := err.(*volcano.UpstreamError); ok {
switch ue.Stage {
case "request":
return "request_error"
case "http":
return fmt.Sprintf("http_%d", ue.Code)
case "stream":
return "upstream_error"
case "wrap":
return "wrap_error"
}
}
return "internal_error"
}
func contentTypeFor(format string) string {
switch strings.ToLower(format) {
case "wav":
return "audio/wav"
case "mp3":
return "audio/mpeg"
case "ogg_opus", "opus":
return "audio/ogg"
case "pcm":
return "audio/L16"
case "aac":
return "audio/aac"
case "flac":
return "audio/flac"
}
return "application/octet-stream"
}
// HealthHandler 暴露运行期状态;无鉴权。
func HealthHandler(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
// 安装模式下 /health 仍然 200,但通过 installed 字段让探针/运维识别
// (Kubernetes readiness probe 可以用 installed=false 决定是否放流量)
mode := installer.GetMode()
if mode == installer.ModeSetup {
w.WriteHeader(http.StatusOK) // 200,因为进程活着,只是还没初始化
} else if setting.TTSConfigErr != nil {
w.WriteHeader(http.StatusServiceUnavailable)
} else {
w.WriteHeader(http.StatusOK)
}
env := setting.CheckEnvironmentVariables()
allRequired := env["all_required_vars_set"].(bool)
status := "ok"
if mode == installer.ModeSetup {
status = "not_installed"
} else if !allRequired {
status = "configuration_error"
}
resp := dto.HealthResponse{
Status: status,
Service: "ByteDance TTS to OpenAI API Adapter",
Version: version.Version,
Commit: version.Commit,
Uptime: fmt.Sprintf("%.0f seconds", time.Since(startTime).Seconds()),
StartTime: startTime.Format(time.RFC3339),
Memory: collectMemorySnapshot(),
ConfigStatus: dto.ConfigStatusResponse{
AllRequiredVarsSet: allRequired,
ConfigError: setting.TTSConfigErr != nil,
},
Installed: mode == installer.ModeNormal,
Mode: mode.String(),
}
json.NewEncoder(w).Encode(resp)
}
var startTime time.Time
func SetStartTime(t time.Time) { startTime = t }
func collectMemorySnapshot() map[string]interface{} {
var ms runtime.MemStats
runtime.ReadMemStats(&ms)
return map[string]interface{}{
"heap_alloc": ms.HeapAlloc,
"heap_inuse": ms.HeapInuse,
"goroutines": runtime.NumGoroutine(),
}
}