From a614ab55c9afacf7d6252a665c9fd4433e10a788 Mon Sep 17 00:00:00 2001 From: tts-stage1 Date: Sun, 11 Oct 2026 15:33:22 +0800 Subject: [PATCH] =?UTF-8?q?feat(route):=20=E9=98=B6=E6=AE=B5=201=20?= =?UTF-8?q?=E5=A4=9A=E6=B8=A0=E9=81=93=E8=B7=AF=E7=94=B1=E5=88=86=E5=8F=91?= =?UTF-8?q?=E5=B1=82=20(Channel=20+=20Router)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 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 个, 全部通过) --- .gitignore | 5 + adapter/route/channel.go | 108 +++++++++++++ adapter/route/router.go | 339 +++++++++++++++++++++++++++++++++++++++ controller/router.go | 20 +++ controller/tts.go | 97 +++++++++-- main.go | 37 +++++ store/channels.go | 207 ++++++++++++++++++++++++ store/db.go | 29 +++- 8 files changed, 830 insertions(+), 12 deletions(-) create mode 100644 adapter/route/channel.go create mode 100644 adapter/route/router.go create mode 100644 controller/router.go create mode 100644 store/channels.go diff --git a/.gitignore b/.gitignore index 59ce6df..1a0ad43 100644 --- a/.gitignore +++ b/.gitignore @@ -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/ diff --git a/adapter/route/channel.go b/adapter/route/channel.go new file mode 100644 index 0000000..dadcd7c --- /dev/null +++ b/adapter/route/channel.go @@ -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 +} diff --git a/adapter/route/router.go b/adapter/route/router.go new file mode 100644 index 0000000..d50514d --- /dev/null +++ b/adapter/route/router.go @@ -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 +} diff --git a/controller/router.go b/controller/router.go new file mode 100644 index 0000000..0bff8f3 --- /dev/null +++ b/controller/router.go @@ -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 } diff --git a/controller/tts.go b/controller/tts.go index 5f89328..c3546c0 100644 --- a/controller/tts.go +++ b/controller/tts.go @@ -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) } } diff --git a/main.go b/main.go index 8d39678..1ee4187 100644 --- a/main.go +++ b/main.go @@ -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) // diff --git a/store/channels.go b/store/channels.go new file mode 100644 index 0000000..9fe9f2c --- /dev/null +++ b/store/channels.go @@ -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, ",") +} diff --git a/store/db.go b/store/db.go index 0bd7c53..b27bd5d 100644 --- a/store/db.go +++ b/store/db.go @@ -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 }