Files
Sakurasan 9f4d631fc4 feat: 网关路由对齐参考实现 + 用量落库 + e2e 全绿
路由与故障转移(参考 openteam 语义)
- channel.Candidates:绑定模型优先(携带 upstream_model 映射),
  未绑定模型回退到权重最低的健康备用渠道;新增 Pick 加权随机与 FilterHealthy 内存健康过滤
- gateway.Dispatch:遍历候选渠道,可重试失败(连接错误/429/5xx)自动故障转移,
  4xx 透传;不再使用单一 SelectChannel
- 修复 gorm default 标签把渠道 weight=0 静默改写为 1 的问题(去掉 default,
  权重 0 语义 = 不参与加权选择,仅作备用承接 unbound 流量)
- RecordFailure 连续 2 次进入 degraded 快速熔断,健康检查成功或冷却过期后复位

网关功能补全
- /v1/models 返回 DB 中启用的模型列表(替换 TODO 存根)
- 请求级 request_id 生成与用量记录接入:流式 SSE 逐块累计 usage、
  非流式从响应提取,按模型定价计算成本后经 usage.Recorder 异步落库
- 流式结束检测:chat 的 [DONE]、messages 的 message_stop、responses 的
  response.completed,避免 keep-alive 上游发完不关连接导致读阻塞到超时
- ResponsesRequest.input 兼容字符串与条目数组两种客户端写法

测试
- 修复 convert_test 对新 input 形态的断言
- 网关 e2e(/tmp/test_gateway.py + mock upstream)72/72 全部通过,连续 3 次稳定
2026-09-01 00:46:05 +08:00

293 lines
7.7 KiB
Go

package channel
import (
"context"
"math/rand"
"opencatd-open/internal/dao"
"opencatd-open/internal/store"
"opencatd-open/internal/pkg/crypto"
"sync"
"time"
)
type Service struct {
channelDAO *dao.ChannelDAO
modelDAO *dao.ModelDAO
// Health tracking
mu sync.RWMutex
healthStatus map[uint64]*channelHealth
// Concurrency control per channel
sems map[uint64]chan struct{}
}
type channelHealth struct {
status string
consecutive int
lastCheck time.Time
cooldown time.Time
}
func NewService(channelDAO *dao.ChannelDAO, modelDAO *dao.ModelDAO) *Service {
return &Service{
channelDAO: channelDAO,
modelDAO: modelDAO,
healthStatus: make(map[uint64]*channelHealth),
sems: make(map[uint64]chan struct{}),
}
}
// SelectedRoute 一次路由决策的完整结果:渠道 + 命中的模型绑定。
// Binding 可能为 nil(渠道经回退路径选中、无绑定记录)。
type SelectedRoute struct {
Channel *store.Channel
Binding *store.ChannelModelBinding
}
// Candidates 返回可用渠道候选:健康 + 启用。
// model 非空时优先取绑定该模型的渠道(携带 upstream_model 映射,权重降序);
// 无绑定则回退到未绑定模型路径:按权重升序(闲置渠道优先探活)。
func (s *Service) Candidates(model string) []Candidate {
if model != "" {
var b []store.ChannelModelBinding
var modelIDs []uint64
s.modelDAO.DB().Model(&store.Model{}).Where("name = ? AND enabled = ?", model, true).Pluck("id", &modelIDs)
if len(modelIDs) > 0 {
s.channelDAO.DB().Where("model_id IN ?", modelIDs).Find(&b)
if cands := s.loadBound(b); len(cands) > 0 {
return cands
}
}
}
// 未绑定模型回退:取优先级最低的空闲健康渠道作为"备用渠道"承接搭车流量
// (排序与绑定候选一致:priority ASC, weight DESC, id ASC,取末位)。
// weight=0 的渠道不被加权随机选中,但可作为最后备用承接 unbound 流量。
var chs []store.Channel
s.channelDAO.DB().Where("enabled = ?", true).
Order("priority ASC, weight DESC, id ASC").Find(&chs)
all := make([]Candidate, 0, len(chs))
for i := range chs {
all = append(all, Candidate{Channel: &chs[i]})
}
healthy := s.FilterHealthy(all)
if len(healthy) == 0 {
return nil
}
return healthy[len(healthy)-1:]
}
// loadBound 按绑定顺序加载渠道候选,过滤健康/启用,携带 upstream_model 映射。
func (s *Service) loadBound(bindings []store.ChannelModelBinding) []Candidate {
if len(bindings) == 0 {
return nil
}
// channel_id -> 绑定(取该渠道对该模型的映射)
byChannel := map[uint64]store.ChannelModelBinding{}
ids := make([]uint64, 0, len(bindings))
for _, b := range bindings {
if _, ok := byChannel[b.ChannelID]; !ok {
ids = append(ids, b.ChannelID)
}
byChannel[b.ChannelID] = b
}
var chs []store.Channel
s.channelDAO.DB().Where("id IN ? AND enabled = ? AND health_status = ?", ids, true, store.ChannelHealthHealthy).
Order("priority ASC, weight DESC, id ASC").Find(&chs)
byID := map[uint64]*store.Channel{}
for i := range chs {
byID[chs[i].ID] = &chs[i]
}
out := make([]Candidate, 0, len(ids))
for _, id := range ids {
if ch, ok := byID[id]; ok {
b := byChannel[id]
out = append(out, Candidate{Channel: ch, Binding: &b})
}
}
return out
}
// GetChannelByKeyID decrypts the API key for a channel
func (s *Service) GetChannelByKeyID(ctx context.Context, channelID uint64) (*store.Channel, error) {
ch, err := s.channelDAO.GetByID(channelID)
if err != nil {
return nil, err
}
return ch, nil
}
// GetAPIKey decrypts the channel's API key
func (s *Service) GetAPIKey(ch *store.Channel) (string, error) {
return crypto.Decrypt(ch.APIKeyEnc)
}
// RecordSuccess records a successful request to a channel
func (s *Service) RecordSuccess(channelID uint64) {
s.mu.Lock()
defer s.mu.Unlock()
h := s.getOrCreateHealth(channelID)
h.consecutive = 0
h.status = store.ChannelHealthHealthy
h.lastCheck = time.Now()
}
// RecordFailure records a failed request to a channel.
// 连续 2 次失败进入 degraded(快速熔断):失败过的渠道让位给健康渠道,
// 健康检查成功或冷却过期后复位。
func (s *Service) RecordFailure(channelID uint64) {
s.mu.Lock()
defer s.mu.Unlock()
h := s.getOrCreateHealth(channelID)
h.consecutive++
h.lastCheck = time.Now()
if h.consecutive >= 2 {
h.status = store.ChannelHealthDegraded
h.cooldown = time.Now().Add(5 * time.Minute)
}
if h.consecutive >= 5 {
h.status = store.ChannelHealthCooldown
h.cooldown = time.Now().Add(15 * time.Minute)
}
}
// RecordTimeout records a timeout to a channel
func (s *Service) RecordTimeout(channelID uint64) {
s.RecordFailure(channelID)
}
func (s *Service) getOrCreateHealth(channelID uint64) *channelHealth {
h, ok := s.healthStatus[channelID]
if !ok {
h = &channelHealth{
status: store.ChannelHealthHealthy,
}
s.healthStatus[channelID] = h
}
return h
}
func (s *Service) filterHealthy(channels []*store.Channel) []*store.Channel {
s.mu.RLock()
defer s.mu.RUnlock()
var healthy []*store.Channel
now := time.Now()
for _, ch := range channels {
h, ok := s.healthStatus[ch.ID]
if !ok {
healthy = append(healthy, ch)
continue
}
// Check if cooldown has expired
if now.After(h.cooldown) && h.cooldown.IsZero() == false {
h.consecutive = 0
h.status = store.ChannelHealthHealthy
healthy = append(healthy, ch)
continue
}
if h.status == store.ChannelHealthHealthy || h.status == store.ChannelHealthDegraded {
healthy = append(healthy, ch)
}
}
return healthy
}
// GetHealthStatus returns the health status of a channel
func (s *Service) GetHealthStatus(channelID uint64) string {
s.mu.RLock()
defer s.mu.RUnlock()
h, ok := s.healthStatus[channelID]
if !ok {
return store.ChannelHealthHealthy
}
return h.status
}
// Candidate 一个候选渠道 + 该模型的映射关系。
type Candidate struct {
Channel *store.Channel
Binding *store.ChannelModelBinding // 全局模型在此渠道的映射(无绑定则 nil)
}
// Pick 按权重加权随机选一个候选渠道(负载均衡;weight<=0 按 1 计)。
func (s *Service) Pick(cands []Candidate) *Candidate {
if len(cands) == 0 {
return nil
}
total := 0
for _, c := range cands {
w := c.Channel.Weight
if w <= 0 {
w = 1
}
total += w
}
r := rand.Intn(total)
acc := 0
for i := range cands {
w := cands[i].Channel.Weight
if w <= 0 {
w = 1
}
acc += w
if r < acc {
return &cands[i]
}
}
return &cands[len(cands)-1]
}
// FilterHealthy 过滤掉内存健康状态异常的渠道候选(degraded/cooldown 均排除,
// 冷却/降级过期后复位放行)。degraded 由单次请求失败触发,作为快速熔断:
// 后续请求先走其他渠道,健康检查成功后恢复。
func (s *Service) FilterHealthy(cands []Candidate) []Candidate {
s.mu.RLock()
defer s.mu.RUnlock()
out := make([]Candidate, 0, len(cands))
now := time.Now()
for _, c := range cands {
h, ok := s.healthStatus[c.Channel.ID]
if !ok || h.status == store.ChannelHealthHealthy {
out = append(out, c)
continue
}
// 冷却/降级已过期:复位并放行
if !h.cooldown.IsZero() && now.After(h.cooldown) {
h.status = store.ChannelHealthHealthy
h.consecutive = 0
out = append(out, c)
}
}
return out
}
// TryAcquire 尝试获取渠道并发槽;渠道满载返回 false(调用方可溢出到其他渠道)。
// MaxConcurrency<=0 视为不限制。
func (s *Service) TryAcquire(ch *store.Channel) (func(), bool) {
if ch.MaxConcurrency <= 0 {
return func() {}, true
}
s.mu.Lock()
sem, ok := s.sems[ch.ID]
if !ok {
sem = make(chan struct{}, ch.MaxConcurrency)
s.sems[ch.ID] = sem
}
s.mu.Unlock()
select {
case sem <- struct{}{}:
return func() { <-sem }, true
default:
return nil, false
}
}