refactor(adapter): 抽出 provider 抽象层(阶段0),controller 不再依赖火山实现
- 新增 adapter/provider:Provider/Capabilities/Request/Credentials/MetricsRecorder/UpstreamError + 注册表 - 火山实现 Provider(provider.go),Synthesis 收敛为私有 synthesize,提供 BuildRequest 过渡映射 - controller 经 provider.Get(name).Synthesize 调用,voice 路由改在 Request 上覆盖 - UpstreamError 上移到 provider 包,火山用类型别名沿用旧名 - setting 新增 GetTTSRequest/GetDefaultFormat/GetDefaultProviderName - main.go 空导入注册火山 - docs/UPSTREAM_ADAPTER_GUIDE.md 更新为 v1.1(阶段0已落地)
This commit is contained in:
+27
-20
@@ -11,7 +11,7 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/volcano-tts/tts-api/adapter/volcano"
|
||||
"github.com/volcano-tts/tts-api/adapter/provider"
|
||||
"github.com/volcano-tts/tts-api/common"
|
||||
"github.com/volcano-tts/tts-api/dto"
|
||||
"github.com/volcano-tts/tts-api/installer"
|
||||
@@ -24,12 +24,12 @@ import (
|
||||
)
|
||||
|
||||
var (
|
||||
volcanoClient *volcano.HTTPClient
|
||||
adapterRec volcano.MetricsRecorder = metrics.AdapterRecorder{}
|
||||
adapterRec provider.MetricsRecorder = metrics.AdapterRecorder{}
|
||||
)
|
||||
|
||||
func InitController() {
|
||||
volcanoClient = volcano.NewHTTPClient()
|
||||
// 上游 provider 在各自包 init() 里已注册(volcano 等),经 provider.Get 取用。
|
||||
// 本函数保留以维持 main.go 的启动调用序列;不再持有具体 client。
|
||||
}
|
||||
|
||||
func truncateForLog(b []byte, max int) string {
|
||||
@@ -49,7 +49,7 @@ func resolveClientFormat(reqFmt string) string {
|
||||
}
|
||||
return strings.ToLower(reqFmt)
|
||||
}
|
||||
return setting.GetTTSOptions().Format
|
||||
return setting.GetDefaultFormat()
|
||||
}
|
||||
|
||||
// OpenaiTTSHandler 是 /v1/audio/speech 的入口。
|
||||
@@ -151,15 +151,15 @@ func OpenaiTTSHandler(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
clientFormat := resolveClientFormat(req.ResponseFormat)
|
||||
|
||||
opts := setting.GetTTSOptions()
|
||||
opts.Text = req.Input
|
||||
ttsReq := setting.GetTTSRequest()
|
||||
ttsReq.Text = req.Input
|
||||
ttsReq.Format = clientFormat
|
||||
ttsReq.Speed = speed
|
||||
|
||||
// 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>"
|
||||
// - voice 为空 → 走 LoadRuntimeConfig 解析过的默认音色(GetTTSRequest 已带)
|
||||
// - voice 非空 → 查 voices 表,替换 VoiceKey / resource_id / Model
|
||||
// - 命中但 enabled=0 → 拒绝;未命中 → 400 "unknown voice: <name>"
|
||||
if req.Voice != "" {
|
||||
s := GetAdminStore()
|
||||
if s == nil {
|
||||
@@ -183,7 +183,7 @@ func OpenaiTTSHandler(w http.ResponseWriter, r *http.Request) {
|
||||
"voice lookup failed", "server_error", "db_read_failed")
|
||||
return
|
||||
}
|
||||
// 覆盖 opts(API key / UID 保留自 setting.GetTTSOptions 快照)
|
||||
// 覆盖音色与厂商凭证(API key / 其它私有参数保留自 GetTTSRequest 快照)
|
||||
if !v.Enabled {
|
||||
log.Printf("警告: voice=%q 已禁用 - 客户端=%s", req.Voice, middleware.GetClientIP(r))
|
||||
middleware.SendJSONError(w, http.StatusForbidden,
|
||||
@@ -191,10 +191,10 @@ func OpenaiTTSHandler(w http.ResponseWriter, r *http.Request) {
|
||||
"invalid_request_error", "voice_disabled")
|
||||
return
|
||||
}
|
||||
opts.Speaker = v.Speaker
|
||||
opts.ResourceID = v.ResourceID
|
||||
ttsReq.VoiceKey = v.Speaker
|
||||
ttsReq.Credentials.Scope["resource_id"] = v.ResourceID
|
||||
if v.Model != "" {
|
||||
opts.Model = 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))
|
||||
@@ -203,7 +203,14 @@ func OpenaiTTSHandler(w http.ResponseWriter, r *http.Request) {
|
||||
ctx, cancel := context.WithTimeout(r.Context(), setting.GetTTSTimeout())
|
||||
defer cancel()
|
||||
|
||||
result, err := volcano.Synthesis(ctx, volcanoClient, opts, req.Input, clientFormat, speed, adapterRec)
|
||||
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)
|
||||
|
||||
finalLabels := telemetry.Labels{
|
||||
@@ -211,8 +218,8 @@ func OpenaiTTSHandler(w http.ResponseWriter, r *http.Request) {
|
||||
// speaker 是火山复刻音色 ID(用户付费资产),不能直接出现在 /metrics label 里
|
||||
//(无鉴权可枚举)。用 sha1[:8] 替代:同 speaker 同 label 保留 per-voice 观测,
|
||||
//但反推不出原值。Admin UI 想要看原名通过 /api/voices 拿 name 字段。
|
||||
"speaker": telemetry.SpeakerLabel(opts.Speaker),
|
||||
"model": opts.Model,
|
||||
"speaker": telemetry.SpeakerLabel(ttsReq.VoiceKey),
|
||||
"model": ttsReq.Model,
|
||||
}
|
||||
if err != nil {
|
||||
finalLabels["status"] = classifyStatus(err)
|
||||
@@ -240,7 +247,7 @@ func OpenaiTTSHandler(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
func classifyStatus(err error) string {
|
||||
if ue, ok := err.(*volcano.UpstreamError); ok {
|
||||
if ue, ok := err.(*provider.UpstreamError); ok {
|
||||
switch ue.Stage {
|
||||
case "request":
|
||||
return "request_error"
|
||||
|
||||
Reference in New Issue
Block a user