mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 16:18:29 +08:00
feat: 按上游计费倍率调度 OpenAI 账号
This commit is contained in:
@@ -6,6 +6,7 @@ import (
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"math"
|
||||
"net/url"
|
||||
"os"
|
||||
"strings"
|
||||
@@ -1052,11 +1053,35 @@ type GatewayOpenAIWSSchedulerScoreWeights struct {
|
||||
Reset float64 `mapstructure:"reset"`
|
||||
// QuotaHeadroom 倾向 7d 剩余额度更健康的账号;默认 0(关闭,不改变原有行为)。
|
||||
QuotaHeadroom float64 `mapstructure:"quota_headroom"`
|
||||
// UpstreamCost 倾向上游声明倍率更低的账号;默认 0(关闭,不改变原有行为)。
|
||||
UpstreamCost float64 `mapstructure:"upstream_cost"`
|
||||
// PreviousResponse/SessionSticky 仅在开启 OpenAI 高级调度的粘性加权时生效。
|
||||
PreviousResponse float64 `mapstructure:"previous_response"`
|
||||
SessionSticky float64 `mapstructure:"session_sticky"`
|
||||
}
|
||||
|
||||
func (w GatewayOpenAIWSSchedulerScoreWeights) BaseWeightSum() float64 {
|
||||
return w.Priority + w.Load + w.Queue + w.ErrorRate + w.TTFT + w.Reset + w.QuotaHeadroom + w.UpstreamCost
|
||||
}
|
||||
|
||||
func (w GatewayOpenAIWSSchedulerScoreWeights) TotalWeightSum() float64 {
|
||||
return w.BaseWeightSum() + w.PreviousResponse + w.SessionSticky
|
||||
}
|
||||
|
||||
func (w GatewayOpenAIWSSchedulerScoreWeights) IsValid() bool {
|
||||
for _, weight := range []float64{
|
||||
w.Priority, w.Load, w.Queue, w.ErrorRate, w.TTFT, w.Reset,
|
||||
w.QuotaHeadroom, w.UpstreamCost, w.PreviousResponse, w.SessionSticky,
|
||||
} {
|
||||
if weight < 0 || math.IsNaN(weight) || math.IsInf(weight, 0) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
baseSum := w.BaseWeightSum()
|
||||
return baseSum > 0 && !math.IsNaN(baseSum) && !math.IsInf(baseSum, 0) &&
|
||||
!math.IsNaN(w.TotalWeightSum()) && !math.IsInf(w.TotalWeightSum(), 0)
|
||||
}
|
||||
|
||||
// GatewayOpenAISchedulerConfig OpenAI 高级调度器配置。
|
||||
type GatewayOpenAISchedulerConfig struct {
|
||||
// StickyEscapeEnabled: 是否允许 session_hash sticky 在账号健康度劣化时临时逃逸
|
||||
@@ -2033,6 +2058,7 @@ func setDefaults() {
|
||||
viper.SetDefault("gateway.openai_ws.scheduler_score_weights.ttft", 0.5)
|
||||
viper.SetDefault("gateway.openai_ws.scheduler_score_weights.reset", 0.0)
|
||||
viper.SetDefault("gateway.openai_ws.scheduler_score_weights.quota_headroom", 0.0)
|
||||
viper.SetDefault("gateway.openai_ws.scheduler_score_weights.upstream_cost", 0.0)
|
||||
viper.SetDefault("gateway.openai_ws.scheduler_score_weights.previous_response", 5.0)
|
||||
viper.SetDefault("gateway.openai_ws.scheduler_score_weights.session_sticky", 3.0)
|
||||
// OpenAI HTTP upstream protocol strategy
|
||||
@@ -2899,25 +2925,26 @@ func (c *Config) Validate() error {
|
||||
if c.Gateway.OpenAIHTTP2.FallbackTTLSeconds < 0 {
|
||||
return fmt.Errorf("gateway.openai_http2.fallback_ttl_seconds must be non-negative")
|
||||
}
|
||||
if c.Gateway.OpenAIWS.SchedulerScoreWeights.Priority < 0 ||
|
||||
c.Gateway.OpenAIWS.SchedulerScoreWeights.Load < 0 ||
|
||||
c.Gateway.OpenAIWS.SchedulerScoreWeights.Queue < 0 ||
|
||||
c.Gateway.OpenAIWS.SchedulerScoreWeights.ErrorRate < 0 ||
|
||||
c.Gateway.OpenAIWS.SchedulerScoreWeights.TTFT < 0 ||
|
||||
c.Gateway.OpenAIWS.SchedulerScoreWeights.QuotaHeadroom < 0 ||
|
||||
c.Gateway.OpenAIWS.SchedulerScoreWeights.PreviousResponse < 0 ||
|
||||
c.Gateway.OpenAIWS.SchedulerScoreWeights.SessionSticky < 0 {
|
||||
return fmt.Errorf("gateway.openai_ws.scheduler_score_weights.* must be non-negative")
|
||||
weights := c.Gateway.OpenAIWS.SchedulerScoreWeights
|
||||
for _, weight := range []float64{
|
||||
weights.Priority, weights.Load, weights.Queue, weights.ErrorRate, weights.TTFT,
|
||||
weights.Reset, weights.QuotaHeadroom, weights.UpstreamCost,
|
||||
weights.PreviousResponse, weights.SessionSticky,
|
||||
} {
|
||||
if weight < 0 || math.IsNaN(weight) || math.IsInf(weight, 0) {
|
||||
return fmt.Errorf("gateway.openai_ws.scheduler_score_weights.* must be non-negative and finite")
|
||||
}
|
||||
}
|
||||
weightSum := c.Gateway.OpenAIWS.SchedulerScoreWeights.Priority +
|
||||
c.Gateway.OpenAIWS.SchedulerScoreWeights.Load +
|
||||
c.Gateway.OpenAIWS.SchedulerScoreWeights.Queue +
|
||||
c.Gateway.OpenAIWS.SchedulerScoreWeights.ErrorRate +
|
||||
c.Gateway.OpenAIWS.SchedulerScoreWeights.TTFT +
|
||||
c.Gateway.OpenAIWS.SchedulerScoreWeights.QuotaHeadroom
|
||||
weightSum := weights.BaseWeightSum()
|
||||
if weightSum <= 0 {
|
||||
return fmt.Errorf("gateway.openai_ws.scheduler_score_weights must not all be zero")
|
||||
}
|
||||
if math.IsNaN(weightSum) || math.IsInf(weightSum, 0) {
|
||||
return fmt.Errorf("gateway.openai_ws.scheduler_score_weights base-weight sum must be finite")
|
||||
}
|
||||
if totalWeightSum := weights.TotalWeightSum(); math.IsNaN(totalWeightSum) || math.IsInf(totalWeightSum, 0) {
|
||||
return fmt.Errorf("gateway.openai_ws.scheduler_score_weights total-weight sum must be finite")
|
||||
}
|
||||
if c.Gateway.OpenAIScheduler.StickyEscapeTTFTMs <= 0 {
|
||||
return fmt.Errorf("gateway.openai_scheduler.sticky_escape_ttft_ms must be positive")
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"math"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
@@ -215,6 +216,9 @@ func TestLoadDefaultOpenAIWSConfig(t *testing.T) {
|
||||
if cfg.Gateway.OpenAIWS.SchedulerScoreWeights.QuotaHeadroom != 0 {
|
||||
t.Fatalf("Gateway.OpenAIWS.SchedulerScoreWeights.QuotaHeadroom = %v, want 0", cfg.Gateway.OpenAIWS.SchedulerScoreWeights.QuotaHeadroom)
|
||||
}
|
||||
if cfg.Gateway.OpenAIWS.SchedulerScoreWeights.UpstreamCost != 0 {
|
||||
t.Fatalf("Gateway.OpenAIWS.SchedulerScoreWeights.UpstreamCost = %v, want 0", cfg.Gateway.OpenAIWS.SchedulerScoreWeights.UpstreamCost)
|
||||
}
|
||||
if !cfg.Gateway.OpenAIWS.StoreDisabledForceNewConn {
|
||||
t.Fatalf("Gateway.OpenAIWS.StoreDisabledForceNewConn = false, want true")
|
||||
}
|
||||
@@ -1866,6 +1870,42 @@ func TestValidateConfig_OpenAIWSRules(t *testing.T) {
|
||||
mutate: func(c *Config) { c.Gateway.OpenAIWS.SchedulerScoreWeights.QuotaHeadroom = -0.1 },
|
||||
wantErr: "gateway.openai_ws.scheduler_score_weights.* must be non-negative",
|
||||
},
|
||||
{
|
||||
name: "scheduler_score_weights upstream_cost 不能为负数",
|
||||
mutate: func(c *Config) { c.Gateway.OpenAIWS.SchedulerScoreWeights.UpstreamCost = -0.1 },
|
||||
wantErr: "gateway.openai_ws.scheduler_score_weights.* must be non-negative",
|
||||
},
|
||||
{
|
||||
name: "scheduler_score_weights reset 不能为负数",
|
||||
mutate: func(c *Config) { c.Gateway.OpenAIWS.SchedulerScoreWeights.Reset = -0.1 },
|
||||
wantErr: "gateway.openai_ws.scheduler_score_weights.* must be non-negative",
|
||||
},
|
||||
{
|
||||
name: "scheduler_score_weights 不能为 NaN",
|
||||
mutate: func(c *Config) { c.Gateway.OpenAIWS.SchedulerScoreWeights.PreviousResponse = math.NaN() },
|
||||
wantErr: "gateway.openai_ws.scheduler_score_weights.* must be non-negative and finite",
|
||||
},
|
||||
{
|
||||
name: "scheduler_score_weights 不能为 Inf",
|
||||
mutate: func(c *Config) { c.Gateway.OpenAIWS.SchedulerScoreWeights.UpstreamCost = math.Inf(1) },
|
||||
wantErr: "gateway.openai_ws.scheduler_score_weights.* must be non-negative and finite",
|
||||
},
|
||||
{
|
||||
name: "scheduler_score_weights 总和不能溢出",
|
||||
mutate: func(c *Config) {
|
||||
c.Gateway.OpenAIWS.SchedulerScoreWeights.Priority = math.MaxFloat64
|
||||
c.Gateway.OpenAIWS.SchedulerScoreWeights.Load = math.MaxFloat64
|
||||
},
|
||||
wantErr: "gateway.openai_ws.scheduler_score_weights base-weight sum must be finite",
|
||||
},
|
||||
{
|
||||
name: "scheduler_score_weights 含 sticky 总和不能溢出",
|
||||
mutate: func(c *Config) {
|
||||
c.Gateway.OpenAIWS.SchedulerScoreWeights.Priority = math.MaxFloat64
|
||||
c.Gateway.OpenAIWS.SchedulerScoreWeights.PreviousResponse = math.MaxFloat64
|
||||
},
|
||||
wantErr: "gateway.openai_ws.scheduler_score_weights total-weight sum must be finite",
|
||||
},
|
||||
{
|
||||
name: "scheduler_score_weights 不能全为 0",
|
||||
mutate: func(c *Config) {
|
||||
@@ -1917,6 +1957,30 @@ func TestValidateConfig_OpenAIWSRules(t *testing.T) {
|
||||
|
||||
require.NoError(t, cfg.Validate())
|
||||
})
|
||||
|
||||
t.Run("upstream_cost 可作为唯一有效调度权重", func(t *testing.T) {
|
||||
cfg := buildValid(t)
|
||||
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Priority = 0
|
||||
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Load = 0
|
||||
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Queue = 0
|
||||
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.ErrorRate = 0
|
||||
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.TTFT = 0
|
||||
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.UpstreamCost = 0.1
|
||||
|
||||
require.NoError(t, cfg.Validate())
|
||||
})
|
||||
|
||||
t.Run("reset 可作为唯一有效调度权重", func(t *testing.T) {
|
||||
cfg := buildValid(t)
|
||||
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Priority = 0
|
||||
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Load = 0
|
||||
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Queue = 0
|
||||
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.ErrorRate = 0
|
||||
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.TTFT = 0
|
||||
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Reset = 0.1
|
||||
|
||||
require.NoError(t, cfg.Validate())
|
||||
})
|
||||
}
|
||||
|
||||
func TestValidateConfig_AutoScaleDisabledIgnoreAutoScaleFields(t *testing.T) {
|
||||
|
||||
@@ -269,6 +269,8 @@ func (h *SettingHandler) GetSettings(c *gin.Context) {
|
||||
PaymentVisibleMethodWxpaySource: settings.PaymentVisibleMethodWxpaySource,
|
||||
PaymentVisibleMethodAlipayEnabled: settings.PaymentVisibleMethodAlipayEnabled,
|
||||
PaymentVisibleMethodWxpayEnabled: settings.PaymentVisibleMethodWxpayEnabled,
|
||||
OpenAILowUpstreamRatePriorityEnabled: settings.OpenAILowUpstreamRatePriorityEnabled,
|
||||
OpenAIOAuthSchedulingRateMultiplier: settings.OpenAIOAuthSchedulingRateMultiplier,
|
||||
OpenAIAdvancedSchedulerEnabled: settings.OpenAIAdvancedSchedulerEnabled,
|
||||
OpenAIAdvancedSchedulerStickyWeightedEnabled: settings.OpenAIAdvancedSchedulerStickyWeightedEnabled,
|
||||
OpenAIAdvancedSchedulerSubscriptionPriorityEnabled: settings.OpenAIAdvancedSchedulerSubscriptionPriorityEnabled,
|
||||
@@ -280,6 +282,7 @@ func (h *SettingHandler) GetSettings(c *gin.Context) {
|
||||
OpenAIAdvancedSchedulerWeightTTFT: settings.OpenAIAdvancedSchedulerWeightTTFT,
|
||||
OpenAIAdvancedSchedulerWeightReset: settings.OpenAIAdvancedSchedulerWeightReset,
|
||||
OpenAIAdvancedSchedulerWeightQuotaHeadroom: settings.OpenAIAdvancedSchedulerWeightQuotaHeadroom,
|
||||
OpenAIAdvancedSchedulerWeightUpstreamCost: settings.OpenAIAdvancedSchedulerWeightUpstreamCost,
|
||||
OpenAIAdvancedSchedulerWeightPreviousResponse: settings.OpenAIAdvancedSchedulerWeightPreviousResponse,
|
||||
OpenAIAdvancedSchedulerWeightSessionSticky: settings.OpenAIAdvancedSchedulerWeightSessionSticky,
|
||||
OpenAIAdvancedSchedulerEffectiveLBTopK: settings.OpenAIAdvancedSchedulerEffectiveLBTopK,
|
||||
@@ -290,6 +293,7 @@ func (h *SettingHandler) GetSettings(c *gin.Context) {
|
||||
OpenAIAdvancedSchedulerEffectiveWeightTTFT: settings.OpenAIAdvancedSchedulerEffectiveWeightTTFT,
|
||||
OpenAIAdvancedSchedulerEffectiveWeightReset: settings.OpenAIAdvancedSchedulerEffectiveWeightReset,
|
||||
OpenAIAdvancedSchedulerEffectiveWeightQuotaHeadroom: settings.OpenAIAdvancedSchedulerEffectiveWeightQuotaHeadroom,
|
||||
OpenAIAdvancedSchedulerEffectiveWeightUpstreamCost: settings.OpenAIAdvancedSchedulerEffectiveWeightUpstreamCost,
|
||||
OpenAIAdvancedSchedulerEffectiveWeightPreviousResponse: settings.OpenAIAdvancedSchedulerEffectiveWeightPreviousResponse,
|
||||
OpenAIAdvancedSchedulerEffectiveWeightSessionSticky: settings.OpenAIAdvancedSchedulerEffectiveWeightSessionSticky,
|
||||
BalanceLowNotifyEnabled: settings.BalanceLowNotifyEnabled,
|
||||
|
||||
@@ -440,6 +440,12 @@ func diffSettings(before *service.SystemSettings, after *service.SystemSettings,
|
||||
if before.PaymentVisibleMethodWxpayEnabled != after.PaymentVisibleMethodWxpayEnabled {
|
||||
changed = append(changed, "payment_visible_method_wxpay_enabled")
|
||||
}
|
||||
if before.OpenAILowUpstreamRatePriorityEnabled != after.OpenAILowUpstreamRatePriorityEnabled {
|
||||
changed = append(changed, "openai_low_upstream_rate_priority_enabled")
|
||||
}
|
||||
if before.OpenAIOAuthSchedulingRateMultiplier != after.OpenAIOAuthSchedulingRateMultiplier {
|
||||
changed = append(changed, "openai_oauth_scheduling_rate_multiplier")
|
||||
}
|
||||
if before.OpenAIAdvancedSchedulerEnabled != after.OpenAIAdvancedSchedulerEnabled {
|
||||
changed = append(changed, "openai_advanced_scheduler_enabled")
|
||||
}
|
||||
@@ -473,6 +479,9 @@ func diffSettings(before *service.SystemSettings, after *service.SystemSettings,
|
||||
if before.OpenAIAdvancedSchedulerWeightQuotaHeadroom != after.OpenAIAdvancedSchedulerWeightQuotaHeadroom {
|
||||
changed = append(changed, "openai_advanced_scheduler_weight_quota_headroom")
|
||||
}
|
||||
if before.OpenAIAdvancedSchedulerWeightUpstreamCost != after.OpenAIAdvancedSchedulerWeightUpstreamCost {
|
||||
changed = append(changed, "openai_advanced_scheduler_weight_upstream_cost")
|
||||
}
|
||||
if before.OpenAIAdvancedSchedulerWeightPreviousResponse != after.OpenAIAdvancedSchedulerWeightPreviousResponse {
|
||||
changed = append(changed, "openai_advanced_scheduler_weight_previous_response")
|
||||
}
|
||||
|
||||
@@ -223,6 +223,7 @@ func TestSettingHandler_UpdateSettings_PersistsPaymentVisibleMethodsAndAdvancedS
|
||||
"payment_visible_method_alipay_enabled": true,
|
||||
"payment_visible_method_wxpay_enabled": false,
|
||||
"openai_advanced_scheduler_enabled": true,
|
||||
"openai_oauth_scheduling_rate_multiplier": 0.05,
|
||||
"openai_advanced_scheduler_subscription_priority_enabled": true,
|
||||
}
|
||||
rawBody, err := json.Marshal(body)
|
||||
@@ -241,6 +242,7 @@ func TestSettingHandler_UpdateSettings_PersistsPaymentVisibleMethodsAndAdvancedS
|
||||
require.Equal(t, "true", repo.values[service.SettingPaymentVisibleMethodAlipayEnabled])
|
||||
require.Equal(t, "false", repo.values[service.SettingPaymentVisibleMethodWxpayEnabled])
|
||||
require.Equal(t, "true", repo.values["openai_advanced_scheduler_enabled"])
|
||||
require.Equal(t, "0.05", repo.values[service.SettingKeyOpenAIOAuthSchedulingRateMultiplier])
|
||||
require.Equal(t, "true", repo.values[service.SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled])
|
||||
|
||||
var resp response.Response
|
||||
@@ -252,6 +254,7 @@ func TestSettingHandler_UpdateSettings_PersistsPaymentVisibleMethodsAndAdvancedS
|
||||
require.Equal(t, true, data["payment_visible_method_alipay_enabled"])
|
||||
require.Equal(t, false, data["payment_visible_method_wxpay_enabled"])
|
||||
require.Equal(t, true, data["openai_advanced_scheduler_enabled"])
|
||||
require.Equal(t, 0.05, data["openai_oauth_scheduling_rate_multiplier"])
|
||||
require.Equal(t, true, data["openai_advanced_scheduler_subscription_priority_enabled"])
|
||||
}
|
||||
|
||||
|
||||
@@ -243,19 +243,22 @@ type UpdateSettingsRequest struct {
|
||||
PaymentVisibleMethodWxpayEnabled *bool `json:"payment_visible_method_wxpay_enabled"`
|
||||
|
||||
// OpenAI account scheduling
|
||||
OpenAIAdvancedSchedulerEnabled *bool `json:"openai_advanced_scheduler_enabled"`
|
||||
OpenAIAdvancedSchedulerStickyWeightedEnabled *bool `json:"openai_advanced_scheduler_sticky_weighted_enabled"`
|
||||
OpenAIAdvancedSchedulerSubscriptionPriorityEnabled *bool `json:"openai_advanced_scheduler_subscription_priority_enabled"`
|
||||
OpenAIAdvancedSchedulerLBTopK *string `json:"openai_advanced_scheduler_lb_top_k"`
|
||||
OpenAIAdvancedSchedulerWeightPriority *string `json:"openai_advanced_scheduler_weight_priority"`
|
||||
OpenAIAdvancedSchedulerWeightLoad *string `json:"openai_advanced_scheduler_weight_load"`
|
||||
OpenAIAdvancedSchedulerWeightQueue *string `json:"openai_advanced_scheduler_weight_queue"`
|
||||
OpenAIAdvancedSchedulerWeightErrorRate *string `json:"openai_advanced_scheduler_weight_error_rate"`
|
||||
OpenAIAdvancedSchedulerWeightTTFT *string `json:"openai_advanced_scheduler_weight_ttft"`
|
||||
OpenAIAdvancedSchedulerWeightReset *string `json:"openai_advanced_scheduler_weight_reset"`
|
||||
OpenAIAdvancedSchedulerWeightQuotaHeadroom *string `json:"openai_advanced_scheduler_weight_quota_headroom"`
|
||||
OpenAIAdvancedSchedulerWeightPreviousResponse *string `json:"openai_advanced_scheduler_weight_previous_response"`
|
||||
OpenAIAdvancedSchedulerWeightSessionSticky *string `json:"openai_advanced_scheduler_weight_session_sticky"`
|
||||
OpenAILowUpstreamRatePriorityEnabled *bool `json:"openai_low_upstream_rate_priority_enabled"`
|
||||
OpenAIOAuthSchedulingRateMultiplier *float64 `json:"openai_oauth_scheduling_rate_multiplier"`
|
||||
OpenAIAdvancedSchedulerEnabled *bool `json:"openai_advanced_scheduler_enabled"`
|
||||
OpenAIAdvancedSchedulerStickyWeightedEnabled *bool `json:"openai_advanced_scheduler_sticky_weighted_enabled"`
|
||||
OpenAIAdvancedSchedulerSubscriptionPriorityEnabled *bool `json:"openai_advanced_scheduler_subscription_priority_enabled"`
|
||||
OpenAIAdvancedSchedulerLBTopK *string `json:"openai_advanced_scheduler_lb_top_k"`
|
||||
OpenAIAdvancedSchedulerWeightPriority *string `json:"openai_advanced_scheduler_weight_priority"`
|
||||
OpenAIAdvancedSchedulerWeightLoad *string `json:"openai_advanced_scheduler_weight_load"`
|
||||
OpenAIAdvancedSchedulerWeightQueue *string `json:"openai_advanced_scheduler_weight_queue"`
|
||||
OpenAIAdvancedSchedulerWeightErrorRate *string `json:"openai_advanced_scheduler_weight_error_rate"`
|
||||
OpenAIAdvancedSchedulerWeightTTFT *string `json:"openai_advanced_scheduler_weight_ttft"`
|
||||
OpenAIAdvancedSchedulerWeightReset *string `json:"openai_advanced_scheduler_weight_reset"`
|
||||
OpenAIAdvancedSchedulerWeightQuotaHeadroom *string `json:"openai_advanced_scheduler_weight_quota_headroom"`
|
||||
OpenAIAdvancedSchedulerWeightUpstreamCost *string `json:"openai_advanced_scheduler_weight_upstream_cost"`
|
||||
OpenAIAdvancedSchedulerWeightPreviousResponse *string `json:"openai_advanced_scheduler_weight_previous_response"`
|
||||
OpenAIAdvancedSchedulerWeightSessionSticky *string `json:"openai_advanced_scheduler_weight_session_sticky"`
|
||||
|
||||
// 余额不足提醒
|
||||
BalanceLowNotifyEnabled *bool `json:"balance_low_notify_enabled"`
|
||||
@@ -1429,6 +1432,18 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
|
||||
}
|
||||
return previousSettings.PaymentVisibleMethodWxpayEnabled
|
||||
}(),
|
||||
OpenAILowUpstreamRatePriorityEnabled: func() bool {
|
||||
if req.OpenAILowUpstreamRatePriorityEnabled != nil {
|
||||
return *req.OpenAILowUpstreamRatePriorityEnabled
|
||||
}
|
||||
return previousSettings.OpenAILowUpstreamRatePriorityEnabled
|
||||
}(),
|
||||
OpenAIOAuthSchedulingRateMultiplier: func() float64 {
|
||||
if req.OpenAIOAuthSchedulingRateMultiplier != nil {
|
||||
return *req.OpenAIOAuthSchedulingRateMultiplier
|
||||
}
|
||||
return previousSettings.OpenAIOAuthSchedulingRateMultiplier
|
||||
}(),
|
||||
OpenAIAdvancedSchedulerEnabled: func() bool {
|
||||
if req.OpenAIAdvancedSchedulerEnabled != nil {
|
||||
return *req.OpenAIAdvancedSchedulerEnabled
|
||||
@@ -1455,6 +1470,7 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
|
||||
OpenAIAdvancedSchedulerWeightTTFT: stringSetting(req.OpenAIAdvancedSchedulerWeightTTFT, previousSettings.OpenAIAdvancedSchedulerWeightTTFT),
|
||||
OpenAIAdvancedSchedulerWeightReset: stringSetting(req.OpenAIAdvancedSchedulerWeightReset, previousSettings.OpenAIAdvancedSchedulerWeightReset),
|
||||
OpenAIAdvancedSchedulerWeightQuotaHeadroom: stringSetting(req.OpenAIAdvancedSchedulerWeightQuotaHeadroom, previousSettings.OpenAIAdvancedSchedulerWeightQuotaHeadroom),
|
||||
OpenAIAdvancedSchedulerWeightUpstreamCost: stringSetting(req.OpenAIAdvancedSchedulerWeightUpstreamCost, previousSettings.OpenAIAdvancedSchedulerWeightUpstreamCost),
|
||||
OpenAIAdvancedSchedulerWeightPreviousResponse: stringSetting(req.OpenAIAdvancedSchedulerWeightPreviousResponse, previousSettings.OpenAIAdvancedSchedulerWeightPreviousResponse),
|
||||
OpenAIAdvancedSchedulerWeightSessionSticky: stringSetting(req.OpenAIAdvancedSchedulerWeightSessionSticky, previousSettings.OpenAIAdvancedSchedulerWeightSessionSticky),
|
||||
BalanceLowNotifyEnabled: func() bool {
|
||||
@@ -1831,6 +1847,8 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
|
||||
PaymentVisibleMethodWxpaySource: updatedSettings.PaymentVisibleMethodWxpaySource,
|
||||
PaymentVisibleMethodAlipayEnabled: updatedSettings.PaymentVisibleMethodAlipayEnabled,
|
||||
PaymentVisibleMethodWxpayEnabled: updatedSettings.PaymentVisibleMethodWxpayEnabled,
|
||||
OpenAILowUpstreamRatePriorityEnabled: updatedSettings.OpenAILowUpstreamRatePriorityEnabled,
|
||||
OpenAIOAuthSchedulingRateMultiplier: updatedSettings.OpenAIOAuthSchedulingRateMultiplier,
|
||||
OpenAIAdvancedSchedulerEnabled: updatedSettings.OpenAIAdvancedSchedulerEnabled,
|
||||
OpenAIAdvancedSchedulerStickyWeightedEnabled: updatedSettings.OpenAIAdvancedSchedulerStickyWeightedEnabled,
|
||||
OpenAIAdvancedSchedulerSubscriptionPriorityEnabled: updatedSettings.OpenAIAdvancedSchedulerSubscriptionPriorityEnabled,
|
||||
@@ -1842,6 +1860,7 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
|
||||
OpenAIAdvancedSchedulerWeightTTFT: updatedSettings.OpenAIAdvancedSchedulerWeightTTFT,
|
||||
OpenAIAdvancedSchedulerWeightReset: updatedSettings.OpenAIAdvancedSchedulerWeightReset,
|
||||
OpenAIAdvancedSchedulerWeightQuotaHeadroom: updatedSettings.OpenAIAdvancedSchedulerWeightQuotaHeadroom,
|
||||
OpenAIAdvancedSchedulerWeightUpstreamCost: updatedSettings.OpenAIAdvancedSchedulerWeightUpstreamCost,
|
||||
OpenAIAdvancedSchedulerWeightPreviousResponse: updatedSettings.OpenAIAdvancedSchedulerWeightPreviousResponse,
|
||||
OpenAIAdvancedSchedulerWeightSessionSticky: updatedSettings.OpenAIAdvancedSchedulerWeightSessionSticky,
|
||||
OpenAIAdvancedSchedulerEffectiveLBTopK: updatedSettings.OpenAIAdvancedSchedulerEffectiveLBTopK,
|
||||
@@ -1852,6 +1871,7 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
|
||||
OpenAIAdvancedSchedulerEffectiveWeightTTFT: updatedSettings.OpenAIAdvancedSchedulerEffectiveWeightTTFT,
|
||||
OpenAIAdvancedSchedulerEffectiveWeightReset: updatedSettings.OpenAIAdvancedSchedulerEffectiveWeightReset,
|
||||
OpenAIAdvancedSchedulerEffectiveWeightQuotaHeadroom: updatedSettings.OpenAIAdvancedSchedulerEffectiveWeightQuotaHeadroom,
|
||||
OpenAIAdvancedSchedulerEffectiveWeightUpstreamCost: updatedSettings.OpenAIAdvancedSchedulerEffectiveWeightUpstreamCost,
|
||||
OpenAIAdvancedSchedulerEffectiveWeightPreviousResponse: updatedSettings.OpenAIAdvancedSchedulerEffectiveWeightPreviousResponse,
|
||||
OpenAIAdvancedSchedulerEffectiveWeightSessionSticky: updatedSettings.OpenAIAdvancedSchedulerEffectiveWeightSessionSticky,
|
||||
BalanceLowNotifyEnabled: updatedSettings.BalanceLowNotifyEnabled,
|
||||
|
||||
@@ -209,29 +209,33 @@ type SystemSettings struct {
|
||||
PaymentVisibleMethodWxpayEnabled bool `json:"payment_visible_method_wxpay_enabled"`
|
||||
|
||||
// OpenAI account scheduling
|
||||
OpenAIAdvancedSchedulerEnabled bool `json:"openai_advanced_scheduler_enabled"`
|
||||
OpenAIAdvancedSchedulerStickyWeightedEnabled bool `json:"openai_advanced_scheduler_sticky_weighted_enabled"`
|
||||
OpenAIAdvancedSchedulerSubscriptionPriorityEnabled bool `json:"openai_advanced_scheduler_subscription_priority_enabled"`
|
||||
OpenAIAdvancedSchedulerLBTopK string `json:"openai_advanced_scheduler_lb_top_k"`
|
||||
OpenAIAdvancedSchedulerWeightPriority string `json:"openai_advanced_scheduler_weight_priority"`
|
||||
OpenAIAdvancedSchedulerWeightLoad string `json:"openai_advanced_scheduler_weight_load"`
|
||||
OpenAIAdvancedSchedulerWeightQueue string `json:"openai_advanced_scheduler_weight_queue"`
|
||||
OpenAIAdvancedSchedulerWeightErrorRate string `json:"openai_advanced_scheduler_weight_error_rate"`
|
||||
OpenAIAdvancedSchedulerWeightTTFT string `json:"openai_advanced_scheduler_weight_ttft"`
|
||||
OpenAIAdvancedSchedulerWeightReset string `json:"openai_advanced_scheduler_weight_reset"`
|
||||
OpenAIAdvancedSchedulerWeightQuotaHeadroom string `json:"openai_advanced_scheduler_weight_quota_headroom"`
|
||||
OpenAIAdvancedSchedulerWeightPreviousResponse string `json:"openai_advanced_scheduler_weight_previous_response"`
|
||||
OpenAIAdvancedSchedulerWeightSessionSticky string `json:"openai_advanced_scheduler_weight_session_sticky"`
|
||||
OpenAIAdvancedSchedulerEffectiveLBTopK string `json:"openai_advanced_scheduler_effective_lb_top_k"`
|
||||
OpenAIAdvancedSchedulerEffectiveWeightPriority string `json:"openai_advanced_scheduler_effective_weight_priority"`
|
||||
OpenAIAdvancedSchedulerEffectiveWeightLoad string `json:"openai_advanced_scheduler_effective_weight_load"`
|
||||
OpenAIAdvancedSchedulerEffectiveWeightQueue string `json:"openai_advanced_scheduler_effective_weight_queue"`
|
||||
OpenAIAdvancedSchedulerEffectiveWeightErrorRate string `json:"openai_advanced_scheduler_effective_weight_error_rate"`
|
||||
OpenAIAdvancedSchedulerEffectiveWeightTTFT string `json:"openai_advanced_scheduler_effective_weight_ttft"`
|
||||
OpenAIAdvancedSchedulerEffectiveWeightReset string `json:"openai_advanced_scheduler_effective_weight_reset"`
|
||||
OpenAIAdvancedSchedulerEffectiveWeightQuotaHeadroom string `json:"openai_advanced_scheduler_effective_weight_quota_headroom"`
|
||||
OpenAIAdvancedSchedulerEffectiveWeightPreviousResponse string `json:"openai_advanced_scheduler_effective_weight_previous_response"`
|
||||
OpenAIAdvancedSchedulerEffectiveWeightSessionSticky string `json:"openai_advanced_scheduler_effective_weight_session_sticky"`
|
||||
OpenAILowUpstreamRatePriorityEnabled bool `json:"openai_low_upstream_rate_priority_enabled"`
|
||||
OpenAIOAuthSchedulingRateMultiplier float64 `json:"openai_oauth_scheduling_rate_multiplier"`
|
||||
OpenAIAdvancedSchedulerEnabled bool `json:"openai_advanced_scheduler_enabled"`
|
||||
OpenAIAdvancedSchedulerStickyWeightedEnabled bool `json:"openai_advanced_scheduler_sticky_weighted_enabled"`
|
||||
OpenAIAdvancedSchedulerSubscriptionPriorityEnabled bool `json:"openai_advanced_scheduler_subscription_priority_enabled"`
|
||||
OpenAIAdvancedSchedulerLBTopK string `json:"openai_advanced_scheduler_lb_top_k"`
|
||||
OpenAIAdvancedSchedulerWeightPriority string `json:"openai_advanced_scheduler_weight_priority"`
|
||||
OpenAIAdvancedSchedulerWeightLoad string `json:"openai_advanced_scheduler_weight_load"`
|
||||
OpenAIAdvancedSchedulerWeightQueue string `json:"openai_advanced_scheduler_weight_queue"`
|
||||
OpenAIAdvancedSchedulerWeightErrorRate string `json:"openai_advanced_scheduler_weight_error_rate"`
|
||||
OpenAIAdvancedSchedulerWeightTTFT string `json:"openai_advanced_scheduler_weight_ttft"`
|
||||
OpenAIAdvancedSchedulerWeightReset string `json:"openai_advanced_scheduler_weight_reset"`
|
||||
OpenAIAdvancedSchedulerWeightQuotaHeadroom string `json:"openai_advanced_scheduler_weight_quota_headroom"`
|
||||
OpenAIAdvancedSchedulerWeightUpstreamCost string `json:"openai_advanced_scheduler_weight_upstream_cost"`
|
||||
OpenAIAdvancedSchedulerWeightPreviousResponse string `json:"openai_advanced_scheduler_weight_previous_response"`
|
||||
OpenAIAdvancedSchedulerWeightSessionSticky string `json:"openai_advanced_scheduler_weight_session_sticky"`
|
||||
OpenAIAdvancedSchedulerEffectiveLBTopK string `json:"openai_advanced_scheduler_effective_lb_top_k"`
|
||||
OpenAIAdvancedSchedulerEffectiveWeightPriority string `json:"openai_advanced_scheduler_effective_weight_priority"`
|
||||
OpenAIAdvancedSchedulerEffectiveWeightLoad string `json:"openai_advanced_scheduler_effective_weight_load"`
|
||||
OpenAIAdvancedSchedulerEffectiveWeightQueue string `json:"openai_advanced_scheduler_effective_weight_queue"`
|
||||
OpenAIAdvancedSchedulerEffectiveWeightErrorRate string `json:"openai_advanced_scheduler_effective_weight_error_rate"`
|
||||
OpenAIAdvancedSchedulerEffectiveWeightTTFT string `json:"openai_advanced_scheduler_effective_weight_ttft"`
|
||||
OpenAIAdvancedSchedulerEffectiveWeightReset string `json:"openai_advanced_scheduler_effective_weight_reset"`
|
||||
OpenAIAdvancedSchedulerEffectiveWeightQuotaHeadroom string `json:"openai_advanced_scheduler_effective_weight_quota_headroom"`
|
||||
OpenAIAdvancedSchedulerEffectiveWeightUpstreamCost string `json:"openai_advanced_scheduler_effective_weight_upstream_cost"`
|
||||
OpenAIAdvancedSchedulerEffectiveWeightPreviousResponse string `json:"openai_advanced_scheduler_effective_weight_previous_response"`
|
||||
OpenAIAdvancedSchedulerEffectiveWeightSessionSticky string `json:"openai_advanced_scheduler_effective_weight_session_sticky"`
|
||||
|
||||
// Payment configuration
|
||||
PaymentEnabled bool `json:"payment_enabled"`
|
||||
|
||||
@@ -189,6 +189,7 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
|
||||
"",
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
service.PlatformGrok,
|
||||
)
|
||||
if err != nil {
|
||||
|
||||
@@ -121,6 +121,7 @@ func (h *OpenAIGatewayHandler) AlphaSearch(c *gin.Context) {
|
||||
service.OpenAIEndpointCapabilityChatCompletions,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
service.PlatformOpenAI,
|
||||
)
|
||||
if err != nil || selection == nil || selection.Account == nil {
|
||||
|
||||
@@ -151,6 +151,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
|
||||
service.OpenAIEndpointCapabilityChatCompletions,
|
||||
false,
|
||||
false,
|
||||
true,
|
||||
requestPlatform,
|
||||
)
|
||||
if err != nil {
|
||||
|
||||
@@ -119,6 +119,7 @@ func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) {
|
||||
service.OpenAIEndpointCapabilityEmbeddings,
|
||||
false,
|
||||
false,
|
||||
true,
|
||||
)
|
||||
if err != nil {
|
||||
if failoverClientGone(c) {
|
||||
|
||||
@@ -111,6 +111,7 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) {
|
||||
service.OpenAIEndpointCapabilityChatCompletions,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
openAICompatibleRequestPlatform(apiKey),
|
||||
)
|
||||
service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds())
|
||||
|
||||
@@ -366,6 +366,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
|
||||
service.OpenAIEndpointCapabilityChatCompletions,
|
||||
requireCompact,
|
||||
false,
|
||||
!imageIntent,
|
||||
requestPlatform,
|
||||
)
|
||||
if err != nil {
|
||||
@@ -909,6 +910,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
|
||||
service.OpenAIEndpointCapabilityChatCompletions,
|
||||
false,
|
||||
false,
|
||||
true,
|
||||
requestPlatform,
|
||||
)
|
||||
if err != nil {
|
||||
@@ -1461,7 +1463,8 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
if service.IsImageGenerationIntent("/v1/responses", reqModel, firstMessage) && !service.GroupAllowsImageGeneration(apiKey.Group) {
|
||||
imageIntent := service.IsImageGenerationIntent("/v1/responses", reqModel, firstMessage)
|
||||
if imageIntent && !service.GroupAllowsImageGeneration(apiKey.Group) {
|
||||
closeOpenAIClientWS(wsConn, coderws.StatusPolicyViolation, service.ImageGenerationPermissionMessage())
|
||||
return
|
||||
}
|
||||
@@ -1604,6 +1607,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
service.OpenAIEndpointCapabilityChatCompletions,
|
||||
false,
|
||||
previousResponseCanMove,
|
||||
!imageIntent,
|
||||
requestPlatform,
|
||||
)
|
||||
if err != nil {
|
||||
|
||||
@@ -863,10 +863,18 @@ func filterSchedulerExtra(extra map[string]any) map[string]any {
|
||||
"auto_pause_5h_disabled",
|
||||
"auto_pause_7d_disabled",
|
||||
"model_rate_limits",
|
||||
service.UpstreamBillingProbeExtraKey,
|
||||
}
|
||||
filtered := make(map[string]any)
|
||||
for _, key := range keys {
|
||||
if value, ok := extra[key]; ok && value != nil {
|
||||
if key == service.UpstreamBillingProbeExtraKey {
|
||||
filteredProbe := filterSchedulerUpstreamBillingProbe(value)
|
||||
if filteredProbe == nil {
|
||||
continue
|
||||
}
|
||||
value = filteredProbe
|
||||
}
|
||||
filtered[key] = value
|
||||
}
|
||||
}
|
||||
@@ -875,3 +883,43 @@ func filterSchedulerExtra(extra map[string]any) map[string]any {
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
func filterSchedulerUpstreamBillingProbe(value any) map[string]any {
|
||||
source, ok := value.(map[string]any)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
|
||||
status, ok := source["status"].(string)
|
||||
if !ok || status == "" {
|
||||
return nil
|
||||
}
|
||||
filtered := map[string]any{"status": status}
|
||||
for _, key := range []string{"received_at", "fresh_until", "next_probe_at"} {
|
||||
if field, exists := source[key]; exists && field != nil {
|
||||
filtered[key] = field
|
||||
}
|
||||
}
|
||||
data, ok := source["data"].(map[string]any)
|
||||
if !ok {
|
||||
return filtered
|
||||
}
|
||||
filteredData := make(map[string]any)
|
||||
for _, key := range []string{
|
||||
"billing_scope",
|
||||
"resolved_rate_multiplier",
|
||||
"peak_rate_enabled",
|
||||
"peak_start",
|
||||
"peak_end",
|
||||
"peak_rate_multiplier",
|
||||
"timezone",
|
||||
} {
|
||||
if field, exists := data[key]; exists && field != nil {
|
||||
filteredData[key] = field
|
||||
}
|
||||
}
|
||||
if len(filteredData) > 0 {
|
||||
filtered["data"] = filteredData
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
@@ -35,3 +36,75 @@ func TestSchedulerMetadataAccountKeepsOpenAISubscriptionIdentity(t *testing.T) {
|
||||
require.True(t, metadata.IsOpenAIChatGPTSubscription())
|
||||
require.Empty(t, metadata.GetCredential("access_token"))
|
||||
}
|
||||
|
||||
func TestSchedulerMetadataAccountProjectsUpstreamBillingProbe(t *testing.T) {
|
||||
lastError := strings.Repeat("upstream diagnostic ", 512)
|
||||
probe := map[string]any{
|
||||
"status": "ok",
|
||||
"data": map[string]any{
|
||||
"billing_scope": "token",
|
||||
"resolved_rate_multiplier": 0.03,
|
||||
"peak_rate_enabled": true,
|
||||
"peak_start": "09:00",
|
||||
"peak_end": "18:00",
|
||||
"peak_rate_multiplier": 2.0,
|
||||
"timezone": "Asia/Shanghai",
|
||||
"effective_rate_multiplier": 0.03,
|
||||
"remote_diagnostic": lastError,
|
||||
},
|
||||
"received_at": "2026-07-13T10:00:00Z",
|
||||
"fresh_until": "2026-07-13T11:00:00Z",
|
||||
"next_probe_at": "2026-07-13T10:30:00Z",
|
||||
"http_status": 502,
|
||||
"last_error": lastError,
|
||||
}
|
||||
account := service.Account{
|
||||
ID: 42,
|
||||
Extra: map[string]any{
|
||||
"upstream_billing_probe": probe,
|
||||
"unused_large_field": "drop-me",
|
||||
},
|
||||
}
|
||||
|
||||
metadata := buildSchedulerMetadataAccount(account)
|
||||
fullPayload, metaPayload, err := marshalSchedulerCacheAccount(account)
|
||||
require.NoError(t, err)
|
||||
|
||||
filtered, ok := metadata.Extra["upstream_billing_probe"].(map[string]any)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, "ok", filtered["status"])
|
||||
require.Equal(t, "2026-07-13T10:00:00Z", filtered["received_at"])
|
||||
require.Equal(t, "2026-07-13T11:00:00Z", filtered["fresh_until"])
|
||||
require.Equal(t, "2026-07-13T10:30:00Z", filtered["next_probe_at"])
|
||||
require.NotContains(t, filtered, "http_status")
|
||||
require.NotContains(t, filtered, "last_error")
|
||||
filteredData, ok := filtered["data"].(map[string]any)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, "token", filteredData["billing_scope"])
|
||||
require.Equal(t, 0.03, filteredData["resolved_rate_multiplier"])
|
||||
require.Equal(t, true, filteredData["peak_rate_enabled"])
|
||||
require.Equal(t, "09:00", filteredData["peak_start"])
|
||||
require.Equal(t, "18:00", filteredData["peak_end"])
|
||||
require.Equal(t, 2.0, filteredData["peak_rate_multiplier"])
|
||||
require.Equal(t, "Asia/Shanghai", filteredData["timezone"])
|
||||
require.NotContains(t, filteredData, "effective_rate_multiplier")
|
||||
require.NotContains(t, filteredData, "remote_diagnostic")
|
||||
require.NotContains(t, metadata.Extra, "unused_large_field")
|
||||
require.Contains(t, string(fullPayload), lastError)
|
||||
require.NotContains(t, string(metaPayload), "last_error")
|
||||
require.Less(t, len(metaPayload)*4, len(fullPayload))
|
||||
}
|
||||
|
||||
func TestSchedulerMetadataAccountDropsInvalidUpstreamBillingProbe(t *testing.T) {
|
||||
for _, probe := range []any{
|
||||
"invalid",
|
||||
map[string]any{},
|
||||
map[string]any{"status": ""},
|
||||
} {
|
||||
metadata := buildSchedulerMetadataAccount(service.Account{
|
||||
Extra: map[string]any{service.UpstreamBillingProbeExtraKey: probe},
|
||||
})
|
||||
|
||||
require.NotContains(t, metadata.Extra, service.UpstreamBillingProbeExtraKey)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -684,6 +684,8 @@ func TestAPIContracts(t *testing.T) {
|
||||
service.SettingPaymentVisibleMethodWxpaySource: service.VisibleMethodSourceOfficialWechat,
|
||||
service.SettingPaymentVisibleMethodAlipayEnabled: "true",
|
||||
service.SettingPaymentVisibleMethodWxpayEnabled: "false",
|
||||
service.SettingKeyOpenAILowUpstreamRatePriorityEnabled: "true",
|
||||
service.SettingKeyOpenAIOAuthSchedulingRateMultiplier: "0.05",
|
||||
"openai_advanced_scheduler_enabled": "true",
|
||||
service.SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled: "false",
|
||||
service.SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled: "false",
|
||||
@@ -878,6 +880,8 @@ func TestAPIContracts(t *testing.T) {
|
||||
"payment_visible_method_wxpay_source": "official_wxpay",
|
||||
"payment_visible_method_alipay_enabled": true,
|
||||
"payment_visible_method_wxpay_enabled": false,
|
||||
"openai_low_upstream_rate_priority_enabled": true,
|
||||
"openai_oauth_scheduling_rate_multiplier": 0.05,
|
||||
"openai_advanced_scheduler_enabled": true,
|
||||
"openai_advanced_scheduler_sticky_weighted_enabled": false,
|
||||
"openai_advanced_scheduler_subscription_priority_enabled": false,
|
||||
@@ -889,6 +893,7 @@ func TestAPIContracts(t *testing.T) {
|
||||
"openai_advanced_scheduler_weight_ttft": "",
|
||||
"openai_advanced_scheduler_weight_reset": "",
|
||||
"openai_advanced_scheduler_weight_quota_headroom": "",
|
||||
"openai_advanced_scheduler_weight_upstream_cost": "",
|
||||
"openai_advanced_scheduler_weight_previous_response": "",
|
||||
"openai_advanced_scheduler_weight_session_sticky": "",
|
||||
"openai_advanced_scheduler_effective_lb_top_k": "7",
|
||||
@@ -899,6 +904,7 @@ func TestAPIContracts(t *testing.T) {
|
||||
"openai_advanced_scheduler_effective_weight_ttft": "0.5",
|
||||
"openai_advanced_scheduler_effective_weight_reset": "0",
|
||||
"openai_advanced_scheduler_effective_weight_quota_headroom": "0",
|
||||
"openai_advanced_scheduler_effective_weight_upstream_cost": "0",
|
||||
"openai_advanced_scheduler_effective_weight_previous_response": "5",
|
||||
"openai_advanced_scheduler_effective_weight_session_sticky": "3",
|
||||
"openai_codex_user_agent": "",
|
||||
@@ -1151,6 +1157,8 @@ func TestAPIContracts(t *testing.T) {
|
||||
"payment_visible_method_wxpay_source": "",
|
||||
"payment_visible_method_alipay_enabled": false,
|
||||
"payment_visible_method_wxpay_enabled": false,
|
||||
"openai_low_upstream_rate_priority_enabled": false,
|
||||
"openai_oauth_scheduling_rate_multiplier": 1,
|
||||
"openai_advanced_scheduler_enabled": false,
|
||||
"openai_advanced_scheduler_sticky_weighted_enabled": false,
|
||||
"openai_advanced_scheduler_subscription_priority_enabled": false,
|
||||
@@ -1162,6 +1170,7 @@ func TestAPIContracts(t *testing.T) {
|
||||
"openai_advanced_scheduler_weight_ttft": "",
|
||||
"openai_advanced_scheduler_weight_reset": "",
|
||||
"openai_advanced_scheduler_weight_quota_headroom": "",
|
||||
"openai_advanced_scheduler_weight_upstream_cost": "",
|
||||
"openai_advanced_scheduler_weight_previous_response": "",
|
||||
"openai_advanced_scheduler_weight_session_sticky": "",
|
||||
"openai_advanced_scheduler_effective_lb_top_k": "7",
|
||||
@@ -1172,6 +1181,7 @@ func TestAPIContracts(t *testing.T) {
|
||||
"openai_advanced_scheduler_effective_weight_ttft": "0.5",
|
||||
"openai_advanced_scheduler_effective_weight_reset": "0",
|
||||
"openai_advanced_scheduler_effective_weight_quota_headroom": "0",
|
||||
"openai_advanced_scheduler_effective_weight_upstream_cost": "0",
|
||||
"openai_advanced_scheduler_effective_weight_previous_response": "5",
|
||||
"openai_advanced_scheduler_effective_weight_session_sticky": "3",
|
||||
"openai_codex_user_agent": "",
|
||||
|
||||
@@ -437,6 +437,10 @@ const (
|
||||
|
||||
// SettingKeyAllowUngroupedKeyScheduling 允许未分组 API Key 调度(默认 false:未分组 Key 返回 403)
|
||||
SettingKeyAllowUngroupedKeyScheduling = "allow_ungrouped_key_scheduling"
|
||||
// SettingKeyOpenAILowUpstreamRatePriorityEnabled 旧调度是否按上游 token 倍率优先。
|
||||
SettingKeyOpenAILowUpstreamRatePriorityEnabled = "openai_low_upstream_rate_priority_enabled"
|
||||
// SettingKeyOpenAIOAuthSchedulingRateMultiplier OAuth 账号参与成本调度时使用的参考倍率。
|
||||
SettingKeyOpenAIOAuthSchedulingRateMultiplier = "openai_oauth_scheduling_rate_multiplier"
|
||||
// SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled OpenAI 高级调度下是否启用粘性加权。
|
||||
SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled = "openai_advanced_scheduler_sticky_weighted_enabled"
|
||||
// SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled OpenAI 高级调度下是否优先使用订阅账号池。
|
||||
@@ -449,6 +453,7 @@ const (
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightTTFT = "openai_advanced_scheduler_weight_ttft"
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightReset = "openai_advanced_scheduler_weight_reset"
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightQuotaHeadroom = "openai_advanced_scheduler_weight_quota_headroom"
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightUpstreamCost = "openai_advanced_scheduler_weight_upstream_cost"
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightPreviousResponse = "openai_advanced_scheduler_weight_previous_response"
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightSessionSticky = "openai_advanced_scheduler_weight_session_sticky"
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -5,6 +5,7 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"math"
|
||||
"strconv"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -88,9 +89,14 @@ type schedulerTestConcurrencyCache struct {
|
||||
acquireResults map[int64]bool
|
||||
waitCounts map[int64]int
|
||||
skipDefaultLoad bool
|
||||
acquiredIDs *[]int64
|
||||
releasedIDs *[]int64
|
||||
}
|
||||
|
||||
func (c schedulerTestConcurrencyCache) AcquireAccountSlot(ctx context.Context, accountID int64, maxConcurrency int, requestID string) (bool, error) {
|
||||
if c.acquiredIDs != nil {
|
||||
*c.acquiredIDs = append(*c.acquiredIDs, accountID)
|
||||
}
|
||||
if c.acquireResults != nil {
|
||||
if result, ok := c.acquireResults[accountID]; ok {
|
||||
return result, nil
|
||||
@@ -100,6 +106,9 @@ func (c schedulerTestConcurrencyCache) AcquireAccountSlot(ctx context.Context, a
|
||||
}
|
||||
|
||||
func (c schedulerTestConcurrencyCache) ReleaseAccountSlot(ctx context.Context, accountID int64, requestID string) error {
|
||||
if c.releasedIDs != nil {
|
||||
*c.releasedIDs = append(*c.releasedIDs, accountID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -303,14 +312,16 @@ func TestOpenAIGatewayService_OpenAIAdvancedSchedulerRuntimeSettings_DBOverrides
|
||||
TTFT: 5,
|
||||
Reset: 6,
|
||||
QuotaHeadroom: 7,
|
||||
PreviousResponse: 8,
|
||||
SessionSticky: 9,
|
||||
UpstreamCost: 8,
|
||||
PreviousResponse: 9,
|
||||
SessionSticky: 10,
|
||||
}
|
||||
repo := &openAIAdvancedSchedulerSettingRepoStub{
|
||||
values: map[string]string{
|
||||
openAIAdvancedSchedulerSettingKey: "true",
|
||||
SettingKeyOpenAIAdvancedSchedulerLBTopK: "3",
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightPriority: "2.5",
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightReset: "0.25",
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightPreviousResponse: "12",
|
||||
},
|
||||
}
|
||||
@@ -324,8 +335,57 @@ func TestOpenAIGatewayService_OpenAIAdvancedSchedulerRuntimeSettings_DBOverrides
|
||||
weights := svc.openAIWSSchedulerWeightsForRequest(ctx)
|
||||
require.Equal(t, 2.5, weights.Priority)
|
||||
require.Equal(t, 2.0, weights.Load)
|
||||
require.Equal(t, 8.0, weights.UpstreamCost)
|
||||
require.Equal(t, 0.25, weights.Reset)
|
||||
require.Equal(t, 12.0, weights.Previous)
|
||||
require.Equal(t, 9.0, weights.SessionSticky)
|
||||
require.Equal(t, 10.0, weights.SessionSticky)
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayService_OpenAIAdvancedSchedulerRuntimeSettings_InvalidWeightSumsFallBackToConfig(t *testing.T) {
|
||||
base := config.GatewayOpenAIWSSchedulerScoreWeights{
|
||||
Priority: 1, Load: 2, Queue: 3, ErrorRate: 4, TTFT: 5, Reset: 6,
|
||||
QuotaHeadroom: 7, UpstreamCost: 8, PreviousResponse: 9, SessionSticky: 10,
|
||||
}
|
||||
maxFloat := strconv.FormatFloat(math.MaxFloat64, 'g', -1, 64)
|
||||
tests := []struct {
|
||||
name string
|
||||
values map[string]string
|
||||
}{
|
||||
{
|
||||
name: "invalid single value",
|
||||
values: map[string]string{
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightPriority: "NaN",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "base sum overflow",
|
||||
values: map[string]string{
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightPriority: maxFloat,
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightLoad: maxFloat,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "sticky total sum overflow",
|
||||
values: map[string]string{
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightPriority: maxFloat,
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightPreviousResponse: maxFloat,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
resetOpenAIAdvancedSchedulerSettingCacheForTest()
|
||||
defer resetOpenAIAdvancedSchedulerSettingCacheForTest()
|
||||
tt.values[openAIAdvancedSchedulerSettingKey] = "true"
|
||||
cfg := &config.Config{}
|
||||
cfg.Gateway.OpenAIWS.SchedulerScoreWeights = base
|
||||
repo := &openAIAdvancedSchedulerSettingRepoStub{values: tt.values}
|
||||
svc := &OpenAIGatewayService{cfg: cfg, rateLimitService: &RateLimitService{settingService: NewSettingService(repo, cfg)}}
|
||||
|
||||
require.Equal(t, (&OpenAIGatewayService{cfg: cfg}).openAIWSSchedulerWeights(), svc.openAIWSSchedulerWeightsForRequest(context.Background()))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayService_SelectAccountWithScheduler_DefaultDisabledUsesLegacyLoadAwareness(t *testing.T) {
|
||||
@@ -530,6 +590,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_DefaultDisabled_Embeddi
|
||||
OpenAIEndpointCapabilityEmbeddings,
|
||||
false,
|
||||
false,
|
||||
true,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, selection)
|
||||
@@ -574,6 +635,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_DefaultDisabled_AllowsG
|
||||
OpenAIEndpointCapabilityChatCompletions,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
PlatformGrok,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
@@ -777,6 +839,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_StickyWeightedPreviousR
|
||||
OpenAIEndpointCapabilityChatCompletions,
|
||||
false,
|
||||
false,
|
||||
true,
|
||||
PlatformOpenAI,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
@@ -800,6 +863,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_StickyWeightedPreviousR
|
||||
OpenAIEndpointCapabilityChatCompletions,
|
||||
false,
|
||||
true,
|
||||
true,
|
||||
PlatformOpenAI,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
@@ -942,6 +1006,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_Enabled_EmbeddingsSkips
|
||||
OpenAIEndpointCapabilityEmbeddings,
|
||||
false,
|
||||
false,
|
||||
true,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, selection)
|
||||
@@ -1016,6 +1081,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_Enabled_EmbeddingsSkips
|
||||
OpenAIEndpointCapabilityEmbeddings,
|
||||
false,
|
||||
false,
|
||||
true,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, selection)
|
||||
@@ -1043,10 +1109,10 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_SessionStickyRateLimite
|
||||
ctx := context.Background()
|
||||
groupID := int64(10101)
|
||||
rateLimitedUntil := time.Now().Add(30 * time.Minute)
|
||||
staleSticky := &Account{ID: 31001, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 0}
|
||||
staleBackup := &Account{ID: 31002, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5}
|
||||
freshSticky := &Account{ID: 31001, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 0, RateLimitResetAt: &rateLimitedUntil}
|
||||
freshBackup := &Account{ID: 31002, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5}
|
||||
staleSticky := &Account{ID: 31001, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 0, GroupIDs: []int64{groupID}}
|
||||
staleBackup := &Account{ID: 31002, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5, GroupIDs: []int64{groupID}}
|
||||
freshSticky := &Account{ID: 31001, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 0, GroupIDs: []int64{groupID}, RateLimitResetAt: &rateLimitedUntil}
|
||||
freshBackup := &Account{ID: 31002, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5, GroupIDs: []int64{groupID}}
|
||||
cache := &schedulerTestGatewayCache{sessionBindings: map[string]int64{"openai:session_hash_rate_limited": 31001}}
|
||||
snapshotCache := &openAISnapshotCacheStub{snapshotAccounts: []*Account{staleSticky, staleBackup}, accountsByID: map[int64]*Account{31001: freshSticky, 31002: freshBackup}}
|
||||
snapshotService := &SchedulerSnapshotService{cache: snapshotCache}
|
||||
@@ -1352,10 +1418,10 @@ func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_SkipsFreshlyRa
|
||||
ctx := context.Background()
|
||||
groupID := int64(10102)
|
||||
rateLimitedUntil := time.Now().Add(30 * time.Minute)
|
||||
stalePrimary := &Account{ID: 32001, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 0}
|
||||
staleSecondary := &Account{ID: 32002, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5}
|
||||
freshPrimary := &Account{ID: 32001, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 0, RateLimitResetAt: &rateLimitedUntil}
|
||||
freshSecondary := &Account{ID: 32002, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5}
|
||||
stalePrimary := &Account{ID: 32001, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 0, GroupIDs: []int64{groupID}}
|
||||
staleSecondary := &Account{ID: 32002, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5, GroupIDs: []int64{groupID}}
|
||||
freshPrimary := &Account{ID: 32001, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 0, GroupIDs: []int64{groupID}, RateLimitResetAt: &rateLimitedUntil}
|
||||
freshSecondary := &Account{ID: 32002, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5, GroupIDs: []int64{groupID}}
|
||||
snapshotCache := &openAISnapshotCacheStub{snapshotAccounts: []*Account{stalePrimary, staleSecondary}, accountsByID: map[int64]*Account{32001: freshPrimary, 32002: freshSecondary}}
|
||||
snapshotService := &SchedulerSnapshotService{cache: snapshotCache}
|
||||
svc := &OpenAIGatewayService{
|
||||
@@ -1419,10 +1485,10 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_SessionStickyDBRuntimeR
|
||||
ctx := context.Background()
|
||||
groupID := int64(10103)
|
||||
rateLimitedUntil := time.Now().Add(30 * time.Minute)
|
||||
staleSticky := &Account{ID: 33001, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 0}
|
||||
staleBackup := &Account{ID: 33002, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5}
|
||||
dbSticky := Account{ID: 33001, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 0, RateLimitResetAt: &rateLimitedUntil}
|
||||
dbBackup := Account{ID: 33002, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5}
|
||||
staleSticky := &Account{ID: 33001, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 0, GroupIDs: []int64{groupID}}
|
||||
staleBackup := &Account{ID: 33002, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5, GroupIDs: []int64{groupID}}
|
||||
dbSticky := Account{ID: 33001, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 0, GroupIDs: []int64{groupID}, RateLimitResetAt: &rateLimitedUntil}
|
||||
dbBackup := Account{ID: 33002, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5, GroupIDs: []int64{groupID}}
|
||||
cache := &schedulerTestGatewayCache{sessionBindings: map[string]int64{"openai:session_hash_db_runtime_recheck": 33001}}
|
||||
snapshotCache := &openAISnapshotCacheStub{
|
||||
snapshotAccounts: []*Account{staleSticky, staleBackup},
|
||||
@@ -1450,10 +1516,10 @@ func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_DBRuntimeReche
|
||||
ctx := context.Background()
|
||||
groupID := int64(10104)
|
||||
rateLimitedUntil := time.Now().Add(30 * time.Minute)
|
||||
stalePrimary := &Account{ID: 34001, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 0}
|
||||
staleSecondary := &Account{ID: 34002, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5}
|
||||
dbPrimary := Account{ID: 34001, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 0, RateLimitResetAt: &rateLimitedUntil}
|
||||
dbSecondary := Account{ID: 34002, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5}
|
||||
stalePrimary := &Account{ID: 34001, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 0, GroupIDs: []int64{groupID}}
|
||||
staleSecondary := &Account{ID: 34002, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5, GroupIDs: []int64{groupID}}
|
||||
dbPrimary := Account{ID: 34001, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 0, GroupIDs: []int64{groupID}, RateLimitResetAt: &rateLimitedUntil}
|
||||
dbSecondary := Account{ID: 34002, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5, GroupIDs: []int64{groupID}}
|
||||
snapshotCache := &openAISnapshotCacheStub{
|
||||
snapshotAccounts: []*Account{stalePrimary, staleSecondary},
|
||||
accountsByID: map[int64]*Account{34001: stalePrimary, 34002: staleSecondary},
|
||||
@@ -1472,6 +1538,101 @@ func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_DBRuntimeReche
|
||||
require.Equal(t, int64(34002), account.ID)
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayService_SelectAccountWithScheduler_DBFreshGroupRecheckReleasesMovedAccount(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
groupID, otherGroupID := int64(10105), int64(10106)
|
||||
stalePrimary := &Account{ID: 34101, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 0, GroupIDs: []int64{groupID}}
|
||||
staleBackup := &Account{ID: 34102, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 10, GroupIDs: []int64{groupID}}
|
||||
dbPrimary := *stalePrimary
|
||||
dbPrimary.GroupIDs = []int64{otherGroupID}
|
||||
dbBackup := *staleBackup
|
||||
snapshotCache := &openAISnapshotCacheStub{
|
||||
snapshotAccounts: []*Account{stalePrimary, staleBackup},
|
||||
accountsByID: map[int64]*Account{stalePrimary.ID: stalePrimary, staleBackup.ID: staleBackup},
|
||||
}
|
||||
acquiredIDs, releasedIDs := []int64{}, []int64{}
|
||||
cfg := &config.Config{RunMode: config.RunModeStandard}
|
||||
svc := &OpenAIGatewayService{
|
||||
accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{dbPrimary, dbBackup}},
|
||||
cfg: cfg,
|
||||
schedulerSnapshot: &SchedulerSnapshotService{cache: snapshotCache},
|
||||
concurrencyService: NewConcurrencyService(schedulerTestConcurrencyCache{
|
||||
acquiredIDs: &acquiredIDs,
|
||||
releasedIDs: &releasedIDs,
|
||||
}),
|
||||
}
|
||||
scheduler := &defaultOpenAIAccountScheduler{service: svc}
|
||||
|
||||
selection, _, err := scheduler.tryAcquireOpenAISelectionOrder(ctx, OpenAIAccountScheduleRequest{
|
||||
GroupID: &groupID, Platform: PlatformOpenAI, RequestedModel: "gpt-5.1",
|
||||
}, []openAIAccountCandidateScore{
|
||||
{account: stalePrimary, loadInfo: &AccountLoadInfo{AccountID: stalePrimary.ID}},
|
||||
{account: staleBackup, loadInfo: &AccountLoadInfo{AccountID: staleBackup.ID}},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, staleBackup.ID, selection.Account.ID)
|
||||
require.Equal(t, []int64{stalePrimary.ID, staleBackup.ID}, acquiredIDs)
|
||||
require.Equal(t, []int64{stalePrimary.ID}, releasedIDs)
|
||||
selection.ReleaseFunc()
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayService_SelectAccountWithLoadAwareness_DBFreshGroupRecheckWaitsOnValidAccount(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
groupID, otherGroupID := int64(10107), int64(10108)
|
||||
stalePrimary := &Account{ID: 34201, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 0, GroupIDs: []int64{groupID}}
|
||||
staleBackup := &Account{ID: 34202, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 10, GroupIDs: []int64{groupID}}
|
||||
dbPrimary := *stalePrimary
|
||||
dbPrimary.GroupIDs = []int64{otherGroupID}
|
||||
dbBackup := *staleBackup
|
||||
snapshotCache := &openAISnapshotCacheStub{
|
||||
snapshotAccounts: []*Account{stalePrimary, staleBackup},
|
||||
accountsByID: map[int64]*Account{stalePrimary.ID: stalePrimary, staleBackup.ID: staleBackup},
|
||||
}
|
||||
cfg := &config.Config{RunMode: config.RunModeStandard}
|
||||
cfg.Gateway.Scheduling.LoadBatchEnabled = true
|
||||
svc := &OpenAIGatewayService{
|
||||
accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{dbPrimary, dbBackup}},
|
||||
cfg: cfg,
|
||||
schedulerSnapshot: &SchedulerSnapshotService{cache: snapshotCache},
|
||||
concurrencyService: NewConcurrencyService(schedulerTestConcurrencyCache{
|
||||
acquireResults: map[int64]bool{staleBackup.ID: false},
|
||||
}),
|
||||
}
|
||||
|
||||
selection, err := svc.SelectAccountWithLoadAwareness(ctx, &groupID, "", "gpt-5.1", nil)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, selection.WaitPlan)
|
||||
require.Equal(t, staleBackup.ID, selection.Account.ID)
|
||||
require.Equal(t, staleBackup.ID, selection.WaitPlan.AccountID)
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayService_RecheckSelectedOpenAIAccountFromDB_SimpleModeUsesFullPool(t *testing.T) {
|
||||
grouped := Account{ID: 34301, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, GroupIDs: []int64{99}}
|
||||
svc := &OpenAIGatewayService{
|
||||
accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{grouped}},
|
||||
cfg: &config.Config{RunMode: config.RunModeSimple},
|
||||
schedulerSnapshot: &SchedulerSnapshotService{cache: &openAISnapshotCacheStub{}},
|
||||
}
|
||||
requestedGroupID := int64(100)
|
||||
|
||||
for _, groupID := range []*int64{nil, &requestedGroupID} {
|
||||
fresh := svc.recheckSelectedOpenAIAccountFromDB(context.Background(), &grouped, groupID, PlatformOpenAI, "gpt-5.1", false, "")
|
||||
require.NotNil(t, fresh)
|
||||
require.Equal(t, grouped.ID, fresh.ID)
|
||||
}
|
||||
|
||||
ungrouped := grouped
|
||||
ungrouped.ID++
|
||||
ungrouped.GroupIDs = nil
|
||||
standardSvc := &OpenAIGatewayService{
|
||||
accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{grouped, ungrouped}},
|
||||
cfg: &config.Config{RunMode: config.RunModeStandard},
|
||||
schedulerSnapshot: &SchedulerSnapshotService{cache: &openAISnapshotCacheStub{}},
|
||||
}
|
||||
require.Nil(t, standardSvc.recheckSelectedOpenAIAccountFromDB(context.Background(), &grouped, nil, PlatformOpenAI, "gpt-5.1", false, ""))
|
||||
require.NotNil(t, standardSvc.recheckSelectedOpenAIAccountFromDB(context.Background(), &ungrouped, nil, PlatformOpenAI, "gpt-5.1", false, ""))
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayService_SelectAccountWithScheduler_PreviousResponseSticky(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
groupID := int64(9)
|
||||
@@ -1871,6 +2032,12 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_SessionStickyEscapeDisa
|
||||
func TestOpenAIGatewayService_SelectAccountWithScheduler_SubscriptionPriorityChoosesSubscriptionPoolFirst(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
groupID := int64(10120)
|
||||
cheapAPIKey := upstreamCostTestAccount(21602, UpstreamBillingProbeStatusOK, 0.01, time.Now().Add(-time.Minute), 30*time.Minute)
|
||||
cheapAPIKey.Status = StatusActive
|
||||
cheapAPIKey.Schedulable = true
|
||||
cheapAPIKey.Concurrency = 1
|
||||
cheapAPIKey.Priority = 0
|
||||
cheapAPIKey.GroupIDs = []int64{groupID}
|
||||
accounts := []Account{
|
||||
{
|
||||
ID: 21601,
|
||||
@@ -1883,16 +2050,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_SubscriptionPriorityCho
|
||||
GroupIDs: []int64{groupID},
|
||||
Credentials: map[string]any{"plan_type": "plus"},
|
||||
},
|
||||
{
|
||||
ID: 21602,
|
||||
Platform: PlatformOpenAI,
|
||||
Type: AccountTypeAPIKey,
|
||||
Status: StatusActive,
|
||||
Schedulable: true,
|
||||
Concurrency: 1,
|
||||
Priority: 0,
|
||||
GroupIDs: []int64{groupID},
|
||||
},
|
||||
*cheapAPIKey,
|
||||
}
|
||||
concurrencyCache := schedulerTestConcurrencyCache{
|
||||
acquireResults: map[int64]bool{21601: true, 21602: true},
|
||||
@@ -1901,10 +2059,12 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_SubscriptionPriorityCho
|
||||
21602: {AccountID: 21602, LoadRate: 0, WaitingCount: 0},
|
||||
},
|
||||
}
|
||||
cfg := newSchedulerTestSubscriptionPriorityConfig()
|
||||
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.UpstreamCost = 100
|
||||
svc := &OpenAIGatewayService{
|
||||
accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts},
|
||||
cache: &schedulerTestGatewayCache{},
|
||||
cfg: newSchedulerTestSubscriptionPriorityConfig(),
|
||||
cfg: cfg,
|
||||
rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true", "", "true"),
|
||||
concurrencyService: NewConcurrencyService(concurrencyCache),
|
||||
}
|
||||
|
||||
@@ -0,0 +1,814 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"math"
|
||||
"strconv"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/config"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type upstreamCostTrackingConcurrencyCache struct {
|
||||
ConcurrencyCache
|
||||
loadMap map[int64]*AccountLoadInfo
|
||||
acquireLimits map[int64][]int
|
||||
releases map[int64]int
|
||||
rejectAcquire bool
|
||||
}
|
||||
|
||||
func (c *upstreamCostTrackingConcurrencyCache) AcquireAccountSlot(_ context.Context, accountID int64, maxConcurrency int, _ string) (bool, error) {
|
||||
if c.acquireLimits == nil {
|
||||
c.acquireLimits = make(map[int64][]int)
|
||||
}
|
||||
c.acquireLimits[accountID] = append(c.acquireLimits[accountID], maxConcurrency)
|
||||
return !c.rejectAcquire, nil
|
||||
}
|
||||
|
||||
func (c *upstreamCostTrackingConcurrencyCache) ReleaseAccountSlot(_ context.Context, accountID int64, _ string) error {
|
||||
if c.releases == nil {
|
||||
c.releases = make(map[int64]int)
|
||||
}
|
||||
c.releases[accountID]++
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *upstreamCostTrackingConcurrencyCache) GetAccountsLoadBatch(_ context.Context, accounts []AccountWithConcurrency) (map[int64]*AccountLoadInfo, error) {
|
||||
out := make(map[int64]*AccountLoadInfo, len(accounts))
|
||||
for _, account := range accounts {
|
||||
if load := c.loadMap[account.ID]; load != nil {
|
||||
copied := *load
|
||||
out[account.ID] = &copied
|
||||
}
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (c *upstreamCostTrackingConcurrencyCache) limits(accountID int64) []int {
|
||||
return append([]int(nil), c.acquireLimits[accountID]...)
|
||||
}
|
||||
|
||||
func (c *upstreamCostTrackingConcurrencyCache) releaseCount(accountID int64) int {
|
||||
return c.releases[accountID]
|
||||
}
|
||||
|
||||
func (c *upstreamCostTrackingConcurrencyCache) totalAcquires() int {
|
||||
total := 0
|
||||
for _, limits := range c.acquireLimits {
|
||||
total += len(limits)
|
||||
}
|
||||
return total
|
||||
}
|
||||
|
||||
type upstreamCostCountingAccountRepo struct {
|
||||
AccountRepository
|
||||
accounts map[int64]*Account
|
||||
getCalls int
|
||||
}
|
||||
|
||||
func (r *upstreamCostCountingAccountRepo) GetByID(_ context.Context, accountID int64) (*Account, error) {
|
||||
r.getCalls++
|
||||
account := r.accounts[accountID]
|
||||
if account == nil {
|
||||
return nil, errors.New("account not found")
|
||||
}
|
||||
cloned := *account
|
||||
return &cloned, nil
|
||||
}
|
||||
|
||||
func (r *upstreamCostCountingAccountRepo) calls() int {
|
||||
return r.getCalls
|
||||
}
|
||||
|
||||
func upstreamCostTestAccount(id int64, status string, rate float64, receivedAt time.Time, interval time.Duration) *Account {
|
||||
return &Account{
|
||||
ID: id,
|
||||
Platform: PlatformOpenAI,
|
||||
Type: AccountTypeAPIKey,
|
||||
Extra: map[string]any{
|
||||
UpstreamBillingProbeExtraKey: map[string]any{
|
||||
"status": status,
|
||||
"data": map[string]any{
|
||||
"billing_scope": "token",
|
||||
"resolved_rate_multiplier": rate,
|
||||
"peak_rate_enabled": false,
|
||||
"effective_rate_multiplier": rate,
|
||||
},
|
||||
"received_at": receivedAt.UTC().Format(time.RFC3339Nano),
|
||||
"fresh_until": receivedAt.Add(2 * interval).UTC().Format(time.RFC3339Nano),
|
||||
"last_attempt_at": receivedAt.UTC().Format(time.RFC3339Nano),
|
||||
"next_probe_at": receivedAt.Add(interval).UTC().Format(time.RFC3339Nano),
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func upstreamCostTestOAuthAccount(id int64) *Account {
|
||||
return &Account{ID: id, Platform: PlatformOpenAI, Type: AccountTypeOAuth}
|
||||
}
|
||||
|
||||
func TestAdvancedCostSchedulerUsesTopKOverflowWhenPreferredAccountIsKnownFull(t *testing.T) {
|
||||
resetOpenAIAdvancedSchedulerSettingCacheForTest()
|
||||
defer resetOpenAIAdvancedSchedulerSettingCacheForTest()
|
||||
|
||||
now := time.Now()
|
||||
cheap := upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 0.03, now.Add(-time.Minute), 30*time.Minute)
|
||||
expensive := upstreamCostTestAccount(2, UpstreamBillingProbeStatusOK, 0.8, now.Add(-time.Minute), 30*time.Minute)
|
||||
for _, account := range []*Account{cheap, expensive} {
|
||||
account.Status = StatusActive
|
||||
account.Schedulable = true
|
||||
account.Concurrency = 1
|
||||
}
|
||||
cache := &upstreamCostTrackingConcurrencyCache{loadMap: map[int64]*AccountLoadInfo{
|
||||
cheap.ID: {AccountID: cheap.ID, CurrentConcurrency: 1, LoadRate: 100},
|
||||
expensive.ID: {AccountID: expensive.ID},
|
||||
}}
|
||||
cfg := &config.Config{}
|
||||
cfg.Gateway.OpenAIWS.LBTopK = 1
|
||||
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.UpstreamCost = 1
|
||||
svc := &OpenAIGatewayService{
|
||||
accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{*cheap, *expensive}},
|
||||
cfg: cfg,
|
||||
rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true"),
|
||||
concurrencyService: NewConcurrencyService(cache),
|
||||
}
|
||||
groupID := int64(1)
|
||||
|
||||
selection, _, err := svc.SelectAccountWithScheduler(context.Background(), &groupID, "", "", "gpt-test", nil, OpenAIUpstreamTransportAny, false)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, expensive.ID, selection.Account.ID)
|
||||
require.Empty(t, cache.limits(cheap.ID))
|
||||
require.Equal(t, []int{1}, cache.limits(expensive.ID))
|
||||
selection.ReleaseFunc()
|
||||
}
|
||||
|
||||
func TestAdvancedSchedulerCapsRejectedCostOverflowAcquires(t *testing.T) {
|
||||
selectionOrder := make([]openAIAccountCandidateScore, 0, 15_000)
|
||||
for id := int64(1); id <= 15_000; id++ {
|
||||
account := &Account{ID: id, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1}
|
||||
selectionOrder = append(selectionOrder, openAIAccountCandidateScore{
|
||||
account: account, loadInfo: &AccountLoadInfo{AccountID: id}, loadKnown: false,
|
||||
})
|
||||
}
|
||||
cache := &upstreamCostTrackingConcurrencyCache{rejectAcquire: true}
|
||||
scheduler := &defaultOpenAIAccountScheduler{service: &OpenAIGatewayService{
|
||||
concurrencyService: NewConcurrencyService(cache),
|
||||
}}
|
||||
|
||||
selection, _, err := scheduler.tryAcquireOpenAISelectionOrder(
|
||||
context.Background(), OpenAIAccountScheduleRequest{Platform: PlatformOpenAI}, selectionOrder,
|
||||
)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, selection)
|
||||
require.Equal(t, openAIAccountSelectionProbeLimit, cache.totalAcquires())
|
||||
}
|
||||
|
||||
func TestOpenAICostOverflowExpandedOnlyWhenCostAddsCandidates(t *testing.T) {
|
||||
candidates := []openAIAccountCandidateScore{
|
||||
{account: &Account{ID: 1, Extra: map[string]any{"openai_compact_supported": true}}},
|
||||
{account: &Account{ID: 2}},
|
||||
}
|
||||
plan := openAIAccountLoadPlan{candidates: candidates, topK: 1, includeOverflowFallback: true}
|
||||
require.True(t, openAICostOverflowExpanded(OpenAIAccountScheduleRequest{}, plan))
|
||||
require.False(t, openAICostOverflowExpanded(OpenAIAccountScheduleRequest{RequireCompact: true}, plan),
|
||||
"one candidate per compact tier does not expand either tier's top-k")
|
||||
plan.topK = len(candidates)
|
||||
require.False(t, openAICostOverflowExpanded(OpenAIAccountScheduleRequest{}, plan))
|
||||
plan.includeOverflowFallback = false
|
||||
plan.topK = 1
|
||||
require.False(t, openAICostOverflowExpanded(OpenAIAccountScheduleRequest{}, plan))
|
||||
}
|
||||
|
||||
func TestAdvancedSchedulerKnownFullOverflowStillFindsAvailableAccount(t *testing.T) {
|
||||
selectionOrder := make([]openAIAccountCandidateScore, 0, openAIAccountSelectionProbeLimit+2)
|
||||
for id := int64(1); id <= openAIAccountSelectionProbeLimit+1; id++ {
|
||||
account := &Account{ID: id, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1}
|
||||
selectionOrder = append(selectionOrder, openAIAccountCandidateScore{
|
||||
account: account,
|
||||
loadInfo: &AccountLoadInfo{AccountID: id, CurrentConcurrency: 1, LoadRate: 100},
|
||||
loadKnown: true,
|
||||
})
|
||||
}
|
||||
availableID := int64(openAIAccountSelectionProbeLimit + 2)
|
||||
available := &Account{ID: availableID, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1}
|
||||
selectionOrder = append(selectionOrder, openAIAccountCandidateScore{
|
||||
account: available, loadInfo: &AccountLoadInfo{AccountID: availableID}, loadKnown: true,
|
||||
})
|
||||
cache := &upstreamCostTrackingConcurrencyCache{}
|
||||
scheduler := &defaultOpenAIAccountScheduler{service: &OpenAIGatewayService{
|
||||
concurrencyService: NewConcurrencyService(cache),
|
||||
}}
|
||||
|
||||
selection, _, err := scheduler.tryAcquireOpenAISelectionOrder(
|
||||
context.Background(), OpenAIAccountScheduleRequest{Platform: PlatformOpenAI}, selectionOrder,
|
||||
)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, selection)
|
||||
require.Equal(t, availableID, selection.Account.ID)
|
||||
require.Equal(t, 1, cache.totalAcquires())
|
||||
selection.ReleaseFunc()
|
||||
}
|
||||
|
||||
func TestAdvancedSchedulerSharesProbeBudgetWithFallbackDBRechecks(t *testing.T) {
|
||||
const size = 15_000
|
||||
latestAccounts := make(map[int64]*Account, size)
|
||||
snapshotAccounts := make(map[int64]*Account, size)
|
||||
selectionOrder := make([]openAIAccountCandidateScore, 0, size)
|
||||
for id := int64(1); id <= size; id++ {
|
||||
stale := &Account{ID: id, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1}
|
||||
latest := *stale
|
||||
latest.Status = StatusDisabled
|
||||
snapshotAccounts[id] = stale
|
||||
latestAccounts[id] = &latest
|
||||
selectionOrder = append(selectionOrder, openAIAccountCandidateScore{
|
||||
account: stale, loadInfo: &AccountLoadInfo{AccountID: id}, loadKnown: false,
|
||||
})
|
||||
}
|
||||
repo := &upstreamCostCountingAccountRepo{accounts: latestAccounts}
|
||||
cache := &upstreamCostTrackingConcurrencyCache{}
|
||||
scheduler := &defaultOpenAIAccountScheduler{service: &OpenAIGatewayService{
|
||||
accountRepo: repo,
|
||||
schedulerSnapshot: &SchedulerSnapshotService{cache: &openAISnapshotCacheStub{accountsByID: snapshotAccounts}},
|
||||
concurrencyService: NewConcurrencyService(cache),
|
||||
}}
|
||||
budget := newOpenAISelectionProbeBudget()
|
||||
budget.enableLimit()
|
||||
req := OpenAIAccountScheduleRequest{Platform: PlatformOpenAI}
|
||||
|
||||
selection, _, err := scheduler.tryAcquireOpenAISelectionOrderWithBudget(context.Background(), req, selectionOrder, budget)
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, selection)
|
||||
selection, _, _, _, err = scheduler.finishLoadBalanceSelectionFallback(
|
||||
context.Background(), req, openAIAccountLoadSelectionAttempt{selectionOrder: selectionOrder}, budget,
|
||||
)
|
||||
|
||||
require.Error(t, err)
|
||||
require.Nil(t, selection)
|
||||
require.Equal(t, openAIAccountSelectionProbeLimit, cache.totalAcquires())
|
||||
require.Equal(t, openAIAccountSelectionProbeLimit, repo.calls())
|
||||
}
|
||||
|
||||
func TestAdvancedCostSchedulerKeepsCompactSupportedOverflowAheadOfUnknown(t *testing.T) {
|
||||
resetOpenAIAdvancedSchedulerSettingCacheForTest()
|
||||
defer resetOpenAIAdvancedSchedulerSettingCacheForTest()
|
||||
|
||||
now := time.Now()
|
||||
preferred := upstreamCostTestAccount(11, UpstreamBillingProbeStatusOK, 0.01, now.Add(-time.Minute), 30*time.Minute)
|
||||
overflow := upstreamCostTestAccount(12, UpstreamBillingProbeStatusOK, 0.1, now.Add(-time.Minute), 30*time.Minute)
|
||||
unknown := upstreamCostTestAccount(13, UpstreamBillingProbeStatusOK, 0.001, now.Add(-time.Minute), 30*time.Minute)
|
||||
preferred.Extra["openai_compact_supported"] = true
|
||||
overflow.Extra["openai_compact_supported"] = true
|
||||
for _, account := range []*Account{preferred, overflow, unknown} {
|
||||
account.Status = StatusActive
|
||||
account.Schedulable = true
|
||||
account.Concurrency = 1
|
||||
}
|
||||
cache := &upstreamCostTrackingConcurrencyCache{loadMap: map[int64]*AccountLoadInfo{
|
||||
preferred.ID: {AccountID: preferred.ID, CurrentConcurrency: 1, LoadRate: 100},
|
||||
overflow.ID: {AccountID: overflow.ID},
|
||||
unknown.ID: {AccountID: unknown.ID},
|
||||
}}
|
||||
cfg := &config.Config{}
|
||||
cfg.Gateway.OpenAIWS.LBTopK = 1
|
||||
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.UpstreamCost = 1
|
||||
svc := &OpenAIGatewayService{
|
||||
accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{*preferred, *overflow, *unknown}},
|
||||
cfg: cfg,
|
||||
rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true"),
|
||||
concurrencyService: NewConcurrencyService(cache),
|
||||
}
|
||||
groupID := int64(1)
|
||||
|
||||
selection, _, err := svc.SelectAccountWithScheduler(context.Background(), &groupID, "", "", "gpt-test", nil, OpenAIUpstreamTransportAny, true)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, overflow.ID, selection.Account.ID)
|
||||
require.Empty(t, cache.limits(preferred.ID))
|
||||
require.Equal(t, []int{1}, cache.limits(overflow.ID))
|
||||
require.Empty(t, cache.limits(unknown.ID))
|
||||
selection.ReleaseFunc()
|
||||
}
|
||||
|
||||
func TestAdvancedSchedulerUnknownLoadFailsOpen(t *testing.T) {
|
||||
account := &Account{ID: 21, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1}
|
||||
cache := &upstreamCostTrackingConcurrencyCache{}
|
||||
scheduler := &defaultOpenAIAccountScheduler{service: &OpenAIGatewayService{concurrencyService: NewConcurrencyService(cache)}}
|
||||
|
||||
selection, _, err := scheduler.tryAcquireOpenAISelectionOrder(context.Background(), OpenAIAccountScheduleRequest{Platform: PlatformOpenAI}, []openAIAccountCandidateScore{{
|
||||
account: account, loadInfo: &AccountLoadInfo{AccountID: account.ID, CurrentConcurrency: 99}, loadKnown: false,
|
||||
}})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, selection)
|
||||
require.Equal(t, []int{1}, cache.limits(account.ID))
|
||||
selection.ReleaseFunc()
|
||||
}
|
||||
|
||||
func TestAdvancedSchedulerReleasesSlotWhenDBDisablesCandidate(t *testing.T) {
|
||||
stale := &Account{ID: 31, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1}
|
||||
backup := &Account{ID: 32, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1}
|
||||
disabled := *stale
|
||||
disabled.Status = StatusDisabled
|
||||
repo := &upstreamCostCountingAccountRepo{accounts: map[int64]*Account{stale.ID: &disabled, backup.ID: backup}}
|
||||
snapshot := &openAISnapshotCacheStub{accountsByID: map[int64]*Account{stale.ID: stale, backup.ID: backup}}
|
||||
cache := &upstreamCostTrackingConcurrencyCache{}
|
||||
scheduler := &defaultOpenAIAccountScheduler{service: &OpenAIGatewayService{
|
||||
accountRepo: repo,
|
||||
schedulerSnapshot: &SchedulerSnapshotService{cache: snapshot},
|
||||
concurrencyService: NewConcurrencyService(cache),
|
||||
}}
|
||||
|
||||
selection, _, err := scheduler.tryAcquireOpenAISelectionOrder(context.Background(), OpenAIAccountScheduleRequest{Platform: PlatformOpenAI}, []openAIAccountCandidateScore{
|
||||
{account: stale, loadInfo: &AccountLoadInfo{AccountID: stale.ID}, loadKnown: true},
|
||||
{account: backup, loadInfo: &AccountLoadInfo{AccountID: backup.ID}, loadKnown: true},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, backup.ID, selection.Account.ID)
|
||||
require.Equal(t, 1, cache.releaseCount(stale.ID))
|
||||
selection.ReleaseFunc()
|
||||
}
|
||||
|
||||
func TestAdvancedSchedulerReacquiresOnceWhenDBConcurrencyChanges(t *testing.T) {
|
||||
stale := &Account{ID: 41, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 10}
|
||||
latest := *stale
|
||||
latest.Concurrency = 1
|
||||
repo := &upstreamCostCountingAccountRepo{accounts: map[int64]*Account{stale.ID: &latest}}
|
||||
snapshot := &openAISnapshotCacheStub{accountsByID: map[int64]*Account{stale.ID: stale}}
|
||||
cache := &upstreamCostTrackingConcurrencyCache{}
|
||||
scheduler := &defaultOpenAIAccountScheduler{service: &OpenAIGatewayService{
|
||||
accountRepo: repo,
|
||||
schedulerSnapshot: &SchedulerSnapshotService{cache: snapshot},
|
||||
concurrencyService: NewConcurrencyService(cache),
|
||||
}}
|
||||
|
||||
selection, _, err := scheduler.tryAcquireOpenAISelectionOrder(context.Background(), OpenAIAccountScheduleRequest{Platform: PlatformOpenAI}, []openAIAccountCandidateScore{{
|
||||
account: stale, loadInfo: &AccountLoadInfo{AccountID: stale.ID}, loadKnown: true,
|
||||
}})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 1, selection.Account.Concurrency)
|
||||
require.Equal(t, []int{10, 1}, cache.limits(stale.ID))
|
||||
require.Equal(t, 1, cache.releaseCount(stale.ID))
|
||||
selection.ReleaseFunc()
|
||||
}
|
||||
|
||||
func TestAdvancedSchedulerKnownFullPoolsDoNotRecheckDB(t *testing.T) {
|
||||
for _, size := range []int{100, 15_000} {
|
||||
t.Run(strconv.Itoa(size), func(t *testing.T) {
|
||||
accounts := make(map[int64]*Account, size)
|
||||
selectionOrder := make([]openAIAccountCandidateScore, 0, size)
|
||||
for i := 1; i <= size; i++ {
|
||||
account := &Account{ID: int64(i), Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1}
|
||||
accounts[account.ID] = account
|
||||
selectionOrder = append(selectionOrder, openAIAccountCandidateScore{
|
||||
account: account,
|
||||
loadInfo: &AccountLoadInfo{AccountID: account.ID, CurrentConcurrency: 1, LoadRate: 100},
|
||||
loadKnown: true,
|
||||
})
|
||||
}
|
||||
repo := &upstreamCostCountingAccountRepo{accounts: accounts}
|
||||
cache := &upstreamCostTrackingConcurrencyCache{}
|
||||
scheduler := &defaultOpenAIAccountScheduler{service: &OpenAIGatewayService{
|
||||
accountRepo: repo,
|
||||
schedulerSnapshot: &SchedulerSnapshotService{cache: &openAISnapshotCacheStub{accountsByID: accounts}},
|
||||
concurrencyService: NewConcurrencyService(cache),
|
||||
}}
|
||||
|
||||
selection, _, err := scheduler.tryAcquireOpenAISelectionOrder(context.Background(), OpenAIAccountScheduleRequest{Platform: PlatformOpenAI}, selectionOrder)
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, selection)
|
||||
require.Zero(t, repo.calls())
|
||||
require.Zero(t, cache.totalAcquires())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAIFreshUpstreamBillingRateRecomputesPeakAtSelectionTime(t *testing.T) {
|
||||
receivedAt := time.Date(2026, 7, 13, 17, 30, 0, 0, time.UTC)
|
||||
account := upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 0.4, receivedAt, time.Hour)
|
||||
snapshot, ok := account.Extra[UpstreamBillingProbeExtraKey].(map[string]any)
|
||||
require.True(t, ok)
|
||||
snapshot["data"] = map[string]any{
|
||||
"billing_scope": "token",
|
||||
"resolved_rate_multiplier": 0.4,
|
||||
"peak_rate_enabled": true,
|
||||
"peak_start": "09:00",
|
||||
"peak_end": "18:00",
|
||||
"peak_rate_multiplier": 2.0,
|
||||
"applied_peak_multiplier": 2.0,
|
||||
"effective_rate_multiplier": 0.8,
|
||||
"timezone": "UTC",
|
||||
}
|
||||
|
||||
duringPeak, ok := openAIFreshUpstreamBillingRate(account, time.Date(2026, 7, 13, 17, 59, 0, 0, time.UTC))
|
||||
require.True(t, ok)
|
||||
require.Equal(t, 0.8, duringPeak)
|
||||
|
||||
afterPeak, ok := openAIFreshUpstreamBillingRate(account, time.Date(2026, 7, 13, 18, 1, 0, 0, time.UTC))
|
||||
require.True(t, ok)
|
||||
require.Equal(t, 0.4, afterPeak)
|
||||
}
|
||||
|
||||
func TestOpenAIUpstreamCostFactorsSparseProbeIsNeutral(t *testing.T) {
|
||||
now := time.Date(2026, 7, 13, 12, 0, 0, 0, time.UTC)
|
||||
accounts := make([]*Account, 0, 10)
|
||||
accounts = append(accounts, upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 1, now.Add(-time.Minute), 30*time.Minute))
|
||||
for id := int64(2); id <= 10; id++ {
|
||||
accounts = append(accounts, &Account{
|
||||
ID: id,
|
||||
Platform: PlatformOpenAI,
|
||||
Type: AccountTypeAPIKey,
|
||||
Extra: map[string]any{
|
||||
UpstreamBillingProbeExtraKey: map[string]any{
|
||||
"status": UpstreamBillingProbeStatusFailed,
|
||||
"last_attempt_at": now.UTC().Format(time.RFC3339Nano),
|
||||
"next_probe_at": now.Add(time.Hour).UTC().Format(time.RFC3339Nano),
|
||||
},
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
factors := openAIUpstreamCostFactors(accounts, now, defaultOpenAIOAuthSchedulingRateMultiplier)
|
||||
for id := int64(1); id <= 10; id++ {
|
||||
require.Equal(t, openAIUpstreamCostNeutralFactor, factors[id])
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAIUpstreamCostFactorsCoverageShrinksSparseSignal(t *testing.T) {
|
||||
now := time.Date(2026, 7, 13, 12, 0, 0, 0, time.UTC)
|
||||
accounts := []*Account{
|
||||
upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 0.03, now.Add(-time.Minute), 30*time.Minute),
|
||||
upstreamCostTestAccount(2, UpstreamBillingProbeStatusOK, 0.8, now.Add(-time.Minute), 30*time.Minute),
|
||||
}
|
||||
for id := int64(3); id <= 10; id++ {
|
||||
accounts = append(accounts, &Account{ID: id, Platform: PlatformOpenAI, Type: AccountTypeAPIKey})
|
||||
}
|
||||
|
||||
factors := openAIUpstreamCostFactors(accounts, now, defaultOpenAIOAuthSchedulingRateMultiplier)
|
||||
center := math.Sqrt(0.03 * 0.8)
|
||||
require.InDelta(t, 0.5+0.2*(1/(1+0.03/center)-0.5), factors[1], 1e-12)
|
||||
require.InDelta(t, 0.5+0.2*(1/(1+0.8/center)-0.5), factors[2], 1e-12)
|
||||
require.Equal(t, openAIUpstreamCostNeutralFactor, factors[3])
|
||||
}
|
||||
|
||||
func TestOpenAIUpstreamCostFactorsUseMedianAgainstOutlier(t *testing.T) {
|
||||
now := time.Date(2026, 7, 13, 12, 0, 0, 0, time.UTC)
|
||||
accounts := []*Account{
|
||||
upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 0.1, now.Add(-time.Minute), 30*time.Minute),
|
||||
upstreamCostTestAccount(2, UpstreamBillingProbeStatusOK, 0.2, now.Add(-time.Minute), 30*time.Minute),
|
||||
upstreamCostTestAccount(3, UpstreamBillingProbeStatusOK, 100, now.Add(-time.Minute), 30*time.Minute),
|
||||
}
|
||||
|
||||
factors := openAIUpstreamCostFactors(accounts, now, defaultOpenAIOAuthSchedulingRateMultiplier)
|
||||
require.InDelta(t, 2.0/3.0, factors[1], 1e-12)
|
||||
require.InDelta(t, 0.5, factors[2], 1e-12)
|
||||
require.InDelta(t, 1/(1+100/0.2), factors[3], 1e-12)
|
||||
}
|
||||
|
||||
func TestOpenAILegacyUpstreamRateOrderRequiresComparableRates(t *testing.T) {
|
||||
now := time.Now()
|
||||
oneKnown := newOpenAILegacyUpstreamRateOrder([]*Account{
|
||||
upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 0.03, now.Add(-time.Minute), 30*time.Minute),
|
||||
{ID: 2, Platform: PlatformOpenAI, Type: AccountTypeAPIKey},
|
||||
}, now, defaultOpenAIOAuthSchedulingRateMultiplier)
|
||||
require.False(t, oneKnown.enabled)
|
||||
|
||||
allEqual := newOpenAILegacyUpstreamRateOrder([]*Account{
|
||||
upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 0.3, now.Add(-time.Minute), 30*time.Minute),
|
||||
upstreamCostTestAccount(2, UpstreamBillingProbeStatusOK, 0.3, now.Add(-time.Minute), 30*time.Minute),
|
||||
}, now, defaultOpenAIOAuthSchedulingRateMultiplier)
|
||||
require.False(t, allEqual.enabled)
|
||||
|
||||
distinct := newOpenAILegacyUpstreamRateOrder([]*Account{
|
||||
upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 0.03, now.Add(-time.Minute), 30*time.Minute),
|
||||
upstreamCostTestAccount(2, UpstreamBillingProbeStatusOK, 0.8, now.Add(-time.Minute), 30*time.Minute),
|
||||
{ID: 3, Platform: PlatformOpenAI, Type: AccountTypeAPIKey},
|
||||
}, now, defaultOpenAIOAuthSchedulingRateMultiplier)
|
||||
require.True(t, distinct.enabled)
|
||||
require.Negative(t, distinct.compare(&Account{ID: 1}, &Account{ID: 2}))
|
||||
require.Negative(t, distinct.compare(&Account{ID: 2}, &Account{ID: 3}))
|
||||
}
|
||||
|
||||
func TestOpenAISchedulingRatePlacesOAuthAtConfiguredReference(t *testing.T) {
|
||||
now := time.Now()
|
||||
cheap := upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 0.02, now.Add(-time.Minute), 30*time.Minute)
|
||||
oauth := upstreamCostTestOAuthAccount(2)
|
||||
expensive := upstreamCostTestAccount(3, UpstreamBillingProbeStatusOK, 0.12, now.Add(-time.Minute), 30*time.Minute)
|
||||
|
||||
order := newOpenAILegacyUpstreamRateOrder([]*Account{cheap, oauth, expensive}, now, 0.05)
|
||||
require.True(t, order.enabled)
|
||||
require.Negative(t, order.compare(cheap, oauth))
|
||||
require.Negative(t, order.compare(oauth, expensive))
|
||||
|
||||
factors := openAIUpstreamCostFactors([]*Account{cheap, oauth, expensive}, now, 0.05)
|
||||
require.Greater(t, factors[cheap.ID], factors[oauth.ID])
|
||||
require.Greater(t, factors[oauth.ID], factors[expensive.ID])
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayServiceLegacyLowRatePriorityUsesConfiguredOAuthReference(t *testing.T) {
|
||||
resetOpenAIAdvancedSchedulerSettingCacheForTest()
|
||||
defer resetOpenAIAdvancedSchedulerSettingCacheForTest()
|
||||
|
||||
now := time.Now()
|
||||
cheap := upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 0.02, now.Add(-time.Minute), 30*time.Minute)
|
||||
oauth := upstreamCostTestOAuthAccount(2)
|
||||
expensive := upstreamCostTestAccount(3, UpstreamBillingProbeStatusOK, 0.12, now.Add(-time.Minute), 30*time.Minute)
|
||||
for _, account := range []*Account{cheap, oauth, expensive} {
|
||||
account.Status = StatusActive
|
||||
account.Schedulable = true
|
||||
account.Concurrency = 1
|
||||
}
|
||||
cheap.Priority, oauth.Priority, expensive.Priority = 20, 10, 0
|
||||
|
||||
settings := &openAIAdvancedSchedulerSettingRepoStub{values: map[string]string{
|
||||
openAIAdvancedSchedulerSettingKey: "false",
|
||||
SettingKeyOpenAILowUpstreamRatePriorityEnabled: "true",
|
||||
SettingKeyOpenAIOAuthSchedulingRateMultiplier: "0.05",
|
||||
}}
|
||||
cfg := &config.Config{}
|
||||
svc := &OpenAIGatewayService{
|
||||
accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{*cheap, *oauth, *expensive}},
|
||||
cache: &schedulerTestGatewayCache{},
|
||||
cfg: cfg,
|
||||
rateLimitService: &RateLimitService{settingService: NewSettingService(settings, cfg)},
|
||||
}
|
||||
groupID := int64(1)
|
||||
|
||||
first, _, err := svc.SelectAccountWithScheduler(context.Background(), &groupID, "", "", "gpt-test", nil, OpenAIUpstreamTransportAny, false)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, cheap.ID, first.Account.ID)
|
||||
|
||||
second, _, err := svc.SelectAccountWithScheduler(context.Background(), &groupID, "", "", "gpt-test", map[int64]struct{}{cheap.ID: {}}, OpenAIUpstreamTransportAny, false)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, oauth.ID, second.Account.ID)
|
||||
}
|
||||
|
||||
func TestOpenAIModelsSelectionIgnoresTokenCostSignal(t *testing.T) {
|
||||
resetOpenAIAdvancedSchedulerSettingCacheForTest()
|
||||
defer resetOpenAIAdvancedSchedulerSettingCacheForTest()
|
||||
|
||||
now := time.Now()
|
||||
cheap := upstreamCostTestAccount(51, UpstreamBillingProbeStatusOK, 0.02, now.Add(-time.Minute), 30*time.Minute)
|
||||
expensive := upstreamCostTestAccount(52, UpstreamBillingProbeStatusOK, 0.8, now.Add(-time.Minute), 30*time.Minute)
|
||||
for _, account := range []*Account{cheap, expensive} {
|
||||
account.Status = StatusActive
|
||||
account.Schedulable = true
|
||||
account.Concurrency = 1
|
||||
}
|
||||
cheap.Priority = 10
|
||||
expensive.Priority = 0
|
||||
settings := &openAIAdvancedSchedulerSettingRepoStub{values: map[string]string{
|
||||
SettingKeyOpenAILowUpstreamRatePriorityEnabled: "true",
|
||||
}}
|
||||
cfg := &config.Config{}
|
||||
svc := &OpenAIGatewayService{
|
||||
accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{*cheap, *expensive}},
|
||||
cfg: cfg,
|
||||
rateLimitService: &RateLimitService{settingService: NewSettingService(settings, cfg)},
|
||||
}
|
||||
|
||||
account, err := svc.SelectAccountForModelWithExclusions(context.Background(), nil, "", "", nil)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, expensive.ID, account.ID)
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayServiceLegacyLowRatePriorityIsIndependentFromAdvancedScheduler(t *testing.T) {
|
||||
now := time.Now()
|
||||
cheap := upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 0.03, now.Add(-time.Minute), 30*time.Minute)
|
||||
cheap.Status, cheap.Schedulable, cheap.Concurrency, cheap.Priority = StatusActive, true, 1, 10
|
||||
expensive := upstreamCostTestAccount(2, UpstreamBillingProbeStatusOK, 0.8, now.Add(-time.Minute), 30*time.Minute)
|
||||
expensive.Status, expensive.Schedulable, expensive.Concurrency, expensive.Priority = StatusActive, true, 1, 0
|
||||
accounts := []Account{*cheap, *expensive}
|
||||
groupID := int64(1)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
enabled bool
|
||||
loadBatch bool
|
||||
loadErr error
|
||||
wantID int64
|
||||
}{
|
||||
{name: "switch off keeps priority first", loadBatch: true, wantID: 2},
|
||||
{name: "load batch", enabled: true, loadBatch: true, wantID: 1},
|
||||
{name: "load batch disabled", enabled: true, wantID: 1},
|
||||
{name: "load lookup failure", enabled: true, loadBatch: true, loadErr: errors.New("load unavailable"), wantID: 1},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
resetOpenAIAdvancedSchedulerSettingCacheForTest()
|
||||
settings := &openAIAdvancedSchedulerSettingRepoStub{values: map[string]string{
|
||||
openAIAdvancedSchedulerSettingKey: "false",
|
||||
SettingKeyOpenAILowUpstreamRatePriorityEnabled: strconv.FormatBool(tt.enabled),
|
||||
}}
|
||||
cfg := &config.Config{}
|
||||
cfg.Gateway.Scheduling.LoadBatchEnabled = tt.loadBatch
|
||||
svc := &OpenAIGatewayService{
|
||||
accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts},
|
||||
cache: &schedulerTestGatewayCache{},
|
||||
cfg: cfg,
|
||||
rateLimitService: &RateLimitService{settingService: NewSettingService(settings, cfg)},
|
||||
concurrencyService: NewConcurrencyService(schedulerTestConcurrencyCache{
|
||||
loadBatchErr: tt.loadErr,
|
||||
loadMap: map[int64]*AccountLoadInfo{
|
||||
1: {AccountID: 1, LoadRate: 90},
|
||||
2: {AccountID: 2, LoadRate: 10},
|
||||
},
|
||||
}),
|
||||
}
|
||||
|
||||
selection, _, err := svc.SelectAccountWithScheduler(context.Background(), &groupID, "", "", "gpt-test", nil, OpenAIUpstreamTransportAny, false)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, tt.wantID, selection.Account.ID)
|
||||
if selection.ReleaseFunc != nil {
|
||||
selection.ReleaseFunc()
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayServiceAdvancedSchedulerIgnoresLegacyLowRateSwitch(t *testing.T) {
|
||||
resetOpenAIAdvancedSchedulerSettingCacheForTest()
|
||||
defer resetOpenAIAdvancedSchedulerSettingCacheForTest()
|
||||
now := time.Now()
|
||||
cheap := upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 0.03, now.Add(-time.Minute), 30*time.Minute)
|
||||
cheap.Status, cheap.Schedulable, cheap.Concurrency, cheap.Priority = StatusActive, true, 1, 10
|
||||
expensive := upstreamCostTestAccount(2, UpstreamBillingProbeStatusOK, 0.8, now.Add(-time.Minute), 30*time.Minute)
|
||||
expensive.Status, expensive.Schedulable, expensive.Concurrency, expensive.Priority = StatusActive, true, 1, 0
|
||||
settings := &openAIAdvancedSchedulerSettingRepoStub{values: map[string]string{
|
||||
openAIAdvancedSchedulerSettingKey: "true",
|
||||
SettingKeyOpenAILowUpstreamRatePriorityEnabled: "true",
|
||||
}}
|
||||
cfg := &config.Config{}
|
||||
cfg.Gateway.OpenAIWS.LBTopK = 1
|
||||
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Priority = 1
|
||||
svc := &OpenAIGatewayService{
|
||||
accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{*cheap, *expensive}},
|
||||
cache: &schedulerTestGatewayCache{},
|
||||
cfg: cfg,
|
||||
rateLimitService: &RateLimitService{settingService: NewSettingService(settings, cfg)},
|
||||
concurrencyService: NewConcurrencyService(schedulerTestConcurrencyCache{}),
|
||||
}
|
||||
groupID := int64(1)
|
||||
|
||||
selection, _, err := svc.SelectAccountWithScheduler(context.Background(), &groupID, "", "", "gpt-test", nil, OpenAIUpstreamTransportAny, false)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(2), selection.Account.ID)
|
||||
if selection.ReleaseFunc != nil {
|
||||
selection.ReleaseFunc()
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayServiceLegacyLowRatePrioritySkipsCooledDownAccount(t *testing.T) {
|
||||
resetOpenAIAdvancedSchedulerSettingCacheForTest()
|
||||
defer resetOpenAIAdvancedSchedulerSettingCacheForTest()
|
||||
|
||||
now := time.Now()
|
||||
cheap := upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 0.03, now.Add(-time.Minute), 30*time.Minute)
|
||||
cheap.Status, cheap.Schedulable, cheap.Concurrency, cheap.Priority = StatusActive, true, 1, 10
|
||||
cooldownUntil := now.Add(time.Minute)
|
||||
cheap.TempUnschedulableUntil = &cooldownUntil
|
||||
expensive := upstreamCostTestAccount(2, UpstreamBillingProbeStatusOK, 0.8, now.Add(-time.Minute), 30*time.Minute)
|
||||
expensive.Status, expensive.Schedulable, expensive.Concurrency, expensive.Priority = StatusActive, true, 1, 0
|
||||
settings := &openAIAdvancedSchedulerSettingRepoStub{values: map[string]string{
|
||||
openAIAdvancedSchedulerSettingKey: "false",
|
||||
SettingKeyOpenAILowUpstreamRatePriorityEnabled: "true",
|
||||
}}
|
||||
cfg := &config.Config{}
|
||||
cfg.Gateway.Scheduling.LoadBatchEnabled = true
|
||||
svc := &OpenAIGatewayService{
|
||||
accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{*cheap, *expensive}},
|
||||
cache: &schedulerTestGatewayCache{},
|
||||
cfg: cfg,
|
||||
rateLimitService: &RateLimitService{settingService: NewSettingService(settings, cfg)},
|
||||
concurrencyService: NewConcurrencyService(schedulerTestConcurrencyCache{loadMap: map[int64]*AccountLoadInfo{
|
||||
1: {AccountID: 1},
|
||||
2: {AccountID: 2},
|
||||
}}),
|
||||
}
|
||||
groupID := int64(1)
|
||||
|
||||
selection, _, err := svc.SelectAccountWithScheduler(context.Background(), &groupID, "", "", "gpt-test", nil, OpenAIUpstreamTransportAny, false)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(2), selection.Account.ID)
|
||||
if selection.ReleaseFunc != nil {
|
||||
selection.ReleaseFunc()
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAIFreshUpstreamBillingRateUsesFreshCachedSuccessOnly(t *testing.T) {
|
||||
now := time.Date(2026, 7, 13, 12, 0, 0, 0, time.UTC)
|
||||
tests := []struct {
|
||||
name string
|
||||
account *Account
|
||||
wantOK bool
|
||||
}{
|
||||
{name: "fresh", account: upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 0.3, now.Add(-time.Minute), 30*time.Minute), wantOK: true},
|
||||
{name: "zero rate", account: upstreamCostTestAccount(2, UpstreamBillingProbeStatusOK, 0, now.Add(-time.Minute), 30*time.Minute), wantOK: true},
|
||||
{name: "transient failure with fresh cache", account: upstreamCostTestAccount(3, UpstreamBillingProbeStatusFailed, 0.3, now.Add(-time.Minute), 30*time.Minute), wantOK: true},
|
||||
{name: "stale", account: upstreamCostTestAccount(4, UpstreamBillingProbeStatusOK, 0.3, now.Add(-61*time.Minute), 30*time.Minute)},
|
||||
{name: "future", account: upstreamCostTestAccount(5, UpstreamBillingProbeStatusOK, 0.3, now.Add(time.Minute), 30*time.Minute)},
|
||||
{name: "unsupported", account: upstreamCostTestAccount(6, UpstreamBillingProbeStatusUnsupported, 0.3, now.Add(-time.Minute), 30*time.Minute)},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
_, ok := openAIFreshUpstreamBillingRate(tt.account, now)
|
||||
require.Equal(t, tt.wantOK, ok)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildOpenAISelectionOrderIncludesOverflowOnlyForCostScheduling(t *testing.T) {
|
||||
scheduler := &defaultOpenAIAccountScheduler{}
|
||||
candidates := []openAIAccountCandidateScore{
|
||||
{account: &Account{ID: 1}, loadInfo: &AccountLoadInfo{}, score: 3},
|
||||
{account: &Account{ID: 2}, loadInfo: &AccountLoadInfo{}, score: 2},
|
||||
{account: &Account{ID: 3}, loadInfo: &AccountLoadInfo{}, score: 1},
|
||||
}
|
||||
|
||||
legacy := scheduler.buildOpenAISelectionOrder(OpenAIAccountScheduleRequest{}, openAIAccountLoadPlan{
|
||||
candidates: candidates,
|
||||
topK: 1,
|
||||
})
|
||||
require.Len(t, legacy, 1)
|
||||
|
||||
costAware := scheduler.buildOpenAISelectionOrder(OpenAIAccountScheduleRequest{}, openAIAccountLoadPlan{
|
||||
candidates: candidates,
|
||||
topK: 1,
|
||||
includeOverflowFallback: true,
|
||||
})
|
||||
require.Equal(t, []int64{1, 2, 3}, []int64{
|
||||
costAware[0].account.ID,
|
||||
costAware[1].account.ID,
|
||||
costAware[2].account.ID,
|
||||
})
|
||||
}
|
||||
|
||||
func TestBuildOpenAIAccountLoadPlanUsesCostOnlyForTokenScope(t *testing.T) {
|
||||
resetOpenAIAdvancedSchedulerSettingCacheForTest()
|
||||
defer resetOpenAIAdvancedSchedulerSettingCacheForTest()
|
||||
|
||||
now := time.Now()
|
||||
accounts := []*Account{
|
||||
upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 0.03, now.Add(-time.Minute), 30*time.Minute),
|
||||
upstreamCostTestOAuthAccount(2),
|
||||
upstreamCostTestAccount(3, UpstreamBillingProbeStatusOK, 0.8, now.Add(-time.Minute), 30*time.Minute),
|
||||
}
|
||||
cfg := &config.Config{}
|
||||
cfg.Gateway.OpenAIWS.LBTopK = 1
|
||||
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.UpstreamCost = 1.5
|
||||
settings := &openAIAdvancedSchedulerSettingRepoStub{values: map[string]string{
|
||||
SettingKeyOpenAIOAuthSchedulingRateMultiplier: "0.05",
|
||||
}}
|
||||
scheduler := &defaultOpenAIAccountScheduler{service: &OpenAIGatewayService{
|
||||
cfg: cfg,
|
||||
rateLimitService: &RateLimitService{settingService: NewSettingService(settings, cfg)},
|
||||
}}
|
||||
loadMap := map[int64]*AccountLoadInfo{
|
||||
1: {AccountID: 1},
|
||||
2: {AccountID: 2},
|
||||
3: {AccountID: 3},
|
||||
}
|
||||
|
||||
tokenPlan := scheduler.buildOpenAIAccountLoadPlan(context.Background(), OpenAIAccountScheduleRequest{UseUpstreamTokenCost: true}, accounts, loadMap)
|
||||
require.Greater(t, tokenPlan.candidates[0].score, tokenPlan.candidates[1].score)
|
||||
require.Greater(t, tokenPlan.candidates[1].score, tokenPlan.candidates[2].score)
|
||||
require.True(t, tokenPlan.includeOverflowFallback)
|
||||
|
||||
otherPlan := scheduler.buildOpenAIAccountLoadPlan(context.Background(), OpenAIAccountScheduleRequest{}, accounts, loadMap)
|
||||
require.Equal(t, otherPlan.candidates[0].score, otherPlan.candidates[1].score)
|
||||
require.Equal(t, otherPlan.candidates[1].score, otherPlan.candidates[2].score)
|
||||
require.False(t, otherPlan.includeOverflowFallback)
|
||||
}
|
||||
|
||||
func TestBuildOpenAIAccountSchedulerScoreSnapshotUpstreamCostIsExactNoOpWithoutSignal(t *testing.T) {
|
||||
accounts := []*Account{
|
||||
{ID: 1, Platform: PlatformOpenAI, Type: AccountTypeAPIKey},
|
||||
{ID: 2, Platform: PlatformOpenAI, Type: AccountTypeAPIKey},
|
||||
}
|
||||
loadMap := map[int64]*AccountLoadInfo{
|
||||
1: {AccountID: 1, LoadRate: 20},
|
||||
2: {AccountID: 2, LoadRate: 80},
|
||||
}
|
||||
weights := GatewayOpenAIWSSchedulerScoreWeightsView{Priority: 1, Load: 1, Queue: 0.7, ErrorRate: 0.8, TTFT: 0.5}
|
||||
baseline := buildOpenAIAccountSchedulerScoreSnapshot(accounts, loadMap, weights, false, defaultOpenAIOAuthSchedulingRateMultiplier)
|
||||
weights.UpstreamCost = 1.5
|
||||
withCost := buildOpenAIAccountSchedulerScoreSnapshot(accounts, loadMap, weights, false, defaultOpenAIOAuthSchedulingRateMultiplier)
|
||||
|
||||
require.Equal(t, baseline, withCost)
|
||||
}
|
||||
|
||||
func TestBuildOpenAIAccountSchedulerScoreSnapshotUsesUpstreamCostSignal(t *testing.T) {
|
||||
now := time.Now()
|
||||
accounts := []*Account{
|
||||
upstreamCostTestAccount(1, UpstreamBillingProbeStatusOK, 0.03, now.Add(-time.Minute), 30*time.Minute),
|
||||
upstreamCostTestAccount(2, UpstreamBillingProbeStatusOK, 0.8, now.Add(-time.Minute), 30*time.Minute),
|
||||
}
|
||||
weights := GatewayOpenAIWSSchedulerScoreWeightsView{UpstreamCost: 1.5}
|
||||
scores := buildOpenAIAccountSchedulerScoreSnapshot(accounts, nil, weights, false, defaultOpenAIOAuthSchedulingRateMultiplier)
|
||||
|
||||
require.Greater(t, scores[1].BaseScore, scores[2].BaseScore)
|
||||
}
|
||||
@@ -20,6 +20,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_UsesWSPassthroughSnapsh
|
||||
Status: StatusActive,
|
||||
Schedulable: true,
|
||||
Concurrency: 10,
|
||||
GroupIDs: []int64{groupID},
|
||||
Extra: map[string]any{
|
||||
"openai_oauth_responses_websockets_v2_mode": OpenAIWSIngressModePassthrough,
|
||||
},
|
||||
|
||||
@@ -159,7 +159,7 @@ func (s *OpenAIGatewayService) SelectAccountForModel(ctx context.Context, groupI
|
||||
// SelectAccountForModelWithExclusions selects an account supporting the requested model while excluding specified accounts.
|
||||
// SelectAccountForModelWithExclusions 选择支持指定模型的账号,同时排除指定的账号。
|
||||
func (s *OpenAIGatewayService) SelectAccountForModelWithExclusions(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}) (*Account, error) {
|
||||
return s.selectAccountForModelWithExclusions(s.withOpenAIQuotaAutoPauseContext(ctx), groupID, PlatformOpenAI, sessionHash, requestedModel, excludedIDs, false, 0, "")
|
||||
return s.selectAccountForModelWithExclusions(s.withOpenAIQuotaAutoPauseContext(ctx), groupID, PlatformOpenAI, sessionHash, requestedModel, excludedIDs, false, 0, "", false)
|
||||
}
|
||||
|
||||
// noAvailableOpenAISelectionError builds the standard "no account available" error
|
||||
@@ -570,7 +570,7 @@ func resolveOpenAIAccountUpstreamModelForRequest(account *Account, requestedMode
|
||||
return upstreamModel
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) selectAccountForModelWithExclusions(ctx context.Context, groupID *int64, platform string, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, stickyAccountID int64, requiredCapability OpenAIEndpointCapability) (*Account, error) {
|
||||
func (s *OpenAIGatewayService) selectAccountForModelWithExclusions(ctx context.Context, groupID *int64, platform string, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, stickyAccountID int64, requiredCapability OpenAIEndpointCapability, preferLowUpstreamRate bool) (*Account, error) {
|
||||
platform = normalizeOpenAICompatiblePlatform(platform)
|
||||
if s.checkChannelPricingRestriction(ctx, groupID, requestedModel) {
|
||||
slog.Warn("channel pricing restriction blocked request",
|
||||
@@ -594,7 +594,7 @@ func (s *OpenAIGatewayService) selectAccountForModelWithExclusions(ctx context.C
|
||||
|
||||
// 3. 按优先级 + LRU 选择最佳账号
|
||||
// Select by priority + LRU
|
||||
selected, compactBlocked := s.selectBestAccount(ctx, groupID, platform, accounts, requestedModel, excludedIDs, requireCompact, requiredCapability)
|
||||
selected, compactBlocked := s.selectBestAccount(ctx, groupID, platform, accounts, requestedModel, excludedIDs, requireCompact, requiredCapability, preferLowUpstreamRate)
|
||||
|
||||
if selected == nil {
|
||||
return nil, noAvailableOpenAISelectionError(requestedModel, compactBlocked)
|
||||
@@ -663,8 +663,8 @@ func (s *OpenAIGatewayService) tryStickySessionHit(ctx context.Context, groupID
|
||||
_ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash)
|
||||
return nil
|
||||
}
|
||||
account = s.recheckSelectedOpenAIAccountFromDB(ctx, account, platform, requestedModel, requireCompact, requiredCapability)
|
||||
if account == nil || !openAIStickyAccountMatchesGroup(account, groupID) {
|
||||
account = s.recheckSelectedOpenAIAccountFromDB(ctx, account, groupID, platform, requestedModel, requireCompact, requiredCapability)
|
||||
if account == nil || !s.openAIAccountMatchesSchedulingGroup(account, groupID) {
|
||||
_ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash)
|
||||
return nil
|
||||
}
|
||||
@@ -687,12 +687,12 @@ func (s *OpenAIGatewayService) tryStickySessionHit(ctx context.Context, groupID
|
||||
// Returns nil if no available account. The second return reports whether at
|
||||
// least one candidate was filtered out solely because it lacks compact support
|
||||
// (only meaningful when requireCompact=true).
|
||||
func (s *OpenAIGatewayService) selectBestAccount(ctx context.Context, groupID *int64, platform string, accounts []Account, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, requiredCapability OpenAIEndpointCapability) (*Account, bool) {
|
||||
func (s *OpenAIGatewayService) selectBestAccount(ctx context.Context, groupID *int64, platform string, accounts []Account, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, requiredCapability OpenAIEndpointCapability, preferLowUpstreamRate bool) (*Account, bool) {
|
||||
platform = normalizeOpenAICompatiblePlatform(platform)
|
||||
var selected *Account
|
||||
selectedCompactTier := -1
|
||||
compactBlocked := false
|
||||
needsUpstreamCheck := s.needsUpstreamChannelRestrictionCheck(ctx, groupID)
|
||||
eligible := make([]*Account, 0, len(accounts))
|
||||
compactTiers := make(map[int64]int, len(accounts))
|
||||
|
||||
for i := range accounts {
|
||||
acc := &accounts[i]
|
||||
@@ -707,7 +707,7 @@ func (s *OpenAIGatewayService) selectBestAccount(ctx context.Context, groupID *i
|
||||
if fresh == nil {
|
||||
continue
|
||||
}
|
||||
fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, platform, requestedModel, false, requiredCapability)
|
||||
fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, groupID, platform, requestedModel, false, requiredCapability)
|
||||
if fresh == nil {
|
||||
continue
|
||||
}
|
||||
@@ -723,30 +723,28 @@ func (s *OpenAIGatewayService) selectBestAccount(ctx context.Context, groupID *i
|
||||
}
|
||||
}
|
||||
|
||||
// 选择优先级最高且最久未使用的账号
|
||||
// Select highest priority and least recently used
|
||||
if selected == nil {
|
||||
selected = fresh
|
||||
selectedCompactTier = compactTier
|
||||
continue
|
||||
}
|
||||
|
||||
// compact 模式下高 tier 优先;同 tier 内才比较 priority/LRU。
|
||||
if requireCompact && compactTier != selectedCompactTier {
|
||||
if compactTier > selectedCompactTier {
|
||||
selected = fresh
|
||||
selectedCompactTier = compactTier
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
if s.isBetterAccount(fresh, selected) {
|
||||
selected = fresh
|
||||
selectedCompactTier = compactTier
|
||||
}
|
||||
eligible = append(eligible, fresh)
|
||||
compactTiers[fresh.ID] = compactTier
|
||||
}
|
||||
|
||||
return selected, compactBlocked
|
||||
if len(eligible) == 0 {
|
||||
return nil, compactBlocked
|
||||
}
|
||||
rateOrder := openAILegacyUpstreamRateOrder{}
|
||||
if preferLowUpstreamRate {
|
||||
rateOrder = newOpenAILegacyUpstreamRateOrder(eligible, time.Now(), s.openAIOAuthSchedulingRateMultiplier(ctx))
|
||||
}
|
||||
sort.SliceStable(eligible, func(i, j int) bool {
|
||||
a, b := eligible[i], eligible[j]
|
||||
if requireCompact && compactTiers[a.ID] != compactTiers[b.ID] {
|
||||
return compactTiers[a.ID] > compactTiers[b.ID]
|
||||
}
|
||||
if rateCmp := rateOrder.compare(a, b); rateCmp != 0 {
|
||||
return rateCmp < 0
|
||||
}
|
||||
return s.isBetterAccount(a, b)
|
||||
})
|
||||
return eligible[0], compactBlocked
|
||||
}
|
||||
|
||||
// isBetterAccount 判断 candidate 是否比 current 更优。
|
||||
@@ -784,10 +782,10 @@ func (s *OpenAIGatewayService) isBetterAccount(candidate, current *Account) bool
|
||||
|
||||
// SelectAccountWithLoadAwareness selects an account with load-awareness and wait plan.
|
||||
func (s *OpenAIGatewayService) SelectAccountWithLoadAwareness(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}) (*AccountSelectionResult, error) {
|
||||
return s.selectAccountWithLoadAwareness(s.withOpenAIQuotaAutoPauseContext(ctx), groupID, PlatformOpenAI, sessionHash, requestedModel, excludedIDs, false, "")
|
||||
return s.selectAccountWithLoadAwareness(s.withOpenAIQuotaAutoPauseContext(ctx), groupID, PlatformOpenAI, sessionHash, requestedModel, excludedIDs, false, "", true)
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Context, groupID *int64, platform string, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, requiredCapability OpenAIEndpointCapability) (*AccountSelectionResult, error) {
|
||||
func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Context, groupID *int64, platform string, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, requiredCapability OpenAIEndpointCapability, useUpstreamTokenCost bool) (*AccountSelectionResult, error) {
|
||||
platform = normalizeOpenAICompatiblePlatform(platform)
|
||||
if s.checkChannelPricingRestriction(ctx, groupID, requestedModel) {
|
||||
slog.Warn("channel pricing restriction blocked request",
|
||||
@@ -797,6 +795,7 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex
|
||||
}
|
||||
|
||||
cfg := s.schedulingConfig()
|
||||
preferLowUpstreamRate := useUpstreamTokenCost && s.isOpenAILowUpstreamRatePriorityEnabled(ctx)
|
||||
needsUpstreamCheck := s.needsUpstreamChannelRestrictionCheck(ctx, groupID)
|
||||
var stickyAccountID int64
|
||||
if sessionHash != "" && s.cache != nil {
|
||||
@@ -805,7 +804,7 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex
|
||||
}
|
||||
}
|
||||
if s.concurrencyService == nil || !cfg.LoadBatchEnabled {
|
||||
account, err := s.selectAccountForModelWithExclusions(ctx, groupID, platform, sessionHash, requestedModel, excludedIDs, requireCompact, stickyAccountID, requiredCapability)
|
||||
account, err := s.selectAccountForModelWithExclusions(ctx, groupID, platform, sessionHash, requestedModel, excludedIDs, requireCompact, stickyAccountID, requiredCapability, preferLowUpstreamRate)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -859,10 +858,10 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex
|
||||
_ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash)
|
||||
}
|
||||
if !clearSticky && isOpenAICompatibleAccountEligibleForRequest(ctx, account, platform, requestedModel, false, requiredCapability) {
|
||||
account = s.recheckSelectedOpenAIAccountFromDB(ctx, account, platform, requestedModel, requireCompact, requiredCapability)
|
||||
account = s.recheckSelectedOpenAIAccountFromDB(ctx, account, groupID, platform, requestedModel, requireCompact, requiredCapability)
|
||||
if account == nil {
|
||||
_ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash)
|
||||
} else if !openAIStickyAccountMatchesGroup(account, groupID) {
|
||||
} else if !s.openAIAccountMatchesSchedulingGroup(account, groupID) {
|
||||
_ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash)
|
||||
} else if s.isOpenAIAccountRuntimeBlocked(account) {
|
||||
_ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash)
|
||||
@@ -940,6 +939,10 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex
|
||||
if len(candidates) == 0 {
|
||||
return nil, ErrNoAvailableAccounts
|
||||
}
|
||||
rateOrder := openAILegacyUpstreamRateOrder{}
|
||||
if preferLowUpstreamRate {
|
||||
rateOrder = newOpenAILegacyUpstreamRateOrder(candidates, time.Now(), s.openAIOAuthSchedulingRateMultiplier(ctx))
|
||||
}
|
||||
|
||||
accountLoads := make([]AccountWithConcurrency, 0, len(candidates))
|
||||
for _, acc := range candidates {
|
||||
@@ -988,6 +991,11 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex
|
||||
}
|
||||
})
|
||||
shuffleWithinSortGroups(available)
|
||||
if rateOrder.enabled {
|
||||
sort.SliceStable(available, func(i, j int) bool {
|
||||
return rateOrder.compare(available[i].account, available[j].account) < 0
|
||||
})
|
||||
}
|
||||
|
||||
selectionOrder := make([]accountWithLoad, 0, len(available))
|
||||
if requireCompact {
|
||||
@@ -1013,7 +1021,7 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex
|
||||
if fresh == nil {
|
||||
continue
|
||||
}
|
||||
fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, platform, requestedModel, requireCompact, requiredCapability)
|
||||
fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, groupID, platform, requestedModel, requireCompact, requiredCapability)
|
||||
if fresh == nil {
|
||||
continue
|
||||
}
|
||||
@@ -1039,6 +1047,11 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex
|
||||
if err != nil {
|
||||
ordered := append([]*Account(nil), candidates...)
|
||||
sortAccountsByPriorityAndLastUsed(ordered, false)
|
||||
if rateOrder.enabled {
|
||||
sort.SliceStable(ordered, func(i, j int) bool {
|
||||
return rateOrder.compare(ordered[i], ordered[j]) < 0
|
||||
})
|
||||
}
|
||||
if requireCompact {
|
||||
ordered = prioritizeOpenAICompactAccounts(ordered)
|
||||
}
|
||||
@@ -1047,7 +1060,7 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex
|
||||
if fresh == nil {
|
||||
continue
|
||||
}
|
||||
fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, platform, requestedModel, requireCompact, requiredCapability)
|
||||
fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, groupID, platform, requestedModel, requireCompact, requiredCapability)
|
||||
if fresh == nil {
|
||||
continue
|
||||
}
|
||||
@@ -1084,6 +1097,11 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex
|
||||
|
||||
// ============ Layer 3: Fallback wait ============
|
||||
sortAccountsByPriorityAndLastUsed(candidates, false)
|
||||
if rateOrder.enabled {
|
||||
sort.SliceStable(candidates, func(i, j int) bool {
|
||||
return rateOrder.compare(candidates[i], candidates[j]) < 0
|
||||
})
|
||||
}
|
||||
if requireCompact {
|
||||
candidates = prioritizeOpenAICompactAccounts(candidates)
|
||||
}
|
||||
@@ -1092,7 +1110,7 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex
|
||||
if fresh == nil {
|
||||
continue
|
||||
}
|
||||
fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, platform, requestedModel, requireCompact, requiredCapability)
|
||||
fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, groupID, platform, requestedModel, requireCompact, requiredCapability)
|
||||
if fresh == nil {
|
||||
continue
|
||||
}
|
||||
@@ -1182,7 +1200,7 @@ func (s *OpenAIGatewayService) parentAccountLookup(ctx context.Context) func(int
|
||||
}
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) recheckSelectedOpenAIAccountFromDB(ctx context.Context, account *Account, platform string, requestedModel string, requireCompact bool, requiredCapability OpenAIEndpointCapability) *Account {
|
||||
func (s *OpenAIGatewayService) recheckSelectedOpenAIAccountFromDB(ctx context.Context, account *Account, groupID *int64, platform string, requestedModel string, requireCompact bool, requiredCapability OpenAIEndpointCapability) *Account {
|
||||
if account == nil {
|
||||
return nil
|
||||
}
|
||||
@@ -1201,6 +1219,9 @@ func (s *OpenAIGatewayService) recheckSelectedOpenAIAccountFromDB(ctx context.Co
|
||||
if err != nil || latest == nil {
|
||||
return nil
|
||||
}
|
||||
if !s.openAIAccountMatchesSchedulingGroup(latest, groupID) {
|
||||
return nil
|
||||
}
|
||||
if !isOpenAICompatibleAccountEligibleForRequest(ctx, latest, platform, requestedModel, requireCompact, requiredCapability) {
|
||||
return nil
|
||||
}
|
||||
@@ -1213,6 +1234,13 @@ func (s *OpenAIGatewayService) recheckSelectedOpenAIAccountFromDB(ctx context.Co
|
||||
return latest
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) openAIAccountMatchesSchedulingGroup(account *Account, groupID *int64) bool {
|
||||
if s != nil && s.cfg != nil && s.cfg.RunMode == config.RunModeSimple {
|
||||
return account != nil
|
||||
}
|
||||
return openAIStickyAccountMatchesGroup(account, groupID)
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) getSchedulableAccount(ctx context.Context, accountID int64) (*Account, error) {
|
||||
var (
|
||||
account *Account
|
||||
|
||||
@@ -205,6 +205,8 @@ func (s *SettingService) InitializeDefaultSettings(ctx context.Context) error {
|
||||
|
||||
// 分组隔离(默认不允许未分组 Key 调度)
|
||||
SettingKeyAllowUngroupedKeyScheduling: "false",
|
||||
SettingKeyOpenAILowUpstreamRatePriorityEnabled: "false",
|
||||
SettingKeyOpenAIOAuthSchedulingRateMultiplier: "1",
|
||||
SettingKeyEnableAnthropicCacheTTL1hInjection: "false",
|
||||
SettingKeyRewriteMessageCacheControl: strconv.FormatBool(s.defaultRewriteMessageCacheControl()),
|
||||
SettingKeyEnableClientDatelineNormalization: "true",
|
||||
@@ -225,6 +227,7 @@ func (s *SettingService) InitializeDefaultSettings(ctx context.Context) error {
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightTTFT: "",
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightReset: "",
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightQuotaHeadroom: "",
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightUpstreamCost: "",
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightPreviousResponse: "",
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightSessionSticky: "",
|
||||
|
||||
@@ -789,6 +792,8 @@ func (s *SettingService) parseSettings(settings map[string]string) *SystemSettin
|
||||
result.PaymentVisibleMethodWxpaySource = NormalizeVisibleMethodSource("wxpay", settings[SettingPaymentVisibleMethodWxpaySource])
|
||||
result.PaymentVisibleMethodAlipayEnabled = settings[SettingPaymentVisibleMethodAlipayEnabled] == "true"
|
||||
result.PaymentVisibleMethodWxpayEnabled = settings[SettingPaymentVisibleMethodWxpayEnabled] == "true"
|
||||
result.OpenAILowUpstreamRatePriorityEnabled = settings[SettingKeyOpenAILowUpstreamRatePriorityEnabled] == "true"
|
||||
result.OpenAIOAuthSchedulingRateMultiplier = parseOpenAIOAuthSchedulingRateMultiplier(settings[SettingKeyOpenAIOAuthSchedulingRateMultiplier])
|
||||
result.OpenAIAdvancedSchedulerEnabled = settings[openAIAdvancedSchedulerSettingKey] == "true"
|
||||
result.OpenAIAdvancedSchedulerStickyWeightedEnabled = settings[SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled] == "true"
|
||||
result.OpenAIAdvancedSchedulerSubscriptionPriorityEnabled = settings[SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled] == "true"
|
||||
@@ -800,6 +805,7 @@ func (s *SettingService) parseSettings(settings map[string]string) *SystemSettin
|
||||
result.OpenAIAdvancedSchedulerWeightTTFT = strings.TrimSpace(settings[SettingKeyOpenAIAdvancedSchedulerWeightTTFT])
|
||||
result.OpenAIAdvancedSchedulerWeightReset = strings.TrimSpace(settings[SettingKeyOpenAIAdvancedSchedulerWeightReset])
|
||||
result.OpenAIAdvancedSchedulerWeightQuotaHeadroom = strings.TrimSpace(settings[SettingKeyOpenAIAdvancedSchedulerWeightQuotaHeadroom])
|
||||
result.OpenAIAdvancedSchedulerWeightUpstreamCost = strings.TrimSpace(settings[SettingKeyOpenAIAdvancedSchedulerWeightUpstreamCost])
|
||||
result.OpenAIAdvancedSchedulerWeightPreviousResponse = strings.TrimSpace(settings[SettingKeyOpenAIAdvancedSchedulerWeightPreviousResponse])
|
||||
result.OpenAIAdvancedSchedulerWeightSessionSticky = strings.TrimSpace(settings[SettingKeyOpenAIAdvancedSchedulerWeightSessionSticky])
|
||||
result.OpenAIAdvancedSchedulerEffectiveLBTopK = s.openAIAdvancedSchedulerEffectiveLBTopK()
|
||||
@@ -811,6 +817,7 @@ func (s *SettingService) parseSettings(settings map[string]string) *SystemSettin
|
||||
result.OpenAIAdvancedSchedulerEffectiveWeightTTFT = formatOpenAIAdvancedSchedulerFloat(effectiveWeights.TTFT)
|
||||
result.OpenAIAdvancedSchedulerEffectiveWeightReset = formatOpenAIAdvancedSchedulerFloat(effectiveWeights.Reset)
|
||||
result.OpenAIAdvancedSchedulerEffectiveWeightQuotaHeadroom = formatOpenAIAdvancedSchedulerFloat(effectiveWeights.QuotaHeadroom)
|
||||
result.OpenAIAdvancedSchedulerEffectiveWeightUpstreamCost = formatOpenAIAdvancedSchedulerFloat(effectiveWeights.UpstreamCost)
|
||||
result.OpenAIAdvancedSchedulerEffectiveWeightPreviousResponse = formatOpenAIAdvancedSchedulerFloat(effectiveWeights.PreviousResponse)
|
||||
result.OpenAIAdvancedSchedulerEffectiveWeightSessionSticky = formatOpenAIAdvancedSchedulerFloat(effectiveWeights.SessionSticky)
|
||||
|
||||
@@ -901,6 +908,7 @@ func (s *SettingService) openAIAdvancedSchedulerEffectiveWeights() config.Gatewa
|
||||
TTFT: 0.5,
|
||||
Reset: 0.0,
|
||||
QuotaHeadroom: 0.0,
|
||||
UpstreamCost: 0.0,
|
||||
PreviousResponse: 5.0,
|
||||
SessionSticky: 3.0,
|
||||
}
|
||||
@@ -909,8 +917,7 @@ func (s *SettingService) openAIAdvancedSchedulerEffectiveWeights() config.Gatewa
|
||||
}
|
||||
|
||||
weights := s.cfg.Gateway.OpenAIWS.SchedulerScoreWeights
|
||||
baseSum := weights.Priority + weights.Load + weights.Queue + weights.ErrorRate + weights.TTFT + weights.QuotaHeadroom
|
||||
if baseSum <= 0 {
|
||||
if !weights.IsValid() {
|
||||
return defaults
|
||||
}
|
||||
return weights
|
||||
@@ -921,6 +928,10 @@ func formatOpenAIAdvancedSchedulerFloat(value float64) string {
|
||||
}
|
||||
|
||||
func (s *SettingService) normalizeOpenAIAdvancedSchedulerOverrides(settings *SystemSettings) error {
|
||||
if rate := settings.OpenAIOAuthSchedulingRateMultiplier; rate < 0 || math.IsNaN(rate) || math.IsInf(rate, 0) {
|
||||
return infraerrors.BadRequest("INVALID_OPENAI_OAUTH_SCHEDULING_RATE_MULTIPLIER", "OpenAI OAuth scheduling rate multiplier must be a finite non-negative number")
|
||||
}
|
||||
|
||||
lbTopK, err := normalizeOptionalPositiveIntString(settings.OpenAIAdvancedSchedulerLBTopK)
|
||||
if err != nil {
|
||||
return infraerrors.BadRequest("INVALID_OPENAI_ADVANCED_SCHEDULER_LB_TOP_K", "openai advanced scheduler TopK must be a positive integer or empty")
|
||||
@@ -935,6 +946,7 @@ func (s *SettingService) normalizeOpenAIAdvancedSchedulerOverrides(settings *Sys
|
||||
&settings.OpenAIAdvancedSchedulerWeightTTFT,
|
||||
&settings.OpenAIAdvancedSchedulerWeightReset,
|
||||
&settings.OpenAIAdvancedSchedulerWeightQuotaHeadroom,
|
||||
&settings.OpenAIAdvancedSchedulerWeightUpstreamCost,
|
||||
&settings.OpenAIAdvancedSchedulerWeightPreviousResponse,
|
||||
&settings.OpenAIAdvancedSchedulerWeightSessionSticky,
|
||||
}
|
||||
@@ -950,18 +962,32 @@ func (s *SettingService) normalizeOpenAIAdvancedSchedulerOverrides(settings *Sys
|
||||
// 覆盖值(空则回退到生效的配置值)叠加后的基础权重和不允许为 0,
|
||||
// 否则调度会静默退化为 TopK 内均匀随机。
|
||||
effective := s.openAIAdvancedSchedulerEffectiveWeights()
|
||||
baseSum := resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightPriority, effective.Priority) +
|
||||
resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightLoad, effective.Load) +
|
||||
resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightQueue, effective.Queue) +
|
||||
resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightErrorRate, effective.ErrorRate) +
|
||||
resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightTTFT, effective.TTFT) +
|
||||
resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightQuotaHeadroom, effective.QuotaHeadroom)
|
||||
if baseSum <= 0 {
|
||||
return infraerrors.BadRequest("INVALID_OPENAI_ADVANCED_SCHEDULER_WEIGHT", "openai advanced scheduler base weights must not all be zero")
|
||||
resolved := config.GatewayOpenAIWSSchedulerScoreWeights{
|
||||
Priority: resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightPriority, effective.Priority),
|
||||
Load: resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightLoad, effective.Load),
|
||||
Queue: resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightQueue, effective.Queue),
|
||||
ErrorRate: resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightErrorRate, effective.ErrorRate),
|
||||
TTFT: resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightTTFT, effective.TTFT),
|
||||
Reset: resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightReset, effective.Reset),
|
||||
QuotaHeadroom: resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightQuotaHeadroom, effective.QuotaHeadroom),
|
||||
UpstreamCost: resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightUpstreamCost, effective.UpstreamCost),
|
||||
PreviousResponse: resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightPreviousResponse, effective.PreviousResponse),
|
||||
SessionSticky: resolveOpenAIAdvancedSchedulerWeight(settings.OpenAIAdvancedSchedulerWeightSessionSticky, effective.SessionSticky),
|
||||
}
|
||||
if !resolved.IsValid() {
|
||||
return infraerrors.BadRequest("INVALID_OPENAI_ADVANCED_SCHEDULER_WEIGHT", "openai advanced scheduler weights must have finite non-zero base and total sums")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func parseOpenAIOAuthSchedulingRateMultiplier(raw string) float64 {
|
||||
value, err := strconv.ParseFloat(strings.TrimSpace(raw), 64)
|
||||
if err != nil || value < 0 || math.IsNaN(value) || math.IsInf(value, 0) {
|
||||
return defaultOpenAIOAuthSchedulingRateMultiplier
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
// resolveOpenAIAdvancedSchedulerWeight 返回覆盖值(已归一化的非空字符串),空则回退默认值。
|
||||
func resolveOpenAIAdvancedSchedulerWeight(normalized string, fallback float64) float64 {
|
||||
if normalized == "" {
|
||||
|
||||
@@ -5,6 +5,8 @@ package service
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"math"
|
||||
"strconv"
|
||||
"testing"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/config"
|
||||
@@ -339,6 +341,8 @@ func TestSettingService_UpdateSettings_PaymentVisibleMethodsAndAdvancedScheduler
|
||||
PaymentVisibleMethodWxpaySource: "easypay",
|
||||
PaymentVisibleMethodAlipayEnabled: true,
|
||||
PaymentVisibleMethodWxpayEnabled: false,
|
||||
OpenAILowUpstreamRatePriorityEnabled: true,
|
||||
OpenAIOAuthSchedulingRateMultiplier: 0.05,
|
||||
OpenAIAdvancedSchedulerEnabled: true,
|
||||
OpenAIAdvancedSchedulerStickyWeightedEnabled: true,
|
||||
OpenAIAdvancedSchedulerSubscriptionPriorityEnabled: true,
|
||||
@@ -350,6 +354,7 @@ func TestSettingService_UpdateSettings_PaymentVisibleMethodsAndAdvancedScheduler
|
||||
OpenAIAdvancedSchedulerWeightTTFT: "0.5",
|
||||
OpenAIAdvancedSchedulerWeightReset: "",
|
||||
OpenAIAdvancedSchedulerWeightQuotaHeadroom: "0.2",
|
||||
OpenAIAdvancedSchedulerWeightUpstreamCost: "1.5",
|
||||
OpenAIAdvancedSchedulerWeightPreviousResponse: "8",
|
||||
OpenAIAdvancedSchedulerWeightSessionSticky: "4",
|
||||
})
|
||||
@@ -358,6 +363,8 @@ func TestSettingService_UpdateSettings_PaymentVisibleMethodsAndAdvancedScheduler
|
||||
require.Equal(t, VisibleMethodSourceEasyPayWechat, repo.updates[SettingPaymentVisibleMethodWxpaySource])
|
||||
require.Equal(t, "true", repo.updates[SettingPaymentVisibleMethodAlipayEnabled])
|
||||
require.Equal(t, "false", repo.updates[SettingPaymentVisibleMethodWxpayEnabled])
|
||||
require.Equal(t, "true", repo.updates[SettingKeyOpenAILowUpstreamRatePriorityEnabled])
|
||||
require.Equal(t, "0.05", repo.updates[SettingKeyOpenAIOAuthSchedulingRateMultiplier])
|
||||
require.Equal(t, "true", repo.updates[openAIAdvancedSchedulerSettingKey])
|
||||
require.Equal(t, "true", repo.updates[SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled])
|
||||
require.Equal(t, "true", repo.updates[SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled])
|
||||
@@ -369,10 +376,81 @@ func TestSettingService_UpdateSettings_PaymentVisibleMethodsAndAdvancedScheduler
|
||||
require.Equal(t, "0.5", repo.updates[SettingKeyOpenAIAdvancedSchedulerWeightTTFT])
|
||||
require.Equal(t, "", repo.updates[SettingKeyOpenAIAdvancedSchedulerWeightReset])
|
||||
require.Equal(t, "0.2", repo.updates[SettingKeyOpenAIAdvancedSchedulerWeightQuotaHeadroom])
|
||||
require.Equal(t, "1.5", repo.updates[SettingKeyOpenAIAdvancedSchedulerWeightUpstreamCost])
|
||||
require.Equal(t, "8", repo.updates[SettingKeyOpenAIAdvancedSchedulerWeightPreviousResponse])
|
||||
require.Equal(t, "4", repo.updates[SettingKeyOpenAIAdvancedSchedulerWeightSessionSticky])
|
||||
}
|
||||
|
||||
func TestSettingService_UpdateSettingsRejectsInvalidOpenAIOAuthSchedulingRateMultiplier(t *testing.T) {
|
||||
repo := &settingUpdateRepoStub{}
|
||||
svc := NewSettingService(repo, &config.Config{})
|
||||
|
||||
for _, rate := range []float64{-0.01, math.NaN(), math.Inf(1)} {
|
||||
err := svc.UpdateSettings(context.Background(), &SystemSettings{OpenAIOAuthSchedulingRateMultiplier: rate})
|
||||
require.Error(t, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSettingService_UpdateSettings_OpenAIAdvancedSchedulerWeightSums(t *testing.T) {
|
||||
maxFloat := strconv.FormatFloat(math.MaxFloat64, 'g', -1, 64)
|
||||
tests := []struct {
|
||||
name string
|
||||
weights SystemSettings
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "reset only base is valid",
|
||||
weights: SystemSettings{
|
||||
OpenAIAdvancedSchedulerWeightPriority: "0",
|
||||
OpenAIAdvancedSchedulerWeightLoad: "0",
|
||||
OpenAIAdvancedSchedulerWeightQueue: "0",
|
||||
OpenAIAdvancedSchedulerWeightErrorRate: "0",
|
||||
OpenAIAdvancedSchedulerWeightTTFT: "0",
|
||||
OpenAIAdvancedSchedulerWeightReset: "1",
|
||||
OpenAIAdvancedSchedulerWeightQuotaHeadroom: "0",
|
||||
OpenAIAdvancedSchedulerWeightUpstreamCost: "0",
|
||||
OpenAIAdvancedSchedulerWeightPreviousResponse: "0",
|
||||
OpenAIAdvancedSchedulerWeightSessionSticky: "0",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "base sum overflow is rejected",
|
||||
weights: SystemSettings{
|
||||
OpenAIAdvancedSchedulerWeightPriority: maxFloat,
|
||||
OpenAIAdvancedSchedulerWeightLoad: maxFloat,
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "sticky total sum overflow is rejected",
|
||||
weights: SystemSettings{
|
||||
OpenAIAdvancedSchedulerWeightPriority: maxFloat,
|
||||
OpenAIAdvancedSchedulerWeightPreviousResponse: maxFloat,
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
svc := NewSettingService(&settingUpdateRepoStub{}, &config.Config{})
|
||||
err := svc.UpdateSettings(context.Background(), &tt.weights)
|
||||
if tt.wantErr {
|
||||
require.Error(t, err)
|
||||
return
|
||||
}
|
||||
require.NoError(t, err)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSettingService_ParseSettingsDefaultsOpenAIOAuthSchedulingRateMultiplier(t *testing.T) {
|
||||
svc := NewSettingService(&settingUpdateRepoStub{}, &config.Config{})
|
||||
|
||||
require.Equal(t, 1.0, svc.parseSettings(map[string]string{}).OpenAIOAuthSchedulingRateMultiplier)
|
||||
require.Equal(t, 0.05, svc.parseSettings(map[string]string{SettingKeyOpenAIOAuthSchedulingRateMultiplier: "0.05"}).OpenAIOAuthSchedulingRateMultiplier)
|
||||
}
|
||||
|
||||
func TestSettingService_GetAllSettings_OpenAIAdvancedSchedulerEffectiveValuesUseConfig(t *testing.T) {
|
||||
cfg := &config.Config{}
|
||||
cfg.Gateway.OpenAIWS.LBTopK = 13
|
||||
@@ -384,8 +462,9 @@ func TestSettingService_GetAllSettings_OpenAIAdvancedSchedulerEffectiveValuesUse
|
||||
TTFT: 6,
|
||||
Reset: 7,
|
||||
QuotaHeadroom: 8,
|
||||
PreviousResponse: 9,
|
||||
SessionSticky: 10,
|
||||
UpstreamCost: 9,
|
||||
PreviousResponse: 10,
|
||||
SessionSticky: 11,
|
||||
}
|
||||
svc := NewSettingService(&settingGetAllRepoStub{values: map[string]string{
|
||||
SettingKeyOpenAIAdvancedSchedulerLBTopK: "3",
|
||||
@@ -401,7 +480,8 @@ func TestSettingService_GetAllSettings_OpenAIAdvancedSchedulerEffectiveValuesUse
|
||||
require.Equal(t, "13", settings.OpenAIAdvancedSchedulerEffectiveLBTopK)
|
||||
require.Equal(t, "2", settings.OpenAIAdvancedSchedulerEffectiveWeightPriority)
|
||||
require.Equal(t, "3", settings.OpenAIAdvancedSchedulerEffectiveWeightLoad)
|
||||
require.Equal(t, "10", settings.OpenAIAdvancedSchedulerEffectiveWeightSessionSticky)
|
||||
require.Equal(t, "9", settings.OpenAIAdvancedSchedulerEffectiveWeightUpstreamCost)
|
||||
require.Equal(t, "11", settings.OpenAIAdvancedSchedulerEffectiveWeightSessionSticky)
|
||||
}
|
||||
|
||||
func TestSettingService_UpdateSettings_AntigravityUserAgentVersion(t *testing.T) {
|
||||
|
||||
@@ -381,6 +381,8 @@ func (s *SettingService) buildSystemSettingsUpdates(ctx context.Context, setting
|
||||
updates[SettingPaymentVisibleMethodWxpaySource] = settings.PaymentVisibleMethodWxpaySource
|
||||
updates[SettingPaymentVisibleMethodAlipayEnabled] = strconv.FormatBool(settings.PaymentVisibleMethodAlipayEnabled)
|
||||
updates[SettingPaymentVisibleMethodWxpayEnabled] = strconv.FormatBool(settings.PaymentVisibleMethodWxpayEnabled)
|
||||
updates[SettingKeyOpenAILowUpstreamRatePriorityEnabled] = strconv.FormatBool(settings.OpenAILowUpstreamRatePriorityEnabled)
|
||||
updates[SettingKeyOpenAIOAuthSchedulingRateMultiplier] = strconv.FormatFloat(settings.OpenAIOAuthSchedulingRateMultiplier, 'f', -1, 64)
|
||||
updates[openAIAdvancedSchedulerSettingKey] = strconv.FormatBool(settings.OpenAIAdvancedSchedulerEnabled)
|
||||
updates[SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled] = strconv.FormatBool(settings.OpenAIAdvancedSchedulerStickyWeightedEnabled)
|
||||
updates[SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled] = strconv.FormatBool(settings.OpenAIAdvancedSchedulerSubscriptionPriorityEnabled)
|
||||
@@ -392,6 +394,7 @@ func (s *SettingService) buildSystemSettingsUpdates(ctx context.Context, setting
|
||||
updates[SettingKeyOpenAIAdvancedSchedulerWeightTTFT] = settings.OpenAIAdvancedSchedulerWeightTTFT
|
||||
updates[SettingKeyOpenAIAdvancedSchedulerWeightReset] = settings.OpenAIAdvancedSchedulerWeightReset
|
||||
updates[SettingKeyOpenAIAdvancedSchedulerWeightQuotaHeadroom] = settings.OpenAIAdvancedSchedulerWeightQuotaHeadroom
|
||||
updates[SettingKeyOpenAIAdvancedSchedulerWeightUpstreamCost] = settings.OpenAIAdvancedSchedulerWeightUpstreamCost
|
||||
updates[SettingKeyOpenAIAdvancedSchedulerWeightPreviousResponse] = settings.OpenAIAdvancedSchedulerWeightPreviousResponse
|
||||
updates[SettingKeyOpenAIAdvancedSchedulerWeightSessionSticky] = settings.OpenAIAdvancedSchedulerWeightSessionSticky
|
||||
|
||||
@@ -541,10 +544,12 @@ func (s *SettingService) refreshCachedSettings(settings *SystemSettings) {
|
||||
})
|
||||
openAIAdvancedSchedulerSettingSF.Forget(openAIAdvancedSchedulerSettingKey)
|
||||
openAIAdvancedSchedulerSettingCache.Store(&cachedOpenAIAdvancedSchedulerSetting{
|
||||
enabled: settings.OpenAIAdvancedSchedulerEnabled,
|
||||
stickyWeightedEnabled: settings.OpenAIAdvancedSchedulerStickyWeightedEnabled,
|
||||
subscriptionPriorityEnabled: settings.OpenAIAdvancedSchedulerSubscriptionPriorityEnabled,
|
||||
lbTopKOverride: parsePositiveIntOverride(settings.OpenAIAdvancedSchedulerLBTopK),
|
||||
lowUpstreamRatePriorityEnabled: settings.OpenAILowUpstreamRatePriorityEnabled,
|
||||
oauthSchedulingRateMultiplier: settings.OpenAIOAuthSchedulingRateMultiplier,
|
||||
enabled: settings.OpenAIAdvancedSchedulerEnabled,
|
||||
stickyWeightedEnabled: settings.OpenAIAdvancedSchedulerStickyWeightedEnabled,
|
||||
subscriptionPriorityEnabled: settings.OpenAIAdvancedSchedulerSubscriptionPriorityEnabled,
|
||||
lbTopKOverride: parsePositiveIntOverride(settings.OpenAIAdvancedSchedulerLBTopK),
|
||||
weightOverrides: parseOpenAIAdvancedSchedulerWeightOverrides(map[string]string{
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightPriority: settings.OpenAIAdvancedSchedulerWeightPriority,
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightLoad: settings.OpenAIAdvancedSchedulerWeightLoad,
|
||||
@@ -553,6 +558,7 @@ func (s *SettingService) refreshCachedSettings(settings *SystemSettings) {
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightTTFT: settings.OpenAIAdvancedSchedulerWeightTTFT,
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightReset: settings.OpenAIAdvancedSchedulerWeightReset,
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightQuotaHeadroom: settings.OpenAIAdvancedSchedulerWeightQuotaHeadroom,
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightUpstreamCost: settings.OpenAIAdvancedSchedulerWeightUpstreamCost,
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightPreviousResponse: settings.OpenAIAdvancedSchedulerWeightPreviousResponse,
|
||||
SettingKeyOpenAIAdvancedSchedulerWeightSessionSticky: settings.OpenAIAdvancedSchedulerWeightSessionSticky,
|
||||
}),
|
||||
|
||||
@@ -219,6 +219,8 @@ type SystemSettings struct {
|
||||
PaymentVisibleMethodWxpayEnabled bool
|
||||
|
||||
// OpenAI 账号调度
|
||||
OpenAILowUpstreamRatePriorityEnabled bool
|
||||
OpenAIOAuthSchedulingRateMultiplier float64
|
||||
OpenAIAdvancedSchedulerEnabled bool
|
||||
OpenAIAdvancedSchedulerStickyWeightedEnabled bool
|
||||
OpenAIAdvancedSchedulerSubscriptionPriorityEnabled bool
|
||||
@@ -230,6 +232,7 @@ type SystemSettings struct {
|
||||
OpenAIAdvancedSchedulerWeightTTFT string
|
||||
OpenAIAdvancedSchedulerWeightReset string
|
||||
OpenAIAdvancedSchedulerWeightQuotaHeadroom string
|
||||
OpenAIAdvancedSchedulerWeightUpstreamCost string
|
||||
OpenAIAdvancedSchedulerWeightPreviousResponse string
|
||||
OpenAIAdvancedSchedulerWeightSessionSticky string
|
||||
OpenAIAdvancedSchedulerEffectiveLBTopK string
|
||||
@@ -240,6 +243,7 @@ type SystemSettings struct {
|
||||
OpenAIAdvancedSchedulerEffectiveWeightTTFT string
|
||||
OpenAIAdvancedSchedulerEffectiveWeightReset string
|
||||
OpenAIAdvancedSchedulerEffectiveWeightQuotaHeadroom string
|
||||
OpenAIAdvancedSchedulerEffectiveWeightUpstreamCost string
|
||||
OpenAIAdvancedSchedulerEffectiveWeightPreviousResponse string
|
||||
OpenAIAdvancedSchedulerEffectiveWeightSessionSticky string
|
||||
|
||||
|
||||
Reference in New Issue
Block a user