Files
Volcano-Engine-TTS-UI/controller/tts.go
T
tts-stage1 a614ab55c9 feat(route): 阶段 1 多渠道路由分发层 (Channel + Router)
- store: 新增 channels 表 (name unique, credentials_json, voices CSV,
  priority, weight, status, auto_ban) + CRUD (List/Get/Insert/SetStatus/Delete)
  + ErrChannelDuplicate 区分 (避免与 ErrDuplicate 互相误报)
- store/db.go: schemaVersion 1 -> 2; channels 表 CREATE IF NOT EXISTS + 复合索引
- adapter/route: 新包,Channel (含 BuildRequest 覆盖 template.Credentials)
  + Router (Select / SelectAndSynthesize / pickByWeight 纯函数)
  + 错误区分 ErrNoChannels (走兜底) vs ErrVoiceNotFound (400)
  + AggregateError 聚合多渠道失败, 供 controller 日志逐个打印
- controller/router.go: SetRouter/GetRouter 句柄 (与 SetAdminStore 同一模式)
- controller/tts.go: 接入 Router, 零 channel 走 setting 兜底 (行为=阶段 0);
  voice 不被任何 channel 接受 -> 400 + voice_not_found; 抽出 finalizeSynth
  收敛合成结果->响应+metrics+日志三件套
- main.go: 启动期从 store 加载 channels, 转 route.Channel, 注入 controller;
  credentials_json 加载失败 -> fail-fast (避免一个错渠道拖崩全部请求)
- .gitignore: 补 .gotmp/ .gomodcache/ .dsh-acl/

验收 (设计文档 §9):
  - 零 channel 兜底, 行为=阶段 0
  - 加 1 个 channel: 走该渠道
  - 加 2+ 个同档不同 weight: pickByWeight weight-1-to-3 分布 ±5% (单测覆盖)
  - 不同 priority: 高优先级优先 (单测覆盖)
  - disabled 不参与选择 (单测覆盖)
  - 失败降级: 同档换下一个 -> 降档, 全失败返 AggregateError (单测覆盖)
  - go build ./... + go vet ./... 全绿
  - go test ./... (adapter/route 15 个, store 2 个, 全部通过)
2026-10-11 15:33:22 +08:00

439 lines
17 KiB
Go

package controller
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"log"
"net/http"
"runtime"
"strings"
"time"
"github.com/volcano-tts/tts-api/adapter/provider"
"github.com/volcano-tts/tts-api/adapter/route"
"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 (
adapterRec provider.MetricsRecorder = metrics.AdapterRecorder{}
)
func InitController() {
// 上游 provider 在各自包 init() 里已注册(volcano 等),经 provider.Get 取用。
// 本函数保留以维持 main.go 的启动调用序列;不再持有具体 client。
}
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.GetTTSOptions().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.GetDefaultFormat()
}
// 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 err := setting.GetTTSConfigErr(); err != nil {
log.Printf("警告: TTS配置未就绪,拒绝请求 - 错误=%v 路径=%s 客户端=%s",
err, 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)
ttsReq := setting.GetTTSRequest()
ttsReq.Text = req.Input
ttsReq.Format = clientFormat
ttsReq.Speed = speed
// M3: voice 路由
// - voice 为空 → 走 LoadRuntimeConfig 解析过的默认音色(GetTTSRequest 已带)
// - voice 非空 → 查 voices 表,替换 VoiceKey / resource_id / Model
// - 命中但 enabled=0 → 拒绝;未命中 → 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=%q - 路径=%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
}
// 覆盖音色与厂商凭证(API key / 其它私有参数保留自 GetTTSRequest 快照)
if !v.Enabled {
log.Printf("警告: voice=%q 已禁用 - 客户端=%s", req.Voice, middleware.GetClientIP(r))
middleware.SendJSONError(w, http.StatusForbidden,
fmt.Sprintf("voice '%s' is disabled", req.Voice),
"invalid_request_error", "voice_disabled")
return
}
ttsReq.VoiceKey = v.Speaker
ttsReq.Credentials.Scope["resource_id"] = v.ResourceID
if v.Model != "" {
ttsReq.Model = v.Model
}
log.Printf("[tts] voice=%s 命中 (speaker=%s resource=%s model=%s) - 客户端=%s",
req.Voice, telemetry.MaskSpeaker(v.Speaker), telemetry.MaskResourceID(v.ResourceID), v.Model, middleware.GetClientIP(r))
}
ctx, cancel := context.WithTimeout(r.Context(), setting.GetTTSTimeout())
defer cancel()
// 阶段 1:多渠道路由分发
// - router 未注入(等于未配置 channels)→ 走 setting 兜底,行为=阶段 0
// - router 注入了但无任何 channel → 同样走 setting 兜底
// - router 注入了且有 channel → 走 SelectAndSynthesize:按 priority/weight 选
// 渠道,Channel.Credentials 整份覆盖 template.Credentials(意味着
// 上面 voice 路由对 Scope["resource_id"] 的覆盖会被 Channel 接管——
// 这是阶段 1 的设计取舍,Channel 自带完整账号凭据,voice 路由仅用于
// 解析上游 speaker ID → ttsReq.VoiceKey,凭据部分由 Channel 决定)
// - voice 不被任一 channel 接受 → router.Select 返 ErrVoiceNotFound,转 400
// - 所有 channel 都失败 → router 返 *route.AggregateError,转 500(日志逐个打印)
rtr := GetRouter()
if rtr == nil || rtr.Empty() {
prov, ok := provider.Get(setting.GetDefaultProviderName())
if !ok {
log.Printf("警告: 未注册上游适配器=%s - 客户端=%s", setting.GetDefaultProviderName(), middleware.GetClientIP(r))
middleware.SendJSONError(w, http.StatusServiceUnavailable,
"no upstream adapter registered", "configuration_error", "provider_unavailable")
return
}
result, err := prov.Synthesize(ctx, ttsReq, adapterRec)
duration := time.Since(start)
finalizeSynth(w, ttsReq, req.Input, clientFormat, result, err, duration, r.URL.Path)
return
}
// 选渠道路径:以"对外 voice 名"(req.Voice,可能为空表示用默认音色)为过滤键。
// 客户端未传 voice 时 req.Voice=="";Channel.Voices 接受一切(空白名单)时也会匹配。
// 这里需要先看 router 是否能接受这个 voice:
if _, selErr := rtr.Select(req.Voice); selErr != nil {
if selErr == route.ErrNoChannels {
// 中途被禁用 / 删空:降级到兜底,避免硬挂
prov, ok := provider.Get(setting.GetDefaultProviderName())
if !ok {
middleware.SendJSONError(w, http.StatusServiceUnavailable,
"no upstream adapter registered", "configuration_error", "provider_unavailable")
return
}
result, err := prov.Synthesize(ctx, ttsReq, adapterRec)
duration := time.Since(start)
finalizeSynth(w, ttsReq, req.Input, clientFormat, result, err, duration, r.URL.Path)
return
}
if errors.Is(selErr, route.ErrVoiceNotFound) {
log.Printf("警告: voice=%q 不被任何渠道支持 - 路径=%s 客户端=%s",
req.Voice, r.URL.Path, middleware.GetClientIP(r))
middleware.SendJSONError(w, http.StatusBadRequest,
fmt.Sprintf("voice '%s' is not supported by any channel", req.Voice),
"invalid_request_error", "voice_not_found")
return
}
log.Printf("警告: 渠道选择失败 - 错误=%v 路径=%s 客户端=%s",
selErr, r.URL.Path, middleware.GetClientIP(r))
middleware.SendJSONError(w, http.StatusInternalServerError,
"channel selection failed", "server_error", "route_error")
return
}
result, ch, err := rtr.SelectAndSynthesize(ctx, req.Voice, ttsReq, 3, adapterRec)
if err == nil && ch != nil {
log.Printf("[tts] 渠道选择命中 channel_id=%d name=%s provider=%s priority=%d weight=%d - 客户端=%s",
ch.ID, ch.Name, ch.Provider, ch.Priority, ch.Weight, middleware.GetClientIP(r))
}
duration := time.Since(start)
finalizeSynth(w, ttsReq, req.Input, clientFormat, result, err, duration, r.URL.Path)
}
// finalizeSynth 把"合成结果 → 响应 + metrics + 日志"三件套收敛到一处。
// 阶段 1 改造后,这个函数同时被 3 个分支调用:
// 1. 零 channel 兜底(直接调 setting 默认 provider)
// 2. Select 报 ErrNoChannels 降级(同上)
// 3. SelectAndSynthesize 完整路径
//
// 不再每次重复写 metrics/响应/日志,降低后续维护成本。
func finalizeSynth(
w http.ResponseWriter,
ttsReq provider.Request,
inputText string,
clientFormat string,
result *dto.SynthesisResult,
err error,
duration time.Duration,
urlPath string,
) {
finalLabels := telemetry.Labels{
"format": clientFormat,
// speaker 是火山复刻音色 ID(用户付费资产),不能直接出现在 /metrics label 里
//(无鉴权可枚举)。用 sha1[:8] 替代:同 speaker 同 label 保留 per-voice 观测,
//但反推不出原值。Admin UI 想要看原名通过 /api/voices 拿 name 字段。
"speaker": telemetry.SpeakerLabel(ttsReq.VoiceKey),
"model": ttsReq.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 文本长度=%d 耗时=%v 错误=%v",
urlPath, len(inputText), 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)
if n, err := w.Write(result.AudioData); err != nil {
// header 已发,无法改 status code;只记日志供排查(常见:客户端中途断开 → broken pipe / connection reset)
log.Printf("警告: 响应写入失败 - 路径=%s 已写=%d/%d 错误=%v",
urlPath, n, len(result.AudioData), err)
}
}
func classifyStatus(err error) string {
if ue, ok := err.(*provider.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 暴露运行期状态;无鉴权。
// HealthzHandler GET /healthz —— 匿名存活探针,**只回 200 与字面量 "ok"**。
//
// 为什么单独做这个:v0.3.0 把详细健康数据(/health)收口到管理鉴权之后,
// 但 K8s liveness/readiness、Docker HEALTHCHECK、负载均衡健康检查默认都不带 Authorization。
// 若把它们继续指向 /health,加鉴权后会一律 401,导致探针失败、Pod 反复重启。
//
// 因此本端点刻意**不返回任何字段**(无版本、无内存、无配置状态、无模式信息),
// 只用于回答"进程还在不在"。运维要细节请走鉴权后的 /health。
func HealthzHandler(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "text/plain; charset=utf-8")
w.Header().Set("Cache-Control", "no-store")
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte("ok"))
}
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.GetTTSConfigErr() != 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.GetTTSConfigErr() != nil,
Error: configErrorMessage(setting.GetTTSConfigErr()),
},
Installed: mode == installer.ModeNormal,
Mode: mode.String(),
}
json.NewEncoder(w).Encode(resp)
}
// configErrorMessage 把运行时配置错误(setting.GetTTSConfigErr())安全地转成可对外暴露的字符串。
// 仅在 normal 模式且有错时调用, error 为 nil 时返 "" (被 omitempty 跳过)。
func configErrorMessage(err error) string {
if err == nil {
return ""
}
return err.Error()
}
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(),
}
}