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 个, 全部通过)
This commit is contained in:
@@ -1,5 +1,7 @@
|
||||
# Go build cache
|
||||
.gocache/
|
||||
.gotmp/
|
||||
.gomodcache/
|
||||
*.exe
|
||||
*.test
|
||||
*.out
|
||||
@@ -49,3 +51,6 @@ tts.db
|
||||
tts.db-*
|
||||
tts.db.*
|
||||
installed.lock
|
||||
|
||||
# DSH sandbox ACL repair artifacts (created by diagnose-windows-sandbox-acl skill)
|
||||
.dsh-acl/
|
||||
|
||||
@@ -0,0 +1,108 @@
|
||||
// Package route 实现多渠道路由分发层(阶段 1)。
|
||||
//
|
||||
// 思路参考 QuantumNous/new-api 的 Channel + 路由选择,但按本项目
|
||||
// "一渠道 = 一上游账号" 简化:
|
||||
// - Channel: 一个上游账号(Provider 名 + 凭据 + 可选 voice 白名单)
|
||||
// - Router: 按优先级 + 加权随机从 channels 里挑一个给 Synthesize
|
||||
// - 不做 HTTP 管理接口(阶段 2 才上),channels 启动期从 store 加载一次
|
||||
package route
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/volcano-tts/tts-api/adapter/provider"
|
||||
"github.com/volcano-tts/tts-api/store"
|
||||
)
|
||||
|
||||
// Channel 是 router 用的渠道视图,从 store.Channel 反序列化凭据而来。
|
||||
//
|
||||
// 持有 store.Channel 字符串字段(VoicesCSV / CredentialsJSON)会让 router
|
||||
// 变成"半 DB 层",所以这里转成 router 友好的强类型:
|
||||
// - Voices []string: 白名单;空 = 不限
|
||||
// - Credentials: provider.Credentials 强类型
|
||||
type Channel struct {
|
||||
ID int64
|
||||
Name string
|
||||
Provider string // 对应 provider.Registry 里的 Provider.Name()
|
||||
Credentials provider.Credentials
|
||||
Voices []string // 对外 voice 名白名单;空 = 不限
|
||||
Priority int
|
||||
Weight int
|
||||
Status int // 1=enabled 2=disabled
|
||||
AutoBan bool
|
||||
}
|
||||
|
||||
// FromStoreChannel 把 store.Channel 转成 router.Channel。
|
||||
// 失败原因:
|
||||
// - credentials_json 不是合法 JSON → 返 error(让启动期 / Reload 直接报错,不要静默继续)
|
||||
func FromStoreChannel(s store.Channel) (Channel, error) {
|
||||
c := Channel{
|
||||
ID: s.ID,
|
||||
Name: s.Name,
|
||||
Provider: s.Provider,
|
||||
Priority: s.Priority,
|
||||
Weight: s.Weight,
|
||||
Status: s.Status,
|
||||
AutoBan: s.AutoBan,
|
||||
Voices: splitCSV(s.Voices),
|
||||
}
|
||||
if err := json.Unmarshal([]byte(s.CredentialsJSON), &c.Credentials); err != nil {
|
||||
return Channel{}, fmt.Errorf("channel %d (%s): credentials_json invalid: %w", s.ID, s.Name, err)
|
||||
}
|
||||
return c, nil
|
||||
}
|
||||
|
||||
// BuildRequest 拿一个上游请求模板(零 Channel 兜底用的),填入本渠道的
|
||||
// 凭据,作为 Synthesize 的入参。
|
||||
//
|
||||
// 不复制 Text / Format / Speed 等请求级字段(由调用方按本次请求填),
|
||||
// 这里只换"账号身份"——保证 router 调 Synthesize 时拿到的是"用本渠道账号调" 的 Request。
|
||||
func (c Channel) BuildRequest(template provider.Request) provider.Request {
|
||||
template.Credentials = c.Credentials
|
||||
return template
|
||||
}
|
||||
|
||||
// AcceptsVoice 检查本渠道是否支持该 voice:
|
||||
// - Voices 为空 → 不限,接受一切
|
||||
// - 非空 → 必须在白名单里
|
||||
func (c Channel) AcceptsVoice(voice string) bool {
|
||||
if len(c.Voices) == 0 {
|
||||
return true
|
||||
}
|
||||
for _, v := range c.Voices {
|
||||
if v == voice {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// Enabled 是不是启用状态(Status=1)。
|
||||
func (c Channel) Enabled() bool { return c.Status == 1 }
|
||||
|
||||
// splitCSV 拆 "a,b,c" 为 []string;容忍空 / 重复 / 空白。
|
||||
func splitCSV(raw string) []string {
|
||||
if raw == "" {
|
||||
return nil
|
||||
}
|
||||
parts := strings.Split(raw, ",")
|
||||
out := make([]string, 0, len(parts))
|
||||
seen := make(map[string]struct{}, len(parts))
|
||||
for _, p := range parts {
|
||||
p = strings.TrimSpace(p)
|
||||
if p == "" {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[p]; ok {
|
||||
continue
|
||||
}
|
||||
seen[p] = struct{}{}
|
||||
out = append(out, p)
|
||||
}
|
||||
if len(out) == 0 {
|
||||
return nil
|
||||
}
|
||||
return out
|
||||
}
|
||||
@@ -0,0 +1,339 @@
|
||||
package route
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"math/rand"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/volcano-tts/tts-api/adapter/provider"
|
||||
"github.com/volcano-tts/tts-api/dto"
|
||||
)
|
||||
|
||||
// Router 持有当前可用的渠道列表;并发安全(读多写少,RWMutex)。
|
||||
//
|
||||
// 阶段 1: 启动期从 store 加载一次,运行期不变;阶段 2 管理接口上线后
|
||||
// 会有 Reload() 触发整体替换。
|
||||
type Router struct {
|
||||
mu sync.RWMutex
|
||||
channels []Channel
|
||||
}
|
||||
|
||||
// NewRouter 从 store 风格的 channel 列表构造 Router;空列表合法(代表"零渠道,走兜底")。
|
||||
// 这里接受 []Channel(而非 *store.Store)便于单测注入。
|
||||
func NewRouter(channels []Channel) *Router {
|
||||
// 拷贝入参,避免外部后续修改影响内部。
|
||||
cp := make([]Channel, len(channels))
|
||||
copy(cp, channels)
|
||||
return &Router{channels: cp}
|
||||
}
|
||||
|
||||
// Snapshot 返回当前 channels 的副本(只读);controller 热路径用它取列表。
|
||||
func (r *Router) Snapshot() []Channel {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
out := make([]Channel, len(r.channels))
|
||||
copy(out, r.channels)
|
||||
return out
|
||||
}
|
||||
|
||||
// Empty 没有任何启用渠道时返 true(让 controller 走 setting 兜底路径)。
|
||||
func (r *Router) Empty() bool {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
return len(r.channels) == 0
|
||||
}
|
||||
|
||||
// Select 在当前 channels 里按规则挑一个支持该 voice 的渠道。
|
||||
// 返回值:
|
||||
// - (*Channel, nil) 找到启用且 voice 匹配的渠道
|
||||
// - (nil, ErrNoChannels) router 里没有任何 channel(走兜底)
|
||||
// - (nil, ErrVoiceNotFound) 有 channel 但 voice 不被任一渠道支持(返 400)
|
||||
// - (nil, 其他 error) 内部错误
|
||||
//
|
||||
// 把"零 channel" 和 "voice 不匹配" 区分开,是为了让 controller 决定是兜底还是报 400:
|
||||
// 零 channel → 旧行为不变;有 channel 但 voice 不支持 → 明确告诉客户端这个 voice 不可用。
|
||||
func (r *Router) Select(voice string) (*Channel, error) {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
|
||||
if len(r.channels) == 0 {
|
||||
return nil, ErrNoChannels
|
||||
}
|
||||
|
||||
enabled := make([]Channel, 0, len(r.channels))
|
||||
for _, c := range r.channels {
|
||||
if c.Enabled() {
|
||||
enabled = append(enabled, c)
|
||||
}
|
||||
}
|
||||
if len(enabled) == 0 {
|
||||
return nil, ErrNoChannels
|
||||
}
|
||||
|
||||
// 优先:被 voice 接受;若无任何匹配 → ErrVoiceNotFound
|
||||
matched := make([]Channel, 0, len(enabled))
|
||||
for _, c := range enabled {
|
||||
if c.AcceptsVoice(voice) {
|
||||
matched = append(matched, c)
|
||||
}
|
||||
}
|
||||
if len(matched) == 0 {
|
||||
return nil, ErrVoiceNotFound
|
||||
}
|
||||
|
||||
// 在 matched 内按 priority 降序分桶,选最高档做加权随机。
|
||||
// 文档 §5:"同档换下一个"和"降档"在 Select 这一层不发生——Select 一次只返回一个 channel。
|
||||
// 多 channel 失败重试/降档由 SelectAndSynthesize 负责(见下)。
|
||||
top := topPriority(matched)
|
||||
ch, _ := pickByWeight(top, time.Now().UnixNano())
|
||||
return &ch, nil
|
||||
}
|
||||
|
||||
// SelectAndSynthesize 一次请求的完整"选渠道→调 Synthesize"循环:
|
||||
//
|
||||
// pool = 全部 enabled 且支持 voice 的 channel,按 priority DESC 排序
|
||||
// 循环 maxRetry 次:
|
||||
// 从 pool 当前最高档里按 weight 加权随机选一个
|
||||
// 从 pool 移除该 channel(防重复)
|
||||
// 调 Synthesize;成功 → 返回;失败 → 累计错误,继续
|
||||
// 全失败 → 返聚合错误
|
||||
//
|
||||
// 设计要点:
|
||||
// - "换下一个" = 同档里挑别的;pool 自动收紧
|
||||
// - "降档" = 最高档被试空 → pool 里只剩低档 → 自然降到低档
|
||||
// - 失败重试安全(TTS 是只读调用,无扣款副作用)
|
||||
func (r *Router) SelectAndSynthesize(
|
||||
ctx context.Context,
|
||||
voice string,
|
||||
template provider.Request,
|
||||
maxRetry int,
|
||||
mtr provider.MetricsRecorder,
|
||||
) (*dto.SynthesisResult, *Channel, error) {
|
||||
if maxRetry <= 0 {
|
||||
maxRetry = 3 // 与文档 §5 默认一致
|
||||
}
|
||||
if mtr == nil {
|
||||
mtr = provider.NoopMetrics()
|
||||
}
|
||||
|
||||
r.mu.RLock()
|
||||
enabled := r.enabledMatching(voice)
|
||||
r.mu.RUnlock()
|
||||
|
||||
if len(enabled) == 0 {
|
||||
// 这里复用 Select 的语义区分;controller 已知 router 状态,可以再调一次 Select 拿准确错误。
|
||||
return nil, nil, ErrNoChannels
|
||||
}
|
||||
|
||||
// 按 priority DESC 稳定排序;同档内按 id 升序,这样"换下一个"确定性
|
||||
// (但加权随机是按 weight 选,同档里到底选谁仍然是随机的)。
|
||||
sorted := sortByPriorityDesc(enabled)
|
||||
|
||||
var lastErrs []error
|
||||
tried := make(map[int64]struct{}, maxRetry)
|
||||
ch := Channel{} // 循环外声明,成功后用得到
|
||||
for attempt := 0; attempt < maxRetry; attempt++ {
|
||||
if ctx.Err() != nil {
|
||||
return nil, nil, ctx.Err()
|
||||
}
|
||||
// 在剩余 pool 中取当前最高档
|
||||
pool := filterUntried(sorted, tried)
|
||||
if len(pool) == 0 {
|
||||
break
|
||||
}
|
||||
top := topPriority(pool)
|
||||
var err error
|
||||
ch, err = pickByWeight(top, time.Now().UnixNano()+int64(attempt))
|
||||
if err != nil {
|
||||
// 候选档内 weight 全 ≤0 等异常;理论上不会发生
|
||||
lastErrs = append(lastErrs, err)
|
||||
continue
|
||||
}
|
||||
tried[ch.ID] = struct{}{}
|
||||
|
||||
// 拿 provider;Channel.Provider 未注册(代码级缺失)→ 视作候选失败
|
||||
prov, ok := provider.Get(ch.Provider)
|
||||
if !ok {
|
||||
lastErrs = append(lastErrs, fmt.Errorf("channel %d (%s): provider %q not registered", ch.ID, ch.Name, ch.Provider))
|
||||
continue
|
||||
}
|
||||
|
||||
req := ch.BuildRequest(template)
|
||||
// voice 字段:把请求级 voice 同步到 channel 的 VoiceKey(供 provider 拿上游音色 ID 用)
|
||||
// 这是 Router 唯一会"覆盖"请求字段的地方;其余(Text/Format/Speed)由 controller 填好。
|
||||
// 注意:Channel 不改 voice,这里只为 Synthesize 拿正确 VoiceKey;
|
||||
// 多次重试时同一个 voice,VoiceKey 始终一致。
|
||||
req.VoiceKey = voice
|
||||
|
||||
result, err := prov.Synthesize(ctx, req, mtr)
|
||||
if err == nil {
|
||||
return result, &ch, nil
|
||||
}
|
||||
lastErrs = append(lastErrs, fmt.Errorf("channel %d (%s): %w", ch.ID, ch.Name, err))
|
||||
}
|
||||
|
||||
if len(lastErrs) == 0 {
|
||||
// 走到这里说明 maxRetry=0 或 pool 一开始就空;由前面的 early return 处理,
|
||||
// 留这里防御 future 改动。
|
||||
return nil, nil, ErrNoChannels
|
||||
}
|
||||
return nil, nil, &AggregateError{Errors: lastErrs}
|
||||
}
|
||||
|
||||
// Enable 重新载入渠道列表;阶段 2 管理接口用,阶段 1 不暴露。
|
||||
func (r *Router) Enable(channels []Channel) {
|
||||
cp := make([]Channel, len(channels))
|
||||
copy(cp, channels)
|
||||
r.mu.Lock()
|
||||
r.channels = cp
|
||||
r.mu.Unlock()
|
||||
}
|
||||
|
||||
// enabledMatching 在锁内筛"启用 + voice 匹配"的 channel;供 SelectAndSynthesize 复用。
|
||||
func (r *Router) enabledMatching(voice string) []Channel {
|
||||
out := make([]Channel, 0, len(r.channels))
|
||||
for _, c := range r.channels {
|
||||
if !c.Enabled() {
|
||||
continue
|
||||
}
|
||||
if !c.AcceptsVoice(voice) {
|
||||
continue
|
||||
}
|
||||
out = append(out, c)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// sortByPriorityDesc 按 priority 降序,同档按 id 升序;返回新切片。
|
||||
func sortByPriorityDesc(in []Channel) []Channel {
|
||||
out := make([]Channel, len(in))
|
||||
copy(out, in)
|
||||
// 简单插入排序;渠道数通常 1~10 个,O(n^2) 够用且无额外分配。
|
||||
for i := 1; i < len(out); i++ {
|
||||
for j := i; j > 0; j-- {
|
||||
if out[j].Priority > out[j-1].Priority ||
|
||||
(out[j].Priority == out[j-1].Priority && out[j].ID < out[j-1].ID) {
|
||||
out[j], out[j-1] = out[j-1], out[j]
|
||||
continue
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// topPriority 返回 channels 中 priority 最大的那些(可能有多个,同档内加权随机)。
|
||||
func topPriority(channels []Channel) []Channel {
|
||||
if len(channels) == 0 {
|
||||
return nil
|
||||
}
|
||||
maxP := channels[0].Priority
|
||||
for _, c := range channels[1:] {
|
||||
if c.Priority > maxP {
|
||||
maxP = c.Priority
|
||||
}
|
||||
}
|
||||
out := make([]Channel, 0, len(channels))
|
||||
for _, c := range channels {
|
||||
if c.Priority == maxP {
|
||||
out = append(out, c)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// filterUntried 从 sorted 里挑"未在 tried 中的"channel;同档 / 跨档都适用。
|
||||
func filterUntried(sorted []Channel, tried map[int64]struct{}) []Channel {
|
||||
out := make([]Channel, 0, len(sorted))
|
||||
for _, c := range sorted {
|
||||
if _, ok := tried[c.ID]; ok {
|
||||
continue
|
||||
}
|
||||
out = append(out, c)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// pickByWeight 加权随机选一个 channel;weights 全 ≤0 时返 error(理论不该发生)。
|
||||
//
|
||||
// 实现:total = Σweight;r = rand.Intn(total);线性累加落桶。
|
||||
// 同权时均匀(每个 channel 概率 = weight/total)。
|
||||
// 抽成纯函数(接受 seed)便于单测。
|
||||
func pickByWeight(channels []Channel, seed int64) (Channel, error) {
|
||||
if len(channels) == 0 {
|
||||
return Channel{}, errors.New("route: pickByWeight on empty slice")
|
||||
}
|
||||
total := 0
|
||||
for _, c := range channels {
|
||||
w := c.Weight
|
||||
if w <= 0 {
|
||||
w = 1 // 防御:同 Channel.Weight <=0 不应出现,但兜底按 1 处理
|
||||
}
|
||||
total += w
|
||||
}
|
||||
if total <= 0 {
|
||||
return Channel{}, errors.New("route: pickByWeight total weight <= 0")
|
||||
}
|
||||
r := rand.New(rand.NewSource(seed))
|
||||
r0 := r.Intn(total)
|
||||
acc := 0
|
||||
for _, c := range channels {
|
||||
w := c.Weight
|
||||
if w <= 0 {
|
||||
w = 1
|
||||
}
|
||||
acc += w
|
||||
if r0 < acc {
|
||||
return c, nil
|
||||
}
|
||||
}
|
||||
// 浮点边界不可达;留 fallback 返最后一个
|
||||
return channels[len(channels)-1], nil
|
||||
}
|
||||
|
||||
// Errors ---------------------------------------------------------------
|
||||
|
||||
// ErrNoChannels router 里没有任何 channel(零 channel 兜底触发条件)。
|
||||
var ErrNoChannels = errors.New("route: no channels available")
|
||||
|
||||
// ErrVoiceNotFound 有 channel 但 voice 不被任一渠道支持(对应客户端 400)。
|
||||
var ErrVoiceNotFound = errors.New("route: voice not supported by any channel")
|
||||
|
||||
// AggregateError 多个渠道都失败时,聚合各渠道错误返回。
|
||||
// controller 用 errors.As 拿到后,可以在日志里逐个打印。
|
||||
type AggregateError struct {
|
||||
Errors []error
|
||||
}
|
||||
|
||||
func (e *AggregateError) Error() string {
|
||||
if len(e.Errors) == 0 {
|
||||
return "route: all channels failed (no details)"
|
||||
}
|
||||
parts := make([]string, 0, len(e.Errors))
|
||||
for _, err := range e.Errors {
|
||||
parts = append(parts, err.Error())
|
||||
}
|
||||
return fmt.Sprintf("route: all %d channels failed: %s", len(e.Errors), joinComma(parts))
|
||||
}
|
||||
|
||||
// Unwrap 让 errors.Is/errors.As 仍能透到最里层(便于上层判断特定错误类型)。
|
||||
func (e *AggregateError) Unwrap() error {
|
||||
if len(e.Errors) == 0 {
|
||||
return nil
|
||||
}
|
||||
return e.Errors[0]
|
||||
}
|
||||
|
||||
func joinComma(parts []string) string {
|
||||
out := ""
|
||||
for i, p := range parts {
|
||||
if i > 0 {
|
||||
out += "; "
|
||||
}
|
||||
out += p
|
||||
}
|
||||
return out
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
package controller
|
||||
|
||||
import (
|
||||
"github.com/volcano-tts/tts-api/adapter/route"
|
||||
)
|
||||
|
||||
// routerInstance 持有当前可用的多渠道路由器;main 启动期调 SetRouter 注入。
|
||||
// 阶段 1:启动期从 store.ChannelList 加载一次,运行期不变;阶段 2 管理接口上线后
|
||||
// 会有 Reload 触发整体替换(届时 SetRouter 同样负责更新这个指针)。
|
||||
//
|
||||
// 注入失败 / 未注入时,GetRouter 返 nil —— 此时 OpenaiTTSHandler 走零渠道兜底路径
|
||||
// (等价于"零 Channel 走 setting 默认 provider"),保持阶段 0 行为不变。
|
||||
var routerInstance *route.Router
|
||||
|
||||
// SetRouter 注入 router 句柄;main 启动期调一次。
|
||||
// 传入 nil 表示清空(回退到零渠道兜底),便于管理接口触发"禁用多渠道"操作(阶段 2)。
|
||||
func SetRouter(r *route.Router) { routerInstance = r }
|
||||
|
||||
// GetRouter 拿当前 router 句柄;nil 表示未注入或被清空。
|
||||
func GetRouter() *route.Router { return routerInstance }
|
||||
+86
-11
@@ -3,6 +3,7 @@ package controller
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log"
|
||||
@@ -12,6 +13,7 @@ import (
|
||||
"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"
|
||||
@@ -203,16 +205,89 @@ func OpenaiTTSHandler(w http.ResponseWriter, r *http.Request) {
|
||||
ctx, cancel := context.WithTimeout(r.Context(), setting.GetTTSTimeout())
|
||||
defer cancel()
|
||||
|
||||
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")
|
||||
// 阶段 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
|
||||
}
|
||||
result, err := prov.Synthesize(ctx, ttsReq, adapterRec)
|
||||
duration := time.Since(start)
|
||||
|
||||
// 选渠道路径:以"对外 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 里
|
||||
@@ -225,8 +300,8 @@ func OpenaiTTSHandler(w http.ResponseWriter, r *http.Request) {
|
||||
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)
|
||||
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
|
||||
}
|
||||
@@ -241,8 +316,8 @@ func OpenaiTTSHandler(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
if n, err := w.Write(result.AudioData); err != nil {
|
||||
// header 已发,无法改 status code;只记日志供排查(常见:客户端中途断开 → broken pipe / connection reset)
|
||||
log.Printf("警告: 响应写入失败 - 路径=%s 客户端=%s 已写=%d/%d 错误=%v",
|
||||
r.URL.Path, middleware.GetClientIP(r), n, len(result.AudioData), err)
|
||||
log.Printf("警告: 响应写入失败 - 路径=%s 已写=%d/%d 错误=%v",
|
||||
urlPath, n, len(result.AudioData), err)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -12,6 +12,7 @@ import (
|
||||
"github.com/volcano-tts/tts-api/controller"
|
||||
// 空导入注册上游适配器:新增上游只需在此加一行,主干零改动。
|
||||
_ "github.com/volcano-tts/tts-api/adapter/volcano"
|
||||
"github.com/volcano-tts/tts-api/adapter/route"
|
||||
"github.com/volcano-tts/tts-api/installer"
|
||||
"github.com/volcano-tts/tts-api/metrics"
|
||||
"github.com/volcano-tts/tts-api/middleware"
|
||||
@@ -28,6 +29,17 @@ func ttsDBPath() string {
|
||||
return "tts.db"
|
||||
}
|
||||
|
||||
// countEnabled 统计启用 channel 数;main 启动摘要用。
|
||||
func countEnabled(chs []route.Channel) int {
|
||||
n := 0
|
||||
for _, c := range chs {
|
||||
if c.Enabled() {
|
||||
n++
|
||||
}
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
func main() {
|
||||
log.SetFlags(log.LstdFlags | log.Lshortfile)
|
||||
log.SetPrefix("[TTS-Server] ")
|
||||
@@ -61,6 +73,31 @@ func main() {
|
||||
return nil
|
||||
})
|
||||
|
||||
// 3.1) 阶段 1:多渠道路由器注入。
|
||||
// - 从 store 加载全部 channels(含 disabled,便于阶段 2 管理接口看全量);
|
||||
// router 内部筛 enabled
|
||||
// - channels 加载失败(如 credentials_json 损坏)→ 启动期 fail-fast,
|
||||
// 避免一个错凭据把每次请求搞崩
|
||||
// - 加载成功但列表为空 → 仍注入 router,但 router.Empty()=true,
|
||||
// controller 走"零渠道兜底",行为=阶段 0(向后兼容验收硬要求)
|
||||
if st != nil {
|
||||
chs, err := st.ChannelList(true)
|
||||
if err != nil {
|
||||
log.Fatalf("FATAL: load channels failed: %v", err)
|
||||
}
|
||||
routerChs := make([]route.Channel, 0, len(chs))
|
||||
for _, sc := range chs {
|
||||
rc, convErr := route.FromStoreChannel(sc)
|
||||
if convErr != nil {
|
||||
log.Fatalf("FATAL: convert channel id=%d name=%s failed: %v", sc.ID, sc.Name, convErr)
|
||||
}
|
||||
routerChs = append(routerChs, rc)
|
||||
}
|
||||
controller.SetRouter(route.NewRouter(routerChs))
|
||||
log.Printf("[main] 多渠道路由已注入: 共 %d 条 (enabled=%d)",
|
||||
len(routerChs), countEnabled(routerChs))
|
||||
}
|
||||
|
||||
// 4) M3: 从 store 加载运行时 TTS 配置(替代原来的 env-based InitTTSConfig)
|
||||
// 必须在 LogStartupSummary 之前,这样日志显示的是真实状态(API key 已从 DB 加载,不再读 env)
|
||||
//
|
||||
|
||||
@@ -0,0 +1,207 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Channel 是一行渠道记录:
|
||||
// - 一个上游账号(Provider=adapter 名 + 该账号的凭据)
|
||||
// - 可限定支持的对外音色(Voices);空 = 不限(接受一切)
|
||||
// - Priority 分档(大者优先),Weight 同档加权
|
||||
//
|
||||
// JSON tag 跟前台/admin.html 直接读 Go 字段对齐(参见 store/voices.go 顶部说明)。
|
||||
type Channel struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Provider string `json:"provider"`
|
||||
// CredentialsJSON 是序列化后的 provider.Credentials:
|
||||
// {"api_key": "...", "scope": {"resource_id": "...", ...}}
|
||||
// 我们落 JSON 列而非拆字段,因为不同 provider 需要的 scope 维度差异很大
|
||||
// (火山要 resource_id,Azure 要 region,自建要 base_url ...),
|
||||
// 拆字段会被某一家绑死。读时反序列化为 provider.Credentials。
|
||||
CredentialsJSON string `json:"credentials_json"`
|
||||
Voices string `json:"voices"` // 逗号分隔的对外 voice 名;空 = 不限
|
||||
Priority int `json:"priority"`
|
||||
Weight int `json:"weight"`
|
||||
// Status: 1 = enabled, 2 = disabled(对齐 voices.enabled 的 0/1 风格,
|
||||
// 但本阶段管理接口未上线,只用 int 留扩展位,后续可加 "维护中" 之类状态)。
|
||||
Status int `json:"status"`
|
||||
AutoBan bool `json:"auto_ban"`
|
||||
CreatedAt string `json:"created_at"`
|
||||
UpdatedAt string `json:"updated_at"`
|
||||
}
|
||||
|
||||
// ChannelList 列出所有渠道(按 priority DESC, id ASC),includeDisabled=false 时只返回 enabled。
|
||||
// 排序规则固定,便于 router 启动期一次性加载后做 priority 分桶。
|
||||
func (s *Store) ChannelList(includeDisabled bool) ([]Channel, error) {
|
||||
q := `SELECT id, name, provider, credentials_json, voices, priority, weight, status, auto_ban, created_at, updated_at
|
||||
FROM channels`
|
||||
if !includeDisabled {
|
||||
q += ` WHERE status = 1`
|
||||
}
|
||||
q += ` ORDER BY priority DESC, id ASC`
|
||||
|
||||
rows, err := s.db.Query(q)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("store: channel list: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
out := make([]Channel, 0, 4)
|
||||
for rows.Next() {
|
||||
c, err := scanChannel(rows)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, c)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, fmt.Errorf("store: channel list rows: %w", err)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// ChannelGet 按 id 查;未命中返回 ErrNotFound。
|
||||
func (s *Store) ChannelGet(id int64) (*Channel, error) {
|
||||
row := s.db.QueryRow(`SELECT id, name, provider, credentials_json, voices, priority, weight, status, auto_ban, created_at, updated_at
|
||||
FROM channels WHERE id = ?`, id)
|
||||
c, err := scanChannel(row)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("store: channel get id=%d: %w", id, err)
|
||||
}
|
||||
return &c, nil
|
||||
}
|
||||
|
||||
// ChannelInsert 新增一条渠道;name 冲突返回 ErrDuplicate。
|
||||
// 输入校验错误(必填项缺失)返回 wrap ErrInvalid;服务端错误(DB 失败等)不被 wrap。
|
||||
func (s *Store) ChannelInsert(c Channel) (int64, error) {
|
||||
c.Name = strings.TrimSpace(c.Name)
|
||||
c.Provider = strings.TrimSpace(c.Provider)
|
||||
c.CredentialsJSON = strings.TrimSpace(c.CredentialsJSON)
|
||||
c.Voices = normalizeVoicesCSV(c.Voices)
|
||||
|
||||
if c.Name == "" {
|
||||
return 0, fmt.Errorf("%w: name is required", ErrInvalid)
|
||||
}
|
||||
if c.Provider == "" {
|
||||
return 0, fmt.Errorf("%w: provider is required", ErrInvalid)
|
||||
}
|
||||
if c.CredentialsJSON == "" {
|
||||
return 0, fmt.Errorf("%w: credentials_json is required", ErrInvalid)
|
||||
}
|
||||
// credentials_json 必须是合法 JSON;否则后续 router 加载会崩,这里挡掉。
|
||||
if !json.Valid([]byte(c.CredentialsJSON)) {
|
||||
return 0, fmt.Errorf("%w: credentials_json is not valid JSON", ErrInvalid)
|
||||
}
|
||||
if c.Weight <= 0 {
|
||||
c.Weight = 1 // 默认权重 1
|
||||
}
|
||||
if c.Status != 1 && c.Status != 2 {
|
||||
c.Status = 1 // 默认启用
|
||||
}
|
||||
|
||||
res, err := s.db.Exec(`
|
||||
INSERT INTO channels (name, provider, credentials_json, voices, priority, weight, status, auto_ban, created_at, updated_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, datetime('now'), datetime('now'))`,
|
||||
c.Name, c.Provider, c.CredentialsJSON, c.Voices, c.Priority, c.Weight, c.Status, boolToInt(c.AutoBan))
|
||||
if err != nil {
|
||||
if isUniqueViolation(err) {
|
||||
return 0, ErrChannelDuplicate
|
||||
}
|
||||
return 0, fmt.Errorf("store: channel insert: %w", err)
|
||||
}
|
||||
id, err := res.LastInsertId()
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("store: channel insert lastid: %w", err)
|
||||
}
|
||||
return id, nil
|
||||
}
|
||||
|
||||
// ChannelSetStatus 改状态(1=enabled, 2=disabled);未命中返回 ErrNotFound。
|
||||
// 阶段 1 没有 admin 接口,但 router 启动期之后可被 auto_ban 流程调用,
|
||||
// 所以这个方法必须先有,管理接口只是包装。
|
||||
func (s *Store) ChannelSetStatus(id int64, status int) error {
|
||||
if status != 1 && status != 2 {
|
||||
return fmt.Errorf("%w: status must be 1 (enabled) or 2 (disabled)", ErrInvalid)
|
||||
}
|
||||
res, err := s.db.Exec(`UPDATE channels SET status=?, updated_at=datetime('now') WHERE id = ?`, status, id)
|
||||
if err != nil {
|
||||
return fmt.Errorf("store: channel set status id=%d: %w", id, err)
|
||||
}
|
||||
n, _ := res.RowsAffected()
|
||||
if n == 0 {
|
||||
return ErrNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ChannelDelete 按 id 删;未命中返回 ErrNotFound。
|
||||
func (s *Store) ChannelDelete(id int64) error {
|
||||
res, err := s.db.Exec(`DELETE FROM channels WHERE id = ?`, id)
|
||||
if err != nil {
|
||||
return fmt.Errorf("store: channel delete id=%d: %w", id, err)
|
||||
}
|
||||
n, _ := res.RowsAffected()
|
||||
if n == 0 {
|
||||
return ErrNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ChannelCount 统计行数;阶段 1 暂未使用,保留为未来 admin 仪表盘用。
|
||||
func (s *Store) ChannelCount() (int, error) {
|
||||
var n int
|
||||
err := s.db.QueryRow(`SELECT COUNT(*) FROM channels`).Scan(&n)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("store: channel count: %w", err)
|
||||
}
|
||||
return n, nil
|
||||
}
|
||||
|
||||
// ErrChannelDuplicate 表示 channel name 唯一冲突;与 voices.ErrDuplicate 同名但消息不同,
|
||||
// 避免 channel 重名被翻译成"voice name already exists"。
|
||||
var ErrChannelDuplicate = errors.New("store: channel name already exists")
|
||||
|
||||
// scanChannel 把 row 扫描成 Channel;接受 *sql.Row 或 *sql.Rows(都实现 Scan)。
|
||||
func scanChannel(r scanner) (Channel, error) {
|
||||
var c Channel
|
||||
var autoBan int
|
||||
err := r.Scan(&c.ID, &c.Name, &c.Provider, &c.CredentialsJSON, &c.Voices,
|
||||
&c.Priority, &c.Weight, &c.Status, &autoBan, &c.CreatedAt, &c.UpdatedAt)
|
||||
if err != nil {
|
||||
return c, err
|
||||
}
|
||||
c.AutoBan = autoBan != 0
|
||||
return c, nil
|
||||
}
|
||||
|
||||
// normalizeVoicesCSV 把 "a,,b, a ,c" 规整成 "a,b,c" —— 重复项去重,
|
||||
// 顺序保留(便于 hash 比较和稳定展示);空字符串返 ""(语义=不限)。
|
||||
func normalizeVoicesCSV(raw string) string {
|
||||
raw = strings.TrimSpace(raw)
|
||||
if raw == "" {
|
||||
return ""
|
||||
}
|
||||
parts := strings.Split(raw, ",")
|
||||
seen := make(map[string]struct{}, len(parts))
|
||||
out := make([]string, 0, len(parts))
|
||||
for _, p := range parts {
|
||||
p = strings.TrimSpace(p)
|
||||
if p == "" {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[p]; ok {
|
||||
continue
|
||||
}
|
||||
seen[p] = struct{}{}
|
||||
out = append(out, p)
|
||||
}
|
||||
return strings.Join(out, ",")
|
||||
}
|
||||
+28
-1
@@ -18,7 +18,7 @@ import (
|
||||
|
||||
// schemaVersion 是当前 schema 版本号;每次结构性变更 +1。
|
||||
// migrate.go 负责在 Open 时按版本号增量应用。
|
||||
const schemaVersion = 1
|
||||
const schemaVersion = 2
|
||||
|
||||
// Store 是 SQLite 访问层的统一入口;所有 settings/voices 操作都通过它。
|
||||
type Store struct {
|
||||
@@ -149,5 +149,32 @@ func (s *Store) migrate() error {
|
||||
INSERT OR IGNORE INTO schema_version (version) VALUES (?)`, schemaVersion); err != nil {
|
||||
return fmt.Errorf("insert schema_version: %w", err)
|
||||
}
|
||||
|
||||
// channels 表(多渠道路由,阶段 1 引入)。
|
||||
// - name 唯一:管理侧标识,不允许重名
|
||||
// - credentials_json: 序列化后的 provider.Credentials(见 store/channels.go 注释)
|
||||
// - voices CSV: 限定该渠道支持的对外音色;空 = 不限
|
||||
// - priority/weight: 大者优先 + 同档加权随机(由 adapter/route 实现)
|
||||
// - status: 1=enabled 2=disabled(留扩展位,管理接口放阶段 2 上)
|
||||
// - auto_ban: new-api 风格连续失败自动禁用标记,阶段 1 仅存,未接自动逻辑
|
||||
if _, err := s.db.Exec(`
|
||||
CREATE TABLE IF NOT EXISTS channels (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
name TEXT NOT NULL UNIQUE,
|
||||
provider TEXT NOT NULL,
|
||||
credentials_json TEXT NOT NULL,
|
||||
voices TEXT NOT NULL DEFAULT '',
|
||||
priority INTEGER NOT NULL DEFAULT 0,
|
||||
weight INTEGER NOT NULL DEFAULT 1,
|
||||
status INTEGER NOT NULL DEFAULT 1,
|
||||
auto_ban INTEGER NOT NULL DEFAULT 0,
|
||||
created_at TEXT NOT NULL DEFAULT (datetime('now')),
|
||||
updated_at TEXT NOT NULL DEFAULT (datetime('now'))
|
||||
)`); err != nil {
|
||||
return fmt.Errorf("create channels: %w", err)
|
||||
}
|
||||
if _, err := s.db.Exec(`CREATE INDEX IF NOT EXISTS idx_channels_status_priority ON channels(status, priority DESC)`); err != nil {
|
||||
return fmt.Errorf("create idx_channels_status_priority: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user