Merge remote-tracking branch 'origin/main' into pr-3762

# Conflicts:
#	backend/internal/repository/concurrency_cache.go
#	backend/internal/repository/concurrency_cache_integration_test.go
This commit is contained in:
shaw
2026-07-07 09:22:44 +08:00
203 changed files with 12442 additions and 1571 deletions
+1 -1
View File
@@ -1 +1 @@
0.1.143
0.1.145
+5 -5
View File
@@ -67,7 +67,11 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
serviceUserPlatformQuotaRepository := repository.NewUserPlatformQuotaServiceAdapter(userPlatformQuotaRepository)
billingCacheService := service.ProvideBillingCacheService(billingCache, userRepository, userSubscriptionRepository, apiKeyRepository, userRPMCache, userGroupRateRepository, configConfig, serviceUserPlatformQuotaRepository)
apiKeyCache := repository.NewAPIKeyCache(redisClient)
apiKeyService := service.ProvideAPIKeyService(apiKeyRepository, userRepository, groupRepository, userSubscriptionRepository, userGroupRateRepository, apiKeyCache, configConfig, billingCacheService)
concurrencyCache := repository.ProvideConcurrencyCache(redisClient, configConfig)
schedulerCache := repository.ProvideSchedulerCache(redisClient, configConfig)
accountRepository := repository.NewAccountRepository(client, db, schedulerCache)
concurrencyService := service.ProvideConcurrencyService(concurrencyCache, accountRepository, configConfig)
apiKeyService := service.ProvideAPIKeyService(apiKeyRepository, userRepository, groupRepository, userSubscriptionRepository, userGroupRateRepository, apiKeyCache, configConfig, billingCacheService, concurrencyService)
apiKeyAuthCacheInvalidator := service.ProvideAPIKeyAuthCacheInvalidator(apiKeyService)
promoService := service.NewPromoService(promoCodeRepository, userRepository, billingCacheService, client, apiKeyAuthCacheInvalidator)
subscriptionService := service.NewSubscriptionService(groupRepository, userSubscriptionRepository, billingCacheService, client, configConfig)
@@ -92,10 +96,6 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
usageLogRepository := repository.NewUsageLogRepository(client, db)
usageService := service.NewUsageService(usageLogRepository, userRepository, client, apiKeyAuthCacheInvalidator)
opsRepository := repository.NewOpsRepository(db)
schedulerCache := repository.ProvideSchedulerCache(redisClient, configConfig)
accountRepository := repository.NewAccountRepository(client, db, schedulerCache)
concurrencyCache := repository.ProvideConcurrencyCache(redisClient, configConfig)
concurrencyService := service.ProvideConcurrencyService(concurrencyCache, accountRepository, configConfig)
usageBillingRepository := repository.NewUsageBillingRepository(client, db)
gatewayCache := repository.NewGatewayCache(redisClient)
schedulerOutboxRepository := repository.NewSchedulerOutboxRepository(db)
+12 -2
View File
@@ -969,6 +969,9 @@ type GatewayOpenAIWSSchedulerScoreWeights struct {
Reset float64 `mapstructure:"reset"`
// QuotaHeadroom 倾向 7d 剩余额度更健康的账号;默认 0(关闭,不改变原有行为)。
QuotaHeadroom float64 `mapstructure:"quota_headroom"`
// PreviousResponse/SessionSticky 仅在开启 OpenAI 高级调度的粘性加权时生效。
PreviousResponse float64 `mapstructure:"previous_response"`
SessionSticky float64 `mapstructure:"session_sticky"`
}
// GatewayOpenAISchedulerConfig OpenAI 高级调度器配置。
@@ -1891,6 +1894,8 @@ 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.previous_response", 5.0)
viper.SetDefault("gateway.openai_ws.scheduler_score_weights.session_sticky", 3.0)
// OpenAI HTTP upstream protocol strategy
viper.SetDefault("gateway.openai_http2.enabled", true)
viper.SetDefault("gateway.openai_http2.allow_proxy_fallback_to_http1", true)
@@ -1945,7 +1950,10 @@ func setDefaults() {
viper.SetDefault("gateway.usage_record.worker_count", 128)
viper.SetDefault("gateway.usage_record.queue_size", 16384)
viper.SetDefault("gateway.usage_record.task_timeout_seconds", 5)
viper.SetDefault("gateway.usage_record.overflow_policy", UsageRecordOverflowPolicySample)
// 默认 sync:队列满时由提交方内联执行(提交点在响应写出之后,不阻塞客户端)。
// sample/drop 会在溢出时静默丢弃计费任务,造成扣费与 usage_logs 对账缺口(issue #3656),
// 仅供显式配置的运维场景使用。
viper.SetDefault("gateway.usage_record.overflow_policy", UsageRecordOverflowPolicySync)
viper.SetDefault("gateway.usage_record.overflow_sample_percent", 10)
viper.SetDefault("gateway.usage_record.auto_scale_enabled", true)
viper.SetDefault("gateway.usage_record.auto_scale_min_workers", 128)
@@ -2671,7 +2679,9 @@ func (c *Config) Validate() error {
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.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")
}
weightSum := c.Gateway.OpenAIWS.SchedulerScoreWeights.Priority +
+2 -2
View File
@@ -1903,8 +1903,8 @@ func TestLoad_DefaultGatewayUsageRecordConfig(t *testing.T) {
if cfg.Gateway.UsageRecord.TaskTimeoutSeconds != 5 {
t.Fatalf("task_timeout_seconds = %d, want 5", cfg.Gateway.UsageRecord.TaskTimeoutSeconds)
}
if cfg.Gateway.UsageRecord.OverflowPolicy != UsageRecordOverflowPolicySample {
t.Fatalf("overflow_policy = %s, want %s", cfg.Gateway.UsageRecord.OverflowPolicy, UsageRecordOverflowPolicySample)
if cfg.Gateway.UsageRecord.OverflowPolicy != UsageRecordOverflowPolicySync {
t.Fatalf("overflow_policy = %s, want %s", cfg.Gateway.UsageRecord.OverflowPolicy, UsageRecordOverflowPolicySync)
}
if cfg.Gateway.UsageRecord.OverflowSamplePercent != 10 {
t.Fatalf("overflow_sample_percent = %d, want 10", cfg.Gateway.UsageRecord.OverflowSamplePercent)
+8 -3
View File
@@ -68,6 +68,9 @@ const (
SubscriptionStatusSuspended = "suspended"
)
// AntigravityGemini31ProAgentModel is the upstream route for Gemini 3.1 Pro High.
const AntigravityGemini31ProAgentModel = "gemini-pro-agent"
// DefaultAntigravityModelMapping 是 Antigravity 平台的默认模型映射
// 当账号未配置 model_mapping 时使用此默认值
// 与前端 useModelWhitelist.ts 中的 antigravityDefaultMappings 保持一致
@@ -103,10 +106,12 @@ var DefaultAntigravityModelMapping = map[string]string{
"gemini-3-flash-preview": "gemini-3-flash",
"gemini-3-pro-preview": "gemini-3-pro-high",
// Gemini 3.1 白名单
"gemini-3.1-pro-high": "gemini-3.1-pro-high",
"gemini-3.1-pro-low": "gemini-3.1-pro-low",
AntigravityGemini31ProAgentModel: AntigravityGemini31ProAgentModel,
"gemini-3.1-pro": AntigravityGemini31ProAgentModel,
"gemini-3.1-pro-high": AntigravityGemini31ProAgentModel,
"gemini-3.1-pro-low": "gemini-3.1-pro-low",
// Gemini 3.1 preview 映射
"gemini-3.1-pro-preview": "gemini-3.1-pro-high",
"gemini-3.1-pro-preview": AntigravityGemini31ProAgentModel,
// Gemini 3.1 image 白名单
"gemini-3.1-flash-image": "gemini-3.1-flash-image",
// Gemini 3.1 image preview 映射
+22
View File
@@ -43,6 +43,28 @@ func TestDefaultAntigravityModelMapping_ContainsNewClaudeModels(t *testing.T) {
}
}
func TestDefaultAntigravityModelMapping_Gemini31ProAliases(t *testing.T) {
t.Parallel()
cases := map[string]string{
AntigravityGemini31ProAgentModel: AntigravityGemini31ProAgentModel,
"gemini-3.1-pro": AntigravityGemini31ProAgentModel,
"gemini-3.1-pro-high": AntigravityGemini31ProAgentModel,
"gemini-3.1-pro-preview": AntigravityGemini31ProAgentModel,
"gemini-3.1-pro-low": "gemini-3.1-pro-low",
}
for from, want := range cases {
got, ok := DefaultAntigravityModelMapping[from]
if !ok {
t.Fatalf("expected mapping for %q to exist", from)
}
if got != want {
t.Fatalf("unexpected mapping for %q: got %q want %q", from, got, want)
}
}
}
func TestDefaultBedrockModelMapping_ContainsNewClaudeModels(t *testing.T) {
t.Parallel()
@@ -253,6 +253,17 @@ func (h *AccountHandler) importCodexSessions(ctx context.Context, req CodexSessi
Message: "已有账号未记录 chatgpt_user_id,已按共享的 chatgpt_account_id 匹配并回填,请确认两者属于同一用户",
})
}
preserveExistingRefresh := item.RefreshToken == "" &&
codexCredentialString(existing.Credentials, "refresh_token") != ""
if preserveExistingRefresh {
result.Warnings = append(result.Warnings, CodexSessionImportMessage{
Index: entry.Index,
Name: accountName,
Message: "已有账号包含 refresh_token,本次 accessToken-only 导入已保留自动续期凭据",
})
effectiveExpiresAt = nil
autoPauseOnExpired = nil
}
mergedCredentials := mergeCodexImportCredentials(existing.Credentials, credentials, item)
mergedExtra := mergeCodexImportMap(existing.Extra, extra)
updateInput := &service.UpdateAccountInput{
@@ -592,7 +603,7 @@ func normalizeCodexImportEntry(entry codexImportEntry) (*codexImportAccount, err
fingerprint := codexTokenFingerprint(item.AccessToken)
item.Extra["access_token_sha256"] = fingerprint
item.IdentityKeys = buildCodexIdentityKeys(item.AccountID, item.UserID, item.Email, item.AccessToken)
item.IdentityKeys = buildCodexImportIdentityKeys(item.AccountID, item.UserID, item.Email, item.AccessToken, item.RefreshToken)
item.Name = buildCodexImportAccountName(item, entry.Index)
return item, nil
@@ -815,13 +826,25 @@ func sanitizeCodexImportCredentialExtras(input map[string]any) map[string]any {
return out
}
// buildCodexIdentityKeys 按身份强度排序生成匹配键:chatgpt_account_id 在同一
// ChatGPT 团队内是共享的,因此 account: 键排在最后,且命中时还需通过
// codexIdentityConflicts 的跨用户校验才生效。
func buildCodexIdentityKeys(accountID, userID, email, accessToken string) []string {
// buildCodexImportIdentityKeys 生成导入条目的匹配键。refresh_token 缺失时
// Codex session 只能作为 accessToken-only 凭据使用,此时以 access token
// 指纹作为唯一稳定身份,避免同 workspace 下共享的 account/user 标识误合并。
func buildCodexImportIdentityKeys(accountID, userID, email, accessToken, refreshToken string) []string {
accessToken = strings.TrimSpace(accessToken)
refreshToken = strings.TrimSpace(refreshToken)
if refreshToken == "" && accessToken != "" {
return []string{"access:" + codexTokenFingerprint(accessToken)}
}
return buildCodexStoredIdentityKeys(accountID, userID, email, accessToken)
}
// buildCodexStoredIdentityKeys 生成存量账号索引键,保留 user/account 维度,
// 让 accessToken-only 账号后续升级为完整 OAuth 时仍能命中并更新原账号。
func buildCodexStoredIdentityKeys(accountID, userID, email, accessToken string) []string {
keys := make([]string, 0, 3)
accountID = strings.TrimSpace(accountID)
userID = strings.TrimSpace(userID)
accessToken = strings.TrimSpace(accessToken)
if userID != "" {
keys = append(keys, "user:"+userID)
}
@@ -830,7 +853,7 @@ func buildCodexIdentityKeys(accountID, userID, email, accessToken string) []stri
keys = append(keys, "email:"+email)
}
}
if accessToken = strings.TrimSpace(accessToken); accessToken != "" {
if accessToken != "" {
keys = append(keys, "access:"+codexTokenFingerprint(accessToken))
}
if accountID != "" {
@@ -854,7 +877,8 @@ func (i *codexAccountIndex) Add(account service.Account) {
if i.accountsByKey == nil {
i.accountsByKey = map[string][]service.Account{}
}
keys := buildCodexIdentityKeys(
i.remove(account.ID)
keys := buildCodexStoredIdentityKeys(
codexCredentialString(account.Credentials, "chatgpt_account_id"),
codexCredentialString(account.Credentials, "chatgpt_user_id"),
codexCredentialString(account.Credentials, "email"),
@@ -865,6 +889,22 @@ func (i *codexAccountIndex) Add(account service.Account) {
}
}
func (i *codexAccountIndex) remove(accountID int64) {
for key, accounts := range i.accountsByKey {
kept := accounts[:0]
for _, account := range accounts {
if account.ID != accountID {
kept = append(kept, account)
}
}
if len(kept) == 0 {
delete(i.accountsByKey, key)
continue
}
i.accountsByKey[key] = kept
}
}
// upsertCodexAccount 保留同一键下的全部候选账号(共享的 account: 键可对应
// 团队内多个账号),同一账号重复 Add 时原位替换为最新状态。
func upsertCodexAccount(accounts []service.Account, account service.Account) []service.Account {
@@ -894,9 +934,9 @@ func (i *codexAccountIndex) Find(keys []string, userID string) (*service.Account
}
// codexIdentityConflicts 判断 account: 键的命中是否把同一 ChatGPT 团队的两个
// 不同成员误连到一起:双方都携带 user id 且不相等时视为冲突。任一侧缺少
// user id 时保留匹配,使早期未记录 chatgpt_user_id 的存量账号仍能被更新
// (并借助凭据合并回填 user id),而不是产生重复账号。
// 不同成员误连到一起:双方都携带 user id 且不相等时视为冲突。存量索引侧
// 仍保留 account 键,任一侧缺少 user id 时允许匹配,使含 refresh_token
// 的常规导入和 accessToken-only 账号升级为完整 OAuth 时仍能更新原账号。
func codexIdentityConflicts(key, userID, storedUserID string) bool {
if !strings.HasPrefix(key, "account:") {
return false
@@ -948,8 +988,15 @@ func mergeCodexImportCredentials(existing, incoming map[string]any, item *codexI
return out
}
if strings.TrimSpace(item.RefreshToken) == "" {
delete(out, "refresh_token")
delete(out, "client_id")
if codexCredentialString(existing, "refresh_token") == "" {
delete(out, "refresh_token")
delete(out, "client_id")
} else {
out["refresh_token"] = existing["refresh_token"]
if clientID, ok := existing["client_id"]; ok {
out["client_id"] = clientID
}
}
}
if strings.TrimSpace(item.IDToken) == "" {
delete(out, "id_token")
@@ -1,6 +1,7 @@
package admin
import (
"context"
"encoding/base64"
"encoding/json"
"fmt"
@@ -144,7 +145,7 @@ func TestNormalizeCodexSessionJSONExtractsCredentialsAndIgnoresSessionToken(t *t
}
}
func TestMergeCodexImportCredentialsClearsStaleRefreshFieldsWhenIncomingHasNoRefreshToken(t *testing.T) {
func TestMergeCodexImportCredentialsPreservesExistingRefreshFieldsWhenIncomingHasNoRefreshToken(t *testing.T) {
existing := map[string]any{
"access_token": "old-access-token",
"refresh_token": "old-refresh-token",
@@ -171,11 +172,11 @@ func TestMergeCodexImportCredentialsClearsStaleRefreshFieldsWhenIncomingHasNoRef
if merged["chatgpt_account_id"] != "acct-new" {
t.Fatalf("chatgpt_account_id = %v, want acct-new", merged["chatgpt_account_id"])
}
if _, ok := merged["refresh_token"]; ok {
t.Fatalf("refresh_token should be cleared")
if merged["refresh_token"] != "old-refresh-token" {
t.Fatalf("refresh_token = %v, want old-refresh-token", merged["refresh_token"])
}
if _, ok := merged["client_id"]; ok {
t.Fatalf("client_id should be cleared")
if merged["client_id"] != "old-client-id" {
t.Fatalf("client_id = %v, want old-client-id", merged["client_id"])
}
if _, ok := merged["id_token"]; ok {
t.Fatalf("id_token should be cleared")
@@ -301,9 +302,9 @@ func TestResolveCodexImportExpiryForNoRefreshTokenUsesEarlierRequestExpiry(t *te
}
func TestCodexIdentityKeysPreferStrongIdentifiers(t *testing.T) {
keys := buildCodexIdentityKeys("acct-1", "user-1", "same@example.com", "token")
keys := buildCodexImportIdentityKeys("acct-1", "user-1", "same@example.com", "token", "refresh")
if len(keys) == 0 || keys[0] != "user:user-1" {
t.Fatalf("user key should have highest priority: %v", keys)
t.Fatalf("user key should have highest priority when refresh token exists: %v", keys)
}
if keys[len(keys)-1] != "account:acct-1" {
t.Fatalf("shared account key should be the last fallback: %v", keys)
@@ -314,7 +315,7 @@ func TestCodexIdentityKeysPreferStrongIdentifiers(t *testing.T) {
}
}
keys = buildCodexIdentityKeys("", "", "same@example.com", "token")
keys = buildCodexImportIdentityKeys("", "", "same@example.com", "token", "refresh")
hasEmail := false
for _, key := range keys {
if key == "email:same@example.com" {
@@ -324,6 +325,11 @@ func TestCodexIdentityKeysPreferStrongIdentifiers(t *testing.T) {
if !hasEmail {
t.Fatalf("weak identity should include email fallback: %v", keys)
}
keys = buildCodexImportIdentityKeys("acct-1", "user-1", "same@example.com", "token", "")
if len(keys) != 1 || !strings.HasPrefix(keys[0], "access:") {
t.Fatalf("accessToken-only identity should use only access fingerprint: %v", keys)
}
}
func TestCodexAccountIndexDoesNotMatchDifferentUsersInSameChatGPTAccount(t *testing.T) {
@@ -333,35 +339,37 @@ func TestCodexAccountIndexDoesNotMatchDifferentUsersInSameChatGPTAccount(t *test
"chatgpt_account_id": "team-1",
"chatgpt_user_id": "user-1",
"access_token": "token-1",
"refresh_token": "refresh-1",
},
}
index := buildCodexAccountIndex([]service.Account{existing})
keys := buildCodexIdentityKeys("team-1", "user-2", "", "token-2")
keys := buildCodexImportIdentityKeys("team-1", "user-2", "", "token-2", "refresh-2")
if got, _ := index.Find(keys, "user-2"); got != nil {
t.Fatalf("Find matched account ID %d for a different chatgpt_user_id in the same team", got.ID)
}
keys = buildCodexIdentityKeys("team-1", "user-1", "", "token-2")
keys = buildCodexImportIdentityKeys("team-1", "user-1", "", "token-2", "refresh-2")
got, _ := index.Find(keys, "user-1")
if got == nil || got.ID != existing.ID {
t.Fatalf("Find by same chatgpt_user_id = %v, want account ID %d", got, existing.ID)
}
}
func TestCodexAccountIndexFallsBackToAccountKeyWhenUserIDMissing(t *testing.T) {
// 存量账号缺少 chatgpt_user_id:携带 user id 的重新导入应命中并更新(回填),
// 而不是创建重复账号。
func TestCodexAccountIndexFallsBackToAccountKeyWhenRefreshTokenExistsAndUserIDMissing(t *testing.T) {
// 含 refresh_token 的常规导入沿用 a5638a4e 的兼容逻辑:存量账号缺少
// chatgpt_user_id 时,携带 user id 的重新导入仍可命中并回填。
legacy := service.Account{
ID: 20,
Credentials: map[string]any{
"chatgpt_account_id": "team-1",
"access_token": "token-old",
"refresh_token": "refresh-old",
},
}
index := buildCodexAccountIndex([]service.Account{legacy})
keys := buildCodexIdentityKeys("team-1", "user-1", "", "token-new")
keys := buildCodexImportIdentityKeys("team-1", "user-1", "", "token-new", "refresh-new")
got, matchedKey := index.Find(keys, "user-1")
if got == nil || got.ID != legacy.ID {
t.Fatalf("Find legacy account without stored user id = %v, want account ID %d", got, legacy.ID)
@@ -370,30 +378,59 @@ func TestCodexAccountIndexFallsBackToAccountKeyWhenUserIDMissing(t *testing.T) {
t.Fatalf("matched key = %q, want account:team-1", matchedKey)
}
// 反向:导入条目无法解析出 user id 时,仍应通过 account 键命中已有账号。
// 反向:含 refresh_token 的导入条目无法解析出 user id 时,仍应通过
// account 键命中已有账号,保持常规导入去重行为。
full := service.Account{
ID: 21,
Credentials: map[string]any{
"chatgpt_account_id": "team-2",
"chatgpt_user_id": "user-9",
"access_token": "token-old",
"refresh_token": "refresh-old",
},
}
index = buildCodexAccountIndex([]service.Account{full})
keys = buildCodexIdentityKeys("team-2", "", "", "token-opaque")
keys = buildCodexImportIdentityKeys("team-2", "", "", "token-opaque", "refresh-new")
got, _ = index.Find(keys, "")
if got == nil || got.ID != full.ID {
t.Fatalf("Find by account key without entry user id = %v, want account ID %d", got, full.ID)
}
}
func TestCodexAccountIndexAccessTokenOnlyUsesTokenFingerprint(t *testing.T) {
existing := service.Account{
ID: 22,
Credentials: map[string]any{
"chatgpt_account_id": "team-1",
"chatgpt_user_id": "user-1",
"access_token": "token-old",
},
}
index := buildCodexAccountIndex([]service.Account{existing})
keys := buildCodexImportIdentityKeys("team-1", "user-1", "", "token-new", "")
if got, matchedKey := index.Find(keys, "user-1"); got != nil {
t.Fatalf("accessToken-only import matched by %q despite different token: account ID %d", matchedKey, got.ID)
}
keys = buildCodexImportIdentityKeys("team-1", "user-1", "", "token-old", "")
got, matchedKey := index.Find(keys, "user-1")
if got == nil || got.ID != existing.ID {
t.Fatalf("Find accessToken-only duplicate by fingerprint = %v, want account ID %d", got, existing.ID)
}
if !strings.HasPrefix(matchedKey, "access:") {
t.Fatalf("matched key = %q, want access fingerprint", matchedKey)
}
}
func TestCodexAccountIndexKeepsAllCandidatesForSharedAccountKey(t *testing.T) {
legacy := service.Account{
ID: 30,
Credentials: map[string]any{
"chatgpt_account_id": "team-1",
"access_token": "token-legacy",
"refresh_token": "refresh-legacy",
},
}
member := service.Account{
@@ -402,10 +439,11 @@ func TestCodexAccountIndexKeepsAllCandidatesForSharedAccountKey(t *testing.T) {
"chatgpt_account_id": "team-1",
"chatgpt_user_id": "user-2",
"access_token": "token-member",
"refresh_token": "refresh-member",
},
}
// 无论索引构建顺序如何,携带新 user id 的条目都应跳过 user-2 的账号、
// 无论索引构建顺序如何,携带新 user id 的条目都应跳过 user-2 的账号,
// 命中缺少 user id 的存量账号,而不是因单一候选被遮蔽而落空。
for _, accounts := range [][]service.Account{
{member, legacy},
@@ -413,7 +451,7 @@ func TestCodexAccountIndexKeepsAllCandidatesForSharedAccountKey(t *testing.T) {
} {
index := buildCodexAccountIndex(accounts)
keys := buildCodexIdentityKeys("team-1", "user-1", "", "token-new")
keys := buildCodexImportIdentityKeys("team-1", "user-1", "", "token-new", "refresh-new")
got, matchedKey := index.Find(keys, "user-1")
if got == nil || got.ID != legacy.ID {
t.Fatalf("Find with shared account key = %v, want legacy account ID %d", got, legacy.ID)
@@ -422,7 +460,7 @@ func TestCodexAccountIndexKeepsAllCandidatesForSharedAccountKey(t *testing.T) {
t.Fatalf("matched key = %q, want account:team-1", matchedKey)
}
keys = buildCodexIdentityKeys("team-1", "user-2", "", "token-new")
keys = buildCodexImportIdentityKeys("team-1", "user-2", "", "token-new", "refresh-new")
got, matchedKey = index.Find(keys, "user-2")
if got == nil || got.ID != member.ID {
t.Fatalf("Find by user key = %v, want member account ID %d", got, member.ID)
@@ -449,18 +487,19 @@ func TestCodexAccountIndexUpsertReplacesSameAccount(t *testing.T) {
"chatgpt_account_id": "team-1",
"chatgpt_user_id": "user-1",
"access_token": "token-new",
"refresh_token": "refresh-new",
},
}
index.Add(backfilled)
// 回填后同一账号在 account 键下应被原位替换而非残留旧副本:
// 其他成员的条目不应再通过旧副本(无 user id)命中该账号。
keys := buildCodexIdentityKeys("team-1", "user-2", "", "token-other")
if got, _ := index.Find(keys, "user-2"); got != nil {
t.Fatalf("stale candidate matched after upsert: account ID %d", got.ID)
keys := buildCodexImportIdentityKeys("team-1", "user-2", "", "token-other", "refresh-other")
if got, matchedKey := index.Find(keys, "user-2"); got != nil {
t.Fatalf("stale candidate matched after upsert by %q: account ID %d", matchedKey, got.ID)
}
keys = buildCodexIdentityKeys("team-1", "user-1", "", "token-other")
keys = buildCodexImportIdentityKeys("team-1", "user-1", "", "token-other", "refresh-other")
got, _ := index.Find(keys, "user-1")
if got == nil || got.ID != backfilled.ID {
t.Fatalf("Find after upsert = %v, want account ID %d", got, backfilled.ID)
@@ -472,28 +511,421 @@ func TestCodexAccountIndexUpsertReplacesSameAccount(t *testing.T) {
func TestCodexIdentitySeenDistinguishesTeamMembers(t *testing.T) {
seen := map[string]codexSeenIdentity{}
member1 := buildCodexIdentityKeys("team-1", "user-1", "", "token-1")
member1 := buildCodexImportIdentityKeys("team-1", "user-1", "", "token-1", "refresh-1")
markCodexIdentitySeen(seen, member1, 1, "user-1")
member2 := buildCodexIdentityKeys("team-1", "user-2", "", "token-2")
member2 := buildCodexImportIdentityKeys("team-1", "user-2", "", "token-2", "refresh-2")
if index, ok := firstSeenCodexIdentity(seen, member2, "user-2"); ok {
t.Fatalf("different team member treated as duplicate of entry %d", index)
}
again := buildCodexIdentityKeys("team-1", "user-1", "", "token-3")
again := buildCodexImportIdentityKeys("team-1", "user-1", "", "token-3", "refresh-3")
index, ok := firstSeenCodexIdentity(seen, again, "user-1")
if !ok || index != 1 {
t.Fatalf("same user re-entry dedup = (%d, %v), want (1, true)", index, ok)
}
// 无 user id 的条目与已见同 account 条目视为重复(保守跳过,与既有行为一致)。
opaque := buildCodexIdentityKeys("team-1", "", "", "token-4")
// 无 user id 的条目不应因共享 account id 与已见团队成员互相去重;
// 只有相同 access token 指纹才视为重复。
opaque := buildCodexImportIdentityKeys("team-1", "", "", "token-4", "")
index, ok = firstSeenCodexIdentity(seen, opaque, "")
if !ok || index != 1 {
t.Fatalf("entry without user id dedup = (%d, %v), want (1, true)", index, ok)
if ok {
t.Fatalf("entry without user id dedup = (%d, %v), want no match", index, ok)
}
}
func TestNormalizeCodexImportUsesJWTSubForAccessTokenOnlyIdentity(t *testing.T) {
accessToken := buildCodexImportTestJWT(t, time.Now().Add(time.Hour), map[string]any{
"sub": "user-from-access-token",
"https://api.openai.com/auth": map[string]any{
"chatgpt_account_id": "workspace-1",
},
})
item, err := normalizeCodexImportEntry(codexImportEntry{Index: 1, Value: accessToken})
if err != nil {
t.Fatalf("normalizeCodexImportEntry error = %v", err)
}
if item.UserID != "user-from-access-token" {
t.Fatalf("UserID = %q, want JWT sub", item.UserID)
}
if len(item.IdentityKeys) != 1 || !strings.HasPrefix(item.IdentityKeys[0], "access:") {
t.Fatalf("IdentityKeys = %v, want access fingerprint only for accessToken-only import", item.IdentityKeys)
}
if got := item.Credentials["chatgpt_user_id"]; got != "user-from-access-token" {
t.Fatalf("credential chatgpt_user_id = %v, want JWT sub", got)
}
}
func TestImportCodexSessionsAccessTokenOnlySameWorkspaceDifferentUsersCreatesTwoAccounts(t *testing.T) {
svc := newCodexImportMemoryAdminService(nil)
handler := NewAccountHandler(svc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
req := CodexSessionImportRequest{SkipDefaultGroupBind: boolPtr(true)}
entries := []codexImportEntry{
{Index: 1, Value: buildCodexAccessOnlyImportValue(t, "workspace-1", "user-1")},
{Index: 2, Value: buildCodexAccessOnlyImportValue(t, "workspace-1", "user-2")},
}
result, err := handler.importCodexSessions(context.Background(), req, entries)
if err != nil {
t.Fatalf("importCodexSessions error = %v", err)
}
if result.Created != 2 || result.Updated != 0 || result.Skipped != 0 || result.Failed != 0 {
t.Fatalf("result = %+v, want two created accounts", result)
}
if len(svc.createdAccounts) != 2 {
t.Fatalf("created accounts = %d, want 2", len(svc.createdAccounts))
}
if svc.createdAccounts[0].Credentials["chatgpt_user_id"] == svc.createdAccounts[1].Credentials["chatgpt_user_id"] {
t.Fatalf("created accounts share user id: %v", svc.createdAccounts)
}
}
func TestImportCodexSessionsAccessTokenOnlySameWorkspaceAndUserDifferentTokensCreatesTwoAccounts(t *testing.T) {
svc := newCodexImportMemoryAdminService(nil)
handler := NewAccountHandler(svc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
req := CodexSessionImportRequest{SkipDefaultGroupBind: boolPtr(true)}
entries := []codexImportEntry{
{Index: 1, Value: map[string]any{
"access_token": buildCodexImportTestJWT(t, time.Now().Add(time.Hour), map[string]any{
"sub": "shared-user",
"jti": "token-1",
"https://api.openai.com/auth": map[string]any{
"chatgpt_account_id": "workspace-1",
},
}),
}},
{Index: 2, Value: map[string]any{
"access_token": buildCodexImportTestJWT(t, time.Now().Add(time.Hour), map[string]any{
"sub": "shared-user",
"jti": "token-2",
"https://api.openai.com/auth": map[string]any{
"chatgpt_account_id": "workspace-1",
},
}),
}},
}
result, err := handler.importCodexSessions(context.Background(), req, entries)
if err != nil {
t.Fatalf("importCodexSessions error = %v", err)
}
if result.Created != 2 || result.Updated != 0 || result.Skipped != 0 || result.Failed != 0 {
t.Fatalf("result = %+v, want two created accounts", result)
}
if len(svc.createdAccounts) != 2 {
t.Fatalf("created accounts = %d, want 2", len(svc.createdAccounts))
}
}
func TestImportCodexSessionsAccessTokenOnlySameUserUpdatesExisting(t *testing.T) {
existingToken := buildCodexAccessToken(t, "workspace-1", "user-1", time.Now().Add(time.Hour))
svc := newCodexImportMemoryAdminService([]service.Account{{
ID: 10,
Name: "existing",
Platform: service.PlatformOpenAI,
Type: service.AccountTypeOAuth,
Credentials: map[string]any{
"chatgpt_account_id": "workspace-1",
"chatgpt_user_id": "user-1",
"access_token": existingToken,
},
}})
handler := NewAccountHandler(svc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
req := CodexSessionImportRequest{SkipDefaultGroupBind: boolPtr(true)}
entries := []codexImportEntry{
{Index: 1, Value: map[string]any{"access_token": existingToken}},
}
result, err := handler.importCodexSessions(context.Background(), req, entries)
if err != nil {
t.Fatalf("importCodexSessions error = %v", err)
}
if result.Created != 0 || result.Updated != 1 || result.Failed != 0 {
t.Fatalf("result = %+v, want one updated account", result)
}
if len(svc.createdAccounts) != 0 {
t.Fatalf("created accounts = %d, want 0", len(svc.createdAccounts))
}
if len(svc.updatedAccounts) != 1 || svc.updatedAccounts[0].id != 10 {
t.Fatalf("updated accounts = %+v, want account 10", svc.updatedAccounts)
}
}
func TestImportCodexSessionsUpgradesAccessTokenOnlyAccountWithRefreshToken(t *testing.T) {
oldToken := buildCodexAccessTokenWithJTI(t, "workspace-1", "user-1", "old-token", time.Now().Add(time.Hour))
newToken := buildCodexAccessTokenWithJTI(t, "workspace-1", "user-1", "new-token", time.Now().Add(time.Hour))
svc := newCodexImportMemoryAdminService([]service.Account{{
ID: 12,
Name: "existing",
Platform: service.PlatformOpenAI,
Type: service.AccountTypeOAuth,
Credentials: map[string]any{
"chatgpt_account_id": "workspace-1",
"chatgpt_user_id": "user-1",
"access_token": oldToken,
},
}})
handler := NewAccountHandler(svc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
req := CodexSessionImportRequest{SkipDefaultGroupBind: boolPtr(true)}
entries := []codexImportEntry{
{Index: 1, Value: map[string]any{
"access_token": newToken,
"refresh_token": "refresh-new",
}},
}
result, err := handler.importCodexSessions(context.Background(), req, entries)
if err != nil {
t.Fatalf("importCodexSessions error = %v", err)
}
if result.Created != 0 || result.Updated != 1 || result.Failed != 0 {
t.Fatalf("result = %+v, want one updated account", result)
}
if len(svc.updatedAccounts) != 1 || svc.updatedAccounts[0].id != 12 {
t.Fatalf("updated accounts = %+v, want account 12", svc.updatedAccounts)
}
if got := svc.updatedAccounts[0].input.Credentials["refresh_token"]; got != "refresh-new" {
t.Fatalf("updated refresh_token = %v, want refresh-new", got)
}
}
func TestImportCodexSessionsAccessTokenOnlyPreservesExistingRefreshToken(t *testing.T) {
existingToken := buildCodexAccessToken(t, "workspace-1", "user-1", time.Now().Add(time.Hour))
svc := newCodexImportMemoryAdminService([]service.Account{{
ID: 13,
Name: "existing",
Platform: service.PlatformOpenAI,
Type: service.AccountTypeOAuth,
Credentials: map[string]any{
"chatgpt_account_id": "workspace-1",
"chatgpt_user_id": "user-1",
"access_token": existingToken,
"refresh_token": "refresh-old",
"client_id": "client-old",
},
}})
handler := NewAccountHandler(svc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
req := CodexSessionImportRequest{SkipDefaultGroupBind: boolPtr(true)}
entries := []codexImportEntry{
{Index: 1, Value: map[string]any{"access_token": existingToken}},
}
result, err := handler.importCodexSessions(context.Background(), req, entries)
if err != nil {
t.Fatalf("importCodexSessions error = %v", err)
}
if result.Created != 0 || result.Updated != 1 || result.Failed != 0 {
t.Fatalf("result = %+v, want one updated account", result)
}
update := svc.updatedAccounts[0].input
if got := update.Credentials["refresh_token"]; got != "refresh-old" {
t.Fatalf("refresh_token = %v, want refresh-old", got)
}
if got := update.Credentials["client_id"]; got != "client-old" {
t.Fatalf("client_id = %v, want client-old", got)
}
if update.ExpiresAt != nil {
t.Fatalf("ExpiresAt = %v, want nil to preserve OAuth account expiry", *update.ExpiresAt)
}
if update.AutoPauseOnExpired != nil {
t.Fatalf("AutoPauseOnExpired = %v, want nil to preserve OAuth account scheduling", *update.AutoPauseOnExpired)
}
}
func TestImportCodexSessionsBatchOldAccessTokenDoesNotRollbackRefreshToken(t *testing.T) {
oldToken := buildCodexAccessTokenWithJTI(t, "workspace-1", "user-1", "old-token", time.Now().Add(time.Hour))
newToken := buildCodexAccessTokenWithJTI(t, "workspace-1", "user-1", "new-token", time.Now().Add(time.Hour))
svc := newCodexImportMemoryAdminService([]service.Account{{
ID: 14,
Name: "existing",
Platform: service.PlatformOpenAI,
Type: service.AccountTypeOAuth,
Credentials: map[string]any{
"chatgpt_account_id": "workspace-1",
"chatgpt_user_id": "user-1",
"access_token": oldToken,
"refresh_token": "refresh-old",
},
}})
handler := NewAccountHandler(svc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
req := CodexSessionImportRequest{SkipDefaultGroupBind: boolPtr(true)}
entries := []codexImportEntry{
{Index: 1, Value: map[string]any{
"access_token": newToken,
"refresh_token": "refresh-new",
}},
{Index: 2, Value: map[string]any{"access_token": oldToken}},
}
result, err := handler.importCodexSessions(context.Background(), req, entries)
if err != nil {
t.Fatalf("importCodexSessions error = %v", err)
}
if result.Updated != 1 || result.Created != 1 || result.Failed != 0 {
t.Fatalf("result = %+v, want first item updated and stale access token created separately", result)
}
if len(svc.updatedAccounts) != 1 || svc.updatedAccounts[0].id != 14 {
t.Fatalf("updated accounts = %+v, want account 14 updated once", svc.updatedAccounts)
}
stored, err := svc.GetAccount(context.Background(), 14)
if err != nil {
t.Fatalf("GetAccount error = %v", err)
}
if got := stored.Credentials["access_token"]; got != newToken {
t.Fatalf("stored access_token rolled back = %v, want new token", got)
}
if got := stored.Credentials["refresh_token"]; got != "refresh-new" {
t.Fatalf("stored refresh_token = %v, want refresh-new", got)
}
}
func TestImportCodexSessionsWithRefreshTokenKeepsExistingDedup(t *testing.T) {
existingToken := buildCodexAccessToken(t, "workspace-1", "user-1", time.Now().Add(time.Hour))
svc := newCodexImportMemoryAdminService([]service.Account{{
ID: 11,
Name: "existing",
Platform: service.PlatformOpenAI,
Type: service.AccountTypeOAuth,
Credentials: map[string]any{
"chatgpt_account_id": "workspace-1",
"chatgpt_user_id": "user-1",
"access_token": existingToken,
"refresh_token": "refresh-old",
},
}})
handler := NewAccountHandler(svc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
req := CodexSessionImportRequest{SkipDefaultGroupBind: boolPtr(true)}
entries := []codexImportEntry{
{Index: 1, Value: buildCodexRefreshImportValue(t, "workspace-1", "user-1", "refresh-new")},
}
result, err := handler.importCodexSessions(context.Background(), req, entries)
if err != nil {
t.Fatalf("importCodexSessions error = %v", err)
}
if result.Created != 0 || result.Updated != 1 || result.Failed != 0 {
t.Fatalf("result = %+v, want one updated account", result)
}
if got := svc.updatedAccounts[0].input.Credentials["refresh_token"]; got != "refresh-new" {
t.Fatalf("updated refresh_token = %v, want refresh-new", got)
}
}
type codexImportMemoryAdminService struct {
*stubAdminService
nextID int64
updatedAccounts []struct {
id int64
input *service.UpdateAccountInput
}
}
func newCodexImportMemoryAdminService(accounts []service.Account) *codexImportMemoryAdminService {
stub := newStubAdminService()
stub.accounts = append([]service.Account(nil), accounts...)
return &codexImportMemoryAdminService{
stubAdminService: stub,
nextID: 100,
}
}
func (s *codexImportMemoryAdminService) CreateAccount(ctx context.Context, input *service.CreateAccountInput) (*service.Account, error) {
s.createdAccounts = append(s.createdAccounts, input)
if s.createAccountErr != nil {
return nil, s.createAccountErr
}
account := service.Account{
ID: s.nextID,
Name: input.Name,
Platform: input.Platform,
Type: input.Type,
Status: service.StatusActive,
Credentials: cloneCodexImportTestMap(input.Credentials),
Extra: cloneCodexImportTestMap(input.Extra),
}
s.nextID++
s.accounts = append(s.accounts, account)
return &account, nil
}
func (s *codexImportMemoryAdminService) UpdateAccount(ctx context.Context, id int64, input *service.UpdateAccountInput) (*service.Account, error) {
s.updatedAccounts = append(s.updatedAccounts, struct {
id int64
input *service.UpdateAccountInput
}{id: id, input: input})
if s.updateAccountErr != nil {
return nil, s.updateAccountErr
}
for idx := range s.accounts {
if s.accounts[idx].ID == id {
s.accounts[idx].Credentials = cloneCodexImportTestMap(input.Credentials)
s.accounts[idx].Extra = cloneCodexImportTestMap(input.Extra)
return &s.accounts[idx], nil
}
}
account := service.Account{ID: id, Status: service.StatusActive, Credentials: cloneCodexImportTestMap(input.Credentials)}
return &account, nil
}
func (s *codexImportMemoryAdminService) GetAccount(ctx context.Context, id int64) (*service.Account, error) {
for idx := range s.accounts {
if s.accounts[idx].ID == id {
return &s.accounts[idx], nil
}
}
return s.stubAdminService.GetAccount(ctx, id)
}
func buildCodexAccessOnlyImportValue(t *testing.T, accountID, userID string) map[string]any {
t.Helper()
return map[string]any{
"access_token": buildCodexAccessToken(t, accountID, userID, time.Now().Add(time.Hour)),
}
}
func buildCodexRefreshImportValue(t *testing.T, accountID, userID, refreshToken string) map[string]any {
t.Helper()
return map[string]any{
"access_token": buildCodexAccessToken(t, accountID, userID, time.Now().Add(time.Hour)),
"refresh_token": refreshToken,
}
}
func buildCodexAccessToken(t *testing.T, accountID, userID string, exp time.Time) string {
t.Helper()
return buildCodexAccessTokenWithJTI(t, accountID, userID, "", exp)
}
func buildCodexAccessTokenWithJTI(t *testing.T, accountID, userID, jti string, exp time.Time) string {
t.Helper()
claims := map[string]any{
"sub": userID,
"https://api.openai.com/auth": map[string]any{
"chatgpt_account_id": accountID,
},
}
if jti != "" {
claims["jti"] = jti
}
return buildCodexImportTestJWT(t, exp, claims)
}
func cloneCodexImportTestMap(input map[string]any) map[string]any {
if input == nil {
return nil
}
out := make(map[string]any, len(input))
for key, value := range input {
out[key] = value
}
return out
}
func boolPtr(v bool) *bool {
return &v
}
func buildCodexImportTestJWT(t *testing.T, exp time.Time, extraClaims map[string]any) string {
t.Helper()
header := map[string]any{
@@ -171,13 +171,29 @@ type CheckMixedChannelRequest struct {
// AccountWithConcurrency extends Account with real-time concurrency info
type AccountWithConcurrency struct {
*dto.Account
CurrentConcurrency int `json:"current_concurrency"`
CurrentConcurrency int `json:"current_concurrency"`
SchedulerScore *AccountSchedulerScore `json:"scheduler_score,omitempty"`
SchedulerScores []AccountSchedulerGroupScore `json:"scheduler_scores,omitempty"`
// 以下字段仅对 Anthropic OAuth/SetupToken 账号有效,且仅在启用相应功能时返回
CurrentWindowCost *float64 `json:"current_window_cost,omitempty"` // 当前窗口费用
ActiveSessions *int `json:"active_sessions,omitempty"` // 当前活跃会话数
CurrentRPM *int `json:"current_rpm,omitempty"` // 当前分钟 RPM 计数
}
type AccountSchedulerScore struct {
BaseScore float64 `json:"base_score"`
StickyScore float64 `json:"sticky_score"`
StickyScoreInfinity bool `json:"sticky_score_infinity"`
StickyWeightedEnabled bool `json:"sticky_weighted_enabled"`
}
type AccountSchedulerGroupScore struct {
GroupID *int64 `json:"group_id"`
GroupName string `json:"group_name,omitempty"`
GroupPriority *int `json:"group_priority,omitempty"`
AccountSchedulerScore
}
const accountListGroupUngroupedQueryValue = "ungrouped"
func (h *AccountHandler) buildAccountResponseWithRuntime(ctx context.Context, account *service.Account) AccountWithConcurrency {
@@ -226,6 +242,232 @@ func (h *AccountHandler) buildAccountResponseWithRuntime(ctx context.Context, ac
return item
}
// scoreOpenAIAccountSchedulerPool 对池内 OpenAI 账号计算调度分数快照。
// loadMap 为共享的账号负载数据(含池内全部账号即可,多余条目无害);传 nil 时自行批查。
func (h *AccountHandler) scoreOpenAIAccountSchedulerPool(ctx context.Context, accounts []service.Account, loadMap map[int64]*service.AccountLoadInfo) map[int64]AccountSchedulerScore {
if len(accounts) == 0 {
return nil
}
openAIAccounts := make([]*service.Account, 0, len(accounts))
for i := range accounts {
account := &accounts[i]
if account.Platform != service.PlatformOpenAI {
continue
}
openAIAccounts = append(openAIAccounts, account)
}
if len(openAIAccounts) == 0 {
return nil
}
if loadMap == nil {
loadMap = h.fetchOpenAIAccountLoadMap(ctx, openAIAccounts)
}
var scores map[int64]service.OpenAIAccountSchedulerScoreSnapshot
if h.rateLimitService != nil {
scores = h.rateLimitService.BuildOpenAIAccountSchedulerScoreSnapshot(ctx, openAIAccounts, loadMap)
} else {
scores = service.BuildOpenAIAccountSchedulerScoreSnapshot(openAIAccounts, loadMap)
}
result := make(map[int64]AccountSchedulerScore, len(scores))
for accountID, score := range scores {
result[accountID] = AccountSchedulerScore{
BaseScore: score.BaseScore,
StickyScore: score.StickyScore,
StickyScoreInfinity: score.StickyScoreInfinity,
StickyWeightedEnabled: score.StickyWeightedEnabled,
}
}
return result
}
// fetchOpenAIAccountLoadMap 一次性批查给定 OpenAI 账号的负载数据;
// 失败时记录日志并返回空表(分数按零负载计算,属可接受降级)。
func (h *AccountHandler) fetchOpenAIAccountLoadMap(ctx context.Context, openAIAccounts []*service.Account) map[int64]*service.AccountLoadInfo {
loadMap := map[int64]*service.AccountLoadInfo{}
if h.concurrencyService == nil || len(openAIAccounts) == 0 {
return loadMap
}
seen := make(map[int64]struct{}, len(openAIAccounts))
loadReq := make([]service.AccountWithConcurrency, 0, len(openAIAccounts))
for _, account := range openAIAccounts {
if account == nil {
continue
}
if _, ok := seen[account.ID]; ok {
continue
}
seen[account.ID] = struct{}{}
loadReq = append(loadReq, service.AccountWithConcurrency{
ID: account.ID,
MaxConcurrency: account.EffectiveLoadFactor(),
})
}
if batchLoad, err := h.concurrencyService.GetAccountsLoadBatch(ctx, loadReq); err != nil {
slog.Warn("openai_scheduler_score_load_batch_failed", "error", err)
} else if batchLoad != nil {
loadMap = batchLoad
}
return loadMap
}
func (h *AccountHandler) buildOpenAIAccountSchedulerScores(
ctx context.Context,
accounts []service.Account,
filterPool []service.Account,
) (map[int64]*AccountSchedulerScore, map[int64][]AccountSchedulerGroupScore) {
if len(accounts) == 0 {
return nil, nil
}
if len(filterPool) == 0 {
filterPool = accounts
}
pageOpenAIAccountIDs := make(map[int64]struct{})
groupIDs := make(map[int64]struct{})
for i := range accounts {
account := &accounts[i]
if account.Platform != service.PlatformOpenAI {
continue
}
pageOpenAIAccountIDs[account.ID] = struct{}{}
if len(account.AccountGroups) == 0 && len(account.GroupIDs) == 0 {
continue
}
for _, accountGroup := range account.AccountGroups {
if accountGroup.GroupID > 0 {
groupIDs[accountGroup.GroupID] = struct{}{}
}
}
for _, groupID := range account.GroupIDs {
if groupID > 0 {
groupIDs[groupID] = struct{}{}
}
}
}
if len(pageOpenAIAccountIDs) == 0 {
return nil, nil
}
// 先取各分组池,再对"过滤池 ∪ 分组池"的账号并集做一次负载批查,
// 避免每个池各查一次 Redis 的 N+1。
groupIDList := make([]int64, 0, len(groupIDs))
for groupID := range groupIDs {
groupIDList = append(groupIDList, groupID)
}
sort.Slice(groupIDList, func(i, j int) bool { return groupIDList[i] < groupIDList[j] })
groupPools := make(map[int64][]service.Account, len(groupIDList))
if h.adminService != nil {
for _, groupID := range groupIDList {
gid := groupID
pool, err := h.adminService.ListOpenAISchedulableAccountsForSchedulerScore(ctx, &gid)
if err != nil {
slog.Warn("openai_scheduler_group_score_pool_failed", "group_id", gid, "error", err)
continue
}
groupPools[gid] = pool
}
}
loadUnion := make([]*service.Account, 0, len(filterPool))
collectOpenAIAccounts := func(pool []service.Account) {
for i := range pool {
if pool[i].Platform == service.PlatformOpenAI {
loadUnion = append(loadUnion, &pool[i])
}
}
}
collectOpenAIAccounts(filterPool)
for _, pool := range groupPools {
collectOpenAIAccounts(pool)
}
loadMap := h.fetchOpenAIAccountLoadMap(ctx, loadUnion)
baseScores := make(map[int64]*AccountSchedulerScore)
for accountID, score := range h.scoreOpenAIAccountSchedulerPool(ctx, filterPool, loadMap) {
copiedScore := score
baseScores[accountID] = &copiedScore
}
groupScoresByAccount := make(map[int64][]AccountSchedulerGroupScore)
scoreGroupPool := func(groupID *int64, groupNameByID map[int64]string, groupPriorityByAccount map[int64]int, pool []service.Account) {
if len(pool) == 0 {
return
}
scores := h.scoreOpenAIAccountSchedulerPool(ctx, pool, loadMap)
for accountID, schedulerScore := range scores {
if _, ok := pageOpenAIAccountIDs[accountID]; !ok {
continue
}
groupScore := AccountSchedulerGroupScore{
GroupID: groupID,
AccountSchedulerScore: schedulerScore,
}
if groupID != nil {
groupScore.GroupName = groupNameByID[*groupID]
if priority, ok := groupPriorityByAccount[accountID]; ok {
groupScore.GroupPriority = &priority
}
}
groupScoresByAccount[accountID] = append(groupScoresByAccount[accountID], groupScore)
}
}
for _, groupID := range groupIDList {
gid := groupID
pool, ok := groupPools[gid]
if !ok {
continue
}
groupNameByID := make(map[int64]string)
groupPriorityByAccount := make(map[int64]int)
for i := range pool {
account := &pool[i]
for _, accountGroup := range account.AccountGroups {
if accountGroup.GroupID != gid {
continue
}
groupPriorityByAccount[account.ID] = accountGroup.Priority
if accountGroup.Group != nil {
groupNameByID[gid] = accountGroup.Group.Name
}
}
}
scoreGroupPool(&gid, groupNameByID, groupPriorityByAccount, pool)
}
for accountID := range groupScoresByAccount {
sort.SliceStable(groupScoresByAccount[accountID], func(i, j int) bool {
left := groupScoresByAccount[accountID][i]
right := groupScoresByAccount[accountID][j]
return *left.GroupID < *right.GroupID
})
}
return baseScores, groupScoresByAccount
}
func (h *AccountHandler) listAccountSchedulerScoreFilterPool(
ctx context.Context,
platform, accountType, status, search string,
groupID int64,
privacyMode string,
) []service.Account {
if h.adminService == nil || (platform != "" && platform != service.PlatformOpenAI) {
return nil
}
// 池只用于 OpenAI 分数计算(非 OpenAI 账号会在打分时被丢弃),
// 无论列表页平台过滤为何,查询一律限定 openai,避免无过滤时全表扫描。
accounts, err := h.adminService.ListAccountsForSchedulerScoreFilter(ctx, service.PlatformOpenAI, accountType, status, search, groupID, privacyMode)
if err != nil {
slog.Warn("openai_scheduler_filter_score_pool_failed", "error", err)
return nil
}
return accounts
}
// List handles listing all accounts with pagination
// GET /api/v1/admin/accounts
func (h *AccountHandler) List(c *gin.Context) {
@@ -278,6 +520,20 @@ func (h *AccountHandler) List(c *gin.Context) {
var windowCosts map[int64]float64
var activeSessions map[int64]int
var rpmCounts map[int64]int
// 仅当前页存在 OpenAI 账号时才计算调度分数,避免为空结果付出池查询开销。
var schedulerScores map[int64]*AccountSchedulerScore
var schedulerGroupScores map[int64][]AccountSchedulerGroupScore
pageHasOpenAIAccounts := false
for i := range accounts {
if accounts[i].Platform == service.PlatformOpenAI {
pageHasOpenAIAccounts = true
break
}
}
if pageHasOpenAIAccounts {
schedulerFilterPool := h.listAccountSchedulerScoreFilterPool(c.Request.Context(), platform, accountType, status, search, groupID, privacyMode)
schedulerScores, schedulerGroupScores = h.buildOpenAIAccountSchedulerScores(c.Request.Context(), accounts, schedulerFilterPool)
}
// 始终获取并发数(Redis ZCARD,极低开销)
if h.concurrencyService != nil {
@@ -358,6 +614,8 @@ func (h *AccountHandler) List(c *gin.Context) {
item := AccountWithConcurrency{
Account: dto.AccountFromService(acc),
CurrentConcurrency: concurrencyCounts[acc.ID],
SchedulerScore: schedulerScores[acc.ID],
SchedulerScores: schedulerGroupScores[acc.ID],
}
// 添加窗口费用(仅当启用时)
@@ -8,6 +8,7 @@ import (
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
@@ -50,3 +51,222 @@ func TestAccountHandlerListIncludesCreatedAt(t *testing.T) {
_, offset := parsed.Zone()
require.Equal(t, 0, offset)
}
func TestAccountHandlerListReturnsSchedulerScoresPerGroup(t *testing.T) {
router, adminSvc := setupAccountListRouter()
now := time.Now().UTC()
groupID := int64(41)
adminSvc.accounts = []service.Account{
{
ID: 101,
Name: "account-high-priority",
Platform: service.PlatformOpenAI,
Type: service.AccountTypeAPIKey,
Status: service.StatusActive,
Schedulable: true,
Concurrency: 10,
Priority: 1,
AccountGroups: []service.AccountGroup{
{AccountID: 101, GroupID: groupID, Priority: 100, Group: &service.Group{ID: groupID, Name: "openai"}},
},
GroupIDs: []int64{groupID},
CreatedAt: now,
UpdatedAt: now,
},
{
ID: 102,
Name: "account-low-priority",
Platform: service.PlatformOpenAI,
Type: service.AccountTypeAPIKey,
Status: service.StatusActive,
Schedulable: true,
Concurrency: 10,
Priority: 100000,
AccountGroups: []service.AccountGroup{
{AccountID: 102, GroupID: groupID, Priority: 1, Group: &service.Group{ID: groupID, Name: "openai"}},
},
GroupIDs: []int64{groupID},
CreatedAt: now,
UpdatedAt: now,
},
}
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/accounts?page=1&page_size=20&platform=openai", nil)
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code)
var payload struct {
Data struct {
Items []struct {
ID int64 `json:"id"`
SchedulerScore struct {
BaseScore float64 `json:"base_score"`
} `json:"scheduler_score"`
SchedulerScores []struct {
GroupID *int64 `json:"group_id"`
GroupName string `json:"group_name"`
GroupPriority *int `json:"group_priority"`
BaseScore float64 `json:"base_score"`
} `json:"scheduler_scores"`
} `json:"items"`
} `json:"data"`
}
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &payload))
require.Len(t, payload.Data.Items, 2)
var high, low *struct {
ID int64 `json:"id"`
SchedulerScore struct {
BaseScore float64 `json:"base_score"`
} `json:"scheduler_score"`
SchedulerScores []struct {
GroupID *int64 `json:"group_id"`
GroupName string `json:"group_name"`
GroupPriority *int `json:"group_priority"`
BaseScore float64 `json:"base_score"`
} `json:"scheduler_scores"`
}
for i := range payload.Data.Items {
item := &payload.Data.Items[i]
switch item.ID {
case 101:
high = item
case 102:
low = item
}
}
require.NotNil(t, high)
require.NotNil(t, low)
require.Len(t, high.SchedulerScores, 1)
require.Len(t, low.SchedulerScores, 1)
require.Equal(t, groupID, *high.SchedulerScores[0].GroupID)
require.Equal(t, "openai", high.SchedulerScores[0].GroupName)
require.Equal(t, 100, *high.SchedulerScores[0].GroupPriority)
require.Equal(t, 1, *low.SchedulerScores[0].GroupPriority)
require.Greater(t, high.SchedulerScores[0].BaseScore, low.SchedulerScores[0].BaseScore)
}
func TestAccountHandlerListKeepsSchedulerScoreScopedToFilter(t *testing.T) {
router, adminSvc := setupAccountListRouter()
now := time.Now().UTC()
groupID := int64(42)
visibleAccount := service.Account{
ID: 201,
Name: "visible-low-priority",
Platform: service.PlatformOpenAI,
Type: service.AccountTypeAPIKey,
Status: service.StatusActive,
Schedulable: true,
Concurrency: 10,
Priority: 100000,
AccountGroups: []service.AccountGroup{
{AccountID: 201, GroupID: groupID, Priority: 1, Group: &service.Group{ID: groupID, Name: "openai"}},
},
GroupIDs: []int64{groupID},
CreatedAt: now,
UpdatedAt: now,
}
hiddenGroupPeer := service.Account{
ID: 202,
Name: "hidden-high-priority",
Platform: service.PlatformOpenAI,
Type: service.AccountTypeAPIKey,
Status: service.StatusActive,
Schedulable: true,
Concurrency: 10,
Priority: 1,
AccountGroups: []service.AccountGroup{
{AccountID: 202, GroupID: groupID, Priority: 2, Group: &service.Group{ID: groupID, Name: "openai"}},
},
GroupIDs: []int64{groupID},
CreatedAt: now,
UpdatedAt: now,
}
adminSvc.accounts = []service.Account{visibleAccount}
adminSvc.accountSchedulerScoreFilterAccounts = []service.Account{visibleAccount, hiddenGroupPeer}
adminSvc.openAISchedulerScorePoolAccounts = []service.Account{visibleAccount, hiddenGroupPeer}
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/accounts?page=1&page_size=1&platform=openai", nil)
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code)
var payload struct {
Data struct {
Items []struct {
ID int64 `json:"id"`
SchedulerScore struct {
BaseScore float64 `json:"base_score"`
} `json:"scheduler_score"`
SchedulerScores []struct {
GroupID *int64 `json:"group_id"`
BaseScore float64 `json:"base_score"`
} `json:"scheduler_scores"`
} `json:"items"`
} `json:"data"`
}
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &payload))
require.Len(t, payload.Data.Items, 1)
item := payload.Data.Items[0]
require.Equal(t, int64(201), item.ID)
require.Len(t, item.SchedulerScores, 1)
require.Equal(t, groupID, *item.SchedulerScores[0].GroupID)
require.Equal(t, item.SchedulerScores[0].BaseScore, item.SchedulerScore.BaseScore)
}
func TestAccountHandlerListSchedulerScoreIgnoresPagination(t *testing.T) {
router, adminSvc := setupAccountListRouter()
now := time.Now().UTC()
visibleAccount := service.Account{
ID: 301,
Name: "visible-low-priority",
Platform: service.PlatformOpenAI,
Type: service.AccountTypeAPIKey,
Status: service.StatusActive,
Schedulable: true,
Concurrency: 10,
Priority: 100000,
CreatedAt: now,
UpdatedAt: now,
}
hiddenFilterPeer := service.Account{
ID: 302,
Name: "hidden-high-priority",
Platform: service.PlatformOpenAI,
Type: service.AccountTypeAPIKey,
Status: service.StatusActive,
Schedulable: true,
Concurrency: 10,
Priority: 1,
CreatedAt: now,
UpdatedAt: now,
}
adminSvc.accounts = []service.Account{visibleAccount}
adminSvc.accountSchedulerScoreFilterAccounts = []service.Account{visibleAccount, hiddenFilterPeer}
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/accounts?page=1&page_size=1&platform=openai", nil)
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code)
var payload struct {
Data struct {
Items []struct {
ID int64 `json:"id"`
SchedulerScore struct {
BaseScore float64 `json:"base_score"`
} `json:"scheduler_score"`
SchedulerScores []struct {
GroupID *int64 `json:"group_id"`
BaseScore float64 `json:"base_score"`
} `json:"scheduler_scores"`
} `json:"items"`
} `json:"data"`
}
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &payload))
require.Len(t, payload.Data.Items, 1)
require.Equal(t, int64(301), payload.Data.Items[0].ID)
require.Less(t, payload.Data.Items[0].SchedulerScore.BaseScore, 3.75)
require.Empty(t, payload.Data.Items[0].SchedulerScores)
}
@@ -10,27 +10,29 @@ import (
)
type stubAdminService struct {
users []service.User
apiKeys []service.APIKey
groups []service.Group
accounts []service.Account
proxies []service.Proxy
proxyCounts []service.ProxyWithAccountCount
redeems []service.RedeemCode
boundAuthIdentity *service.AdminBindAuthIdentityInput
boundAuthIdentityFor int64
createdAccounts []*service.CreateAccountInput
createdProxies []*service.CreateProxyInput
updatedProxyIDs []int64
updatedProxies []*service.UpdateProxyInput
testedProxyIDs []int64
getUserErr error
createAccountErr error
createSparkShadowErr error
updateAccountErr error
bulkUpdateAccountErr error
checkMixedErr error
lastMixedCheck struct {
users []service.User
apiKeys []service.APIKey
groups []service.Group
accounts []service.Account
accountSchedulerScoreFilterAccounts []service.Account
openAISchedulerScorePoolAccounts []service.Account
proxies []service.Proxy
proxyCounts []service.ProxyWithAccountCount
redeems []service.RedeemCode
boundAuthIdentity *service.AdminBindAuthIdentityInput
boundAuthIdentityFor int64
createdAccounts []*service.CreateAccountInput
createdProxies []*service.CreateProxyInput
updatedProxyIDs []int64
updatedProxies []*service.UpdateProxyInput
testedProxyIDs []int64
getUserErr error
createAccountErr error
createSparkShadowErr error
updateAccountErr error
bulkUpdateAccountErr error
checkMixedErr error
lastMixedCheck struct {
accountID int64
platform string
groupIDs []int64
@@ -329,7 +331,56 @@ func (s *stubAdminService) ListAccounts(ctx context.Context, page, pageSize int,
s.lastListAccounts.sortBy = sortBy
s.lastListAccounts.sortOrder = sortOrder
s.lastListAccounts.calls++
return s.accounts, int64(len(s.accounts)), nil
accounts := s.accounts
total := len(accounts)
if page < 1 {
page = 1
}
if pageSize < 1 {
pageSize = total
}
start := (page - 1) * pageSize
if start >= total {
return []service.Account{}, int64(total), nil
}
end := start + pageSize
if end > total {
end = total
}
return accounts[start:end], int64(total), nil
}
func (s *stubAdminService) ListAccountsForSchedulerScoreFilter(_ context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]service.Account, error) {
if s.accountSchedulerScoreFilterAccounts != nil {
return s.accountSchedulerScoreFilterAccounts, nil
}
return s.accounts, nil
}
func (s *stubAdminService) ListOpenAISchedulableAccountsForSchedulerScore(_ context.Context, groupID *int64) ([]service.Account, error) {
accounts := s.openAISchedulerScorePoolAccounts
if accounts == nil {
accounts = s.accounts
}
out := make([]service.Account, 0, len(accounts))
for _, account := range accounts {
if account.Platform != service.PlatformOpenAI || !account.IsSchedulable() {
continue
}
if groupID == nil {
if len(account.AccountGroups) == 0 && len(account.GroupIDs) == 0 {
out = append(out, account)
}
continue
}
for _, accountGroup := range account.AccountGroups {
if accountGroup.GroupID == *groupID {
out = append(out, account)
break
}
}
}
return out, nil
}
func (s *stubAdminService) GetAccount(ctx context.Context, id int64) (*service.Account, error) {
+41 -8
View File
@@ -73,6 +73,13 @@ func NewOpsHandler(opsService *service.OpsService) *OpsHandler {
}
// GetErrorLogs lists ops error logs.
// applyOpsErrorSortParams reads sort_by/sort_order query params into the filter.
// Column whitelist and order normalization live in the repository; unknown
// values degrade to the default (created_at DESC), mirroring the usage list.
func applyOpsErrorSortParams(c *gin.Context, filter *service.OpsErrorLogFilter) {
filter.SetSort(c.Query("sort_by"), c.Query("sort_order"))
}
// GET /api/v1/admin/ops/errors
func (h *OpsHandler) GetErrorLogs(c *gin.Context) {
if h.opsService == nil {
@@ -114,10 +121,17 @@ func (h *OpsHandler) GetErrorLogs(c *gin.Context) {
// buildOpsErrorLogsWhere 以 COALESCE(requested_model, model) 比对。
filter.Model = strings.TrimSpace(c.Query("model"))
// Force request errors: client-visible status >= 400.
// buildOpsErrorLogsWhere already applies this for non-upstream phase.
if strings.EqualFold(strings.TrimSpace(filter.Phase), "upstream") {
filter.Phase = ""
// 请求错误语义:client-visible status>=400 守卫恒生效(未设
// IncludeRecoveredUpstream 时 phase=upstream 不再绕过守卫),故
// phase=upstream 作为普通过滤条件保留——此前这里清空该值,导致
// 错误类型下拉选「上游」等于不过滤。
// 分类(用户侧粗分类码)→ phase/type ANY 条件,与用户端 /usage/errors 同一映射;
// 未知分类返回空切片 = 不过滤。与 phase 参数可同时设置(AND 语义)。
if cat := strings.TrimSpace(c.Query("category")); cat != "" {
phases, types := service.CategoryToFilter(cat)
filter.ErrorPhasesAny = phases
filter.ErrorTypesAny = types
}
if platform := strings.TrimSpace(c.Query("platform")); platform != "" {
@@ -187,6 +201,8 @@ func (h *OpsHandler) GetErrorLogs(c *gin.Context) {
filter.StatusCodes = out
}
applyOpsErrorSortParams(c, filter)
result, err := h.opsService.GetErrorLogs(c.Request.Context(), filter)
if err != nil {
response.ErrorFrom(c, err)
@@ -234,10 +250,17 @@ func (h *OpsHandler) ListRequestErrors(c *gin.Context) {
// buildOpsErrorLogsWhere 以 COALESCE(requested_model, model) 比对。
filter.Model = strings.TrimSpace(c.Query("model"))
// Force request errors: client-visible status >= 400.
// buildOpsErrorLogsWhere already applies this for non-upstream phase.
if strings.EqualFold(strings.TrimSpace(filter.Phase), "upstream") {
filter.Phase = ""
// 请求错误语义:client-visible status>=400 守卫恒生效(未设
// IncludeRecoveredUpstream 时 phase=upstream 不再绕过守卫),故
// phase=upstream 作为普通过滤条件保留——此前这里清空该值,导致
// 错误类型下拉选「上游」等于不过滤。
// 分类(用户侧粗分类码)→ phase/type ANY 条件,与用户端 /usage/errors 同一映射;
// 未知分类返回空切片 = 不过滤。与 phase 参数可同时设置(AND 语义)。
if cat := strings.TrimSpace(c.Query("category")); cat != "" {
phases, types := service.CategoryToFilter(cat)
filter.ErrorPhasesAny = phases
filter.ErrorTypesAny = types
}
if platform := strings.TrimSpace(c.Query("platform")); platform != "" {
@@ -291,6 +314,8 @@ func (h *OpsHandler) ListRequestErrors(c *gin.Context) {
filter.StatusCodes = out
}
applyOpsErrorSortParams(c, filter)
result, err := h.opsService.GetErrorLogs(c.Request.Context(), filter)
if err != nil {
response.ErrorFrom(c, err)
@@ -362,6 +387,8 @@ func (h *OpsHandler) ListRequestErrorUpstreamErrors(c *gin.Context) {
}
filter.View = "all"
filter.Phase = "upstream"
// 上游错误列表需含 status<400 的 recovered 行,显式豁免客户端可见守卫。
filter.IncludeRecoveredUpstream = true
filter.Owner = "provider"
filter.Source = strings.TrimSpace(c.Query("error_source"))
filter.Query = strings.TrimSpace(c.Query("q"))
@@ -377,6 +404,8 @@ func (h *OpsHandler) ListRequestErrorUpstreamErrors(c *gin.Context) {
filter.ClientRequestID = clientRequestID
}
applyOpsErrorSortParams(c, filter)
result, err := h.opsService.GetErrorLogs(c.Request.Context(), filter)
if err != nil {
response.ErrorFrom(c, err)
@@ -442,6 +471,8 @@ func (h *OpsHandler) ListUpstreamErrors(c *gin.Context) {
filter.View = parseOpsViewParam(c)
filter.Phase = "upstream"
// 上游错误列表需含 status<400 的 recovered 行,显式豁免客户端可见守卫。
filter.IncludeRecoveredUpstream = true
filter.Owner = "provider"
filter.Source = strings.TrimSpace(c.Query("error_source"))
filter.Query = strings.TrimSpace(c.Query("q"))
@@ -497,6 +528,8 @@ func (h *OpsHandler) ListUpstreamErrors(c *gin.Context) {
filter.StatusCodes = out
}
applyOpsErrorSortParams(c, filter)
result, err := h.opsService.GetErrorLogs(c.Request.Context(), filter)
if err != nil {
response.ErrorFrom(c, err)
+488 -362
View File
@@ -119,188 +119,211 @@ func (h *SettingHandler) GetSettings(c *gin.Context) {
}
payload := dto.SystemSettings{
RegistrationEnabled: settings.RegistrationEnabled,
EmailVerifyEnabled: settings.EmailVerifyEnabled,
RegistrationEmailSuffixWhitelist: settings.RegistrationEmailSuffixWhitelist,
PromoCodeEnabled: settings.PromoCodeEnabled,
PasswordResetEnabled: settings.PasswordResetEnabled,
FrontendURL: settings.FrontendURL,
InvitationCodeEnabled: settings.InvitationCodeEnabled,
TotpEnabled: settings.TotpEnabled,
TotpEncryptionKeyConfigured: h.settingService.IsTotpEncryptionKeyConfigured(),
LoginAgreementEnabled: settings.LoginAgreementEnabled,
LoginAgreementMode: settings.LoginAgreementMode,
LoginAgreementUpdatedAt: settings.LoginAgreementUpdatedAt,
LoginAgreementDocuments: loginAgreementDocumentsToDTO(settings.LoginAgreementDocuments),
SMTPHost: settings.SMTPHost,
SMTPPort: settings.SMTPPort,
SMTPUsername: settings.SMTPUsername,
SMTPPasswordConfigured: settings.SMTPPasswordConfigured,
SMTPFrom: settings.SMTPFrom,
SMTPFromName: settings.SMTPFromName,
SMTPUseTLS: settings.SMTPUseTLS,
TurnstileEnabled: settings.TurnstileEnabled,
TurnstileSiteKey: settings.TurnstileSiteKey,
TurnstileSecretKeyConfigured: settings.TurnstileSecretKeyConfigured,
APIKeyACLTrustForwardedIP: settings.APIKeyACLTrustForwardedIP,
LinuxDoConnectEnabled: settings.LinuxDoConnectEnabled,
LinuxDoConnectClientID: settings.LinuxDoConnectClientID,
LinuxDoConnectClientSecretConfigured: settings.LinuxDoConnectClientSecretConfigured,
LinuxDoConnectRedirectURL: settings.LinuxDoConnectRedirectURL,
DingTalkConnectEnabled: settings.DingTalkConnectEnabled,
DingTalkConnectClientID: settings.DingTalkConnectClientID,
DingTalkConnectClientSecretConfigured: settings.DingTalkConnectClientSecretConfigured,
DingTalkConnectRedirectURL: settings.DingTalkConnectRedirectURL,
DingTalkConnectCorpRestrictionPolicy: settings.DingTalkConnectCorpRestrictionPolicy,
DingTalkConnectInternalCorpID: settings.DingTalkConnectInternalCorpID,
DingTalkConnectBypassRegistration: settings.DingTalkConnectBypassRegistration,
DingTalkConnectSyncCorpEmail: settings.DingTalkConnectSyncCorpEmail,
DingTalkConnectSyncDisplayName: settings.DingTalkConnectSyncDisplayName,
DingTalkConnectSyncDept: settings.DingTalkConnectSyncDept,
DingTalkConnectSyncCorpEmailAttrKey: settings.DingTalkConnectSyncCorpEmailAttrKey,
DingTalkConnectSyncDisplayNameAttrKey: settings.DingTalkConnectSyncDisplayNameAttrKey,
DingTalkConnectSyncDeptAttrKey: settings.DingTalkConnectSyncDeptAttrKey,
DingTalkConnectSyncCorpEmailAttrName: settings.DingTalkConnectSyncCorpEmailAttrName,
DingTalkConnectSyncDisplayNameAttrName: settings.DingTalkConnectSyncDisplayNameAttrName,
DingTalkConnectSyncDeptAttrName: settings.DingTalkConnectSyncDeptAttrName,
WeChatConnectEnabled: settings.WeChatConnectEnabled,
WeChatConnectAppID: settings.WeChatConnectAppID,
WeChatConnectAppSecretConfigured: settings.WeChatConnectAppSecretConfigured,
WeChatConnectOpenAppID: settings.WeChatConnectOpenAppID,
WeChatConnectOpenAppSecretConfigured: settings.WeChatConnectOpenAppSecretConfigured,
WeChatConnectMPAppID: settings.WeChatConnectMPAppID,
WeChatConnectMPAppSecretConfigured: settings.WeChatConnectMPAppSecretConfigured,
WeChatConnectMobileAppID: settings.WeChatConnectMobileAppID,
WeChatConnectMobileAppSecretConfigured: settings.WeChatConnectMobileAppSecretConfigured,
WeChatConnectOpenEnabled: settings.WeChatConnectOpenEnabled,
WeChatConnectMPEnabled: settings.WeChatConnectMPEnabled,
WeChatConnectMobileEnabled: settings.WeChatConnectMobileEnabled,
WeChatConnectMode: settings.WeChatConnectMode,
WeChatConnectScopes: settings.WeChatConnectScopes,
WeChatConnectRedirectURL: settings.WeChatConnectRedirectURL,
WeChatConnectFrontendRedirectURL: settings.WeChatConnectFrontendRedirectURL,
OIDCConnectEnabled: settings.OIDCConnectEnabled,
OIDCConnectProviderName: settings.OIDCConnectProviderName,
OIDCConnectClientID: settings.OIDCConnectClientID,
OIDCConnectClientSecretConfigured: settings.OIDCConnectClientSecretConfigured,
OIDCConnectIssuerURL: settings.OIDCConnectIssuerURL,
OIDCConnectDiscoveryURL: settings.OIDCConnectDiscoveryURL,
OIDCConnectAuthorizeURL: settings.OIDCConnectAuthorizeURL,
OIDCConnectTokenURL: settings.OIDCConnectTokenURL,
OIDCConnectUserInfoURL: settings.OIDCConnectUserInfoURL,
OIDCConnectJWKSURL: settings.OIDCConnectJWKSURL,
OIDCConnectScopes: settings.OIDCConnectScopes,
OIDCConnectRedirectURL: settings.OIDCConnectRedirectURL,
OIDCConnectFrontendRedirectURL: settings.OIDCConnectFrontendRedirectURL,
OIDCConnectTokenAuthMethod: settings.OIDCConnectTokenAuthMethod,
OIDCConnectUsePKCE: settings.OIDCConnectUsePKCE,
OIDCConnectValidateIDToken: settings.OIDCConnectValidateIDToken,
OIDCConnectAllowedSigningAlgs: settings.OIDCConnectAllowedSigningAlgs,
OIDCConnectClockSkewSeconds: settings.OIDCConnectClockSkewSeconds,
OIDCConnectRequireEmailVerified: settings.OIDCConnectRequireEmailVerified,
OIDCConnectUserInfoEmailPath: settings.OIDCConnectUserInfoEmailPath,
OIDCConnectUserInfoIDPath: settings.OIDCConnectUserInfoIDPath,
OIDCConnectUserInfoUsernamePath: settings.OIDCConnectUserInfoUsernamePath,
GitHubOAuthEnabled: settings.GitHubOAuthEnabled,
GitHubOAuthClientID: settings.GitHubOAuthClientID,
GitHubOAuthClientSecretConfigured: settings.GitHubOAuthClientSecretConfigured,
GitHubOAuthRedirectURL: settings.GitHubOAuthRedirectURL,
GitHubOAuthFrontendRedirectURL: settings.GitHubOAuthFrontendRedirectURL,
GoogleOAuthEnabled: settings.GoogleOAuthEnabled,
GoogleOAuthClientID: settings.GoogleOAuthClientID,
GoogleOAuthClientSecretConfigured: settings.GoogleOAuthClientSecretConfigured,
GoogleOAuthRedirectURL: settings.GoogleOAuthRedirectURL,
GoogleOAuthFrontendRedirectURL: settings.GoogleOAuthFrontendRedirectURL,
SiteName: settings.SiteName,
SiteLogo: settings.SiteLogo,
SiteSubtitle: settings.SiteSubtitle,
APIBaseURL: settings.APIBaseURL,
ContactInfo: settings.ContactInfo,
DocURL: settings.DocURL,
HomeContent: settings.HomeContent,
HideCcsImportButton: settings.HideCcsImportButton,
PurchaseSubscriptionEnabled: settings.PurchaseSubscriptionEnabled,
PurchaseSubscriptionURL: settings.PurchaseSubscriptionURL,
TableDefaultPageSize: settings.TableDefaultPageSize,
TablePageSizeOptions: settings.TablePageSizeOptions,
CustomMenuItems: dto.ParseCustomMenuItems(settings.CustomMenuItems),
CustomEndpoints: dto.ParseCustomEndpoints(settings.CustomEndpoints),
DefaultConcurrency: settings.DefaultConcurrency,
DefaultBalance: settings.DefaultBalance,
RiskControlEnabled: settings.RiskControlEnabled,
CyberSessionBlockEnabled: settings.CyberSessionBlockEnabled,
CyberSessionBlockTTLSeconds: settings.CyberSessionBlockTTLSeconds,
AffiliateRebateRate: settings.AffiliateRebateRate,
AffiliateRebateFreezeHours: settings.AffiliateRebateFreezeHours,
AffiliateRebateDurationDays: settings.AffiliateRebateDurationDays,
AffiliateRebatePerInviteeCap: settings.AffiliateRebatePerInviteeCap,
DefaultUserRPMLimit: settings.DefaultUserRPMLimit,
DefaultSubscriptions: defaultSubscriptions,
EnableModelFallback: settings.EnableModelFallback,
FallbackModelAnthropic: settings.FallbackModelAnthropic,
FallbackModelOpenAI: settings.FallbackModelOpenAI,
FallbackModelGemini: settings.FallbackModelGemini,
FallbackModelAntigravity: settings.FallbackModelAntigravity,
EnableIdentityPatch: settings.EnableIdentityPatch,
IdentityPatchPrompt: settings.IdentityPatchPrompt,
OpsMonitoringEnabled: opsEnabled && settings.OpsMonitoringEnabled,
OpsRealtimeMonitoringEnabled: settings.OpsRealtimeMonitoringEnabled,
OpsQueryModeDefault: settings.OpsQueryModeDefault,
OpsMetricsIntervalSeconds: settings.OpsMetricsIntervalSeconds,
MinClaudeCodeVersion: settings.MinClaudeCodeVersion,
MaxClaudeCodeVersion: settings.MaxClaudeCodeVersion,
AllowUngroupedKeyScheduling: settings.AllowUngroupedKeyScheduling,
BackendModeEnabled: settings.BackendModeEnabled,
EnableFingerprintUnification: settings.EnableFingerprintUnification,
EnableMetadataPassthrough: settings.EnableMetadataPassthrough,
EnableCCHSigning: settings.EnableCCHSigning,
EnableClaudeOAuthSystemPromptInjection: settings.EnableClaudeOAuthSystemPromptInjection,
ClaudeOAuthSystemPrompt: settings.ClaudeOAuthSystemPrompt,
ClaudeOAuthSystemPromptBlocks: settings.ClaudeOAuthSystemPromptBlocks,
EnableAnthropicCacheTTL1hInjection: settings.EnableAnthropicCacheTTL1hInjection,
RewriteMessageCacheControl: settings.RewriteMessageCacheControl,
EnableClientDatelineNormalization: settings.EnableClientDatelineNormalization,
AntigravityUserAgentVersion: settings.AntigravityUserAgentVersion,
OpenAICodexUserAgent: settings.OpenAICodexUserAgent,
MinCodexVersion: settings.MinCodexVersion,
MaxCodexVersion: settings.MaxCodexVersion,
CodexCLIOnlyBlacklist: settings.CodexCLIOnlyBlacklist,
CodexCLIOnlyWhitelist: settings.CodexCLIOnlyWhitelist,
CodexCLIOnlyAllowAppServerClients: settings.CodexCLIOnlyAllowAppServerClients,
CodexCLIOnlyEngineFingerprintSignals: settings.CodexCLIOnlyEngineFingerprintSignals,
WebSearchEmulationEnabled: settings.WebSearchEmulationEnabled,
PaymentVisibleMethodAlipaySource: settings.PaymentVisibleMethodAlipaySource,
PaymentVisibleMethodWxpaySource: settings.PaymentVisibleMethodWxpaySource,
PaymentVisibleMethodAlipayEnabled: settings.PaymentVisibleMethodAlipayEnabled,
PaymentVisibleMethodWxpayEnabled: settings.PaymentVisibleMethodWxpayEnabled,
OpenAIAdvancedSchedulerEnabled: settings.OpenAIAdvancedSchedulerEnabled,
BalanceLowNotifyEnabled: settings.BalanceLowNotifyEnabled,
BalanceLowNotifyThreshold: settings.BalanceLowNotifyThreshold,
BalanceLowNotifyRechargeURL: settings.BalanceLowNotifyRechargeURL,
SubscriptionExpiryNotifyEnabled: settings.SubscriptionExpiryNotifyEnabled,
AccountQuotaNotifyEnabled: settings.AccountQuotaNotifyEnabled,
AccountQuotaNotifyEmails: dto.NotifyEmailEntriesFromService(settings.AccountQuotaNotifyEmails),
PaymentEnabled: paymentCfg.Enabled,
PaymentMinAmount: paymentCfg.MinAmount,
PaymentMaxAmount: paymentCfg.MaxAmount,
PaymentDailyLimit: paymentCfg.DailyLimit,
PaymentOrderTimeoutMin: paymentCfg.OrderTimeoutMin,
PaymentMaxPendingOrders: paymentCfg.MaxPendingOrders,
PaymentEnabledTypes: paymentCfg.EnabledTypes,
PaymentBalanceDisabled: paymentCfg.BalanceDisabled,
PaymentBalanceRechargeMultiplier: paymentCfg.BalanceRechargeMultiplier,
PaymentRechargeFeeRate: paymentCfg.RechargeFeeRate,
PaymentLoadBalanceStrat: paymentCfg.LoadBalanceStrategy,
PaymentProductNamePrefix: paymentCfg.ProductNamePrefix,
PaymentProductNameSuffix: paymentCfg.ProductNameSuffix,
PaymentHelpImageURL: paymentCfg.HelpImageURL,
PaymentHelpText: paymentCfg.HelpText,
PaymentCancelRateLimitEnabled: paymentCfg.CancelRateLimitEnabled,
PaymentCancelRateLimitMax: paymentCfg.CancelRateLimitMax,
PaymentCancelRateLimitWindow: paymentCfg.CancelRateLimitWindow,
PaymentCancelRateLimitUnit: paymentCfg.CancelRateLimitUnit,
PaymentCancelRateLimitMode: paymentCfg.CancelRateLimitMode,
PaymentAlipayForceQRCode: paymentCfg.AlipayForceQRCode,
RegistrationEnabled: settings.RegistrationEnabled,
EmailVerifyEnabled: settings.EmailVerifyEnabled,
RegistrationEmailSuffixWhitelist: settings.RegistrationEmailSuffixWhitelist,
PromoCodeEnabled: settings.PromoCodeEnabled,
PasswordResetEnabled: settings.PasswordResetEnabled,
FrontendURL: settings.FrontendURL,
InvitationCodeEnabled: settings.InvitationCodeEnabled,
TotpEnabled: settings.TotpEnabled,
TotpEncryptionKeyConfigured: h.settingService.IsTotpEncryptionKeyConfigured(),
LoginAgreementEnabled: settings.LoginAgreementEnabled,
LoginAgreementMode: settings.LoginAgreementMode,
LoginAgreementUpdatedAt: settings.LoginAgreementUpdatedAt,
LoginAgreementDocuments: loginAgreementDocumentsToDTO(settings.LoginAgreementDocuments),
SMTPHost: settings.SMTPHost,
SMTPPort: settings.SMTPPort,
SMTPUsername: settings.SMTPUsername,
SMTPPasswordConfigured: settings.SMTPPasswordConfigured,
SMTPFrom: settings.SMTPFrom,
SMTPFromName: settings.SMTPFromName,
SMTPUseTLS: settings.SMTPUseTLS,
TurnstileEnabled: settings.TurnstileEnabled,
TurnstileSiteKey: settings.TurnstileSiteKey,
TurnstileSecretKeyConfigured: settings.TurnstileSecretKeyConfigured,
APIKeyACLTrustForwardedIP: settings.APIKeyACLTrustForwardedIP,
LinuxDoConnectEnabled: settings.LinuxDoConnectEnabled,
LinuxDoConnectClientID: settings.LinuxDoConnectClientID,
LinuxDoConnectClientSecretConfigured: settings.LinuxDoConnectClientSecretConfigured,
LinuxDoConnectRedirectURL: settings.LinuxDoConnectRedirectURL,
DingTalkConnectEnabled: settings.DingTalkConnectEnabled,
DingTalkConnectClientID: settings.DingTalkConnectClientID,
DingTalkConnectClientSecretConfigured: settings.DingTalkConnectClientSecretConfigured,
DingTalkConnectRedirectURL: settings.DingTalkConnectRedirectURL,
DingTalkConnectCorpRestrictionPolicy: settings.DingTalkConnectCorpRestrictionPolicy,
DingTalkConnectInternalCorpID: settings.DingTalkConnectInternalCorpID,
DingTalkConnectBypassRegistration: settings.DingTalkConnectBypassRegistration,
DingTalkConnectSyncCorpEmail: settings.DingTalkConnectSyncCorpEmail,
DingTalkConnectSyncDisplayName: settings.DingTalkConnectSyncDisplayName,
DingTalkConnectSyncDept: settings.DingTalkConnectSyncDept,
DingTalkConnectSyncCorpEmailAttrKey: settings.DingTalkConnectSyncCorpEmailAttrKey,
DingTalkConnectSyncDisplayNameAttrKey: settings.DingTalkConnectSyncDisplayNameAttrKey,
DingTalkConnectSyncDeptAttrKey: settings.DingTalkConnectSyncDeptAttrKey,
DingTalkConnectSyncCorpEmailAttrName: settings.DingTalkConnectSyncCorpEmailAttrName,
DingTalkConnectSyncDisplayNameAttrName: settings.DingTalkConnectSyncDisplayNameAttrName,
DingTalkConnectSyncDeptAttrName: settings.DingTalkConnectSyncDeptAttrName,
WeChatConnectEnabled: settings.WeChatConnectEnabled,
WeChatConnectAppID: settings.WeChatConnectAppID,
WeChatConnectAppSecretConfigured: settings.WeChatConnectAppSecretConfigured,
WeChatConnectOpenAppID: settings.WeChatConnectOpenAppID,
WeChatConnectOpenAppSecretConfigured: settings.WeChatConnectOpenAppSecretConfigured,
WeChatConnectMPAppID: settings.WeChatConnectMPAppID,
WeChatConnectMPAppSecretConfigured: settings.WeChatConnectMPAppSecretConfigured,
WeChatConnectMobileAppID: settings.WeChatConnectMobileAppID,
WeChatConnectMobileAppSecretConfigured: settings.WeChatConnectMobileAppSecretConfigured,
WeChatConnectOpenEnabled: settings.WeChatConnectOpenEnabled,
WeChatConnectMPEnabled: settings.WeChatConnectMPEnabled,
WeChatConnectMobileEnabled: settings.WeChatConnectMobileEnabled,
WeChatConnectMode: settings.WeChatConnectMode,
WeChatConnectScopes: settings.WeChatConnectScopes,
WeChatConnectRedirectURL: settings.WeChatConnectRedirectURL,
WeChatConnectFrontendRedirectURL: settings.WeChatConnectFrontendRedirectURL,
OIDCConnectEnabled: settings.OIDCConnectEnabled,
OIDCConnectProviderName: settings.OIDCConnectProviderName,
OIDCConnectClientID: settings.OIDCConnectClientID,
OIDCConnectClientSecretConfigured: settings.OIDCConnectClientSecretConfigured,
OIDCConnectIssuerURL: settings.OIDCConnectIssuerURL,
OIDCConnectDiscoveryURL: settings.OIDCConnectDiscoveryURL,
OIDCConnectAuthorizeURL: settings.OIDCConnectAuthorizeURL,
OIDCConnectTokenURL: settings.OIDCConnectTokenURL,
OIDCConnectUserInfoURL: settings.OIDCConnectUserInfoURL,
OIDCConnectJWKSURL: settings.OIDCConnectJWKSURL,
OIDCConnectScopes: settings.OIDCConnectScopes,
OIDCConnectRedirectURL: settings.OIDCConnectRedirectURL,
OIDCConnectFrontendRedirectURL: settings.OIDCConnectFrontendRedirectURL,
OIDCConnectTokenAuthMethod: settings.OIDCConnectTokenAuthMethod,
OIDCConnectUsePKCE: settings.OIDCConnectUsePKCE,
OIDCConnectValidateIDToken: settings.OIDCConnectValidateIDToken,
OIDCConnectAllowedSigningAlgs: settings.OIDCConnectAllowedSigningAlgs,
OIDCConnectClockSkewSeconds: settings.OIDCConnectClockSkewSeconds,
OIDCConnectRequireEmailVerified: settings.OIDCConnectRequireEmailVerified,
OIDCConnectUserInfoEmailPath: settings.OIDCConnectUserInfoEmailPath,
OIDCConnectUserInfoIDPath: settings.OIDCConnectUserInfoIDPath,
OIDCConnectUserInfoUsernamePath: settings.OIDCConnectUserInfoUsernamePath,
GitHubOAuthEnabled: settings.GitHubOAuthEnabled,
GitHubOAuthClientID: settings.GitHubOAuthClientID,
GitHubOAuthClientSecretConfigured: settings.GitHubOAuthClientSecretConfigured,
GitHubOAuthRedirectURL: settings.GitHubOAuthRedirectURL,
GitHubOAuthFrontendRedirectURL: settings.GitHubOAuthFrontendRedirectURL,
GoogleOAuthEnabled: settings.GoogleOAuthEnabled,
GoogleOAuthClientID: settings.GoogleOAuthClientID,
GoogleOAuthClientSecretConfigured: settings.GoogleOAuthClientSecretConfigured,
GoogleOAuthRedirectURL: settings.GoogleOAuthRedirectURL,
GoogleOAuthFrontendRedirectURL: settings.GoogleOAuthFrontendRedirectURL,
SiteName: settings.SiteName,
SiteLogo: settings.SiteLogo,
SiteSubtitle: settings.SiteSubtitle,
APIBaseURL: settings.APIBaseURL,
ContactInfo: settings.ContactInfo,
DocURL: settings.DocURL,
HomeContent: settings.HomeContent,
HideCcsImportButton: settings.HideCcsImportButton,
PurchaseSubscriptionEnabled: settings.PurchaseSubscriptionEnabled,
PurchaseSubscriptionURL: settings.PurchaseSubscriptionURL,
TableDefaultPageSize: settings.TableDefaultPageSize,
TablePageSizeOptions: settings.TablePageSizeOptions,
CustomMenuItems: dto.ParseCustomMenuItems(settings.CustomMenuItems),
CustomEndpoints: dto.ParseCustomEndpoints(settings.CustomEndpoints),
DefaultConcurrency: settings.DefaultConcurrency,
DefaultBalance: settings.DefaultBalance,
RiskControlEnabled: settings.RiskControlEnabled,
CyberSessionBlockEnabled: settings.CyberSessionBlockEnabled,
CyberSessionBlockTTLSeconds: settings.CyberSessionBlockTTLSeconds,
AffiliateRebateRate: settings.AffiliateRebateRate,
AffiliateRebateFreezeHours: settings.AffiliateRebateFreezeHours,
AffiliateRebateDurationDays: settings.AffiliateRebateDurationDays,
AffiliateRebatePerInviteeCap: settings.AffiliateRebatePerInviteeCap,
DefaultUserRPMLimit: settings.DefaultUserRPMLimit,
DefaultSubscriptions: defaultSubscriptions,
EnableModelFallback: settings.EnableModelFallback,
FallbackModelAnthropic: settings.FallbackModelAnthropic,
FallbackModelOpenAI: settings.FallbackModelOpenAI,
FallbackModelGemini: settings.FallbackModelGemini,
FallbackModelAntigravity: settings.FallbackModelAntigravity,
EnableIdentityPatch: settings.EnableIdentityPatch,
IdentityPatchPrompt: settings.IdentityPatchPrompt,
OpsMonitoringEnabled: opsEnabled && settings.OpsMonitoringEnabled,
OpsRealtimeMonitoringEnabled: settings.OpsRealtimeMonitoringEnabled,
OpsQueryModeDefault: settings.OpsQueryModeDefault,
OpsMetricsIntervalSeconds: settings.OpsMetricsIntervalSeconds,
MinClaudeCodeVersion: settings.MinClaudeCodeVersion,
MaxClaudeCodeVersion: settings.MaxClaudeCodeVersion,
AllowUngroupedKeyScheduling: settings.AllowUngroupedKeyScheduling,
BackendModeEnabled: settings.BackendModeEnabled,
EnableFingerprintUnification: settings.EnableFingerprintUnification,
EnableMetadataPassthrough: settings.EnableMetadataPassthrough,
EnableCCHSigning: settings.EnableCCHSigning,
EnableClaudeOAuthSystemPromptInjection: settings.EnableClaudeOAuthSystemPromptInjection,
ClaudeOAuthSystemPrompt: settings.ClaudeOAuthSystemPrompt,
ClaudeOAuthSystemPromptBlocks: settings.ClaudeOAuthSystemPromptBlocks,
EnableAnthropicCacheTTL1hInjection: settings.EnableAnthropicCacheTTL1hInjection,
RewriteMessageCacheControl: settings.RewriteMessageCacheControl,
EnableClientDatelineNormalization: settings.EnableClientDatelineNormalization,
AntigravityUserAgentVersion: settings.AntigravityUserAgentVersion,
OpenAICodexUserAgent: settings.OpenAICodexUserAgent,
MinCodexVersion: settings.MinCodexVersion,
MaxCodexVersion: settings.MaxCodexVersion,
CodexCLIOnlyBlacklist: settings.CodexCLIOnlyBlacklist,
CodexCLIOnlyWhitelist: settings.CodexCLIOnlyWhitelist,
CodexCLIOnlyAllowAppServerClients: settings.CodexCLIOnlyAllowAppServerClients,
CodexCLIOnlyEngineFingerprintSignals: settings.CodexCLIOnlyEngineFingerprintSignals,
WebSearchEmulationEnabled: settings.WebSearchEmulationEnabled,
PaymentVisibleMethodAlipaySource: settings.PaymentVisibleMethodAlipaySource,
PaymentVisibleMethodWxpaySource: settings.PaymentVisibleMethodWxpaySource,
PaymentVisibleMethodAlipayEnabled: settings.PaymentVisibleMethodAlipayEnabled,
PaymentVisibleMethodWxpayEnabled: settings.PaymentVisibleMethodWxpayEnabled,
OpenAIAdvancedSchedulerEnabled: settings.OpenAIAdvancedSchedulerEnabled,
OpenAIAdvancedSchedulerStickyWeightedEnabled: settings.OpenAIAdvancedSchedulerStickyWeightedEnabled,
OpenAIAdvancedSchedulerSubscriptionPriorityEnabled: settings.OpenAIAdvancedSchedulerSubscriptionPriorityEnabled,
OpenAIAdvancedSchedulerLBTopK: settings.OpenAIAdvancedSchedulerLBTopK,
OpenAIAdvancedSchedulerWeightPriority: settings.OpenAIAdvancedSchedulerWeightPriority,
OpenAIAdvancedSchedulerWeightLoad: settings.OpenAIAdvancedSchedulerWeightLoad,
OpenAIAdvancedSchedulerWeightQueue: settings.OpenAIAdvancedSchedulerWeightQueue,
OpenAIAdvancedSchedulerWeightErrorRate: settings.OpenAIAdvancedSchedulerWeightErrorRate,
OpenAIAdvancedSchedulerWeightTTFT: settings.OpenAIAdvancedSchedulerWeightTTFT,
OpenAIAdvancedSchedulerWeightReset: settings.OpenAIAdvancedSchedulerWeightReset,
OpenAIAdvancedSchedulerWeightQuotaHeadroom: settings.OpenAIAdvancedSchedulerWeightQuotaHeadroom,
OpenAIAdvancedSchedulerWeightPreviousResponse: settings.OpenAIAdvancedSchedulerWeightPreviousResponse,
OpenAIAdvancedSchedulerWeightSessionSticky: settings.OpenAIAdvancedSchedulerWeightSessionSticky,
OpenAIAdvancedSchedulerEffectiveLBTopK: settings.OpenAIAdvancedSchedulerEffectiveLBTopK,
OpenAIAdvancedSchedulerEffectiveWeightPriority: settings.OpenAIAdvancedSchedulerEffectiveWeightPriority,
OpenAIAdvancedSchedulerEffectiveWeightLoad: settings.OpenAIAdvancedSchedulerEffectiveWeightLoad,
OpenAIAdvancedSchedulerEffectiveWeightQueue: settings.OpenAIAdvancedSchedulerEffectiveWeightQueue,
OpenAIAdvancedSchedulerEffectiveWeightErrorRate: settings.OpenAIAdvancedSchedulerEffectiveWeightErrorRate,
OpenAIAdvancedSchedulerEffectiveWeightTTFT: settings.OpenAIAdvancedSchedulerEffectiveWeightTTFT,
OpenAIAdvancedSchedulerEffectiveWeightReset: settings.OpenAIAdvancedSchedulerEffectiveWeightReset,
OpenAIAdvancedSchedulerEffectiveWeightQuotaHeadroom: settings.OpenAIAdvancedSchedulerEffectiveWeightQuotaHeadroom,
OpenAIAdvancedSchedulerEffectiveWeightPreviousResponse: settings.OpenAIAdvancedSchedulerEffectiveWeightPreviousResponse,
OpenAIAdvancedSchedulerEffectiveWeightSessionSticky: settings.OpenAIAdvancedSchedulerEffectiveWeightSessionSticky,
BalanceLowNotifyEnabled: settings.BalanceLowNotifyEnabled,
BalanceLowNotifyThreshold: settings.BalanceLowNotifyThreshold,
BalanceLowNotifyRechargeURL: settings.BalanceLowNotifyRechargeURL,
SubscriptionExpiryNotifyEnabled: settings.SubscriptionExpiryNotifyEnabled,
AccountQuotaNotifyEnabled: settings.AccountQuotaNotifyEnabled,
AccountQuotaNotifyEmails: dto.NotifyEmailEntriesFromService(settings.AccountQuotaNotifyEmails),
PaymentEnabled: paymentCfg.Enabled,
PaymentMinAmount: paymentCfg.MinAmount,
PaymentMaxAmount: paymentCfg.MaxAmount,
PaymentDailyLimit: paymentCfg.DailyLimit,
PaymentOrderTimeoutMin: paymentCfg.OrderTimeoutMin,
PaymentMaxPendingOrders: paymentCfg.MaxPendingOrders,
PaymentEnabledTypes: paymentCfg.EnabledTypes,
PaymentBalanceDisabled: paymentCfg.BalanceDisabled,
PaymentBalanceRechargeMultiplier: paymentCfg.BalanceRechargeMultiplier,
PaymentSubscriptionUSDToCNYRate: paymentCfg.SubscriptionUSDToCNYRate,
PaymentRechargeFeeRate: paymentCfg.RechargeFeeRate,
PaymentLoadBalanceStrat: paymentCfg.LoadBalanceStrategy,
PaymentProductNamePrefix: paymentCfg.ProductNamePrefix,
PaymentProductNameSuffix: paymentCfg.ProductNameSuffix,
PaymentHelpImageURL: paymentCfg.HelpImageURL,
PaymentHelpText: paymentCfg.HelpText,
PaymentCancelRateLimitEnabled: paymentCfg.CancelRateLimitEnabled,
PaymentCancelRateLimitMax: paymentCfg.CancelRateLimitMax,
PaymentCancelRateLimitWindow: paymentCfg.CancelRateLimitWindow,
PaymentCancelRateLimitUnit: paymentCfg.CancelRateLimitUnit,
PaymentCancelRateLimitMode: paymentCfg.CancelRateLimitMode,
PaymentAlipayForceQRCode: paymentCfg.AlipayForceQRCode,
ChannelMonitorEnabled: settings.ChannelMonitorEnabled,
ChannelMonitorDefaultIntervalSeconds: settings.ChannelMonitorDefaultIntervalSeconds,
@@ -618,7 +641,19 @@ type UpdateSettingsRequest struct {
PaymentVisibleMethodWxpayEnabled *bool `json:"payment_visible_method_wxpay_enabled"`
// OpenAI account scheduling
OpenAIAdvancedSchedulerEnabled *bool `json:"openai_advanced_scheduler_enabled"`
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"`
// 余额不足提醒
BalanceLowNotifyEnabled *bool `json:"balance_low_notify_enabled"`
@@ -638,6 +673,7 @@ type UpdateSettingsRequest struct {
PaymentEnabledTypes []string `json:"payment_enabled_types"`
PaymentBalanceDisabled *bool `json:"payment_balance_disabled"`
PaymentBalanceRechargeMultiplier *float64 `json:"payment_balance_recharge_multiplier"`
PaymentSubscriptionUSDToCNYRate *float64 `json:"payment_subscription_usd_to_cny_rate"`
PaymentRechargeFeeRate *float64 `json:"payment_recharge_fee_rate"`
PaymentLoadBalanceStrat *string `json:"payment_load_balance_strategy"`
PaymentProductNamePrefix *string `json:"payment_product_name_prefix"`
@@ -1792,6 +1828,28 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
}
return previousSettings.OpenAIAdvancedSchedulerEnabled
}(),
OpenAIAdvancedSchedulerStickyWeightedEnabled: func() bool {
if req.OpenAIAdvancedSchedulerStickyWeightedEnabled != nil {
return *req.OpenAIAdvancedSchedulerStickyWeightedEnabled
}
return previousSettings.OpenAIAdvancedSchedulerStickyWeightedEnabled
}(),
OpenAIAdvancedSchedulerSubscriptionPriorityEnabled: func() bool {
if req.OpenAIAdvancedSchedulerSubscriptionPriorityEnabled != nil {
return *req.OpenAIAdvancedSchedulerSubscriptionPriorityEnabled
}
return previousSettings.OpenAIAdvancedSchedulerSubscriptionPriorityEnabled
}(),
OpenAIAdvancedSchedulerLBTopK: stringSetting(req.OpenAIAdvancedSchedulerLBTopK, previousSettings.OpenAIAdvancedSchedulerLBTopK),
OpenAIAdvancedSchedulerWeightPriority: stringSetting(req.OpenAIAdvancedSchedulerWeightPriority, previousSettings.OpenAIAdvancedSchedulerWeightPriority),
OpenAIAdvancedSchedulerWeightLoad: stringSetting(req.OpenAIAdvancedSchedulerWeightLoad, previousSettings.OpenAIAdvancedSchedulerWeightLoad),
OpenAIAdvancedSchedulerWeightQueue: stringSetting(req.OpenAIAdvancedSchedulerWeightQueue, previousSettings.OpenAIAdvancedSchedulerWeightQueue),
OpenAIAdvancedSchedulerWeightErrorRate: stringSetting(req.OpenAIAdvancedSchedulerWeightErrorRate, previousSettings.OpenAIAdvancedSchedulerWeightErrorRate),
OpenAIAdvancedSchedulerWeightTTFT: stringSetting(req.OpenAIAdvancedSchedulerWeightTTFT, previousSettings.OpenAIAdvancedSchedulerWeightTTFT),
OpenAIAdvancedSchedulerWeightReset: stringSetting(req.OpenAIAdvancedSchedulerWeightReset, previousSettings.OpenAIAdvancedSchedulerWeightReset),
OpenAIAdvancedSchedulerWeightQuotaHeadroom: stringSetting(req.OpenAIAdvancedSchedulerWeightQuotaHeadroom, previousSettings.OpenAIAdvancedSchedulerWeightQuotaHeadroom),
OpenAIAdvancedSchedulerWeightPreviousResponse: stringSetting(req.OpenAIAdvancedSchedulerWeightPreviousResponse, previousSettings.OpenAIAdvancedSchedulerWeightPreviousResponse),
OpenAIAdvancedSchedulerWeightSessionSticky: stringSetting(req.OpenAIAdvancedSchedulerWeightSessionSticky, previousSettings.OpenAIAdvancedSchedulerWeightSessionSticky),
BalanceLowNotifyEnabled: func() bool {
if req.BalanceLowNotifyEnabled != nil {
return *req.BalanceLowNotifyEnabled
@@ -1959,6 +2017,7 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
EnabledTypes: req.PaymentEnabledTypes,
BalanceDisabled: req.PaymentBalanceDisabled,
BalanceRechargeMultiplier: req.PaymentBalanceRechargeMultiplier,
SubscriptionUSDToCNYRate: req.PaymentSubscriptionUSDToCNYRate,
RechargeFeeRate: req.PaymentRechargeFeeRate,
LoadBalanceStrategy: req.PaymentLoadBalanceStrat,
ProductNamePrefix: req.PaymentProductNamePrefix,
@@ -2014,184 +2073,207 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
}
payload := dto.SystemSettings{
RegistrationEnabled: updatedSettings.RegistrationEnabled,
EmailVerifyEnabled: updatedSettings.EmailVerifyEnabled,
RegistrationEmailSuffixWhitelist: updatedSettings.RegistrationEmailSuffixWhitelist,
PromoCodeEnabled: updatedSettings.PromoCodeEnabled,
PasswordResetEnabled: updatedSettings.PasswordResetEnabled,
FrontendURL: updatedSettings.FrontendURL,
InvitationCodeEnabled: updatedSettings.InvitationCodeEnabled,
TotpEnabled: updatedSettings.TotpEnabled,
TotpEncryptionKeyConfigured: h.settingService.IsTotpEncryptionKeyConfigured(),
LoginAgreementEnabled: updatedSettings.LoginAgreementEnabled,
LoginAgreementMode: updatedSettings.LoginAgreementMode,
LoginAgreementUpdatedAt: updatedSettings.LoginAgreementUpdatedAt,
LoginAgreementDocuments: loginAgreementDocumentsToDTO(updatedSettings.LoginAgreementDocuments),
SMTPHost: updatedSettings.SMTPHost,
SMTPPort: updatedSettings.SMTPPort,
SMTPUsername: updatedSettings.SMTPUsername,
SMTPPasswordConfigured: updatedSettings.SMTPPasswordConfigured,
SMTPFrom: updatedSettings.SMTPFrom,
SMTPFromName: updatedSettings.SMTPFromName,
SMTPUseTLS: updatedSettings.SMTPUseTLS,
TurnstileEnabled: updatedSettings.TurnstileEnabled,
TurnstileSiteKey: updatedSettings.TurnstileSiteKey,
TurnstileSecretKeyConfigured: updatedSettings.TurnstileSecretKeyConfigured,
APIKeyACLTrustForwardedIP: updatedSettings.APIKeyACLTrustForwardedIP,
LinuxDoConnectEnabled: updatedSettings.LinuxDoConnectEnabled,
LinuxDoConnectClientID: updatedSettings.LinuxDoConnectClientID,
LinuxDoConnectClientSecretConfigured: updatedSettings.LinuxDoConnectClientSecretConfigured,
LinuxDoConnectRedirectURL: updatedSettings.LinuxDoConnectRedirectURL,
DingTalkConnectEnabled: updatedSettings.DingTalkConnectEnabled,
DingTalkConnectClientID: updatedSettings.DingTalkConnectClientID,
DingTalkConnectClientSecretConfigured: updatedSettings.DingTalkConnectClientSecretConfigured,
DingTalkConnectRedirectURL: updatedSettings.DingTalkConnectRedirectURL,
DingTalkConnectCorpRestrictionPolicy: updatedSettings.DingTalkConnectCorpRestrictionPolicy,
DingTalkConnectInternalCorpID: updatedSettings.DingTalkConnectInternalCorpID,
DingTalkConnectBypassRegistration: updatedSettings.DingTalkConnectBypassRegistration,
DingTalkConnectSyncCorpEmail: updatedSettings.DingTalkConnectSyncCorpEmail,
DingTalkConnectSyncDisplayName: updatedSettings.DingTalkConnectSyncDisplayName,
DingTalkConnectSyncDept: updatedSettings.DingTalkConnectSyncDept,
DingTalkConnectSyncCorpEmailAttrKey: updatedSettings.DingTalkConnectSyncCorpEmailAttrKey,
DingTalkConnectSyncDisplayNameAttrKey: updatedSettings.DingTalkConnectSyncDisplayNameAttrKey,
DingTalkConnectSyncDeptAttrKey: updatedSettings.DingTalkConnectSyncDeptAttrKey,
DingTalkConnectSyncCorpEmailAttrName: updatedSettings.DingTalkConnectSyncCorpEmailAttrName,
DingTalkConnectSyncDisplayNameAttrName: updatedSettings.DingTalkConnectSyncDisplayNameAttrName,
DingTalkConnectSyncDeptAttrName: updatedSettings.DingTalkConnectSyncDeptAttrName,
WeChatConnectEnabled: updatedSettings.WeChatConnectEnabled,
WeChatConnectAppID: updatedSettings.WeChatConnectAppID,
WeChatConnectAppSecretConfigured: updatedSettings.WeChatConnectAppSecretConfigured,
WeChatConnectOpenAppID: updatedSettings.WeChatConnectOpenAppID,
WeChatConnectOpenAppSecretConfigured: updatedSettings.WeChatConnectOpenAppSecretConfigured,
WeChatConnectMPAppID: updatedSettings.WeChatConnectMPAppID,
WeChatConnectMPAppSecretConfigured: updatedSettings.WeChatConnectMPAppSecretConfigured,
WeChatConnectMobileAppID: updatedSettings.WeChatConnectMobileAppID,
WeChatConnectMobileAppSecretConfigured: updatedSettings.WeChatConnectMobileAppSecretConfigured,
WeChatConnectOpenEnabled: updatedSettings.WeChatConnectOpenEnabled,
WeChatConnectMPEnabled: updatedSettings.WeChatConnectMPEnabled,
WeChatConnectMobileEnabled: updatedSettings.WeChatConnectMobileEnabled,
WeChatConnectMode: updatedSettings.WeChatConnectMode,
WeChatConnectScopes: updatedSettings.WeChatConnectScopes,
WeChatConnectRedirectURL: updatedSettings.WeChatConnectRedirectURL,
WeChatConnectFrontendRedirectURL: updatedSettings.WeChatConnectFrontendRedirectURL,
OIDCConnectEnabled: updatedSettings.OIDCConnectEnabled,
OIDCConnectProviderName: updatedSettings.OIDCConnectProviderName,
OIDCConnectClientID: updatedSettings.OIDCConnectClientID,
OIDCConnectClientSecretConfigured: updatedSettings.OIDCConnectClientSecretConfigured,
OIDCConnectIssuerURL: updatedSettings.OIDCConnectIssuerURL,
OIDCConnectDiscoveryURL: updatedSettings.OIDCConnectDiscoveryURL,
OIDCConnectAuthorizeURL: updatedSettings.OIDCConnectAuthorizeURL,
OIDCConnectTokenURL: updatedSettings.OIDCConnectTokenURL,
OIDCConnectUserInfoURL: updatedSettings.OIDCConnectUserInfoURL,
OIDCConnectJWKSURL: updatedSettings.OIDCConnectJWKSURL,
OIDCConnectScopes: updatedSettings.OIDCConnectScopes,
OIDCConnectRedirectURL: updatedSettings.OIDCConnectRedirectURL,
OIDCConnectFrontendRedirectURL: updatedSettings.OIDCConnectFrontendRedirectURL,
OIDCConnectTokenAuthMethod: updatedSettings.OIDCConnectTokenAuthMethod,
OIDCConnectUsePKCE: updatedSettings.OIDCConnectUsePKCE,
OIDCConnectValidateIDToken: updatedSettings.OIDCConnectValidateIDToken,
OIDCConnectAllowedSigningAlgs: updatedSettings.OIDCConnectAllowedSigningAlgs,
OIDCConnectClockSkewSeconds: updatedSettings.OIDCConnectClockSkewSeconds,
OIDCConnectRequireEmailVerified: updatedSettings.OIDCConnectRequireEmailVerified,
OIDCConnectUserInfoEmailPath: updatedSettings.OIDCConnectUserInfoEmailPath,
OIDCConnectUserInfoIDPath: updatedSettings.OIDCConnectUserInfoIDPath,
OIDCConnectUserInfoUsernamePath: updatedSettings.OIDCConnectUserInfoUsernamePath,
GitHubOAuthEnabled: updatedSettings.GitHubOAuthEnabled,
GitHubOAuthClientID: updatedSettings.GitHubOAuthClientID,
GitHubOAuthClientSecretConfigured: updatedSettings.GitHubOAuthClientSecretConfigured,
GitHubOAuthRedirectURL: updatedSettings.GitHubOAuthRedirectURL,
GitHubOAuthFrontendRedirectURL: updatedSettings.GitHubOAuthFrontendRedirectURL,
GoogleOAuthEnabled: updatedSettings.GoogleOAuthEnabled,
GoogleOAuthClientID: updatedSettings.GoogleOAuthClientID,
GoogleOAuthClientSecretConfigured: updatedSettings.GoogleOAuthClientSecretConfigured,
GoogleOAuthRedirectURL: updatedSettings.GoogleOAuthRedirectURL,
GoogleOAuthFrontendRedirectURL: updatedSettings.GoogleOAuthFrontendRedirectURL,
SiteName: updatedSettings.SiteName,
SiteLogo: updatedSettings.SiteLogo,
SiteSubtitle: updatedSettings.SiteSubtitle,
APIBaseURL: updatedSettings.APIBaseURL,
ContactInfo: updatedSettings.ContactInfo,
DocURL: updatedSettings.DocURL,
HomeContent: updatedSettings.HomeContent,
HideCcsImportButton: updatedSettings.HideCcsImportButton,
PurchaseSubscriptionEnabled: updatedSettings.PurchaseSubscriptionEnabled,
PurchaseSubscriptionURL: updatedSettings.PurchaseSubscriptionURL,
TableDefaultPageSize: updatedSettings.TableDefaultPageSize,
TablePageSizeOptions: updatedSettings.TablePageSizeOptions,
CustomMenuItems: dto.ParseCustomMenuItems(updatedSettings.CustomMenuItems),
CustomEndpoints: dto.ParseCustomEndpoints(updatedSettings.CustomEndpoints),
DefaultConcurrency: updatedSettings.DefaultConcurrency,
DefaultBalance: updatedSettings.DefaultBalance,
AffiliateRebateRate: updatedSettings.AffiliateRebateRate,
AffiliateRebateFreezeHours: updatedSettings.AffiliateRebateFreezeHours,
AffiliateRebateDurationDays: updatedSettings.AffiliateRebateDurationDays,
AffiliateRebatePerInviteeCap: updatedSettings.AffiliateRebatePerInviteeCap,
DefaultUserRPMLimit: updatedSettings.DefaultUserRPMLimit,
DefaultSubscriptions: updatedDefaultSubscriptions,
EnableModelFallback: updatedSettings.EnableModelFallback,
FallbackModelAnthropic: updatedSettings.FallbackModelAnthropic,
FallbackModelOpenAI: updatedSettings.FallbackModelOpenAI,
FallbackModelGemini: updatedSettings.FallbackModelGemini,
FallbackModelAntigravity: updatedSettings.FallbackModelAntigravity,
EnableIdentityPatch: updatedSettings.EnableIdentityPatch,
IdentityPatchPrompt: updatedSettings.IdentityPatchPrompt,
OpsMonitoringEnabled: updatedSettings.OpsMonitoringEnabled,
OpsRealtimeMonitoringEnabled: updatedSettings.OpsRealtimeMonitoringEnabled,
OpsQueryModeDefault: updatedSettings.OpsQueryModeDefault,
OpsMetricsIntervalSeconds: updatedSettings.OpsMetricsIntervalSeconds,
MinClaudeCodeVersion: updatedSettings.MinClaudeCodeVersion,
MaxClaudeCodeVersion: updatedSettings.MaxClaudeCodeVersion,
AllowUngroupedKeyScheduling: updatedSettings.AllowUngroupedKeyScheduling,
BackendModeEnabled: updatedSettings.BackendModeEnabled,
EnableFingerprintUnification: updatedSettings.EnableFingerprintUnification,
EnableMetadataPassthrough: updatedSettings.EnableMetadataPassthrough,
EnableCCHSigning: updatedSettings.EnableCCHSigning,
EnableClaudeOAuthSystemPromptInjection: updatedSettings.EnableClaudeOAuthSystemPromptInjection,
ClaudeOAuthSystemPrompt: updatedSettings.ClaudeOAuthSystemPrompt,
ClaudeOAuthSystemPromptBlocks: updatedSettings.ClaudeOAuthSystemPromptBlocks,
EnableAnthropicCacheTTL1hInjection: updatedSettings.EnableAnthropicCacheTTL1hInjection,
RewriteMessageCacheControl: updatedSettings.RewriteMessageCacheControl,
EnableClientDatelineNormalization: updatedSettings.EnableClientDatelineNormalization,
AntigravityUserAgentVersion: updatedSettings.AntigravityUserAgentVersion,
OpenAICodexUserAgent: updatedSettings.OpenAICodexUserAgent,
MinCodexVersion: updatedSettings.MinCodexVersion,
MaxCodexVersion: updatedSettings.MaxCodexVersion,
CodexCLIOnlyBlacklist: updatedSettings.CodexCLIOnlyBlacklist,
CodexCLIOnlyWhitelist: updatedSettings.CodexCLIOnlyWhitelist,
CodexCLIOnlyAllowAppServerClients: updatedSettings.CodexCLIOnlyAllowAppServerClients,
CodexCLIOnlyEngineFingerprintSignals: updatedSettings.CodexCLIOnlyEngineFingerprintSignals,
PaymentVisibleMethodAlipaySource: updatedSettings.PaymentVisibleMethodAlipaySource,
PaymentVisibleMethodWxpaySource: updatedSettings.PaymentVisibleMethodWxpaySource,
PaymentVisibleMethodAlipayEnabled: updatedSettings.PaymentVisibleMethodAlipayEnabled,
PaymentVisibleMethodWxpayEnabled: updatedSettings.PaymentVisibleMethodWxpayEnabled,
OpenAIAdvancedSchedulerEnabled: updatedSettings.OpenAIAdvancedSchedulerEnabled,
BalanceLowNotifyEnabled: updatedSettings.BalanceLowNotifyEnabled,
BalanceLowNotifyThreshold: updatedSettings.BalanceLowNotifyThreshold,
BalanceLowNotifyRechargeURL: updatedSettings.BalanceLowNotifyRechargeURL,
SubscriptionExpiryNotifyEnabled: updatedSettings.SubscriptionExpiryNotifyEnabled,
AccountQuotaNotifyEnabled: updatedSettings.AccountQuotaNotifyEnabled,
AccountQuotaNotifyEmails: dto.NotifyEmailEntriesFromService(updatedSettings.AccountQuotaNotifyEmails),
PaymentEnabled: updatedPaymentCfg.Enabled,
PaymentMinAmount: updatedPaymentCfg.MinAmount,
PaymentMaxAmount: updatedPaymentCfg.MaxAmount,
PaymentDailyLimit: updatedPaymentCfg.DailyLimit,
PaymentOrderTimeoutMin: updatedPaymentCfg.OrderTimeoutMin,
PaymentMaxPendingOrders: updatedPaymentCfg.MaxPendingOrders,
PaymentEnabledTypes: updatedPaymentCfg.EnabledTypes,
PaymentBalanceDisabled: updatedPaymentCfg.BalanceDisabled,
PaymentBalanceRechargeMultiplier: updatedPaymentCfg.BalanceRechargeMultiplier,
PaymentRechargeFeeRate: updatedPaymentCfg.RechargeFeeRate,
PaymentLoadBalanceStrat: updatedPaymentCfg.LoadBalanceStrategy,
PaymentProductNamePrefix: updatedPaymentCfg.ProductNamePrefix,
PaymentProductNameSuffix: updatedPaymentCfg.ProductNameSuffix,
PaymentHelpImageURL: updatedPaymentCfg.HelpImageURL,
PaymentHelpText: updatedPaymentCfg.HelpText,
PaymentCancelRateLimitEnabled: updatedPaymentCfg.CancelRateLimitEnabled,
PaymentCancelRateLimitMax: updatedPaymentCfg.CancelRateLimitMax,
PaymentCancelRateLimitWindow: updatedPaymentCfg.CancelRateLimitWindow,
PaymentCancelRateLimitUnit: updatedPaymentCfg.CancelRateLimitUnit,
PaymentCancelRateLimitMode: updatedPaymentCfg.CancelRateLimitMode,
PaymentAlipayForceQRCode: updatedPaymentCfg.AlipayForceQRCode,
RegistrationEnabled: updatedSettings.RegistrationEnabled,
EmailVerifyEnabled: updatedSettings.EmailVerifyEnabled,
RegistrationEmailSuffixWhitelist: updatedSettings.RegistrationEmailSuffixWhitelist,
PromoCodeEnabled: updatedSettings.PromoCodeEnabled,
PasswordResetEnabled: updatedSettings.PasswordResetEnabled,
FrontendURL: updatedSettings.FrontendURL,
InvitationCodeEnabled: updatedSettings.InvitationCodeEnabled,
TotpEnabled: updatedSettings.TotpEnabled,
TotpEncryptionKeyConfigured: h.settingService.IsTotpEncryptionKeyConfigured(),
LoginAgreementEnabled: updatedSettings.LoginAgreementEnabled,
LoginAgreementMode: updatedSettings.LoginAgreementMode,
LoginAgreementUpdatedAt: updatedSettings.LoginAgreementUpdatedAt,
LoginAgreementDocuments: loginAgreementDocumentsToDTO(updatedSettings.LoginAgreementDocuments),
SMTPHost: updatedSettings.SMTPHost,
SMTPPort: updatedSettings.SMTPPort,
SMTPUsername: updatedSettings.SMTPUsername,
SMTPPasswordConfigured: updatedSettings.SMTPPasswordConfigured,
SMTPFrom: updatedSettings.SMTPFrom,
SMTPFromName: updatedSettings.SMTPFromName,
SMTPUseTLS: updatedSettings.SMTPUseTLS,
TurnstileEnabled: updatedSettings.TurnstileEnabled,
TurnstileSiteKey: updatedSettings.TurnstileSiteKey,
TurnstileSecretKeyConfigured: updatedSettings.TurnstileSecretKeyConfigured,
APIKeyACLTrustForwardedIP: updatedSettings.APIKeyACLTrustForwardedIP,
LinuxDoConnectEnabled: updatedSettings.LinuxDoConnectEnabled,
LinuxDoConnectClientID: updatedSettings.LinuxDoConnectClientID,
LinuxDoConnectClientSecretConfigured: updatedSettings.LinuxDoConnectClientSecretConfigured,
LinuxDoConnectRedirectURL: updatedSettings.LinuxDoConnectRedirectURL,
DingTalkConnectEnabled: updatedSettings.DingTalkConnectEnabled,
DingTalkConnectClientID: updatedSettings.DingTalkConnectClientID,
DingTalkConnectClientSecretConfigured: updatedSettings.DingTalkConnectClientSecretConfigured,
DingTalkConnectRedirectURL: updatedSettings.DingTalkConnectRedirectURL,
DingTalkConnectCorpRestrictionPolicy: updatedSettings.DingTalkConnectCorpRestrictionPolicy,
DingTalkConnectInternalCorpID: updatedSettings.DingTalkConnectInternalCorpID,
DingTalkConnectBypassRegistration: updatedSettings.DingTalkConnectBypassRegistration,
DingTalkConnectSyncCorpEmail: updatedSettings.DingTalkConnectSyncCorpEmail,
DingTalkConnectSyncDisplayName: updatedSettings.DingTalkConnectSyncDisplayName,
DingTalkConnectSyncDept: updatedSettings.DingTalkConnectSyncDept,
DingTalkConnectSyncCorpEmailAttrKey: updatedSettings.DingTalkConnectSyncCorpEmailAttrKey,
DingTalkConnectSyncDisplayNameAttrKey: updatedSettings.DingTalkConnectSyncDisplayNameAttrKey,
DingTalkConnectSyncDeptAttrKey: updatedSettings.DingTalkConnectSyncDeptAttrKey,
DingTalkConnectSyncCorpEmailAttrName: updatedSettings.DingTalkConnectSyncCorpEmailAttrName,
DingTalkConnectSyncDisplayNameAttrName: updatedSettings.DingTalkConnectSyncDisplayNameAttrName,
DingTalkConnectSyncDeptAttrName: updatedSettings.DingTalkConnectSyncDeptAttrName,
WeChatConnectEnabled: updatedSettings.WeChatConnectEnabled,
WeChatConnectAppID: updatedSettings.WeChatConnectAppID,
WeChatConnectAppSecretConfigured: updatedSettings.WeChatConnectAppSecretConfigured,
WeChatConnectOpenAppID: updatedSettings.WeChatConnectOpenAppID,
WeChatConnectOpenAppSecretConfigured: updatedSettings.WeChatConnectOpenAppSecretConfigured,
WeChatConnectMPAppID: updatedSettings.WeChatConnectMPAppID,
WeChatConnectMPAppSecretConfigured: updatedSettings.WeChatConnectMPAppSecretConfigured,
WeChatConnectMobileAppID: updatedSettings.WeChatConnectMobileAppID,
WeChatConnectMobileAppSecretConfigured: updatedSettings.WeChatConnectMobileAppSecretConfigured,
WeChatConnectOpenEnabled: updatedSettings.WeChatConnectOpenEnabled,
WeChatConnectMPEnabled: updatedSettings.WeChatConnectMPEnabled,
WeChatConnectMobileEnabled: updatedSettings.WeChatConnectMobileEnabled,
WeChatConnectMode: updatedSettings.WeChatConnectMode,
WeChatConnectScopes: updatedSettings.WeChatConnectScopes,
WeChatConnectRedirectURL: updatedSettings.WeChatConnectRedirectURL,
WeChatConnectFrontendRedirectURL: updatedSettings.WeChatConnectFrontendRedirectURL,
OIDCConnectEnabled: updatedSettings.OIDCConnectEnabled,
OIDCConnectProviderName: updatedSettings.OIDCConnectProviderName,
OIDCConnectClientID: updatedSettings.OIDCConnectClientID,
OIDCConnectClientSecretConfigured: updatedSettings.OIDCConnectClientSecretConfigured,
OIDCConnectIssuerURL: updatedSettings.OIDCConnectIssuerURL,
OIDCConnectDiscoveryURL: updatedSettings.OIDCConnectDiscoveryURL,
OIDCConnectAuthorizeURL: updatedSettings.OIDCConnectAuthorizeURL,
OIDCConnectTokenURL: updatedSettings.OIDCConnectTokenURL,
OIDCConnectUserInfoURL: updatedSettings.OIDCConnectUserInfoURL,
OIDCConnectJWKSURL: updatedSettings.OIDCConnectJWKSURL,
OIDCConnectScopes: updatedSettings.OIDCConnectScopes,
OIDCConnectRedirectURL: updatedSettings.OIDCConnectRedirectURL,
OIDCConnectFrontendRedirectURL: updatedSettings.OIDCConnectFrontendRedirectURL,
OIDCConnectTokenAuthMethod: updatedSettings.OIDCConnectTokenAuthMethod,
OIDCConnectUsePKCE: updatedSettings.OIDCConnectUsePKCE,
OIDCConnectValidateIDToken: updatedSettings.OIDCConnectValidateIDToken,
OIDCConnectAllowedSigningAlgs: updatedSettings.OIDCConnectAllowedSigningAlgs,
OIDCConnectClockSkewSeconds: updatedSettings.OIDCConnectClockSkewSeconds,
OIDCConnectRequireEmailVerified: updatedSettings.OIDCConnectRequireEmailVerified,
OIDCConnectUserInfoEmailPath: updatedSettings.OIDCConnectUserInfoEmailPath,
OIDCConnectUserInfoIDPath: updatedSettings.OIDCConnectUserInfoIDPath,
OIDCConnectUserInfoUsernamePath: updatedSettings.OIDCConnectUserInfoUsernamePath,
GitHubOAuthEnabled: updatedSettings.GitHubOAuthEnabled,
GitHubOAuthClientID: updatedSettings.GitHubOAuthClientID,
GitHubOAuthClientSecretConfigured: updatedSettings.GitHubOAuthClientSecretConfigured,
GitHubOAuthRedirectURL: updatedSettings.GitHubOAuthRedirectURL,
GitHubOAuthFrontendRedirectURL: updatedSettings.GitHubOAuthFrontendRedirectURL,
GoogleOAuthEnabled: updatedSettings.GoogleOAuthEnabled,
GoogleOAuthClientID: updatedSettings.GoogleOAuthClientID,
GoogleOAuthClientSecretConfigured: updatedSettings.GoogleOAuthClientSecretConfigured,
GoogleOAuthRedirectURL: updatedSettings.GoogleOAuthRedirectURL,
GoogleOAuthFrontendRedirectURL: updatedSettings.GoogleOAuthFrontendRedirectURL,
SiteName: updatedSettings.SiteName,
SiteLogo: updatedSettings.SiteLogo,
SiteSubtitle: updatedSettings.SiteSubtitle,
APIBaseURL: updatedSettings.APIBaseURL,
ContactInfo: updatedSettings.ContactInfo,
DocURL: updatedSettings.DocURL,
HomeContent: updatedSettings.HomeContent,
HideCcsImportButton: updatedSettings.HideCcsImportButton,
PurchaseSubscriptionEnabled: updatedSettings.PurchaseSubscriptionEnabled,
PurchaseSubscriptionURL: updatedSettings.PurchaseSubscriptionURL,
TableDefaultPageSize: updatedSettings.TableDefaultPageSize,
TablePageSizeOptions: updatedSettings.TablePageSizeOptions,
CustomMenuItems: dto.ParseCustomMenuItems(updatedSettings.CustomMenuItems),
CustomEndpoints: dto.ParseCustomEndpoints(updatedSettings.CustomEndpoints),
DefaultConcurrency: updatedSettings.DefaultConcurrency,
DefaultBalance: updatedSettings.DefaultBalance,
AffiliateRebateRate: updatedSettings.AffiliateRebateRate,
AffiliateRebateFreezeHours: updatedSettings.AffiliateRebateFreezeHours,
AffiliateRebateDurationDays: updatedSettings.AffiliateRebateDurationDays,
AffiliateRebatePerInviteeCap: updatedSettings.AffiliateRebatePerInviteeCap,
DefaultUserRPMLimit: updatedSettings.DefaultUserRPMLimit,
DefaultSubscriptions: updatedDefaultSubscriptions,
EnableModelFallback: updatedSettings.EnableModelFallback,
FallbackModelAnthropic: updatedSettings.FallbackModelAnthropic,
FallbackModelOpenAI: updatedSettings.FallbackModelOpenAI,
FallbackModelGemini: updatedSettings.FallbackModelGemini,
FallbackModelAntigravity: updatedSettings.FallbackModelAntigravity,
EnableIdentityPatch: updatedSettings.EnableIdentityPatch,
IdentityPatchPrompt: updatedSettings.IdentityPatchPrompt,
OpsMonitoringEnabled: updatedSettings.OpsMonitoringEnabled,
OpsRealtimeMonitoringEnabled: updatedSettings.OpsRealtimeMonitoringEnabled,
OpsQueryModeDefault: updatedSettings.OpsQueryModeDefault,
OpsMetricsIntervalSeconds: updatedSettings.OpsMetricsIntervalSeconds,
MinClaudeCodeVersion: updatedSettings.MinClaudeCodeVersion,
MaxClaudeCodeVersion: updatedSettings.MaxClaudeCodeVersion,
AllowUngroupedKeyScheduling: updatedSettings.AllowUngroupedKeyScheduling,
BackendModeEnabled: updatedSettings.BackendModeEnabled,
EnableFingerprintUnification: updatedSettings.EnableFingerprintUnification,
EnableMetadataPassthrough: updatedSettings.EnableMetadataPassthrough,
EnableCCHSigning: updatedSettings.EnableCCHSigning,
EnableClaudeOAuthSystemPromptInjection: updatedSettings.EnableClaudeOAuthSystemPromptInjection,
ClaudeOAuthSystemPrompt: updatedSettings.ClaudeOAuthSystemPrompt,
ClaudeOAuthSystemPromptBlocks: updatedSettings.ClaudeOAuthSystemPromptBlocks,
EnableAnthropicCacheTTL1hInjection: updatedSettings.EnableAnthropicCacheTTL1hInjection,
RewriteMessageCacheControl: updatedSettings.RewriteMessageCacheControl,
EnableClientDatelineNormalization: updatedSettings.EnableClientDatelineNormalization,
AntigravityUserAgentVersion: updatedSettings.AntigravityUserAgentVersion,
OpenAICodexUserAgent: updatedSettings.OpenAICodexUserAgent,
MinCodexVersion: updatedSettings.MinCodexVersion,
MaxCodexVersion: updatedSettings.MaxCodexVersion,
CodexCLIOnlyBlacklist: updatedSettings.CodexCLIOnlyBlacklist,
CodexCLIOnlyWhitelist: updatedSettings.CodexCLIOnlyWhitelist,
CodexCLIOnlyAllowAppServerClients: updatedSettings.CodexCLIOnlyAllowAppServerClients,
CodexCLIOnlyEngineFingerprintSignals: updatedSettings.CodexCLIOnlyEngineFingerprintSignals,
PaymentVisibleMethodAlipaySource: updatedSettings.PaymentVisibleMethodAlipaySource,
PaymentVisibleMethodWxpaySource: updatedSettings.PaymentVisibleMethodWxpaySource,
PaymentVisibleMethodAlipayEnabled: updatedSettings.PaymentVisibleMethodAlipayEnabled,
PaymentVisibleMethodWxpayEnabled: updatedSettings.PaymentVisibleMethodWxpayEnabled,
OpenAIAdvancedSchedulerEnabled: updatedSettings.OpenAIAdvancedSchedulerEnabled,
OpenAIAdvancedSchedulerStickyWeightedEnabled: updatedSettings.OpenAIAdvancedSchedulerStickyWeightedEnabled,
OpenAIAdvancedSchedulerSubscriptionPriorityEnabled: updatedSettings.OpenAIAdvancedSchedulerSubscriptionPriorityEnabled,
OpenAIAdvancedSchedulerLBTopK: updatedSettings.OpenAIAdvancedSchedulerLBTopK,
OpenAIAdvancedSchedulerWeightPriority: updatedSettings.OpenAIAdvancedSchedulerWeightPriority,
OpenAIAdvancedSchedulerWeightLoad: updatedSettings.OpenAIAdvancedSchedulerWeightLoad,
OpenAIAdvancedSchedulerWeightQueue: updatedSettings.OpenAIAdvancedSchedulerWeightQueue,
OpenAIAdvancedSchedulerWeightErrorRate: updatedSettings.OpenAIAdvancedSchedulerWeightErrorRate,
OpenAIAdvancedSchedulerWeightTTFT: updatedSettings.OpenAIAdvancedSchedulerWeightTTFT,
OpenAIAdvancedSchedulerWeightReset: updatedSettings.OpenAIAdvancedSchedulerWeightReset,
OpenAIAdvancedSchedulerWeightQuotaHeadroom: updatedSettings.OpenAIAdvancedSchedulerWeightQuotaHeadroom,
OpenAIAdvancedSchedulerWeightPreviousResponse: updatedSettings.OpenAIAdvancedSchedulerWeightPreviousResponse,
OpenAIAdvancedSchedulerWeightSessionSticky: updatedSettings.OpenAIAdvancedSchedulerWeightSessionSticky,
OpenAIAdvancedSchedulerEffectiveLBTopK: updatedSettings.OpenAIAdvancedSchedulerEffectiveLBTopK,
OpenAIAdvancedSchedulerEffectiveWeightPriority: updatedSettings.OpenAIAdvancedSchedulerEffectiveWeightPriority,
OpenAIAdvancedSchedulerEffectiveWeightLoad: updatedSettings.OpenAIAdvancedSchedulerEffectiveWeightLoad,
OpenAIAdvancedSchedulerEffectiveWeightQueue: updatedSettings.OpenAIAdvancedSchedulerEffectiveWeightQueue,
OpenAIAdvancedSchedulerEffectiveWeightErrorRate: updatedSettings.OpenAIAdvancedSchedulerEffectiveWeightErrorRate,
OpenAIAdvancedSchedulerEffectiveWeightTTFT: updatedSettings.OpenAIAdvancedSchedulerEffectiveWeightTTFT,
OpenAIAdvancedSchedulerEffectiveWeightReset: updatedSettings.OpenAIAdvancedSchedulerEffectiveWeightReset,
OpenAIAdvancedSchedulerEffectiveWeightQuotaHeadroom: updatedSettings.OpenAIAdvancedSchedulerEffectiveWeightQuotaHeadroom,
OpenAIAdvancedSchedulerEffectiveWeightPreviousResponse: updatedSettings.OpenAIAdvancedSchedulerEffectiveWeightPreviousResponse,
OpenAIAdvancedSchedulerEffectiveWeightSessionSticky: updatedSettings.OpenAIAdvancedSchedulerEffectiveWeightSessionSticky,
BalanceLowNotifyEnabled: updatedSettings.BalanceLowNotifyEnabled,
BalanceLowNotifyThreshold: updatedSettings.BalanceLowNotifyThreshold,
BalanceLowNotifyRechargeURL: updatedSettings.BalanceLowNotifyRechargeURL,
SubscriptionExpiryNotifyEnabled: updatedSettings.SubscriptionExpiryNotifyEnabled,
AccountQuotaNotifyEnabled: updatedSettings.AccountQuotaNotifyEnabled,
AccountQuotaNotifyEmails: dto.NotifyEmailEntriesFromService(updatedSettings.AccountQuotaNotifyEmails),
PaymentEnabled: updatedPaymentCfg.Enabled,
PaymentMinAmount: updatedPaymentCfg.MinAmount,
PaymentMaxAmount: updatedPaymentCfg.MaxAmount,
PaymentDailyLimit: updatedPaymentCfg.DailyLimit,
PaymentOrderTimeoutMin: updatedPaymentCfg.OrderTimeoutMin,
PaymentMaxPendingOrders: updatedPaymentCfg.MaxPendingOrders,
PaymentEnabledTypes: updatedPaymentCfg.EnabledTypes,
PaymentBalanceDisabled: updatedPaymentCfg.BalanceDisabled,
PaymentBalanceRechargeMultiplier: updatedPaymentCfg.BalanceRechargeMultiplier,
PaymentSubscriptionUSDToCNYRate: updatedPaymentCfg.SubscriptionUSDToCNYRate,
PaymentRechargeFeeRate: updatedPaymentCfg.RechargeFeeRate,
PaymentLoadBalanceStrat: updatedPaymentCfg.LoadBalanceStrategy,
PaymentProductNamePrefix: updatedPaymentCfg.ProductNamePrefix,
PaymentProductNameSuffix: updatedPaymentCfg.ProductNameSuffix,
PaymentHelpImageURL: updatedPaymentCfg.HelpImageURL,
PaymentHelpText: updatedPaymentCfg.HelpText,
PaymentCancelRateLimitEnabled: updatedPaymentCfg.CancelRateLimitEnabled,
PaymentCancelRateLimitMax: updatedPaymentCfg.CancelRateLimitMax,
PaymentCancelRateLimitWindow: updatedPaymentCfg.CancelRateLimitWindow,
PaymentCancelRateLimitUnit: updatedPaymentCfg.CancelRateLimitUnit,
PaymentCancelRateLimitMode: updatedPaymentCfg.CancelRateLimitMode,
PaymentAlipayForceQRCode: updatedPaymentCfg.AlipayForceQRCode,
ChannelMonitorEnabled: updatedSettings.ChannelMonitorEnabled,
ChannelMonitorDefaultIntervalSeconds: updatedSettings.ChannelMonitorDefaultIntervalSeconds,
@@ -2238,7 +2320,8 @@ func hasPaymentFields(req UpdateSettingsRequest) bool {
req.PaymentMaxAmount != nil || req.PaymentDailyLimit != nil ||
req.PaymentOrderTimeoutMin != nil || req.PaymentMaxPendingOrders != nil ||
req.PaymentEnabledTypes != nil || req.PaymentBalanceDisabled != nil ||
req.PaymentBalanceRechargeMultiplier != nil || req.PaymentRechargeFeeRate != nil ||
req.PaymentBalanceRechargeMultiplier != nil || req.PaymentSubscriptionUSDToCNYRate != nil ||
req.PaymentRechargeFeeRate != nil ||
req.PaymentLoadBalanceStrat != nil || req.PaymentProductNamePrefix != nil ||
req.PaymentProductNameSuffix != nil || req.PaymentHelpImageURL != nil ||
req.PaymentHelpText != nil || req.PaymentCancelRateLimitEnabled != nil ||
@@ -2677,6 +2760,42 @@ func diffSettings(before *service.SystemSettings, after *service.SystemSettings,
if before.OpenAIAdvancedSchedulerEnabled != after.OpenAIAdvancedSchedulerEnabled {
changed = append(changed, "openai_advanced_scheduler_enabled")
}
if before.OpenAIAdvancedSchedulerStickyWeightedEnabled != after.OpenAIAdvancedSchedulerStickyWeightedEnabled {
changed = append(changed, "openai_advanced_scheduler_sticky_weighted_enabled")
}
if before.OpenAIAdvancedSchedulerSubscriptionPriorityEnabled != after.OpenAIAdvancedSchedulerSubscriptionPriorityEnabled {
changed = append(changed, "openai_advanced_scheduler_subscription_priority_enabled")
}
if before.OpenAIAdvancedSchedulerLBTopK != after.OpenAIAdvancedSchedulerLBTopK {
changed = append(changed, "openai_advanced_scheduler_lb_top_k")
}
if before.OpenAIAdvancedSchedulerWeightPriority != after.OpenAIAdvancedSchedulerWeightPriority {
changed = append(changed, "openai_advanced_scheduler_weight_priority")
}
if before.OpenAIAdvancedSchedulerWeightLoad != after.OpenAIAdvancedSchedulerWeightLoad {
changed = append(changed, "openai_advanced_scheduler_weight_load")
}
if before.OpenAIAdvancedSchedulerWeightQueue != after.OpenAIAdvancedSchedulerWeightQueue {
changed = append(changed, "openai_advanced_scheduler_weight_queue")
}
if before.OpenAIAdvancedSchedulerWeightErrorRate != after.OpenAIAdvancedSchedulerWeightErrorRate {
changed = append(changed, "openai_advanced_scheduler_weight_error_rate")
}
if before.OpenAIAdvancedSchedulerWeightTTFT != after.OpenAIAdvancedSchedulerWeightTTFT {
changed = append(changed, "openai_advanced_scheduler_weight_ttft")
}
if before.OpenAIAdvancedSchedulerWeightReset != after.OpenAIAdvancedSchedulerWeightReset {
changed = append(changed, "openai_advanced_scheduler_weight_reset")
}
if before.OpenAIAdvancedSchedulerWeightQuotaHeadroom != after.OpenAIAdvancedSchedulerWeightQuotaHeadroom {
changed = append(changed, "openai_advanced_scheduler_weight_quota_headroom")
}
if before.OpenAIAdvancedSchedulerWeightPreviousResponse != after.OpenAIAdvancedSchedulerWeightPreviousResponse {
changed = append(changed, "openai_advanced_scheduler_weight_previous_response")
}
if before.OpenAIAdvancedSchedulerWeightSessionSticky != after.OpenAIAdvancedSchedulerWeightSessionSticky {
changed = append(changed, "openai_advanced_scheduler_weight_session_sticky")
}
// 余额、订阅到期与账号限额通知
if before.BalanceLowNotifyEnabled != after.BalanceLowNotifyEnabled {
changed = append(changed, "balance_low_notify_enabled")
@@ -3829,3 +3948,10 @@ func equalPlatformQuotaSettings(before, after map[string]*service.DefaultPlatfor
}
return true
}
func stringSetting(value *string, fallback string) string {
if value == nil {
return fallback
}
return *value
}
@@ -217,12 +217,13 @@ func TestSettingHandler_UpdateSettings_PersistsPaymentVisibleMethodsAndAdvancedS
handler := NewSettingHandler(svc, nil, nil, nil, nil, nil, nil)
body := map[string]any{
"promo_code_enabled": true,
"payment_visible_method_alipay_source": "easypay",
"payment_visible_method_wxpay_source": "wxpay",
"payment_visible_method_alipay_enabled": true,
"payment_visible_method_wxpay_enabled": false,
"openai_advanced_scheduler_enabled": true,
"promo_code_enabled": true,
"payment_visible_method_alipay_source": "easypay",
"payment_visible_method_wxpay_source": "wxpay",
"payment_visible_method_alipay_enabled": true,
"payment_visible_method_wxpay_enabled": false,
"openai_advanced_scheduler_enabled": true,
"openai_advanced_scheduler_subscription_priority_enabled": true,
}
rawBody, err := json.Marshal(body)
require.NoError(t, err)
@@ -240,6 +241,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, "true", repo.values[service.SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled])
var resp response.Response
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp))
@@ -250,6 +252,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, true, data["openai_advanced_scheduler_subscription_priority_enabled"])
}
func TestSettingHandler_UpdateSettings_PreservesLegacyBlankPaymentVisibleMethodSource(t *testing.T) {
@@ -11,18 +11,20 @@ import (
func TestAPIKeyFromService_MapsLastUsedAt(t *testing.T) {
lastUsed := time.Now().UTC().Truncate(time.Second)
src := &service.APIKey{
ID: 1,
UserID: 2,
Key: "sk-map-last-used",
Name: "Mapper",
Status: service.StatusActive,
LastUsedAt: &lastUsed,
ID: 1,
UserID: 2,
Key: "sk-map-last-used",
Name: "Mapper",
Status: service.StatusActive,
LastUsedAt: &lastUsed,
CurrentConcurrency: 3,
}
out := APIKeyFromService(src)
require.NotNil(t, out)
require.NotNil(t, out.LastUsedAt)
require.WithinDuration(t, lastUsed, *out.LastUsedAt, time.Second)
require.Equal(t, 3, out.CurrentConcurrency)
}
func TestAPIKeyFromService_MapsNilLastUsedAt(t *testing.T) {
+26 -25
View File
@@ -79,31 +79,32 @@ func APIKeyFromService(k *service.APIKey) *APIKey {
return nil
}
out := &APIKey{
ID: k.ID,
UserID: k.UserID,
Key: k.Key,
Name: k.Name,
GroupID: k.GroupID,
Status: k.Status,
IPWhitelist: k.IPWhitelist,
IPBlacklist: k.IPBlacklist,
LastUsedAt: k.LastUsedAt,
Quota: k.Quota,
QuotaUsed: k.QuotaUsed,
ExpiresAt: k.ExpiresAt,
CreatedAt: k.CreatedAt,
UpdatedAt: k.UpdatedAt,
RateLimit5h: k.RateLimit5h,
RateLimit1d: k.RateLimit1d,
RateLimit7d: k.RateLimit7d,
Usage5h: k.EffectiveUsage5h(),
Usage1d: k.EffectiveUsage1d(),
Usage7d: k.EffectiveUsage7d(),
Window5hStart: k.Window5hStart,
Window1dStart: k.Window1dStart,
Window7dStart: k.Window7dStart,
User: UserFromServiceShallow(k.User),
Group: GroupFromServiceShallow(k.Group),
ID: k.ID,
UserID: k.UserID,
Key: k.Key,
Name: k.Name,
GroupID: k.GroupID,
Status: k.Status,
IPWhitelist: k.IPWhitelist,
IPBlacklist: k.IPBlacklist,
LastUsedAt: k.LastUsedAt,
Quota: k.Quota,
QuotaUsed: k.QuotaUsed,
ExpiresAt: k.ExpiresAt,
CreatedAt: k.CreatedAt,
UpdatedAt: k.UpdatedAt,
CurrentConcurrency: k.CurrentConcurrency,
RateLimit5h: k.RateLimit5h,
RateLimit1d: k.RateLimit1d,
RateLimit7d: k.RateLimit7d,
Usage5h: k.EffectiveUsage5h(),
Usage1d: k.EffectiveUsage1d(),
Usage7d: k.EffectiveUsage7d(),
Window5hStart: k.Window5hStart,
Window1dStart: k.Window1dStart,
Window7dStart: k.Window7dStart,
User: UserFromServiceShallow(k.User),
Group: GroupFromServiceShallow(k.Group),
}
if k.Window5hStart != nil && !service.IsWindowExpired(k.Window5hStart, service.RateLimitWindow5h) {
t := k.Window5hStart.Add(service.RateLimitWindow5h)
+24 -1
View File
@@ -208,7 +208,29 @@ type SystemSettings struct {
PaymentVisibleMethodWxpayEnabled bool `json:"payment_visible_method_wxpay_enabled"`
// OpenAI account scheduling
OpenAIAdvancedSchedulerEnabled bool `json:"openai_advanced_scheduler_enabled"`
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"`
// Payment configuration
PaymentEnabled bool `json:"payment_enabled"`
@@ -220,6 +242,7 @@ type SystemSettings struct {
PaymentEnabledTypes []string `json:"payment_enabled_types"`
PaymentBalanceDisabled bool `json:"payment_balance_disabled"`
PaymentBalanceRechargeMultiplier float64 `json:"payment_balance_recharge_multiplier"`
PaymentSubscriptionUSDToCNYRate float64 `json:"payment_subscription_usd_to_cny_rate"`
PaymentRechargeFeeRate float64 `json:"payment_recharge_fee_rate"`
PaymentLoadBalanceStrat string `json:"payment_load_balance_strategy"`
PaymentProductNamePrefix string `json:"payment_product_name_prefix"`
+2
View File
@@ -63,6 +63,8 @@ type APIKey struct {
ExpiresAt *time.Time `json:"expires_at"` // Expiration time (nil = never expires)
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
// CurrentConcurrency is the real-time active request count for this API key.
CurrentConcurrency int `json:"current_concurrency"`
// Rate limit fields
RateLimit5h float64 `json:"rate_limit_5h"`
+37 -1
View File
@@ -1006,7 +1006,8 @@ func (h *GatewayHandler) Models(c *gin.Context) {
// Get available models from account configurations for the selected group platform.
availableModels := h.gatewayService.GetAvailableModels(c.Request.Context(), groupID, platform)
if apiKey != nil && apiKey.Group != nil && apiKey.Group.CustomModelsListEnabled() {
availableModels = filterModelsByCustomList(availableModels, defaultModelIDsForPlatform(platform), apiKey.Group.ModelsListConfig.Models)
fallbackModels := defaultModelIDsForPlatform(platform)
availableModels = filterModelsByCustomList(customModelsListSource(platform, availableModels, fallbackModels), fallbackModels, apiKey.Group.ModelsListConfig.Models)
writeCustomModelsList(c, platform, availableModels)
return
}
@@ -1090,6 +1091,13 @@ func writeOpenAIModelsList(c *gin.Context, modelIDs []string) {
})
}
func customModelsListSource(platform string, availableModels, fallbackModels []string) []string {
if platform == service.PlatformAnthropic && len(availableModels) > 0 {
return mergeModelIDs(availableModels, fallbackModels)
}
return availableModels
}
func filterModelsByCustomList(availableModels, fallbackModels, selectedModels []string) []string {
if len(selectedModels) == 0 {
return availableModels
@@ -1158,6 +1166,15 @@ func defaultModelIDsForPlatform(platform string) []string {
ids = append(ids, model.ID)
}
return ids
case service.PlatformAnthropic:
ids := make([]string, 0, len(claude.DefaultModels)+len(antigravity.DefaultModels()))
for _, model := range claude.DefaultModels {
ids = append(ids, model.ID)
}
for _, model := range antigravity.DefaultModels() {
ids = append(ids, model.ID)
}
return mergeModelIDs(ids, nil)
case service.PlatformGrok:
return xai.DefaultModelIDs()
default:
@@ -1169,6 +1186,25 @@ func defaultModelIDsForPlatform(platform string) []string {
}
}
func mergeModelIDs(primary, secondary []string) []string {
seen := make(map[string]struct{}, len(primary)+len(secondary))
merged := make([]string, 0, len(primary)+len(secondary))
for _, models := range [][]string{primary, secondary} {
for _, model := range models {
model = strings.TrimSpace(model)
if model == "" {
continue
}
if _, ok := seen[model]; ok {
continue
}
seen[model] = struct{}{}
merged = append(merged, model)
}
}
return merged
}
// AntigravityModels 返回 Antigravity 支持的全部模型
// GET /antigravity/models
func (h *GatewayHandler) AntigravityModels(c *gin.Context) {
+41 -2
View File
@@ -10,6 +10,7 @@ import (
"sync"
"time"
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
@@ -211,6 +212,14 @@ func (h *ConcurrencyHelper) TryAcquireUserSlot(ctx context.Context, userID int64
return result.ReleaseFunc, true, nil
}
func (h *ConcurrencyHelper) TryAcquireUserSlotForAPIKey(ctx context.Context, userID int64, maxConcurrency int, apiKeyID int64) (func(), bool, error) {
releaseFunc, acquired, err := h.TryAcquireUserSlot(ctx, userID, maxConcurrency)
if err != nil || !acquired {
return releaseFunc, acquired, err
}
return h.withAPIKeySlot(ctx, apiKeyID, releaseFunc), true, nil
}
// TryAcquireAccountSlot 尝试立即获取账号并发槽位。
// 返回值: (releaseFunc, acquired, error)
func (h *ConcurrencyHelper) TryAcquireAccountSlot(ctx context.Context, accountID int64, maxConcurrency int) (func(), bool, error) {
@@ -241,7 +250,7 @@ func (h *ConcurrencyHelper) acquireUserSlotWithWaitTimeout(c *gin.Context, userI
}
if acquired {
return releaseFunc, nil
return h.withAPIKeySlotFromGin(c, releaseFunc), nil
}
queueLimit := service.CalculateMaxWait(maxConcurrency) - maxConcurrency
@@ -258,7 +267,37 @@ func (h *ConcurrencyHelper) acquireUserSlotWithWaitTimeout(c *gin.Context, userI
defer h.DecrementWaitCount(ctx, userID)
// Need to wait - handle streaming ping if needed
return h.waitForSlotWithPingTimeout(c, "user", userID, maxConcurrency, timeout, isStream, streamStarted, false)
releaseFunc, err = h.waitForSlotWithPingTimeout(c, "user", userID, maxConcurrency, timeout, isStream, streamStarted, false)
if err != nil {
return nil, err
}
return h.withAPIKeySlotFromGin(c, releaseFunc), nil
}
func (h *ConcurrencyHelper) withAPIKeySlotFromGin(c *gin.Context, releaseFunc func()) func() {
if c == nil {
return releaseFunc
}
apiKey, ok := middleware2.GetAPIKeyFromContext(c)
if !ok || apiKey == nil {
return releaseFunc
}
return h.withAPIKeySlot(c.Request.Context(), apiKey.ID, releaseFunc)
}
func (h *ConcurrencyHelper) withAPIKeySlot(ctx context.Context, apiKeyID int64, releaseFunc func()) func() {
if h == nil || h.concurrencyService == nil || apiKeyID <= 0 {
return releaseFunc
}
apiKeyReleaseFunc := h.concurrencyService.TrackAPIKeySlot(ctx, apiKeyID)
return func() {
if releaseFunc != nil {
releaseFunc()
}
if apiKeyReleaseFunc != nil {
apiKeyReleaseFunc()
}
}
}
// AcquireAccountSlotWithWait acquires an account concurrency slot, waiting if necessary.
@@ -9,6 +9,7 @@ import (
"testing"
"time"
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
@@ -29,6 +30,9 @@ type helperConcurrencyCacheStub struct {
waitDecrementCalls int
waitMaxWait int
waitIncrementHook func()
apiKeyTrackCalls int
apiKeyReleaseCalls int
apiKeyTrackIDs []int64
}
func (s *helperConcurrencyCacheStub) AcquireAccountSlot(ctx context.Context, accountID int64, maxConcurrency int, requestID string) (bool, error) {
@@ -97,6 +101,29 @@ func (s *helperConcurrencyCacheStub) GetUserConcurrency(ctx context.Context, use
return 0, nil
}
func (s *helperConcurrencyCacheStub) TrackAPIKeySlot(ctx context.Context, apiKeyID int64, requestID string) error {
s.mu.Lock()
defer s.mu.Unlock()
s.apiKeyTrackCalls++
s.apiKeyTrackIDs = append(s.apiKeyTrackIDs, apiKeyID)
return nil
}
func (s *helperConcurrencyCacheStub) ReleaseAPIKeySlot(ctx context.Context, apiKeyID int64, requestID string) error {
s.mu.Lock()
defer s.mu.Unlock()
s.apiKeyReleaseCalls++
return nil
}
func (s *helperConcurrencyCacheStub) GetAPIKeyConcurrencyBatch(ctx context.Context, apiKeyIDs []int64) (map[int64]int, error) {
out := make(map[int64]int, len(apiKeyIDs))
for _, apiKeyID := range apiKeyIDs {
out[apiKeyID] = 0
}
return out, nil
}
func (s *helperConcurrencyCacheStub) IncrementWaitCount(ctx context.Context, userID int64, maxWait int) (bool, error) {
s.mu.Lock()
s.waitIncrementCalls++
@@ -270,6 +297,48 @@ func TestAcquireUserSlotWithWait_ImmediateAcquireSkipsWaitQueue(t *testing.T) {
require.Equal(t, 1, cache.userReleaseCalls)
}
func TestAcquireUserSlotWithWait_TracksAPIKeySlot(t *testing.T) {
cache := &helperConcurrencyCacheStub{
userSeq: []bool{true},
}
concurrency := service.NewConcurrencyService(cache)
helper := NewConcurrencyHelper(concurrency, SSEPingFormatNone, 5*time.Millisecond)
c, _ := newHelperTestContext(http.MethodPost, "/v1/messages")
c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{ID: 77})
streamStarted := false
release, err := helper.acquireUserSlotWithWaitTimeout(c, 202, 3, time.Second, false, &streamStarted)
require.NoError(t, err)
require.NotNil(t, release)
require.Equal(t, 1, cache.apiKeyTrackCalls)
require.Equal(t, []int64{77}, cache.apiKeyTrackIDs)
release()
require.Equal(t, 1, cache.userReleaseCalls)
require.Equal(t, 1, cache.apiKeyReleaseCalls)
}
func TestTryAcquireUserSlotForAPIKey_TracksAPIKeySlot(t *testing.T) {
cache := &helperConcurrencyCacheStub{
userSeq: []bool{true},
}
concurrency := service.NewConcurrencyService(cache)
helper := NewConcurrencyHelper(concurrency, SSEPingFormatNone, 5*time.Millisecond)
release, acquired, err := helper.TryAcquireUserSlotForAPIKey(context.Background(), 202, 3, 77)
require.NoError(t, err)
require.True(t, acquired)
require.NotNil(t, release)
require.Equal(t, 1, cache.apiKeyTrackCalls)
require.Equal(t, []int64{77}, cache.apiKeyTrackIDs)
release()
require.Equal(t, 1, cache.userReleaseCalls)
require.Equal(t, 1, cache.apiKeyReleaseCalls)
}
func TestAcquireUserSlotWithWait_WaitSuccessDecrementsBeforeReturn(t *testing.T) {
cache := &helperConcurrencyCacheStub{
userSeq: []bool{false, true},
@@ -269,6 +269,149 @@ func TestGatewayModels_CustomModelsListKeepsConcreteModelAllowedByWildcardMappin
require.Equal(t, []string{"claude-sonnet-4-6"}, modelIDsForTest(got.Data))
}
func TestGatewayModels_AnthropicCustomModelsListIncludesOAuthClaudeAndMappedDeepSeek(t *testing.T) {
gin.SetMode(gin.TestMode)
groupID := int64(28)
h := newGatewayModelsHandlerForTest(
&gatewayModelsAccountRepoStub{
byGroup: map[int64][]service.Account{
groupID: {
{
ID: 1,
Platform: service.PlatformAnthropic,
Type: service.AccountTypeOAuth,
},
{
ID: 2,
Platform: service.PlatformAnthropic,
Type: service.AccountTypeAPIKey,
Credentials: map[string]any{
"model_mapping": map[string]any{
"deepseek-v4-pro": "deepseek-v4-pro",
},
},
},
},
},
},
)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodGet, "/v1/models", nil)
c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
Group: &service.Group{
ID: groupID,
Platform: service.PlatformAnthropic,
ModelsListConfig: service.GroupModelsListConfig{
Enabled: true,
Models: []string{"claude-fable-5", "claude-opus-4-8", "deepseek-v4-pro"},
},
},
})
h.Models(c)
require.Equal(t, http.StatusOK, rec.Code)
var got gatewayModelsResponseForTest
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
require.Equal(t, []string{"claude-fable-5", "claude-opus-4-8", "deepseek-v4-pro"}, modelIDsForTest(got.Data))
}
func TestGatewayModels_AnthropicCustomModelsListDisabledKeepsMappedModelList(t *testing.T) {
gin.SetMode(gin.TestMode)
groupID := int64(29)
h := newGatewayModelsHandlerForTest(
&gatewayModelsAccountRepoStub{
byGroup: map[int64][]service.Account{
groupID: {
{
ID: 1,
Platform: service.PlatformAnthropic,
Type: service.AccountTypeOAuth,
},
{
ID: 2,
Platform: service.PlatformAnthropic,
Type: service.AccountTypeAPIKey,
Credentials: map[string]any{
"model_mapping": map[string]any{
"deepseek-v4-pro": "deepseek-v4-pro",
},
},
},
},
},
},
)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodGet, "/v1/models", nil)
c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
Group: &service.Group{
ID: groupID,
Platform: service.PlatformAnthropic,
ModelsListConfig: service.GroupModelsListConfig{
Enabled: false,
Models: []string{"claude-fable-5", "deepseek-v4-pro"},
},
},
})
h.Models(c)
require.Equal(t, http.StatusOK, rec.Code)
var got gatewayModelsResponseForTest
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
require.Equal(t, []string{"deepseek-v4-pro"}, modelIDsForTest(got.Data))
}
func TestGatewayModels_AnthropicCustomModelsListIncludesOAuthClaudeWithoutMappings(t *testing.T) {
gin.SetMode(gin.TestMode)
groupID := int64(30)
h := newGatewayModelsHandlerForTest(
&gatewayModelsAccountRepoStub{
byGroup: map[int64][]service.Account{
groupID: {
{
ID: 1,
Platform: service.PlatformAnthropic,
Type: service.AccountTypeOAuth,
},
},
},
},
)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodGet, "/v1/models", nil)
c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
Group: &service.Group{
ID: groupID,
Platform: service.PlatformAnthropic,
ModelsListConfig: service.GroupModelsListConfig{
Enabled: true,
Models: []string{"claude-opus-4-6-thinking", "claude-sonnet-4-5"},
},
},
})
h.Models(c)
require.Equal(t, http.StatusOK, rec.Code)
var got gatewayModelsResponseForTest
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
require.Equal(t, []string{"claude-opus-4-6-thinking", "claude-sonnet-4-5"}, modelIDsForTest(got.Data))
}
func TestGatewayModels_CustomModelsListCanReturnEmptyWhenSelectionsUnavailable(t *testing.T) {
gin.SetMode(gin.TestMode)
+1
View File
@@ -174,6 +174,7 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
service.OpenAIUpstreamTransportHTTPSSE,
"",
false,
false,
service.PlatformGrok,
)
if err != nil {
@@ -145,6 +145,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
service.OpenAIUpstreamTransportAny,
service.OpenAIEndpointCapabilityChatCompletions,
false,
false,
requestPlatform,
)
if err != nil {
@@ -117,6 +117,7 @@ func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) {
service.OpenAIUpstreamTransportHTTPSSE,
service.OpenAIEndpointCapabilityEmbeddings,
false,
false,
)
if err != nil {
reqLog.Warn("openai_embeddings.account_select_failed",
@@ -110,6 +110,7 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) {
service.OpenAIUpstreamTransportAny,
service.OpenAIEndpointCapabilityChatCompletions,
false,
false,
openAICompatibleRequestPlatform(apiKey),
)
service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds())
@@ -350,6 +350,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
service.OpenAIUpstreamTransportAny,
service.OpenAIEndpointCapabilityChatCompletions,
requireCompact,
false,
requestPlatform,
)
if err != nil {
@@ -783,6 +784,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
service.OpenAIUpstreamTransportAny,
service.OpenAIEndpointCapabilityChatCompletions,
false,
false,
requestPlatform,
)
if err != nil {
@@ -1266,6 +1268,8 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
closeOpenAIClientWS(wsConn, coderws.StatusPolicyViolation, "previous_response_id must be a response.id (resp_*), not a message id")
return
}
firstMessageToolCoverage := service.AnalyzeToolCallOutputContextCoverageBytes(firstMessage)
previousResponseCanMove := !firstMessageToolCoverage.HasFunctionCallOutput || firstMessageToolCoverage.ContextCoversAllCallIDs
reqLog = reqLog.With(
zap.Bool("ws_ingress", true),
zap.String("model", reqModel),
@@ -1318,7 +1322,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
// 必须尽早注册,确保任何 early return 都能释放已获取的并发槽位。
defer releaseTurnSlots()
userReleaseFunc, userAcquired, err := h.concurrencyHelper.TryAcquireUserSlot(ctx, subject.UserID, subject.Concurrency)
userReleaseFunc, userAcquired, err := h.concurrencyHelper.TryAcquireUserSlotForAPIKey(ctx, subject.UserID, subject.Concurrency, apiKey.ID)
if err != nil {
reqLog.Warn("openai.websocket_user_slot_acquire_failed", zap.Error(err))
closeOpenAIClientWS(wsConn, coderws.StatusInternalError, "failed to acquire user concurrency slot")
@@ -1333,7 +1337,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
if currentUserRelease != nil {
return true
}
userReleaseFunc, userAcquired, err := h.concurrencyHelper.TryAcquireUserSlot(ctx, subject.UserID, subject.Concurrency)
userReleaseFunc, userAcquired, err := h.concurrencyHelper.TryAcquireUserSlotForAPIKey(ctx, subject.UserID, subject.Concurrency, apiKey.ID)
if err != nil {
reqLog.Warn("openai.websocket_user_slot_reacquire_failed", zap.Error(err))
closeOpenAIClientWS(wsConn, coderws.StatusInternalError, "failed to acquire user concurrency slot")
@@ -1381,6 +1385,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
requiredTransport,
service.OpenAIEndpointCapabilityChatCompletions,
false,
previousResponseCanMove,
requestPlatform,
)
if err != nil {
@@ -1484,7 +1489,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
// 防御式清理:避免异常路径下旧槽位覆盖导致泄漏。
releaseTurnSlots()
// 非首轮 turn 需要重新抢占并发槽位,避免长连接空闲占槽。
userReleaseFunc, userAcquired, err := h.concurrencyHelper.TryAcquireUserSlot(ctx, subject.UserID, subject.Concurrency)
userReleaseFunc, userAcquired, err := h.concurrencyHelper.TryAcquireUserSlotForAPIKey(ctx, subject.UserID, subject.Concurrency, apiKey.ID)
if err != nil {
return service.NewOpenAIWSClientCloseError(coderws.StatusInternalError, "failed to acquire user concurrency slot", err)
}
@@ -1581,8 +1586,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
// 说明该会话链不属于本次调度到的账号,原样转发会触发上游会话链鉴权失败(“鉴权失败,请检查 API Key”)。
// 故剥离首包里的 previous_response_id,改用首包内 input 重建上下文;带 function_call_output 的
// 工具续链无法重建,保持原样。仅作用于首轮首包,后续 turn 的续链由 WS 转发层既有逻辑处理。
if previousResponseID != "" && !scheduleDecision.StickyPreviousHit &&
!service.ValidateFunctionCallOutputContextBytes(wsFirstMessage).HasFunctionCallOutput {
if previousResponseID != "" && !scheduleDecision.StickyPreviousHit && previousResponseCanMove {
wsFirstMessage = service.RemovePreviousResponseIDFromBody(wsFirstMessage)
reqLog.Debug("openai.websocket_previous_response_id_stripped_cross_group",
zap.Int64("account_id", account.ID),
@@ -150,6 +150,7 @@ func (h *PaymentHandler) GetCheckoutInfo(c *gin.Context) {
Plans: planList,
BalanceDisabled: cfg.BalanceDisabled,
BalanceRechargeMultiplier: cfg.BalanceRechargeMultiplier,
SubscriptionUSDToCNYRate: cfg.SubscriptionUSDToCNYRate,
RechargeFeeRate: cfg.RechargeFeeRate,
HelpText: cfg.HelpText,
HelpImageURL: cfg.HelpImageURL,
@@ -165,6 +166,7 @@ type checkoutInfoResponse struct {
Plans []checkoutPlan `json:"plans"`
BalanceDisabled bool `json:"balance_disabled"`
BalanceRechargeMultiplier float64 `json:"balance_recharge_multiplier"`
SubscriptionUSDToCNYRate float64 `json:"subscription_usd_to_cny_rate"`
RechargeFeeRate float64 `json:"recharge_fee_rate"`
HelpText string `json:"help_text"`
HelpImageURL string `json:"help_image_url"`
@@ -322,6 +322,9 @@ func (h *UsageHandler) ListErrors(c *gin.Context) {
filter.ErrorTypesAny = types
}
// 排序对齐用量明细:列白名单与方向归一在 repo 层,非法值回退 created_at DESC。
filter.SetSort(c.Query("sort_by"), c.Query("sort_order"))
result, err := h.opsService.ListUserErrorRequests(c.Request.Context(), subject.UserID, filter)
if err != nil {
response.ErrorFrom(c, err)
+54 -5
View File
@@ -39,6 +39,12 @@ type EasyPay struct {
httpClient *http.Client
}
type easyPayCustomMethod struct {
Type string `json:"type"`
UpstreamType string `json:"upstreamType"`
DisplayName string `json:"displayName"`
}
// NewEasyPay creates a new EasyPay provider.
// config keys: pid, pkey, apiBase, notifyUrl, returnUrl, cid, cidAlipay, cidWxpay
func NewEasyPay(instanceID string, config map[string]string) (*EasyPay, error) {
@@ -95,7 +101,13 @@ func (e *EasyPay) apiBase() string {
func (e *EasyPay) Name() string { return "EasyPay" }
func (e *EasyPay) ProviderKey() string { return payment.TypeEasyPay }
func (e *EasyPay) SupportedTypes() []payment.PaymentType {
return []payment.PaymentType{payment.TypeAlipay, payment.TypeWxpay}
types := []payment.PaymentType{payment.TypeAlipay, payment.TypeWxpay}
for _, method := range e.customMethods() {
if method.Type != "" {
types = append(types, method.Type)
}
}
return types
}
func (e *EasyPay) MerchantIdentityMetadata() map[string]string {
@@ -124,13 +136,14 @@ func (e *EasyPay) CreatePayment(ctx context.Context, req payment.CreatePaymentRe
// TradeNo is empty; it arrives via the notify callback after payment.
func (e *EasyPay) createRedirectPayment(req payment.CreatePaymentRequest) (*payment.CreatePaymentResponse, error) {
notifyURL, returnURL := e.resolveURLs(req)
paymentType := e.upstreamPaymentType(req.PaymentType)
params := map[string]string{
"pid": e.config["pid"], "type": req.PaymentType,
"pid": e.config["pid"], "type": paymentType,
"out_trade_no": req.OrderID, "notify_url": notifyURL,
"return_url": returnURL, "name": req.Subject,
"money": req.Amount,
}
if cid := e.resolveCID(req.PaymentType); cid != "" {
if cid := e.resolveCID(paymentType); cid != "" {
params["cid"] = cid
}
if req.IsMobile {
@@ -150,13 +163,14 @@ func (e *EasyPay) createRedirectPayment(req payment.CreatePaymentRequest) (*paym
// createAPIPayment calls mapi.php to get payurl/qrcode (existing behavior).
func (e *EasyPay) createAPIPayment(ctx context.Context, req payment.CreatePaymentRequest) (*payment.CreatePaymentResponse, error) {
notifyURL, returnURL := e.resolveURLs(req)
paymentType := e.upstreamPaymentType(req.PaymentType)
params := map[string]string{
"pid": e.config["pid"], "type": req.PaymentType,
"pid": e.config["pid"], "type": paymentType,
"out_trade_no": req.OrderID, "notify_url": notifyURL,
"return_url": returnURL, "name": req.Subject,
"money": req.Amount, "clientip": req.ClientIP,
}
if cid := e.resolveCID(req.PaymentType); cid != "" {
if cid := e.resolveCID(paymentType); cid != "" {
params["cid"] = cid
}
if req.IsMobile {
@@ -204,6 +218,41 @@ func (e *EasyPay) resolveURLs(req payment.CreatePaymentRequest) (string, string)
return notifyURL, returnURL
}
func (e *EasyPay) customMethods() []easyPayCustomMethod {
if e == nil {
return nil
}
raw := strings.TrimSpace(e.config["customMethods"])
if raw == "" {
return nil
}
var methods []easyPayCustomMethod
if err := json.Unmarshal([]byte(raw), &methods); err != nil {
return nil
}
result := make([]easyPayCustomMethod, 0, len(methods))
for _, method := range methods {
method.Type = strings.TrimSpace(method.Type)
method.UpstreamType = strings.TrimSpace(method.UpstreamType)
method.DisplayName = strings.TrimSpace(method.DisplayName)
if method.Type == "" || method.UpstreamType == "" {
continue
}
result = append(result, method)
}
return result
}
func (e *EasyPay) upstreamPaymentType(paymentType string) string {
paymentType = strings.TrimSpace(paymentType)
for _, method := range e.customMethods() {
if paymentType == method.Type {
return method.UpstreamType
}
}
return paymentType
}
func (e *EasyPay) QueryOrder(ctx context.Context, tradeNo string) (*payment.QueryOrderResponse, error) {
params := map[string]string{
"act": "order", "pid": e.config["pid"],
@@ -179,6 +179,102 @@ func TestEasyPayRefundResponseErrors(t *testing.T) {
}
}
func TestEasyPayCustomMethodsUseConfiguredUpstreamType(t *testing.T) {
t.Parallel()
provider, err := NewEasyPay("test-instance", map[string]string{
"pid": "pid-1",
"pkey": "pkey-1",
"apiBase": "https://pay.example.com",
"notifyUrl": "https://example.com/notify",
"returnUrl": "https://example.com/return",
"paymentMode": paymentModePopup,
"customMethods": `[{"type":"ldc","upstreamType":"epay","displayName":"LDC"},{"type":"usdt_trc20","upstreamType":"usdt","displayName":"USDT-TRC20"}]`,
})
if err != nil {
t.Fatalf("NewEasyPay: %v", err)
}
resp, err := provider.CreatePayment(context.Background(), payment.CreatePaymentRequest{
OrderID: "sub2-custom-1",
Amount: "1.00",
PaymentType: "usdt_trc20",
Subject: "Custom EasyPay",
})
if err != nil {
t.Fatalf("CreatePayment: %v", err)
}
payURL, err := url.Parse(resp.PayURL)
if err != nil {
t.Fatalf("parse pay url: %v", err)
}
if got := payURL.Query().Get("type"); got != "usdt" {
t.Fatalf("pay url type = %q, want usdt (%s)", got, resp.PayURL)
}
}
func TestEasyPayCustomMethodsResolveCIDFromConfiguredUpstreamType(t *testing.T) {
t.Parallel()
provider, err := NewEasyPay("test-instance", map[string]string{
"pid": "pid-1",
"pkey": "pkey-1",
"apiBase": "https://pay.example.com",
"notifyUrl": "https://example.com/notify",
"returnUrl": "https://example.com/return",
"paymentMode": paymentModePopup,
"cidAlipay": "cid-alipay",
"cidWxpay": "cid-wxpay",
"customMethods": `[{"type":"ldc","upstreamType":"alipay","displayName":"LDC"}]`,
})
if err != nil {
t.Fatalf("NewEasyPay: %v", err)
}
resp, err := provider.CreatePayment(context.Background(), payment.CreatePaymentRequest{
OrderID: "sub2-custom-cid",
Amount: "1.00",
PaymentType: "ldc",
Subject: "Custom EasyPay CID",
})
if err != nil {
t.Fatalf("CreatePayment: %v", err)
}
payURL, err := url.Parse(resp.PayURL)
if err != nil {
t.Fatalf("parse pay url: %v", err)
}
if got := payURL.Query().Get("type"); got != "alipay" {
t.Fatalf("pay url type = %q, want alipay (%s)", got, resp.PayURL)
}
if got := payURL.Query().Get("cid"); got != "cid-alipay" {
t.Fatalf("pay url cid = %q, want cid-alipay (%s)", got, resp.PayURL)
}
}
func TestEasyPaySupportedTypesIncludeCustomMethods(t *testing.T) {
t.Parallel()
provider, err := NewEasyPay("test-instance", map[string]string{
"pid": "pid-1",
"pkey": "pkey-1",
"apiBase": "https://pay.example.com",
"notifyUrl": "https://example.com/notify",
"returnUrl": "https://example.com/return",
"customMethods": `[{"type":"ldc","upstreamType":"epay","displayName":"LDC"},{"type":"usdt_trc20","upstreamType":"usdt","displayName":"USDT-TRC20"}]`,
})
if err != nil {
t.Fatalf("NewEasyPay: %v", err)
}
got := strings.Join(provider.SupportedTypes(), ",")
for _, want := range []string{"alipay", "wxpay", "ldc", "usdt_trc20"} {
if !strings.Contains(got, want) {
t.Fatalf("SupportedTypes() = %q, want it to include %q", got, want)
}
}
}
func newTestEasyPay(t *testing.T, apiBase string) *EasyPay {
t.Helper()
+3
View File
@@ -18,6 +18,9 @@ type Model struct {
// DefaultModels OpenAI models list
var DefaultModels = []Model{
{ID: "gpt-5.6-sol", Object: "model", Created: 1780876800, OwnedBy: "openai", Type: "model", DisplayName: "GPT-5.6 Sol"},
{ID: "gpt-5.6-terra", Object: "model", Created: 1780876800, OwnedBy: "openai", Type: "model", DisplayName: "GPT-5.6 Terra"},
{ID: "gpt-5.6-luna", Object: "model", Created: 1780876800, OwnedBy: "openai", Type: "model", DisplayName: "GPT-5.6 Luna"},
{ID: "gpt-5.5", Object: "model", Created: 1776873600, OwnedBy: "openai", Type: "model", DisplayName: "GPT-5.5"},
{ID: "gpt-5.4", Object: "model", Created: 1738368000, OwnedBy: "openai", Type: "model", DisplayName: "GPT-5.4"},
{ID: "gpt-5.4-mini", Object: "model", Created: 1738368000, OwnedBy: "openai", Type: "model", DisplayName: "GPT-5.4 Mini"},
+98 -1
View File
@@ -482,7 +482,7 @@ func (r *accountRepository) List(ctx context.Context, params pagination.Paginati
return r.ListWithFilters(ctx, params, "", "", "", "", 0, "")
}
func (r *accountRepository) ListWithFilters(ctx context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]service.Account, *pagination.PaginationResult, error) {
func (r *accountRepository) accountListFilteredQuery(platform, accountType, status, search string, groupID int64, privacyMode string) *dbent.AccountQuery {
q := r.client.Account.Query()
if platform != "" {
@@ -575,6 +575,11 @@ func (r *accountRepository) ListWithFilters(ctx context.Context, params paginati
}))
}
return q
}
func (r *accountRepository) ListWithFilters(ctx context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]service.Account, *pagination.PaginationResult, error) {
q := r.accountListFilteredQuery(platform, accountType, status, search, groupID, privacyMode)
// Clone before Count so interceptor-appended predicates (SoftDeleteMixin's
// deleted_at IS NULL) don't accumulate on the shared builder and pollute the
// subsequent list query. Same pattern used in group_repo/promo_code_repo/user_repo
@@ -603,6 +608,14 @@ func (r *accountRepository) ListWithFilters(ctx context.Context, params paginati
return outAccounts, paginationResultFromTotal(int64(total), params), nil
}
func (r *accountRepository) ListAllWithFilters(ctx context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]service.Account, error) {
accounts, err := r.accountListFilteredQuery(platform, accountType, status, search, groupID, privacyMode).All(ctx)
if err != nil {
return nil, err
}
return r.accountsToService(ctx, accounts)
}
func (r *accountRepository) ListOpsAccountsForStats(ctx context.Context, platformFilter string, groupIDFilter *int64) ([]service.Account, error) {
if r == nil || r.client == nil {
return []service.Account{}, nil
@@ -1061,6 +1074,90 @@ func (r *accountRepository) ListSchedulableByGroupID(ctx context.Context, groupI
})
}
func (r *accountRepository) ListSchedulableCapacityByGroupIDs(ctx context.Context, groupIDs []int64) ([]service.GroupAccountCapacityRow, error) {
groupIDs = uniquePositiveInt64s(groupIDs)
if len(groupIDs) == 0 {
return []service.GroupAccountCapacityRow{}, nil
}
if r.sql == nil {
rows := make([]service.GroupAccountCapacityRow, 0)
for _, groupID := range groupIDs {
accounts, err := r.ListSchedulableByGroupID(ctx, groupID)
if err != nil {
return nil, err
}
for i := range accounts {
acc := &accounts[i]
rows = append(rows, service.GroupAccountCapacityRow{
GroupID: groupID,
AccountID: acc.ID,
Concurrency: acc.Concurrency,
Extra: copyJSONMap(acc.Extra),
SessionWindowStart: acc.SessionWindowStart,
SessionWindowEnd: acc.SessionWindowEnd,
SessionWindowStatus: acc.SessionWindowStatus,
})
}
}
return rows, nil
}
rows, err := r.sql.QueryContext(ctx, `
SELECT
ag.group_id,
a.id AS account_id,
a.concurrency,
COALESCE(a.extra, '{}'::jsonb)::text AS extra,
a.session_window_start,
a.session_window_end,
COALESCE(a.session_window_status, '') AS session_window_status
FROM account_groups ag
JOIN accounts a ON a.id = ag.account_id
WHERE ag.group_id = ANY($1)
AND a.deleted_at IS NULL
AND a.status = $2
AND a.schedulable = TRUE
AND (a.temp_unschedulable_until IS NULL OR a.temp_unschedulable_until <= $3)
AND (a.expires_at IS NULL OR a.expires_at > $3 OR a.auto_pause_on_expired = FALSE)
AND (a.overload_until IS NULL OR a.overload_until <= $3)
AND (a.rate_limit_reset_at IS NULL OR a.rate_limit_reset_at <= $3)
ORDER BY ag.group_id ASC, ag.priority ASC, a.priority ASC, a.id ASC
`, pq.Array(groupIDs), service.StatusActive, time.Now())
if err != nil {
return nil, err
}
defer func() { _ = rows.Close() }()
out := make([]service.GroupAccountCapacityRow, 0)
for rows.Next() {
var row service.GroupAccountCapacityRow
var extraRaw string
if err := rows.Scan(
&row.GroupID,
&row.AccountID,
&row.Concurrency,
&extraRaw,
&row.SessionWindowStart,
&row.SessionWindowEnd,
&row.SessionWindowStatus,
); err != nil {
return nil, err
}
if extraRaw != "" && extraRaw != "null" {
var extra map[string]any
if err := json.Unmarshal([]byte(extraRaw), &extra); err != nil {
return nil, err
}
row.Extra = extra
}
out = append(out, row)
}
if err := rows.Err(); err != nil {
return nil, err
}
return out, nil
}
func (r *accountRepository) ListSchedulableByPlatform(ctx context.Context, platform string) ([]service.Account, error) {
now := time.Now()
accounts, err := r.client.Account.Query().
@@ -27,6 +27,8 @@ const (
accountSlotKeyPrefix = "concurrency:account:"
// 格式: concurrency:user:{userID}
userSlotKeyPrefix = "concurrency:user:"
// 格式: concurrency:api_key:{apiKeyID}
apiKeySlotKeyPrefix = "concurrency:api_key:"
// 等待队列计数器格式: concurrency:wait:{userID}
waitQueueKeyPrefix = "concurrency:wait:"
// 账号级等待队列计数器格式: wait:account:{accountID}
@@ -108,6 +110,28 @@ var (
return redis.call('ZCARD', key)
`)
// trackSlotScript 记录 stats-only 槽位,不做并发上限判断。
// KEYS[1] = 有序集合键
// ARGV[1] = TTL(秒)
// ARGV[2] = requestID
trackSlotScript = redis.NewScript(`
-- Redis 3.2-4.x compat: opt into effects replication so redis.call('TIME')
-- replicates correctly. No-op on Redis 5.0+ (effects replication is default).
redis.replicate_commands()
local key = KEYS[1]
local ttl = tonumber(ARGV[1])
local requestID = ARGV[2]
local timeResult = redis.call('TIME')
local now = tonumber(timeResult[1])
local expireBefore = now - ttl
redis.call('ZREMRANGEBYSCORE', key, '-inf', expireBefore)
redis.call('ZADD', key, now, requestID)
redis.call('EXPIRE', key, ttl)
return 1
`)
// incrementWaitScript - refreshes TTL on each increment to keep queue depth accurate
// KEYS[1] = wait queue key
// ARGV[1] = maxWait
@@ -237,6 +261,10 @@ func userSlotKey(userID int64) string {
return fmt.Sprintf("%s%d", userSlotKeyPrefix, userID)
}
func apiKeySlotKey(apiKeyID int64) string {
return fmt.Sprintf("%s%d", apiKeySlotKeyPrefix, apiKeyID)
}
func waitQueueKey(userID int64) string {
return fmt.Sprintf("%s%d", waitQueueKeyPrefix, userID)
}
@@ -546,6 +574,54 @@ func (c *concurrencyCache) GetUserConcurrency(ctx context.Context, userID int64)
return result, nil
}
func (c *concurrencyCache) TrackAPIKeySlot(ctx context.Context, apiKeyID int64, requestID string) error {
key := apiKeySlotKey(apiKeyID)
_, err := trackSlotScript.Run(ctx, c.rdb, []string{key}, c.slotTTLSeconds, requestID).Result()
return err
}
func (c *concurrencyCache) ReleaseAPIKeySlot(ctx context.Context, apiKeyID int64, requestID string) error {
key := apiKeySlotKey(apiKeyID)
return c.rdb.ZRem(ctx, key, requestID).Err()
}
func (c *concurrencyCache) GetAPIKeyConcurrencyBatch(ctx context.Context, apiKeyIDs []int64) (map[int64]int, error) {
if len(apiKeyIDs) == 0 {
return map[int64]int{}, nil
}
now, err := c.rdb.Time(ctx).Result()
if err != nil {
return nil, fmt.Errorf("redis TIME: %w", err)
}
cutoffTime := now.Unix() - int64(c.slotTTLSeconds)
pipe := c.rdb.Pipeline()
type apiKeyCmd struct {
apiKeyID int64
zcardCmd *redis.IntCmd
}
cmds := make([]apiKeyCmd, 0, len(apiKeyIDs))
for _, apiKeyID := range apiKeyIDs {
slotKey := apiKeySlotKeyPrefix + strconv.FormatInt(apiKeyID, 10)
pipe.ZRemRangeByScore(ctx, slotKey, "-inf", strconv.FormatInt(cutoffTime, 10))
cmds = append(cmds, apiKeyCmd{
apiKeyID: apiKeyID,
zcardCmd: pipe.ZCard(ctx, slotKey),
})
}
if _, err := pipe.Exec(ctx); err != nil && !errors.Is(err, redis.Nil) {
return nil, fmt.Errorf("pipeline exec: %w", err)
}
result := make(map[int64]int, len(apiKeyIDs))
for _, cmd := range cmds {
result[cmd.apiKeyID] = int(cmd.zcardCmd.Val())
}
return result, nil
}
// Wait queue operations
func (c *concurrencyCache) IncrementWaitCount(ctx context.Context, userID int64, maxWait int) (bool, error) {
@@ -814,6 +890,8 @@ func (c *concurrencyCache) CleanupExpiredAccountSlotKeys(ctx context.Context) er
// CleanupStaleProcessSlots 启动时清理非当前进程前缀的槽位。
// 清理范围来自活跃索引,避免在 Redis 上 SCAN 全部 concurrency:* 键。
// API Key 槽位(concurrency:api_key:*)是 stats-only 数据:每次 Track/读取都会按分数
// 裁剪过期成员,key 自带 TTL,可在一个 slot TTL 内自愈,因此不参与启动清理。
func (c *concurrencyCache) CleanupStaleProcessSlots(ctx context.Context, activeRequestPrefix string) error {
if activeRequestPrefix == "" {
return nil
@@ -3,6 +3,7 @@
package repository
import (
"context"
"errors"
"fmt"
"strconv"
@@ -37,6 +38,18 @@ func (s *ConcurrencyCacheSuite) SetupTest() {
s.cache = s.rawCache
}
type apiKeyConcurrencyCacheForTest interface {
TrackAPIKeySlot(ctx context.Context, apiKeyID int64, requestID string) error
ReleaseAPIKeySlot(ctx context.Context, apiKeyID int64, requestID string) error
GetAPIKeyConcurrencyBatch(ctx context.Context, apiKeyIDs []int64) (map[int64]int, error)
}
func (s *ConcurrencyCacheSuite) apiKeyConcurrencyCache() apiKeyConcurrencyCacheForTest {
cache, ok := s.cache.(apiKeyConcurrencyCacheForTest)
require.True(s.T(), ok)
return cache
}
func (s *ConcurrencyCacheSuite) TestAccountSlot_AcquireAndRelease() {
accountID := int64(10)
reqID1, reqID2, reqID3 := "req1", "req2", "req3"
@@ -218,6 +231,34 @@ func (s *ConcurrencyCacheSuite) TestUserSlot_TTL() {
s.AssertTTLWithin(ttl, 1*time.Second, testSlotTTL)
}
func (s *ConcurrencyCacheSuite) TestAPIKeySlot_TrackReleaseAndBatchCount() {
cache := s.apiKeyConcurrencyCache()
apiKeyID := int64(300)
emptyAPIKeyID := int64(301)
slotKey := fmt.Sprintf("%s%d", apiKeySlotKeyPrefix, apiKeyID)
require.NoError(s.T(), cache.TrackAPIKeySlot(s.ctx, apiKeyID, "req1"))
require.NoError(s.T(), cache.TrackAPIKeySlot(s.ctx, apiKeyID, "req2"))
counts, err := cache.GetAPIKeyConcurrencyBatch(s.ctx, []int64{apiKeyID, emptyAPIKeyID})
require.NoError(s.T(), err)
require.Equal(s.T(), map[int64]int{apiKeyID: 2, emptyAPIKeyID: 0}, counts)
ttl, err := s.rdb.TTL(s.ctx, slotKey).Result()
require.NoError(s.T(), err, "TTL")
s.AssertTTLWithin(ttl, 1*time.Second, testSlotTTL)
require.NoError(s.T(), cache.ReleaseAPIKeySlot(s.ctx, apiKeyID, "req1"))
counts, err = cache.GetAPIKeyConcurrencyBatch(s.ctx, []int64{apiKeyID})
require.NoError(s.T(), err)
require.Equal(s.T(), 1, counts[apiKeyID])
require.NoError(s.T(), cache.ReleaseAPIKeySlot(s.ctx, apiKeyID, "req2"))
counts, err = cache.GetAPIKeyConcurrencyBatch(s.ctx, []int64{apiKeyID})
require.NoError(s.T(), err)
require.Equal(s.T(), 0, counts[apiKeyID])
}
func (s *ConcurrencyCacheSuite) TestWaitQueue_IncrementAndDecrement() {
userID := int64(20)
waitKey := fmt.Sprintf("%s%d", waitQueueKeyPrefix, userID)
@@ -312,9 +353,11 @@ func (s *ConcurrencyCacheSuite) TestAccountWaitQueue_IncrementAndDecrement() {
func (s *ConcurrencyCacheSuite) TestCleanupStaleProcessSlots() {
accountID := int64(901)
userID := int64(902)
apiKeyID := int64(903)
unindexedAccountID := int64(1901)
accountKey := fmt.Sprintf("%s%d", accountSlotKeyPrefix, accountID)
userKey := fmt.Sprintf("%s%d", userSlotKeyPrefix, userID)
apiKeyKey := fmt.Sprintf("%s%d", apiKeySlotKeyPrefix, apiKeyID)
unindexedAccountKey := fmt.Sprintf("%s%d", accountSlotKeyPrefix, unindexedAccountID)
userWaitKey := fmt.Sprintf("%s%d", waitQueueKeyPrefix, userID)
accountWaitKey := fmt.Sprintf("%s%d", accountWaitKeyPrefix, accountID)
@@ -333,6 +376,10 @@ func (s *ConcurrencyCacheSuite) TestCleanupStaleProcessSlots() {
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, unindexedAccountKey,
redis.Z{Score: float64(now), Member: "oldproc-unindexed"},
).Err())
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, apiKeyKey,
redis.Z{Score: float64(now), Member: "oldproc-3"},
redis.Z{Score: float64(now), Member: "keep-3"},
).Err())
require.NoError(s.T(), s.rdb.Set(s.ctx, userWaitKey, 3, time.Minute).Err())
require.NoError(s.T(), s.rdb.Set(s.ctx, accountWaitKey, 2, time.Minute).Err())
require.NoError(s.T(), s.rdb.Set(s.ctx, unindexedAccountWaitKey, 2, time.Minute).Err())
@@ -355,6 +402,11 @@ func (s *ConcurrencyCacheSuite) TestCleanupStaleProcessSlots() {
require.NoError(s.T(), err)
require.Equal(s.T(), []string{"keep-2"}, userMembers)
// API Key 槽位(stats-only)不在启动清理范围内,靠分数裁剪与 key TTL 自愈。
apiKeyMembers, err := s.rdb.ZRange(s.ctx, apiKeyKey, 0, -1).Result()
require.NoError(s.T(), err)
require.ElementsMatch(s.T(), []string{"keep-3", "oldproc-3"}, apiKeyMembers)
_, err = s.rdb.Get(s.ctx, userWaitKey).Result()
require.True(s.T(), errors.Is(err, redis.Nil))
+43
View File
@@ -466,6 +466,49 @@ func (r *groupRepository) ListActive(ctx context.Context) ([]service.Group, erro
return outGroups, nil
}
func (r *groupRepository) ListActiveIDs(ctx context.Context) ([]int64, error) {
if r.sql != nil {
rows, err := r.sql.QueryContext(ctx, `
SELECT id
FROM groups
WHERE status = $1
AND deleted_at IS NULL
ORDER BY sort_order ASC, id ASC
`, service.StatusActive)
if err != nil {
return nil, err
}
defer func() { _ = rows.Close() }()
ids := make([]int64, 0)
for rows.Next() {
var id int64
if err := rows.Scan(&id); err != nil {
return nil, err
}
ids = append(ids, id)
}
if err := rows.Err(); err != nil {
return nil, err
}
return ids, nil
}
groups, err := r.client.Group.Query().
Where(group.StatusEQ(service.StatusActive)).
Select(group.FieldID).
Order(dbent.Asc(group.FieldSortOrder), dbent.Asc(group.FieldID)).
All(ctx)
if err != nil {
return nil, err
}
ids := make([]int64, 0, len(groups))
for i := range groups {
ids = append(ids, groups[i].ID)
}
return ids, nil
}
func (r *groupRepository) ListActiveByPlatform(ctx context.Context, platform string) ([]service.Group, error) {
groups, err := r.client.Group.Query().
Where(group.StatusEQ(service.StatusActive), group.PlatformEQ(platform)).
@@ -85,10 +85,21 @@ func TestBuildOpsErrorLogsWhere_CyberPolicyStatusExemption(t *testing.T) {
t.Fatalf("default filter must still include the status >= 400 guard for non-cyber rows\nfull: %s", where)
}
// phase=upstream skips the status guard entirely — exemption is irrelevant there.
// phase=upstream WITHOUT the recovered-upstream opt-in keeps the status guard:
// request-error list endpoints filter by phase=upstream as a plain condition.
whereUpstream, _ := buildOpsErrorLogsWhere(&service.OpsErrorLogFilter{Phase: "upstream"})
if strings.Contains(whereUpstream, "status_code") {
t.Fatalf("upstream phase filter must not add any status_code clause\nfull: %s", whereUpstream)
if !strings.Contains(whereUpstream, "COALESCE(e.status_code, 0) >= 400") {
t.Fatalf("upstream phase without IncludeRecoveredUpstream must keep the status guard\nfull: %s", whereUpstream)
}
if !strings.Contains(whereUpstream, "e.error_phase = $") {
t.Fatalf("upstream phase filter must emit the error_phase condition\nfull: %s", whereUpstream)
}
// phase=upstream WITH IncludeRecoveredUpstream (ops 上游列表) skips the guard,
// exposing recovered (<400) upstream rows.
whereRecovered, _ := buildOpsErrorLogsWhere(&service.OpsErrorLogFilter{Phase: "upstream", IncludeRecoveredUpstream: true})
if strings.Contains(whereRecovered, "status_code") {
t.Fatalf("upstream phase with IncludeRecoveredUpstream must not add any status_code clause\nfull: %s", whereRecovered)
}
}
+54 -6
View File
@@ -177,6 +177,37 @@ func opsInsertErrorLogArgs(input *service.OpsInsertErrorLogInput) []any {
}
}
// opsErrorLogsOrderBy builds the ORDER BY clause from a whitelist, mirroring
// usageLogOrderBy semantics. Unknown SortBy falls back to created_at; e.id is
// always appended as tiebreaker for stable pagination.
func opsErrorLogsOrderBy(filter *service.OpsErrorLogFilter) string {
sortBy := ""
sortOrder := ""
if filter != nil {
sortBy = strings.ToLower(strings.TrimSpace(filter.SortBy))
sortOrder = strings.ToLower(strings.TrimSpace(filter.SortOrder))
}
var column string
switch sortBy {
case "model":
column = "COALESCE(NULLIF(TRIM(e.requested_model), ''), e.model)"
case "status_code":
// 与展示列/过滤保持同义:列表展示 COALESCE(upstream_status_code, status_code, 0),
// status_code 过滤也用同一表达式,故排序必须一致——否则 recovered upstream 行
//(status_code<400 但展示上游 5xx)排序键与显示值/分页切分不符。
column = "COALESCE(e.upstream_status_code, e.status_code, 0)"
default:
column = "e.created_at"
}
dir := "DESC"
if sortOrder == "asc" {
dir = "ASC"
}
return fmt.Sprintf("%s %s, e.id %s", column, dir, dir)
}
func (r *opsRepository) ListErrorLogs(ctx context.Context, filter *service.OpsErrorLogFilter) (*service.OpsErrorLogList, error) {
if r == nil || r.db == nil {
return nil, fmt.Errorf("nil ops repository")
@@ -233,25 +264,29 @@ SELECT
COALESCE(a.name, ''),
e.group_id,
COALESCE(g.name, ''),
CASE WHEN e.client_ip IS NULL THEN NULL ELSE e.client_ip::text END,
CASE WHEN e.client_ip IS NULL THEN NULL ELSE host(e.client_ip) END,
COALESCE(e.request_path, ''),
e.stream,
COALESCE(e.inbound_endpoint, ''),
COALESCE(e.upstream_endpoint, ''),
COALESCE(e.requested_model, ''),
COALESCE(e.upstream_model, ''),
COALESCE(e.user_agent, ''),
e.request_type,
COALESCE(ak.name, ''),
ak.deleted_at,
COALESCE(e.deleted_key_name, '')
COALESCE(e.deleted_key_name, ''),
e.deleted_key_owner_user_id,
COALESCE(du.email, '')
FROM ops_error_logs e
LEFT JOIN accounts a ON e.account_id = a.id
LEFT JOIN groups g ON e.group_id = g.id
LEFT JOIN users u ON e.user_id = u.id
LEFT JOIN users u2 ON e.resolved_by_user_id = u2.id
LEFT JOIN users du ON e.deleted_key_owner_user_id = du.id
LEFT JOIN api_keys ak ON ak.id = e.api_key_id
` + where + `
ORDER BY e.created_at DESC
ORDER BY ` + opsErrorLogsOrderBy(filter) + `
LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2)
rows, err := r.db.QueryContext(ctx, selectSQL, argsWithLimit...)
@@ -279,6 +314,8 @@ LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2)
var apiKeyName string
var apiKeyDeletedAt sql.NullTime
var deletedKeyName string
var deletedKeyOwnerID sql.NullInt64
var deletedKeyOwnerEmail string
if err := rows.Scan(
&item.ID,
&item.CreatedAt,
@@ -311,10 +348,13 @@ LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2)
&item.UpstreamEndpoint,
&item.RequestedModel,
&item.UpstreamModel,
&item.UserAgent,
&requestType,
&apiKeyName,
&apiKeyDeletedAt,
&deletedKeyName,
&deletedKeyOwnerID,
&deletedKeyOwnerEmail,
); err != nil {
return nil, err
}
@@ -364,6 +404,12 @@ LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2)
}
// 已删除:ak.deleted_at 非空(软删),或仅命中 deleted_key_name 兜底。
item.APIKeyDeleted = apiKeyDeletedAt.Valid || (apiKeyName == "" && deletedKeyName != "")
// 已删除 KEY 所有者快照:认证失败行 user_id 为空,列表用户列以此回退。
if deletedKeyOwnerID.Valid {
v := deletedKeyOwnerID.Int64
item.DeletedKeyOwnerUserID = &v
item.DeletedKeyOwnerEmail = deletedKeyOwnerEmail
}
out = append(out, &item)
}
if err := rows.Err(); err != nil {
@@ -417,7 +463,7 @@ SELECT
COALESCE(a.name, ''),
e.group_id,
COALESCE(g.name, ''),
CASE WHEN e.client_ip IS NULL THEN NULL ELSE e.client_ip::text END,
CASE WHEN e.client_ip IS NULL THEN NULL ELSE host(e.client_ip) END,
COALESCE(e.request_path, ''),
e.stream,
COALESCE(e.inbound_endpoint, ''),
@@ -927,12 +973,14 @@ func buildOpsErrorLogsWhere(filter *service.OpsErrorLogFilter) (string, []any) {
if filter != nil {
resolvedFilter = filter.Resolved
}
// Keep list endpoints scoped to client errors unless explicitly filtering upstream phase.
// Keep list endpoints scoped to client errors unless the caller explicitly opts
// into recovered upstream rows (Phase=="upstream" + IncludeRecoveredUpstream,
// ops 专用上游列表)。请求错误语义的端点即便过滤 phase=upstream 也保留该守卫。
// cyber_policy is exempt from the status >= 400 guard: streaming cyber hits arrive with
// status 200 (the SSE stream opened successfully before upstream returned response.failed),
// but they are always client-visible blocked requests that belong in admin + user error
// lists. Without the exemption the entire streaming-path cyber sink would be invisible.
if phaseFilter != "upstream" {
if phaseFilter != "upstream" || filter == nil || !filter.IncludeRecoveredUpstream {
clauses = append(clauses, "(COALESCE(e.status_code, 0) >= 400 OR e.error_type = 'cyber_policy')")
}
@@ -518,7 +518,7 @@ func filterSchedulerCredentials(credentials map[string]any) map[string]any {
if len(credentials) == 0 {
return nil
}
keys := []string{"model_mapping", "compact_model_mapping", "api_key", "project_id", "oauth_type"}
keys := []string{"model_mapping", "compact_model_mapping", "api_key", "project_id", "oauth_type", "plan_type"}
filtered := make(map[string]any)
for _, key := range keys {
if value, ok := credentials[key]; ok && value != nil {
@@ -0,0 +1,37 @@
package repository
import (
"testing"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/stretchr/testify/require"
)
func TestFilterSchedulerCredentialsKeepsSubscriptionPlanType(t *testing.T) {
filtered := filterSchedulerCredentials(map[string]any{
"plan_type": "plus",
"access_token": "secret-access-token",
"refresh_token": "secret-refresh-token",
})
require.Equal(t, "plus", filtered["plan_type"])
require.NotContains(t, filtered, "access_token")
require.NotContains(t, filtered, "refresh_token")
}
func TestSchedulerMetadataAccountKeepsOpenAISubscriptionIdentity(t *testing.T) {
account := service.Account{
ID: 24,
Platform: service.PlatformOpenAI,
Type: service.AccountTypeOAuth,
Credentials: map[string]any{
"plan_type": "plus",
"access_token": "secret-access-token",
},
}
metadata := buildSchedulerMetadataAccount(account)
require.True(t, metadata.IsOpenAIChatGPTSubscription())
require.Empty(t, metadata.GetCredential("access_token"))
}
+17 -10
View File
@@ -372,12 +372,13 @@ func (r *usageLogRepository) CreateBestEffort(ctx context.Context, log *service.
}
}
// 队列满时阻塞等待而非立即丢弃:批处理器持续排空队列,短暂等待即可入队。
// 立即丢弃会造成“已扣费但无 usage_log”的永久数据缺口(issue #3656);
// 阻塞上限由调用方 ctx 期限约束,超时后由上层同步兜底。
select {
case r.bestEffortBatchCh <- req:
case <-ctx.Done():
return service.MarkUsageLogCreateDropped(ctx.Err())
default:
return service.MarkUsageLogCreateDropped(errors.New("usage log best-effort queue full"))
}
select {
@@ -493,12 +494,12 @@ func (r *usageLogRepository) createBatched(ctx context.Context, log *service.Usa
resultCh: make(chan usageLogCreateResult, 1),
}
// 队列满时阻塞等待而非立即报错:本路径是 best-effort 丢弃后的最后兜底,
// 立即失败会让日志永久丢失;阻塞上限由调用方 ctx 期限约束。
select {
case r.createBatchCh <- req:
case <-ctx.Done():
return false, service.MarkUsageLogCreateNotPersisted(ctx.Err())
default:
return false, service.MarkUsageLogCreateNotPersisted(errors.New("usage log create batch queue full"))
}
select {
@@ -520,22 +521,28 @@ func (r *usageLogRepository) createBatched(ctx context.Context, log *service.Usa
}
func (r *usageLogRepository) ensureCreateBatcher() {
if r == nil || r.db == nil || r.createBatchCh != nil {
if r == nil || r.db == nil {
return
}
// nil 检查必须在 Once 内部:在外层做无同步快路径读会与 Once 内的写构成数据竞争。
r.createBatchOnce.Do(func() {
r.createBatchCh = make(chan usageLogCreateRequest, usageLogCreateBatchQueueCap)
go r.runCreateBatcher(r.db)
if r.createBatchCh == nil {
r.createBatchCh = make(chan usageLogCreateRequest, usageLogCreateBatchQueueCap)
go r.runCreateBatcher(r.db)
}
})
}
func (r *usageLogRepository) ensureBestEffortBatcher() {
if r == nil || r.db == nil || r.bestEffortBatchCh != nil {
if r == nil || r.db == nil {
return
}
// 同 ensureCreateBatcher:nil 检查放在 Once 内部以避免数据竞争。
r.bestEffortBatchOnce.Do(func() {
r.bestEffortBatchCh = make(chan usageLogBestEffortRequest, usageLogBestEffortBatchQueueCap)
go r.runBestEffortBatcher(r.db)
if r.bestEffortBatchCh == nil {
r.bestEffortBatchCh = make(chan usageLogBestEffortRequest, usageLogBestEffortBatchQueueCap)
go r.runBestEffortBatcher(r.db)
}
})
}
@@ -288,21 +288,21 @@ func TestUsageLogRepositoryCreateBestEffort_BatchPathDuplicateRequestID(t *testi
}, 3*time.Second, 20*time.Millisecond)
}
func TestUsageLogRepositoryCreateBestEffort_QueueFullReturnsDropped(t *testing.T) {
ctx := context.Background()
func TestUsageLogRepositoryCreateBestEffort_QueueFullBlocksUntilCtxDeadline(t *testing.T) {
// 队列满时不再立即丢弃:阻塞等待入队,直到调用方 ctx 到期才标记 dropped(issue #3656)。
client := testEntClient(t)
repo := newUsageLogRepositoryWithSQL(client, integrationDB)
repo.bestEffortBatchCh = make(chan usageLogBestEffortRequest, 1)
repo.bestEffortBatchCh <- usageLogBestEffortRequest{}
user := mustCreateUser(t, client, &service.User{Email: fmt.Sprintf("usage-best-effort-full-%d@example.com", time.Now().UnixNano())})
apiKey := mustCreateApiKey(t, client, &service.APIKey{UserID: user.ID, Key: "sk-usage-best-effort-full-" + uuid.NewString(), Name: "k"})
account := mustCreateAccount(t, client, &service.Account{Name: "acc-usage-best-effort-full-" + uuid.NewString()})
ctx, cancel := context.WithTimeout(context.Background(), 200*time.Millisecond)
defer cancel()
start := time.Now()
err := repo.CreateBestEffort(ctx, &service.UsageLog{
UserID: user.ID,
APIKeyID: apiKey.ID,
AccountID: account.ID,
UserID: 1,
APIKeyID: 2,
AccountID: 3,
RequestID: uuid.NewString(),
Model: "claude-3",
InputTokens: 10,
@@ -314,6 +314,40 @@ func TestUsageLogRepositoryCreateBestEffort_QueueFullReturnsDropped(t *testing.T
require.Error(t, err)
require.True(t, service.IsUsageLogCreateDropped(err))
require.GreaterOrEqual(t, time.Since(start), 150*time.Millisecond)
}
func TestUsageLogRepositoryCreateBestEffort_QueueFullWaitsForDrain(t *testing.T) {
// 队列满但批处理器随后排空时,阻塞的入队应成功完成而非丢弃。
client := testEntClient(t)
repo := newUsageLogRepositoryWithSQL(client, integrationDB)
repo.bestEffortBatchCh = make(chan usageLogBestEffortRequest, 1)
repo.bestEffortBatchCh <- usageLogBestEffortRequest{}
go func() {
time.Sleep(100 * time.Millisecond)
<-repo.bestEffortBatchCh // 排空占位请求,为阻塞中的入队腾出空间
req := <-repo.bestEffortBatchCh
sendUsageLogBestEffortResult(req.resultCh, nil)
}()
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
err := repo.CreateBestEffort(ctx, &service.UsageLog{
UserID: 1,
APIKeyID: 2,
AccountID: 3,
RequestID: uuid.NewString(),
Model: "claude-3",
InputTokens: 10,
OutputTokens: 20,
TotalCost: 0.5,
ActualCost: 0.5,
CreatedAt: time.Now().UTC(),
})
require.NoError(t, err)
}
func TestUsageLogRepositoryCreate_BatchPathCanceledContextMarksNotPersisted(t *testing.T) {
@@ -346,7 +380,7 @@ func TestUsageLogRepositoryCreate_BatchPathCanceledContextMarksNotPersisted(t *t
}
func TestUsageLogRepositoryCreate_BatchPathQueueFullMarksNotPersisted(t *testing.T) {
ctx := context.Background()
// 队列满时阻塞等待入队,直到调用方 ctx 到期才标记 not persisted(issue #3656)。
client := testEntClient(t)
repo := newUsageLogRepositoryWithSQL(client, integrationDB)
repo.createBatchCh = make(chan usageLogCreateRequest, 1)
@@ -356,6 +390,10 @@ func TestUsageLogRepositoryCreate_BatchPathQueueFullMarksNotPersisted(t *testing
apiKey := mustCreateApiKey(t, client, &service.APIKey{UserID: user.ID, Key: "sk-usage-create-full-" + uuid.NewString(), Name: "k"})
account := mustCreateAccount(t, client, &service.Account{Name: "acc-usage-create-full-" + uuid.NewString()})
ctx, cancel := context.WithTimeout(context.Background(), 200*time.Millisecond)
defer cancel()
start := time.Now()
inserted, err := repo.Create(ctx, &service.UsageLog{
UserID: user.ID,
APIKeyID: apiKey.ID,
@@ -372,6 +410,7 @@ func TestUsageLogRepositoryCreate_BatchPathQueueFullMarksNotPersisted(t *testing
require.False(t, inserted)
require.Error(t, err)
require.True(t, service.IsUsageLogCreateNotPersisted(err))
require.GreaterOrEqual(t, time.Since(start), 150*time.Millisecond)
}
func TestUsageLogRepositoryCreate_BatchPathCanceledAfterQueueMarksNotPersisted(t *testing.T) {
+63 -9
View File
@@ -233,6 +233,7 @@ func TestAPIContracts(t *testing.T) {
"ip_whitelist": null,
"ip_blacklist": null,
"last_used_at": null,
"current_concurrency": 0,
"quota": 0,
"quota_used": 0,
"rate_limit_5h": 0,
@@ -282,6 +283,7 @@ func TestAPIContracts(t *testing.T) {
"ip_whitelist": null,
"ip_blacklist": null,
"last_used_at": null,
"current_concurrency": 0,
"quota": 0,
"quota_used": 0,
"rate_limit_5h": 0,
@@ -661,15 +663,17 @@ func TestAPIContracts(t *testing.T) {
service.SettingKeyTableDefaultPageSize: "20",
service.SettingKeyTablePageSizeOptions: "[10,20,50,100]",
service.SettingKeyOpsMonitoringEnabled: "false",
service.SettingKeyOpsRealtimeMonitoringEnabled: "true",
service.SettingKeyOpsQueryModeDefault: "auto",
service.SettingKeyOpsMetricsIntervalSeconds: "60",
service.SettingPaymentVisibleMethodAlipaySource: service.VisibleMethodSourceEasyPayAlipay,
service.SettingPaymentVisibleMethodWxpaySource: service.VisibleMethodSourceOfficialWechat,
service.SettingPaymentVisibleMethodAlipayEnabled: "true",
service.SettingPaymentVisibleMethodWxpayEnabled: "false",
"openai_advanced_scheduler_enabled": "true",
service.SettingKeyOpsMonitoringEnabled: "false",
service.SettingKeyOpsRealtimeMonitoringEnabled: "true",
service.SettingKeyOpsQueryModeDefault: "auto",
service.SettingKeyOpsMetricsIntervalSeconds: "60",
service.SettingPaymentVisibleMethodAlipaySource: service.VisibleMethodSourceEasyPayAlipay,
service.SettingPaymentVisibleMethodWxpaySource: service.VisibleMethodSourceOfficialWechat,
service.SettingPaymentVisibleMethodAlipayEnabled: "true",
service.SettingPaymentVisibleMethodWxpayEnabled: "false",
"openai_advanced_scheduler_enabled": "true",
service.SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled: "false",
service.SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled: "false",
})
},
method: http.MethodGet,
@@ -861,6 +865,28 @@ func TestAPIContracts(t *testing.T) {
"payment_visible_method_alipay_enabled": true,
"payment_visible_method_wxpay_enabled": false,
"openai_advanced_scheduler_enabled": true,
"openai_advanced_scheduler_sticky_weighted_enabled": false,
"openai_advanced_scheduler_subscription_priority_enabled": false,
"openai_advanced_scheduler_lb_top_k": "",
"openai_advanced_scheduler_weight_priority": "",
"openai_advanced_scheduler_weight_load": "",
"openai_advanced_scheduler_weight_queue": "",
"openai_advanced_scheduler_weight_error_rate": "",
"openai_advanced_scheduler_weight_ttft": "",
"openai_advanced_scheduler_weight_reset": "",
"openai_advanced_scheduler_weight_quota_headroom": "",
"openai_advanced_scheduler_weight_previous_response": "",
"openai_advanced_scheduler_weight_session_sticky": "",
"openai_advanced_scheduler_effective_lb_top_k": "7",
"openai_advanced_scheduler_effective_weight_priority": "1",
"openai_advanced_scheduler_effective_weight_load": "1",
"openai_advanced_scheduler_effective_weight_queue": "0.7",
"openai_advanced_scheduler_effective_weight_error_rate": "0.8",
"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_previous_response": "5",
"openai_advanced_scheduler_effective_weight_session_sticky": "3",
"openai_codex_user_agent": "",
"openai_fast_policy_settings": {
"rules": []
@@ -875,6 +901,7 @@ func TestAPIContracts(t *testing.T) {
"payment_max_pending_orders": 0,
"payment_balance_disabled": false,
"payment_balance_recharge_multiplier": 0,
"payment_subscription_usd_to_cny_rate": 0,
"payment_recharge_fee_rate": 0,
"payment_load_balance_strategy": "",
"payment_product_name_prefix": "",
@@ -1110,6 +1137,28 @@ func TestAPIContracts(t *testing.T) {
"payment_visible_method_alipay_enabled": false,
"payment_visible_method_wxpay_enabled": false,
"openai_advanced_scheduler_enabled": false,
"openai_advanced_scheduler_sticky_weighted_enabled": false,
"openai_advanced_scheduler_subscription_priority_enabled": false,
"openai_advanced_scheduler_lb_top_k": "",
"openai_advanced_scheduler_weight_priority": "",
"openai_advanced_scheduler_weight_load": "",
"openai_advanced_scheduler_weight_queue": "",
"openai_advanced_scheduler_weight_error_rate": "",
"openai_advanced_scheduler_weight_ttft": "",
"openai_advanced_scheduler_weight_reset": "",
"openai_advanced_scheduler_weight_quota_headroom": "",
"openai_advanced_scheduler_weight_previous_response": "",
"openai_advanced_scheduler_weight_session_sticky": "",
"openai_advanced_scheduler_effective_lb_top_k": "7",
"openai_advanced_scheduler_effective_weight_priority": "1",
"openai_advanced_scheduler_effective_weight_load": "1",
"openai_advanced_scheduler_effective_weight_queue": "0.7",
"openai_advanced_scheduler_effective_weight_error_rate": "0.8",
"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_previous_response": "5",
"openai_advanced_scheduler_effective_weight_session_sticky": "3",
"openai_codex_user_agent": "",
"openai_fast_policy_settings": {
"rules": []
@@ -1123,6 +1172,7 @@ func TestAPIContracts(t *testing.T) {
"payment_enabled_types": null,
"payment_balance_disabled": false,
"payment_balance_recharge_multiplier": 0,
"payment_subscription_usd_to_cny_rate": 0,
"payment_recharge_fee_rate": 0,
"payment_load_balance_strategy": "",
"payment_product_name_prefix": "",
@@ -1690,6 +1740,10 @@ func (s *stubAccountRepo) List(ctx context.Context, params pagination.Pagination
return nil, nil, errors.New("not implemented")
}
func (s *stubAccountRepo) ListAllWithFilters(context.Context, string, string, string, string, int64, string) ([]service.Account, error) {
return nil, nil
}
func (s *stubAccountRepo) ListWithFilters(ctx context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]service.Account, *pagination.PaginationResult, error) {
return nil, nil, errors.New("not implemented")
}
+76
View File
@@ -70,6 +70,14 @@ type Account struct {
modelMappingCacheRawPtr uintptr
modelMappingCacheRawLen int
modelMappingCacheRawSig uint64
// header_overrides 热路径缓存(非持久化字段,同 model_mapping 缓存先例)
headerOverrideCache map[string]string
headerOverrideCacheReady bool
headerOverrideCacheCredentialsPtr uintptr
headerOverrideCacheRawPtr uintptr
headerOverrideCacheRawLen int
headerOverrideCacheRawSig uint64
}
type OpenAIEndpointCapability string
@@ -580,6 +588,7 @@ func (a *Account) resolveModelMapping(rawMapping map[string]any) map[string]stri
"gemini-3.1-pro-high",
"gemini-3.1-pro-low",
})
applyAntigravityGemini31ProAliases(result)
}
return result
}
@@ -646,6 +655,61 @@ func ensureAntigravityDefaultPassthroughs(mapping map[string]string, models []st
}
}
func applyAntigravityGemini31ProAliases(mapping map[string]string) {
target := strings.TrimSpace(mapping[domain.AntigravityGemini31ProAgentModel])
if target == "" {
return
}
aliases := []struct {
model string
legacyTargets map[string]struct{}
}{
{
model: "gemini-3.1-pro",
legacyTargets: map[string]struct{}{
"gemini-3.1-pro": {},
},
},
{
model: "gemini-3.1-pro-high",
legacyTargets: map[string]struct{}{
"gemini-3.1-pro-high": {},
},
},
{
model: "gemini-3.1-pro-preview",
legacyTargets: map[string]struct{}{
"gemini-3.1-pro-preview": {},
"gemini-3.1-pro-high": {},
},
},
}
for _, alias := range aliases {
current, exists := mapping[alias.model]
if exists {
if _, legacy := alias.legacyTargets[current]; legacy {
mapping[alias.model] = target
}
continue
}
if mappingHasWildcardForModel(mapping, alias.model) {
continue
}
mapping[alias.model] = target
}
}
func mappingHasWildcardForModel(mapping map[string]string, model string) bool {
for pattern := range mapping {
if matchWildcard(pattern, model) {
return true
}
}
return false
}
func normalizeRequestedModelForLookup(platform, requestedModel string) string {
trimmed := strings.TrimSpace(requestedModel)
if trimmed == "" {
@@ -1126,6 +1190,18 @@ func (a *Account) IsOpenAIOAuth() bool {
return a.IsOpenAI() && a.Type == AccountTypeOAuth
}
func (a *Account) IsOpenAIChatGPTSubscription() bool {
if !a.IsOpenAIOAuth() {
return false
}
switch strings.ToLower(strings.TrimSpace(a.GetCredential("plan_type"))) {
case "", "free", "abnormal":
return false
default:
return true
}
}
func (a *Account) IsOpenAIPersonalAccessToken() bool {
if !a.IsOpenAIOAuth() {
return false
@@ -0,0 +1,280 @@
package service
import (
"net/http"
"strings"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
"golang.org/x/net/http/httpguts"
)
// 请求头覆写(header override):仅对 Anthropic / OpenAI 平台的 api_key 账号生效。
// 管理员在账号上配置一组 header name -> value,转发到上游前用配置值覆盖同名请求头
// (匹配不区分大小写);value 为空的条目视为"未填写",不参与覆盖。
const (
credKeyHeaderOverrideEnabled = "header_override_enabled"
credKeyHeaderOverrides = "header_overrides"
maxHeaderOverrideEntries = 64
maxHeaderOverrideNameLength = 200
maxHeaderOverrideValueLength = 8192
)
// headerOverrideBlockedNames 禁止覆写的请求头(小写)。
// - 连接控制/逐跳头:由 HTTP 栈管理,覆写会破坏请求传输;
// - host/content-length:由 Go 的 Request.Host / ContentLength 字段管理,header 覆写不生效或产生冲突;
// - content-type:承载报文框架信息(multipart boundary 为每请求随机值),静态覆写必然与 body 不匹配;
// - authorization/x-api-key/cookie 等:上游认证头由账号凭据统一注入,禁止通过覆写篡改或重新引入;
// - accept-encoding:强制压缩会破坏网关对上游流式响应(SSE/usage)的解析;
// - sec-websocket-*:WebSocket 握手头由拨号器管理(OpenAI WS 模式);
// - session_id/x-claude-code-session-id 等:逐请求会话隔离头,固定值会造成会话串扰。
var headerOverrideBlockedNames = map[string]struct{}{
"host": {},
"content-length": {},
"content-type": {},
"transfer-encoding": {},
"connection": {},
"keep-alive": {},
"proxy-authenticate": {},
"proxy-authorization": {},
"proxy-connection": {},
"te": {},
"trailer": {},
"upgrade": {},
"authorization": {},
"x-api-key": {},
"x-goog-api-key": {},
"cookie": {},
"accept-encoding": {},
"sec-websocket-key": {},
"sec-websocket-version": {},
"sec-websocket-extensions": {},
"sec-websocket-protocol": {},
"sec-websocket-accept": {},
"session_id": {},
"conversation_id": {},
"x-codex-turn-state": {},
"x-codex-turn-metadata": {},
"chatgpt-account-id": {},
"x-claude-code-session-id": {},
"x-client-request-id": {},
}
func isHeaderOverrideBlockedName(lowerName string) bool {
_, blocked := headerOverrideBlockedNames[lowerName]
return blocked
}
// IsHeaderOverrideEligible 报告账号类型是否支持请求头覆写。
// 目前仅开放 Anthropic / OpenAI 两个平台的 api_key 账号。
func (a *Account) IsHeaderOverrideEligible() bool {
if a == nil || a.Type != AccountTypeAPIKey {
return false
}
return a.Platform == PlatformAnthropic || a.Platform == PlatformOpenAI
}
// IsHeaderOverrideEnabled 报告账号是否启用了请求头覆写。
func (a *Account) IsHeaderOverrideEnabled() bool {
if !a.IsHeaderOverrideEligible() || a.Credentials == nil {
return false
}
enabled, ok := a.Credentials[credKeyHeaderOverrideEnabled].(bool)
return ok && enabled
}
// GetHeaderOverrides 返回生效的请求头覆写表(key 统一小写)。
// 未启用、不符合平台/类型条件或配置为空时返回 nil。
// 空 value 的条目(模板占位)与非法/禁止的 header 名会被跳过。
// 结果带热路径缓存(同 GetModelMapping 先例):同一 credentials 映射在
// 一次请求 / 一条 WS 会话内的多次调用只做一次解析与校验。
func (a *Account) GetHeaderOverrides() map[string]string {
if !a.IsHeaderOverrideEnabled() {
return nil
}
rawMapping, rawIsAnyMap := a.Credentials[credKeyHeaderOverrides].(map[string]any)
if !rawIsAnyMap {
// 非 JSON 反序列化产物(如直接注入的 map[string]string):直接解析,不缓存
return resolveHeaderOverrides(stringMappingFromRaw(a.Credentials[credKeyHeaderOverrides]))
}
credentialsPtr := mapPtr(a.Credentials)
rawPtr := mapPtr(rawMapping)
rawLen := len(rawMapping)
rawSig := uint64(0)
rawSigReady := false
if a.headerOverrideCacheReady &&
a.headerOverrideCacheCredentialsPtr == credentialsPtr &&
a.headerOverrideCacheRawPtr == rawPtr &&
a.headerOverrideCacheRawLen == rawLen {
rawSig = modelMappingSignature(rawMapping)
rawSigReady = true
if a.headerOverrideCacheRawSig == rawSig {
return a.headerOverrideCache
}
}
overrides := resolveHeaderOverrides(stringMappingFromRaw(rawMapping))
if !rawSigReady {
rawSig = modelMappingSignature(rawMapping)
}
a.headerOverrideCache = overrides
a.headerOverrideCacheReady = true
a.headerOverrideCacheCredentialsPtr = credentialsPtr
a.headerOverrideCacheRawPtr = rawPtr
a.headerOverrideCacheRawLen = rawLen
a.headerOverrideCacheRawSig = rawSig
return overrides
}
// resolveHeaderOverrides 解析并防御性过滤原始覆写表:保存路径已做校验,
// 这里兜底未经 Normalize 落库的数据(含名单扩充前保存的旧配置),非法条目直接跳过。
func resolveHeaderOverrides(raw map[string]string) map[string]string {
if len(raw) == 0 {
return nil
}
result := make(map[string]string, len(raw))
for name, value := range raw {
lowerName, value, err := normalizeHeaderOverrideEntry(name, value)
if err != nil || lowerName == "" || value == "" {
continue
}
result[lowerName] = value
}
if len(result) == 0 {
return nil
}
return result
}
// HeaderOverrideValue 返回指定 header(小写名)的生效覆写值。
// 供转发链路在 header 写入前感知覆写结果(如 anthropic-beta 需要参与 body 净化)。
func (a *Account) HeaderOverrideValue(lowerName string) (string, bool) {
value, ok := a.GetHeaderOverrides()[lowerName]
return value, ok
}
// ApplyHeaderOverrides 将账号配置的请求头覆写应用到出站请求头。
// 对每个覆写条目:先删除所有大小写变体(转发链路会以 wire casing 直接写入 map,
// 可能存在非 canonical key),再按已知 wire casing 写入,避免产生重复头。
// 账号未启用或不符合条件时为 no-op,可安全地在 OAuth/api_key 共用的构建器中调用。
func (a *Account) ApplyHeaderOverrides(h http.Header) {
if h == nil {
return
}
overrides := a.GetHeaderOverrides()
if len(overrides) == 0 {
return
}
// 覆写名两两不同(大小写不敏感)且各自只操作同名键,应用顺序不影响结果。
// 全量 EqualFold 扫描兜底删除任意 casing 的既有键:透传链路可能保留客户端
// 原始 casing,非 canonical/wire casing 的键 deleteHeaderAllForms 覆盖不到。
for name, value := range overrides {
for existing := range h {
if strings.EqualFold(existing, name) {
delete(h, existing)
}
}
h[resolveWireCasing(name)] = []string{value}
}
}
// NormalizeHeaderOverrideCredentials 校验并原地规范化 credentials 中的请求头覆写字段。
// 供账号创建/更新/批量更新的保存路径调用;credentials 未携带相关字段时为 no-op。
// 规范化内容:header 名转小写并去除首尾空白,value 去除首尾空白,丢弃名和值均为空的条目。
func NormalizeHeaderOverrideCredentials(credentials map[string]any) error {
if credentials == nil {
return nil
}
if raw, ok := credentials[credKeyHeaderOverrideEnabled]; ok && raw != nil {
if _, isBool := raw.(bool); !isBool {
return infraerrors.New(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE",
"header_override_enabled must be a boolean")
}
}
raw, ok := credentials[credKeyHeaderOverrides]
if !ok || raw == nil {
return nil
}
var entries map[string]any
switch m := raw.(type) {
case map[string]any:
entries = m
case map[string]string:
entries = make(map[string]any, len(m))
for k, v := range m {
entries[k] = v
}
default:
return infraerrors.New(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE",
"header_overrides must be an object of header name to string value")
}
if len(entries) > maxHeaderOverrideEntries {
return infraerrors.Newf(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE",
"header_overrides supports at most %d entries", maxHeaderOverrideEntries)
}
normalized := make(map[string]any, len(entries))
for name, rawValue := range entries {
value, isString := rawValue.(string)
if !isString {
return infraerrors.Newf(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE",
"header %q value must be a string", name)
}
lowerName, value, err := normalizeHeaderOverrideEntry(name, value)
if err != nil {
return err
}
if lowerName == "" {
continue // 丢弃完全为空的占位行
}
if _, dup := normalized[lowerName]; dup {
return infraerrors.Newf(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE",
"duplicate header name %q (matching is case-insensitive)", lowerName)
}
normalized[lowerName] = value
}
credentials[credKeyHeaderOverrides] = normalized
return nil
}
// normalizeHeaderOverrideEntry 校验并规范化单个覆写条目,保存路径(Normalize,err → 400)
// 与应用路径(resolveHeaderOverrides,err → 跳过)共用同一套规则,避免两处校验漂移。
// 名和值均为空表示空占位行,返回 ("", "", nil);空 value 的具名条目合法(模板占位)。
func normalizeHeaderOverrideEntry(name, value string) (string, string, error) {
lowerName := strings.ToLower(strings.TrimSpace(name))
value = strings.TrimSpace(value)
if lowerName == "" {
if value == "" {
return "", "", nil
}
return "", "", infraerrors.New(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE",
"header name must not be empty")
}
if len(lowerName) > maxHeaderOverrideNameLength {
return "", "", infraerrors.Newf(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE",
"header name %q exceeds %d characters", lowerName, maxHeaderOverrideNameLength)
}
if !httpguts.ValidHeaderFieldName(lowerName) {
return "", "", infraerrors.Newf(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE",
"invalid header name %q", lowerName)
}
if isHeaderOverrideBlockedName(lowerName) {
return "", "", infraerrors.Newf(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE",
"header %q is not allowed to be overridden", lowerName)
}
if len(value) > maxHeaderOverrideValueLength {
return "", "", infraerrors.Newf(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE",
"header %q value exceeds %d characters", lowerName, maxHeaderOverrideValueLength)
}
if !httpguts.ValidHeaderFieldValue(value) {
return "", "", infraerrors.Newf(http.StatusBadRequest, "INVALID_HEADER_OVERRIDE",
"header %q has an invalid value", lowerName)
}
return lowerName, value, nil
}
@@ -0,0 +1,339 @@
//go:build unit
package service
import (
"net/http"
"strings"
"testing"
"github.com/stretchr/testify/require"
)
func headerOverrideTestAccount(platform, accountType string, credentials map[string]any) *Account {
return &Account{
Platform: platform,
Type: accountType,
Credentials: credentials,
}
}
func TestIsHeaderOverrideEligible(t *testing.T) {
tests := []struct {
name string
platform string
accType string
want bool
}{
{"anthropic apikey", PlatformAnthropic, AccountTypeAPIKey, true},
{"openai apikey", PlatformOpenAI, AccountTypeAPIKey, true},
{"anthropic oauth", PlatformAnthropic, AccountTypeOAuth, false},
{"openai oauth", PlatformOpenAI, AccountTypeOAuth, false},
{"gemini apikey", PlatformGemini, AccountTypeAPIKey, false},
{"grok apikey", PlatformGrok, AccountTypeAPIKey, false},
{"antigravity apikey", PlatformAntigravity, AccountTypeAPIKey, false},
{"anthropic bedrock", PlatformAnthropic, AccountTypeBedrock, false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
acc := headerOverrideTestAccount(tt.platform, tt.accType, nil)
require.Equal(t, tt.want, acc.IsHeaderOverrideEligible())
})
}
var nilAccount *Account
require.False(t, nilAccount.IsHeaderOverrideEligible())
require.False(t, nilAccount.IsHeaderOverrideEnabled())
require.Nil(t, nilAccount.GetHeaderOverrides())
}
func TestIsHeaderOverrideEnabled(t *testing.T) {
acc := headerOverrideTestAccount(PlatformAnthropic, AccountTypeAPIKey, map[string]any{
credKeyHeaderOverrideEnabled: true,
})
require.True(t, acc.IsHeaderOverrideEnabled())
// 未配置 / 非 bool / false 均视为未启用
require.False(t, headerOverrideTestAccount(PlatformAnthropic, AccountTypeAPIKey, nil).IsHeaderOverrideEnabled())
require.False(t, headerOverrideTestAccount(PlatformAnthropic, AccountTypeAPIKey, map[string]any{
credKeyHeaderOverrideEnabled: "true",
}).IsHeaderOverrideEnabled())
require.False(t, headerOverrideTestAccount(PlatformAnthropic, AccountTypeAPIKey, map[string]any{
credKeyHeaderOverrideEnabled: false,
}).IsHeaderOverrideEnabled())
// 不符合平台/类型条件时即使配置了 true 也不启用
require.False(t, headerOverrideTestAccount(PlatformAnthropic, AccountTypeOAuth, map[string]any{
credKeyHeaderOverrideEnabled: true,
}).IsHeaderOverrideEnabled())
require.False(t, headerOverrideTestAccount(PlatformGemini, AccountTypeAPIKey, map[string]any{
credKeyHeaderOverrideEnabled: true,
}).IsHeaderOverrideEnabled())
}
func TestGetHeaderOverrides(t *testing.T) {
acc := headerOverrideTestAccount(PlatformOpenAI, AccountTypeAPIKey, map[string]any{
credKeyHeaderOverrideEnabled: true,
credKeyHeaderOverrides: map[string]any{
"User-Agent": "my-agent/1.0", // 大写 key 归一化为小写
" X-App ": "cli", // 名称去空白
"x-empty": "", // 空 value(模板占位)跳过
"authorization": "Bearer leaked", // 禁止覆写的头跳过
"bad name": "value", // 非法 header 名跳过
"x-padded": " padded ", // value 去空白
},
})
overrides := acc.GetHeaderOverrides()
require.Equal(t, map[string]string{
"user-agent": "my-agent/1.0",
"x-app": "cli",
"x-padded": "padded",
}, overrides)
// 未启用时返回 nil
disabled := headerOverrideTestAccount(PlatformOpenAI, AccountTypeAPIKey, map[string]any{
credKeyHeaderOverrides: map[string]any{"user-agent": "x"},
})
require.Nil(t, disabled.GetHeaderOverrides())
// 启用但全部为空 value 时返回 nil
empty := headerOverrideTestAccount(PlatformOpenAI, AccountTypeAPIKey, map[string]any{
credKeyHeaderOverrideEnabled: true,
credKeyHeaderOverrides: map[string]any{"user-agent": ""},
})
require.Nil(t, empty.GetHeaderOverrides())
// 未经 Normalize 落库的超长数据 / WebSocket 握手头在应用时被防御性跳过
oversizedValue := strings.Repeat("a", maxHeaderOverrideValueLength+1)
defensive := headerOverrideTestAccount(PlatformOpenAI, AccountTypeAPIKey, map[string]any{
credKeyHeaderOverrideEnabled: true,
credKeyHeaderOverrides: map[string]any{
"x-big": oversizedValue,
"sec-websocket-key": "forged",
"content-type": "application/json", // 名单扩充前落库的数据也要被拦截
"x-claude-code-session-id": "pinned-session",
"x-ok": "ok",
},
})
require.Equal(t, map[string]string{"x-ok": "ok"}, defensive.GetHeaderOverrides())
}
func TestApplyHeaderOverrides(t *testing.T) {
acc := headerOverrideTestAccount(PlatformAnthropic, AccountTypeAPIKey, map[string]any{
credKeyHeaderOverrideEnabled: true,
credKeyHeaderOverrides: map[string]any{
"user-agent": "override-agent/2.0",
"anthropic-beta": "custom-beta-1",
"x-custom": "custom-value",
},
})
h := http.Header{}
// 模拟转发链路:canonical key 与 wire casing 原样 key 混合存在
h.Set("User-Agent", "claude-cli/2.1.161 (external, cli)")
h["anthropic-beta"] = []string{"claude-code-20250219,oauth-2025-04-20"} // 非 canonical 原样 key
h.Set("Content-Type", "application/json")
acc.ApplyHeaderOverrides(h)
// user-agent 覆盖且只有一个值(已知头恢复 wire casing)
require.Equal(t, []string{"override-agent/2.0"}, h["User-Agent"])
// anthropic-beta:非 canonical 旧值被清除,写入 wire casing(小写)
require.Equal(t, []string{"custom-beta-1"}, h["anthropic-beta"])
require.Empty(t, h["Anthropic-Beta"])
// 新增头(未知头以小写原样键写入,与转发链路 wire casing 约定一致)
require.Equal(t, []string{"custom-value"}, h["x-custom"])
require.Equal(t, "custom-value", getHeaderRaw(h, "x-custom"))
// 未覆写的头不受影响
require.Equal(t, "application/json", h.Get("Content-Type"))
// 覆盖后不存在任何大小写重复
count := 0
for k := range h {
if k == "anthropic-beta" || k == "Anthropic-Beta" {
count++
}
}
require.Equal(t, 1, count)
}
func TestApplyHeaderOverridesNoOpPaths(t *testing.T) {
baseline := func() http.Header {
h := http.Header{}
h.Set("User-Agent", "orig")
return h
}
// OAuth 账号:即使配置了覆写也不生效
oauth := headerOverrideTestAccount(PlatformAnthropic, AccountTypeOAuth, map[string]any{
credKeyHeaderOverrideEnabled: true,
credKeyHeaderOverrides: map[string]any{"user-agent": "hacked"},
})
h := baseline()
oauth.ApplyHeaderOverrides(h)
require.Equal(t, "orig", h.Get("User-Agent"))
// 未启用开关
off := headerOverrideTestAccount(PlatformAnthropic, AccountTypeAPIKey, map[string]any{
credKeyHeaderOverrides: map[string]any{"user-agent": "hacked"},
})
h = baseline()
off.ApplyHeaderOverrides(h)
require.Equal(t, "orig", h.Get("User-Agent"))
// 禁止覆写的头(authorization / x-api-key / host 等)不会被应用
blocked := headerOverrideTestAccount(PlatformOpenAI, AccountTypeAPIKey, map[string]any{
credKeyHeaderOverrideEnabled: true,
credKeyHeaderOverrides: map[string]any{
"Authorization": "Bearer evil",
"X-Api-Key": "evil",
"Host": "evil.example.com",
"Content-Length": "0",
},
})
h = http.Header{}
h.Set("Authorization", "Bearer real-key")
blocked.ApplyHeaderOverrides(h)
require.Equal(t, "Bearer real-key", h.Get("Authorization"))
require.Empty(t, h.Get("X-Api-Key"))
require.Empty(t, h.Get("Host"))
// nil header 不 panic
blocked.ApplyHeaderOverrides(nil)
}
func TestNormalizeHeaderOverrideCredentials(t *testing.T) {
t.Run("nil credentials no-op", func(t *testing.T) {
require.NoError(t, NormalizeHeaderOverrideCredentials(nil))
})
t.Run("missing keys no-op", func(t *testing.T) {
creds := map[string]any{"api_key": "sk-xxx"}
require.NoError(t, NormalizeHeaderOverrideCredentials(creds))
_, exists := creds[credKeyHeaderOverrides]
require.False(t, exists)
})
t.Run("normalizes names and values", func(t *testing.T) {
creds := map[string]any{
credKeyHeaderOverrideEnabled: true,
credKeyHeaderOverrides: map[string]any{
" User-Agent ": " my-agent ",
"X-App": "",
"": "", // 完全空行被丢弃
},
}
require.NoError(t, NormalizeHeaderOverrideCredentials(creds))
require.Equal(t, map[string]any{
"user-agent": "my-agent",
"x-app": "",
}, creds[credKeyHeaderOverrides])
})
t.Run("accepts map[string]string input", func(t *testing.T) {
creds := map[string]any{
credKeyHeaderOverrides: map[string]string{"X-App": "cli"},
}
require.NoError(t, NormalizeHeaderOverrideCredentials(creds))
require.Equal(t, map[string]any{"x-app": "cli"}, creds[credKeyHeaderOverrides])
})
t.Run("rejects non-bool enabled", func(t *testing.T) {
err := NormalizeHeaderOverrideCredentials(map[string]any{
credKeyHeaderOverrideEnabled: "yes",
})
require.Error(t, err)
})
t.Run("rejects non-object overrides", func(t *testing.T) {
err := NormalizeHeaderOverrideCredentials(map[string]any{
credKeyHeaderOverrides: []any{"user-agent"},
})
require.Error(t, err)
})
t.Run("rejects non-string value", func(t *testing.T) {
err := NormalizeHeaderOverrideCredentials(map[string]any{
credKeyHeaderOverrides: map[string]any{"x-app": 123},
})
require.Error(t, err)
})
t.Run("rejects invalid header name", func(t *testing.T) {
for _, name := range []string{"bad name", "bad:name", "bad\nname", "值"} {
err := NormalizeHeaderOverrideCredentials(map[string]any{
credKeyHeaderOverrides: map[string]any{name: "v"},
})
require.Error(t, err, "name %q should be rejected", name)
}
})
t.Run("rejects empty name with value", func(t *testing.T) {
err := NormalizeHeaderOverrideCredentials(map[string]any{
credKeyHeaderOverrides: map[string]any{" ": "v"},
})
require.Error(t, err)
})
t.Run("rejects blocked headers", func(t *testing.T) {
for _, name := range []string{
"Authorization", "x-api-key", "Host", "content-length", "Transfer-Encoding",
"connection", "accept-encoding", "Sec-WebSocket-Key", "session_id",
"conversation_id", "x-codex-turn-state", "chatgpt-account-id",
"Content-Type", "Cookie", "x-goog-api-key",
"X-Claude-Code-Session-Id", "x-client-request-id",
} {
err := NormalizeHeaderOverrideCredentials(map[string]any{
credKeyHeaderOverrides: map[string]any{name: "v"},
})
require.Error(t, err, "blocked header %q should be rejected", name)
}
})
t.Run("allows tab inside value", func(t *testing.T) {
creds := map[string]any{
credKeyHeaderOverrides: map[string]any{"x-app": "a\tb"},
}
require.NoError(t, NormalizeHeaderOverrideCredentials(creds))
require.Equal(t, map[string]any{"x-app": "a\tb"}, creds[credKeyHeaderOverrides])
})
t.Run("rejects invalid value", func(t *testing.T) {
err := NormalizeHeaderOverrideCredentials(map[string]any{
credKeyHeaderOverrides: map[string]any{"x-app": "bad\nvalue"},
})
require.Error(t, err)
})
t.Run("rejects duplicate names case-insensitively", func(t *testing.T) {
err := NormalizeHeaderOverrideCredentials(map[string]any{
credKeyHeaderOverrides: map[string]any{
"User-Agent": "a",
"user-agent": "b",
},
})
require.Error(t, err)
})
t.Run("rejects too many entries", func(t *testing.T) {
entries := make(map[string]any, maxHeaderOverrideEntries+1)
for i := 0; i <= maxHeaderOverrideEntries; i++ {
entries["x-h-"+string(rune('a'+i%26))+string(rune('a'+(i/26)%26))+string(rune('a'+(i/676)%26))] = "v"
}
err := NormalizeHeaderOverrideCredentials(map[string]any{
credKeyHeaderOverrides: entries,
})
require.Error(t, err)
})
t.Run("rejects oversized value", func(t *testing.T) {
big := make([]byte, maxHeaderOverrideValueLength+1)
for i := range big {
big[i] = 'a'
}
err := NormalizeHeaderOverrideCredentials(map[string]any{
credKeyHeaderOverrides: map[string]any{"x-app": string(big)},
})
require.Error(t, err)
})
}
@@ -39,6 +39,9 @@ type AccountRepository interface {
List(ctx context.Context, params pagination.PaginationParams) ([]Account, *pagination.PaginationResult, error)
ListWithFilters(ctx context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, *pagination.PaginationResult, error)
// ListAllWithFilters 返回符合过滤条件的全部账号(不分页),用于账号列表页
// 计算 OpenAI 调度分数的过滤范围池。
ListAllWithFilters(ctx context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, error)
ListByGroup(ctx context.Context, groupID int64) ([]Account, error)
ListActive(ctx context.Context) ([]Account, error)
ListOAuthRefreshCandidates(ctx context.Context) ([]Account, error)
@@ -79,6 +79,10 @@ func (s *accountRepoStub) List(ctx context.Context, params pagination.Pagination
panic("unexpected List call")
}
func (s *accountRepoStub) ListAllWithFilters(context.Context, string, string, string, string, int64, string) ([]Account, error) {
return nil, nil
}
func (s *accountRepoStub) ListWithFilters(ctx context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, *pagination.PaginationResult, error) {
panic("unexpected ListWithFilters call")
}
@@ -295,6 +295,9 @@ func (s *AccountTestService) testClaudeAccountConnection(c *gin.Context, account
setAnthropicAPIKeyAuthHeader(req.Header, account, authToken)
}
// 账号级请求头覆写:测试请求与真实转发保持一致的最终头
account.ApplyHeaderOverrides(req.Header)
// Get proxy URL
proxyURL := ""
if account.ProxyID != nil && account.Proxy != nil {
@@ -600,9 +603,19 @@ func (s *AccountTestService) testOpenAIAccountConnection(c *gin.Context, account
if isOAuth {
req.Host = "chatgpt.com"
req.Header.Set("accept", "text/event-stream")
req.Header.Set("OpenAI-Beta", "responses=experimental")
req.Header.Set("Originator", "codex_cli_rs")
if customUA := strings.TrimSpace(credentialAccount.GetOpenAIUserAgent()); customUA != "" {
req.Header.Set("User-Agent", customUA)
} else {
req.Header.Set("User-Agent", codexCLIUserAgent)
}
setOpenAIChatGPTAccountHeaders(req.Header, credentialAccount)
}
// 账号级请求头覆写:测试请求与真实转发保持一致的最终头
credentialAccount.ApplyHeaderOverrides(req.Header)
// Get proxy URL
proxyURL := ""
if account.ProxyID != nil && account.Proxy != nil {
@@ -756,6 +769,9 @@ func (s *AccountTestService) testOpenAIChatCompletionsConnection(
req.Header.Set("Accept", "text/event-stream")
req.Header.Set("Authorization", "Bearer "+authToken)
// 账号级请求头覆写:测试请求与真实转发保持一致的最终头
account.ApplyHeaderOverrides(req.Header)
proxyURL := ""
if account.ProxyID != nil && account.Proxy != nil {
proxyURL = account.Proxy.URL()
@@ -848,6 +864,9 @@ func (s *AccountTestService) testOpenAICompactConnection(c *gin.Context, account
setOpenAIChatGPTAccountHeaders(req.Header, account)
}
// 账号级请求头覆写:测试请求与真实转发保持一致的最终头
account.ApplyHeaderOverrides(req.Header)
proxyURL := ""
if account.ProxyID != nil && account.Proxy != nil {
proxyURL = account.Proxy.URL()
@@ -1599,6 +1618,9 @@ func (s *AccountTestService) testOpenAIImageAPIKey(c *gin.Context, ctx context.C
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+authToken)
// 账号级请求头覆写:测试请求与真实转发保持一致的最终头
account.ApplyHeaderOverrides(req.Header)
proxyURL := ""
if account.ProxyID != nil && account.Proxy != nil {
proxyURL = account.Proxy.URL()
@@ -184,6 +184,7 @@ type UsageInfo struct {
FiveHour *UsageProgress `json:"five_hour"` // 5小时窗口
SevenDay *UsageProgress `json:"seven_day,omitempty"` // 7天窗口
SevenDaySonnet *UsageProgress `json:"seven_day_sonnet,omitempty"` // 7天Sonnet窗口
SevenDayFable *UsageProgress `json:"seven_day_fable,omitempty"` // 7天Fable窗口(响应头 7d_oi)
GeminiSharedDaily *UsageProgress `json:"gemini_shared_daily,omitempty"` // Gemini shared pool RPD (Google One / Code Assist)
GeminiProDaily *UsageProgress `json:"gemini_pro_daily,omitempty"` // Gemini Pro 日配额
GeminiFlashDaily *UsageProgress `json:"gemini_flash_daily,omitempty"` // Gemini Flash 日配额
@@ -236,6 +237,12 @@ type UsageInfo struct {
Error string `json:"error,omitempty"`
}
// ClaudeUsageWindow Anthropic /api/oauth/usage 返回的单个用量窗口
type ClaudeUsageWindow struct {
Utilization float64 `json:"utilization"`
ResetsAt string `json:"resets_at"`
}
// ClaudeUsageResponse Anthropic API返回的usage结构
type ClaudeUsageResponse struct {
FiveHour struct {
@@ -250,6 +257,10 @@ type ClaudeUsageResponse struct {
Utilization float64 `json:"utilization"`
ResetsAt string `json:"resets_at"`
} `json:"seven_day_sonnet"`
// Fable 专属 7d 窗口(对应响应头 7d_oi,claim 名为 seven_day_overage_included,
// 见 anthropic-ratelimit-unified-representative-claim 头)。上游 usage API
// 若不下发该字段,GetUsage 会用被动采样数据回填。
SevenDayOverageIncluded ClaudeUsageWindow `json:"seven_day_overage_included"`
}
// ClaudeUsageFetchOptions 包含获取 Claude 用量数据所需的所有选项
@@ -429,6 +440,12 @@ func (s *AccountUsageService) GetUsage(ctx context.Context, accountID int64, for
// 5. 将主动查询结果同步到被动缓存,下次 passive 加载即为最新值
s.syncActiveToPassive(ctx, account.ID, usage)
// 6. 上游 usage API 目前不一定下发 Fable 7d 窗口;缺失时回填被动采样
// (7d_oi 响应头)的数据,避免主动查询后 7d F 进度条丢失。
if usage.SevenDayFable == nil {
usage.SevenDayFable = buildPassiveUsageWindow(account.Extra, "passive_usage_7d_oi_utilization", "passive_usage_7d_oi_reset")
}
s.tryClearRecoverableAccountError(ctx, account)
return usage, nil
}
@@ -471,25 +488,10 @@ func (s *AccountUsageService) GetPassiveUsage(ctx context.Context, accountID int
}
// 构建 7d 窗口(从被动采样数据)
util7d := parseExtraFloat64(account.Extra["passive_usage_7d_utilization"])
reset7dRaw := parseExtraFloat64(account.Extra["passive_usage_7d_reset"])
if util7d > 0 || reset7dRaw > 0 {
var resetAt *time.Time
var remaining int
if reset7dRaw > 0 {
t := time.Unix(int64(reset7dRaw), 0)
resetAt = &t
remaining = int(time.Until(t).Seconds())
if remaining < 0 {
remaining = 0
}
}
info.SevenDay = &UsageProgress{
Utilization: util7d * 100,
ResetsAt: resetAt,
RemainingSeconds: remaining,
}
}
info.SevenDay = buildPassiveUsageWindow(account.Extra, "passive_usage_7d_utilization", "passive_usage_7d_reset")
// 构建 7d Fable 窗口(从被动采样的 7d_oi 响应头数据)
info.SevenDayFable = buildPassiveUsageWindow(account.Extra, "passive_usage_7d_oi_utilization", "passive_usage_7d_oi_reset")
// 添加窗口统计
s.addWindowStats(ctx, account, info)
@@ -497,6 +499,31 @@ func (s *AccountUsageService) GetPassiveUsage(ctx context.Context, accountID int
return info, nil
}
// buildPassiveUsageWindow 从 Extra 中的被动采样数据(utilization 为 0-1 小数、reset 为 Unix 秒)
// 构建用量窗口,无数据时返回 nil。
func buildPassiveUsageWindow(extra map[string]any, utilKey, resetKey string) *UsageProgress {
util := parseExtraFloat64(extra[utilKey])
resetRaw := parseExtraFloat64(extra[resetKey])
if util <= 0 && resetRaw <= 0 {
return nil
}
var resetAt *time.Time
var remaining int
if resetRaw > 0 {
t := time.Unix(int64(resetRaw), 0)
resetAt = &t
remaining = int(time.Until(t).Seconds())
if remaining < 0 {
remaining = 0
}
}
return &UsageProgress{
Utilization: util * 100,
ResetsAt: resetAt,
RemainingSeconds: remaining,
}
}
// syncActiveToPassive 将主动查询的最新数据回写到 Extra 被动缓存,
// 这样下次被动加载时能看到最新值。
func (s *AccountUsageService) syncActiveToPassive(ctx context.Context, accountID int64, usage *UsageInfo) {
@@ -511,6 +538,12 @@ func (s *AccountUsageService) syncActiveToPassive(ctx context.Context, accountID
extraUpdates["passive_usage_7d_reset"] = usage.SevenDay.ResetsAt.Unix()
}
}
if usage.SevenDayFable != nil {
extraUpdates["passive_usage_7d_oi_utilization"] = usage.SevenDayFable.Utilization / 100
if usage.SevenDayFable.ResetsAt != nil {
extraUpdates["passive_usage_7d_oi_reset"] = usage.SevenDayFable.ResetsAt.Unix()
}
}
if len(extraUpdates) > 0 {
extraUpdates["passive_usage_sampled_at"] = time.Now().UTC().Format(time.RFC3339)
@@ -1010,8 +1043,8 @@ func enrichUsageWithAccountError(info *UsageInfo, account *Account) {
// 使用独立缓存(1 分钟),与 API 缓存分离
func (s *AccountUsageService) addWindowStats(ctx context.Context, account *Account, usage *UsageInfo) {
// 修复:即使 FiveHour 为 nil,也要尝试获取统计数据
// 因为 SevenDay/SevenDaySonnet 可能需要
if usage.FiveHour == nil && usage.SevenDay == nil && usage.SevenDaySonnet == nil {
// 因为 SevenDay/SevenDaySonnet/SevenDayFable 可能需要
if usage.FiveHour == nil && usage.SevenDay == nil && usage.SevenDaySonnet == nil && usage.SevenDayFable == nil {
return
}
@@ -1347,6 +1380,22 @@ func (s *AccountUsageService) buildUsageInfo(resp *ClaudeUsageResponse, updatedA
}
}
// 7天Fable窗口(响应头 7d_oi 对应的窗口)
if fable := resp.SevenDayOverageIncluded; fable.ResetsAt != "" {
if fableReset, err := parseTime(fable.ResetsAt); err == nil {
info.SevenDayFable = &UsageProgress{
Utilization: fable.Utilization,
ResetsAt: &fableReset,
RemainingSeconds: int(time.Until(fableReset).Seconds()),
}
} else {
log.Printf("Failed to parse SevenDayFable.ResetsAt: %s, error: %v", fable.ResetsAt, err)
info.SevenDayFable = &UsageProgress{
Utilization: fable.Utilization,
}
}
}
return info
}
@@ -0,0 +1,119 @@
package service
import (
"encoding/json"
"testing"
"time"
"github.com/stretchr/testify/require"
)
func TestClaudeUsageResponse_FableWindowDecoding(t *testing.T) {
t.Run("seven_day_overage_included", func(t *testing.T) {
raw := `{
"five_hour": {"utilization": 12.0, "resets_at": "2026-07-03T10:00:00Z"},
"seven_day": {"utilization": 34.0, "resets_at": "2026-07-08T00:00:00Z"},
"seven_day_overage_included": {"utilization": 56.0, "resets_at": "2026-07-08T03:00:00Z"}
}`
var resp ClaudeUsageResponse
require.NoError(t, json.Unmarshal([]byte(raw), &resp))
require.Equal(t, 56.0, resp.SevenDayOverageIncluded.Utilization)
require.Equal(t, "2026-07-08T03:00:00Z", resp.SevenDayOverageIncluded.ResetsAt)
})
t.Run("absent", func(t *testing.T) {
raw := `{"five_hour": {"utilization": 12.0, "resets_at": "2026-07-03T10:00:00Z"}}`
var resp ClaudeUsageResponse
require.NoError(t, json.Unmarshal([]byte(raw), &resp))
require.Zero(t, resp.SevenDayOverageIncluded.Utilization)
require.Empty(t, resp.SevenDayOverageIncluded.ResetsAt)
})
}
func TestBuildUsageInfo_SevenDayFable(t *testing.T) {
svc := &AccountUsageService{}
now := time.Now()
resetAt := now.Add(72 * time.Hour).UTC().Truncate(time.Second)
var resp ClaudeUsageResponse
resp.FiveHour.Utilization = 10
resp.SevenDayOverageIncluded = ClaudeUsageWindow{
Utilization: 88,
ResetsAt: resetAt.Format(time.RFC3339),
}
info := svc.buildUsageInfo(&resp, &now)
require.NotNil(t, info.SevenDayFable)
require.Equal(t, 88.0, info.SevenDayFable.Utilization)
require.NotNil(t, info.SevenDayFable.ResetsAt)
require.True(t, info.SevenDayFable.ResetsAt.Equal(resetAt))
require.Greater(t, info.SevenDayFable.RemainingSeconds, 0)
// 无 Fable 数据时不应创建窗口
var empty ClaudeUsageResponse
empty.FiveHour.Utilization = 10
info = svc.buildUsageInfo(&empty, &now)
require.Nil(t, info.SevenDayFable)
}
func TestBuildPassiveUsageWindow(t *testing.T) {
future := time.Now().Add(48 * time.Hour).Unix()
t.Run("utilization and reset", func(t *testing.T) {
window := buildPassiveUsageWindow(map[string]any{
"passive_usage_7d_oi_utilization": 0.87,
"passive_usage_7d_oi_reset": float64(future),
}, "passive_usage_7d_oi_utilization", "passive_usage_7d_oi_reset")
require.NotNil(t, window)
require.InDelta(t, 87.0, window.Utilization, 1e-9)
require.NotNil(t, window.ResetsAt)
require.Equal(t, future, window.ResetsAt.Unix())
require.Greater(t, window.RemainingSeconds, 0)
})
t.Run("no data returns nil", func(t *testing.T) {
require.Nil(t, buildPassiveUsageWindow(nil, "u", "r"))
require.Nil(t, buildPassiveUsageWindow(map[string]any{}, "u", "r"))
})
t.Run("expired reset clamps remaining to zero", func(t *testing.T) {
past := time.Now().Add(-time.Hour).Unix()
window := buildPassiveUsageWindow(map[string]any{
"u": 0.5,
"r": float64(past),
}, "u", "r")
require.NotNil(t, window)
require.Equal(t, 0, window.RemainingSeconds)
})
t.Run("utilization only", func(t *testing.T) {
window := buildPassiveUsageWindow(map[string]any{"u": 0.25}, "u", "r")
require.NotNil(t, window)
require.InDelta(t, 25.0, window.Utilization, 1e-9)
require.Nil(t, window.ResetsAt)
})
}
func TestSyncActiveToPassive_WritesFableExtras(t *testing.T) {
repo := &accountUsageCodexProbeRepo{updateExtraCh: make(chan map[string]any, 1)}
svc := &AccountUsageService{accountRepo: repo}
resetAt := time.Now().Add(72 * time.Hour).Truncate(time.Second)
usage := &UsageInfo{
SevenDayFable: &UsageProgress{
Utilization: 87,
ResetsAt: &resetAt,
},
}
svc.syncActiveToPassive(t.Context(), 1, usage)
select {
case updates := <-repo.updateExtraCh:
require.InDelta(t, 0.87, updates["passive_usage_7d_oi_utilization"], 1e-9)
require.Equal(t, resetAt.Unix(), updates["passive_usage_7d_oi_reset"])
require.Contains(t, updates, "passive_usage_sampled_at")
default:
t.Fatal("expected UpdateExtra to be called with fable extras")
}
}
@@ -4,6 +4,8 @@ package service
import (
"testing"
"github.com/Wei-Shaw/sub2api/internal/domain"
)
func TestMatchWildcard(t *testing.T) {
@@ -320,6 +322,86 @@ func TestAccountGetMappedModel(t *testing.T) {
}
}
func TestAccountGetModelMapping_AntigravityNormalizesGemini31ProAliases(t *testing.T) {
t.Parallel()
account := &Account{
Platform: PlatformAntigravity,
Credentials: map[string]any{
"model_mapping": map[string]any{
domain.AntigravityGemini31ProAgentModel: domain.AntigravityGemini31ProAgentModel,
"gemini-3.1-pro-high": "gemini-3.1-pro-high",
"gemini-3.1-pro-preview": "gemini-3.1-pro-high",
},
},
}
mapping := account.GetModelMapping()
if got := mapping["gemini-3.1-pro"]; got != domain.AntigravityGemini31ProAgentModel {
t.Fatalf("expected gemini-3.1-pro to map to %q, got %q", domain.AntigravityGemini31ProAgentModel, got)
}
if got := mapping["gemini-3.1-pro-high"]; got != domain.AntigravityGemini31ProAgentModel {
t.Fatalf("expected gemini-3.1-pro-high to map to %q, got %q", domain.AntigravityGemini31ProAgentModel, got)
}
if got := mapping["gemini-3.1-pro-preview"]; got != domain.AntigravityGemini31ProAgentModel {
t.Fatalf("expected gemini-3.1-pro-preview to map to %q, got %q", domain.AntigravityGemini31ProAgentModel, got)
}
}
func TestAccountGetModelMapping_AntigravityPreservesGemini31ProOverrides(t *testing.T) {
t.Parallel()
account := &Account{
Platform: PlatformAntigravity,
Credentials: map[string]any{
"model_mapping": map[string]any{
domain.AntigravityGemini31ProAgentModel: domain.AntigravityGemini31ProAgentModel,
"gemini-3.1-pro-high": "custom-high",
"gemini-3.1-pro-preview": "custom-preview",
},
},
}
mapping := account.GetModelMapping()
if got := mapping["gemini-3.1-pro-high"]; got != "custom-high" {
t.Fatalf("expected gemini-3.1-pro-high override to be preserved, got %q", got)
}
if got := mapping["gemini-3.1-pro-preview"]; got != "custom-preview" {
t.Fatalf("expected gemini-3.1-pro-preview override to be preserved, got %q", got)
}
if got := mapping["gemini-3.1-pro"]; got != domain.AntigravityGemini31ProAgentModel {
t.Fatalf("expected gemini-3.1-pro alias to default to %q, got %q", domain.AntigravityGemini31ProAgentModel, got)
}
}
func TestAccountGetModelMapping_AntigravityGemini31ProAliasesRespectWildcard(t *testing.T) {
t.Parallel()
account := &Account{
Platform: PlatformAntigravity,
Credentials: map[string]any{
"model_mapping": map[string]any{
domain.AntigravityGemini31ProAgentModel: domain.AntigravityGemini31ProAgentModel,
"gemini-3.1-*": "custom-wildcard",
},
},
}
mapping := account.GetModelMapping()
if got := mapping["gemini-3.1-pro"]; got != "" {
t.Fatalf("expected gemini-3.1-pro exact alias to stay unset when wildcard exists, got %q", got)
}
if got := mapping["gemini-3.1-pro-high"]; got != "" {
t.Fatalf("expected gemini-3.1-pro-high exact alias to stay unset when wildcard exists, got %q", got)
}
if got := mapping["gemini-3.1-pro-preview"]; got != "" {
t.Fatalf("expected gemini-3.1-pro-preview exact alias to stay unset when wildcard exists, got %q", got)
}
}
func TestAccountResolveMappedModel(t *testing.T) {
tests := []struct {
name string
@@ -5,23 +5,16 @@ package service
import (
"testing"
"github.com/Wei-Shaw/sub2api/internal/pkg/xai"
"github.com/stretchr/testify/require"
)
func TestNormalizeAccountConcurrencyCapsGrokOAuthUnlessUnsafe(t *testing.T) {
t.Setenv(xai.EnvUnsafeAllowHighConcurrency, "")
func TestNormalizeAccountConcurrencyDefaultsInvalidGrokOAuthToOne(t *testing.T) {
require.Equal(t, 1, normalizeAccountConcurrency(PlatformGrok, AccountTypeOAuth, 0))
require.Equal(t, 1, normalizeAccountConcurrency(PlatformGrok, AccountTypeOAuth, -5))
require.Equal(t, 1, normalizeAccountConcurrency(PlatformGrok, AccountTypeOAuth, 50))
}
func TestNormalizeAccountConcurrencyPreservesExplicitValues(t *testing.T) {
require.Equal(t, 50, normalizeAccountConcurrency(PlatformGrok, AccountTypeOAuth, 50))
require.Equal(t, 2, normalizeAccountConcurrency(PlatformOpenAI, AccountTypeOAuth, 2))
require.Equal(t, 2, normalizeAccountConcurrency(PlatformGrok, AccountTypeAPIKey, 2))
}
func TestNormalizeAccountConcurrencyAllowsGrokOAuthUnsafeOverride(t *testing.T) {
t.Setenv(xai.EnvUnsafeAllowHighConcurrency, "true")
require.Equal(t, 50, normalizeAccountConcurrency(PlatformGrok, AccountTypeOAuth, 50))
require.Equal(t, 1, normalizeAccountConcurrency(PlatformGrok, AccountTypeOAuth, 0))
}
+37 -3
View File
@@ -78,6 +78,12 @@ type AdminService interface {
// Account management
ListAccounts(ctx context.Context, page, pageSize int, platform, accountType, status, search string, groupID int64, privacyMode string, sortBy, sortOrder string) ([]Account, int64, error)
// ListAccountsForSchedulerScoreFilter 返回符合过滤条件的全部账号(不分页),
// 作为账号列表页计算 OpenAI 调度分数的过滤范围池。
ListAccountsForSchedulerScoreFilter(ctx context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, error)
// ListOpenAISchedulableAccountsForSchedulerScore 返回指定分组(nil 为未分组)内
// 可调度的 OpenAI 账号,用于按组计算调度分数。
ListOpenAISchedulableAccountsForSchedulerScore(ctx context.Context, groupID *int64) ([]Account, error)
GetAccount(ctx context.Context, id int64) (*Account, error)
GetAccountsByIDs(ctx context.Context, ids []int64) ([]*Account, error)
CreateAccount(ctx context.Context, input *CreateAccountInput) (*Account, error)
@@ -2618,6 +2624,23 @@ func (s *adminServiceImpl) ListAccounts(ctx context.Context, page, pageSize int,
return accounts, result.Total, nil
}
func (s *adminServiceImpl) ListAccountsForSchedulerScoreFilter(ctx context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, error) {
if s == nil || s.accountRepo == nil {
return nil, nil
}
return s.accountRepo.ListAllWithFilters(ctx, platform, accountType, status, search, groupID, privacyMode)
}
func (s *adminServiceImpl) ListOpenAISchedulableAccountsForSchedulerScore(ctx context.Context, groupID *int64) ([]Account, error) {
if s == nil || s.accountRepo == nil {
return nil, nil
}
if groupID != nil {
return s.accountRepo.ListSchedulableByGroupIDAndPlatform(ctx, *groupID, PlatformOpenAI)
}
return s.accountRepo.ListSchedulableUngroupedByPlatform(ctx, PlatformOpenAI)
}
func (s *adminServiceImpl) GetAccount(ctx context.Context, id int64) (*Account, error) {
return s.accountRepo.GetByID(ctx, id)
}
@@ -2640,9 +2663,6 @@ func normalizeAccountConcurrency(platform, accountType string, concurrency int)
if concurrency <= 0 {
return 1
}
if concurrency > 1 && !xai.AllowUnsafeHighConcurrency() {
return 1
}
}
return concurrency
}
@@ -2671,6 +2691,11 @@ func (s *adminServiceImpl) CreateAccount(ctx context.Context, input *CreateAccou
}
}
// 校验并规范化请求头覆写配置(header 名小写化、格式检查)
if err := NormalizeHeaderOverrideCredentials(input.Credentials); err != nil {
return nil, err
}
account := &Account{
Name: input.Name,
Notes: normalizeAccountNotes(input.Notes),
@@ -2801,6 +2826,10 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U
// 敏感子键采用"incoming 没提供就保留"的合并语义:前端响应已脱敏,
// 全对象 PUT 编辑时不会再带回 token,避免覆盖时清空已有凭证。
account.Credentials = MergePreservingSensitiveCreds(account.Credentials, input.Credentials)
// 校验并规范化请求头覆写配置(header 名小写化、格式检查)
if err := NormalizeHeaderOverrideCredentials(account.Credentials); err != nil {
return nil, err
}
}
// Extra 使用 map:需要区分“未提供(nil)”与“显式清空({})”。
// 关闭配额限制时前端会删除 quota_* 键并提交 extra:{},此时也必须落库。
@@ -3019,6 +3048,11 @@ func (s *adminServiceImpl) BulkUpdateAccounts(ctx context.Context, input *BulkUp
}
}
// 校验并规范化请求头覆写配置(批量路径为 JSONB 顶层 key 合并,直接校验增量即可)
if err := NormalizeHeaderOverrideCredentials(input.Credentials); err != nil {
return nil, err
}
// Prepare bulk updates for columns and JSONB fields.
repoUpdates := AccountBulkUpdate{
Credentials: input.Credentials,
@@ -88,6 +88,10 @@ func (s *accountRepoStubForBulkUpdate) ListByGroup(_ context.Context, groupID in
return nil, nil
}
func (s *accountRepoStubForBulkUpdate) ListAllWithFilters(context.Context, string, string, string, string, int64, string) ([]Account, error) {
return nil, nil
}
func (s *accountRepoStubForBulkUpdate) ListWithFilters(_ context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, *pagination.PaginationResult, error) {
s.listCalled = true
s.lastListParams = params
@@ -25,6 +25,10 @@ type accountRepoStubForAdminList struct {
listWithFiltersErr error
}
func (s *accountRepoStubForAdminList) ListAllWithFilters(context.Context, string, string, string, string, int64, string) ([]Account, error) {
return nil, nil
}
func (s *accountRepoStubForAdminList) ListWithFilters(_ context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, *pagination.PaginationResult, error) {
s.listWithFiltersCalls++
s.listWithFiltersParams = params
@@ -12,6 +12,10 @@ const (
// antigravityRefreshWindow Antigravity token 提前刷新窗口:15分钟
// Google OAuth token 有效期55分钟,提前15分钟刷新
antigravityRefreshWindow = 15 * time.Minute
antigravityForceTokenRefreshExtraKey = "antigravity_force_token_refresh"
antigravityForceTokenRefreshReasonExtraKey = "antigravity_force_token_refresh_reason"
antigravityForceTokenRefreshAtExtraKey = "antigravity_force_token_refresh_at"
)
// AntigravityTokenRefresher 实现 TokenRefresher 接口
@@ -41,6 +45,9 @@ func (r *AntigravityTokenRefresher) NeedsRefresh(account *Account, _ time.Durati
if !r.CanRefresh(account) {
return false
}
if accountNeedsAntigravityForceTokenRefresh(account) {
return true
}
expiresAt := account.GetCredentialAsTime("expires_at")
if expiresAt == nil {
return false
@@ -54,6 +61,29 @@ func (r *AntigravityTokenRefresher) NeedsRefresh(account *Account, _ time.Durati
return needsRefresh
}
func accountNeedsAntigravityForceTokenRefresh(account *Account) bool {
return account != nil &&
account.Platform == PlatformAntigravity &&
account.Type == AccountTypeOAuth &&
account.getExtraBool(antigravityForceTokenRefreshExtraKey)
}
func antigravityForceTokenRefreshExtra(reason string) map[string]any {
return map[string]any{
antigravityForceTokenRefreshExtraKey: true,
antigravityForceTokenRefreshReasonExtraKey: reason,
antigravityForceTokenRefreshAtExtraKey: time.Now().UTC().Format(time.RFC3339),
}
}
func clearAntigravityForceTokenRefreshExtra() map[string]any {
return map[string]any{
antigravityForceTokenRefreshExtraKey: false,
antigravityForceTokenRefreshReasonExtraKey: "",
antigravityForceTokenRefreshAtExtraKey: "",
}
}
// Refresh 执行 token 刷新
func (r *AntigravityTokenRefresher) Refresh(ctx context.Context, account *Account) (map[string]any, error) {
tokenInfo, err := r.antigravityOAuthService.RefreshAccountToken(ctx, account)
+1
View File
@@ -44,6 +44,7 @@ type APIKey struct {
UpdatedAt time.Time
User *User
Group *Group
CurrentConcurrency int
// Quota fields
Quota float64 // Quota limit in USD (0 = unlimited)
@@ -203,6 +203,7 @@ type APIKeyService struct {
userGroupRateRepo UserGroupRateRepository
cache APIKeyCache
rateLimitCacheInvalid RateLimitCacheInvalidator // optional: invalidate Redis rate limit cache
concurrencyService *ConcurrencyService
cfg *config.Config
authCacheL1 *ristretto.Cache
authCfg apiKeyAuthCacheConfig
@@ -240,6 +241,10 @@ func (s *APIKeyService) SetRateLimitCacheInvalidator(inv RateLimitCacheInvalidat
s.rateLimitCacheInvalid = inv
}
func (s *APIKeyService) SetConcurrencyService(concurrencyService *ConcurrencyService) {
s.concurrencyService = concurrencyService
}
func (s *APIKeyService) compileAPIKeyIPRules(apiKey *APIKey) {
if apiKey == nil {
return
@@ -436,9 +441,40 @@ func (s *APIKeyService) List(ctx context.Context, userID int64, params paginatio
if err != nil {
return nil, nil, fmt.Errorf("list api keys: %w", err)
}
s.fillCurrentConcurrency(ctx, keys)
return keys, pagination, nil
}
func (s *APIKeyService) fillCurrentConcurrency(ctx context.Context, keys []APIKey) {
if s == nil || s.concurrencyService == nil || len(keys) == 0 {
return
}
ids := make([]int64, 0, len(keys))
for i := range keys {
if keys[i].ID > 0 {
ids = append(ids, keys[i].ID)
}
}
counts, err := s.concurrencyService.GetAPIKeyConcurrencyBatch(ctx, ids)
if err != nil {
return
}
for i := range keys {
keys[i].CurrentConcurrency = counts[keys[i].ID]
}
}
func (s *APIKeyService) currentConcurrencyForAPIKey(ctx context.Context, apiKeyID int64) int {
if s == nil || s.concurrencyService == nil || apiKeyID <= 0 {
return 0
}
counts, err := s.concurrencyService.GetAPIKeyConcurrencyBatch(ctx, []int64{apiKeyID})
if err != nil {
return 0
}
return counts[apiKeyID]
}
func (s *APIKeyService) VerifyOwnership(ctx context.Context, userID int64, apiKeyIDs []int64) ([]int64, error) {
if len(apiKeyIDs) == 0 {
return []int64{}, nil
@@ -458,6 +494,9 @@ func (s *APIKeyService) GetByID(ctx context.Context, id int64) (*APIKey, error)
return nil, fmt.Errorf("get api key: %w", err)
}
s.compileAPIKeyIPRules(apiKey)
if apiKey != nil {
apiKey.CurrentConcurrency = s.currentConcurrencyForAPIKey(ctx, apiKey.ID)
}
return apiKey, nil
}
@@ -300,6 +300,40 @@ func TestApiKeyService_Delete_NotFound(t *testing.T) {
require.Empty(t, cache.deleteAuthKeys)
}
func TestAPIKeyService_List_FillsCurrentConcurrency(t *testing.T) {
repo := &apiKeyRepoStub{
allowListByUserID: true,
listByUserIDKeys: []APIKey{
{ID: 10, UserID: 7, Key: "sk-10", Name: "key-10"},
{ID: 11, UserID: 7, Key: "sk-11", Name: "key-11"},
},
}
concurrency := NewConcurrencyService(&stubConcurrencyCacheForTest{
apiKeyConcurrency: map[int64]int{10: 2, 11: 0},
})
svc := &APIKeyService{apiKeyRepo: repo, concurrencyService: concurrency}
keys, _, err := svc.List(context.Background(), 7, pagination.PaginationParams{Page: 1, PageSize: 20}, APIKeyListFilters{})
require.NoError(t, err)
require.Len(t, keys, 2)
require.Equal(t, 2, keys[0].CurrentConcurrency)
require.Equal(t, 0, keys[1].CurrentConcurrency)
}
func TestAPIKeyService_GetByID_FillsCurrentConcurrency(t *testing.T) {
repo := &apiKeyRepoStub{
apiKey: &APIKey{ID: 10, UserID: 7, Key: "sk-10", Name: "key-10"},
}
concurrency := NewConcurrencyService(&stubConcurrencyCacheForTest{
apiKeyConcurrency: map[int64]int{10: 4},
})
svc := &APIKeyService{apiKeyRepo: repo, concurrencyService: concurrency}
key, err := svc.GetByID(context.Background(), 10)
require.NoError(t, err)
require.Equal(t, 4, key.CurrentConcurrency)
}
// TestApiKeyService_Delete_DeleteFails 测试删除操作失败时的错误处理。
// 预期行为:
// - GetKeyAndOwnerID 返回正确的所有者 ID
+13 -1
View File
@@ -280,6 +280,11 @@ func (s *BillingService) initFallbackPricing() {
s.fallbackPrices["gpt-5.5"] = s.fallbackPrices["gpt-5.4"]
s.fallbackPrices["gpt-5.5-pro"] = s.fallbackPrices["gpt-5.4"]
// GPT-5.6(sol / terra / luna)暂无独立定价,回退到 GPT-5.4。
s.fallbackPrices["gpt-5.6-sol"] = s.fallbackPrices["gpt-5.4"]
s.fallbackPrices["gpt-5.6-terra"] = s.fallbackPrices["gpt-5.4"]
s.fallbackPrices["gpt-5.6-luna"] = s.fallbackPrices["gpt-5.4"]
s.fallbackPrices["gpt-5.4-mini"] = &ModelPricing{
InputPricePerToken: 7.5e-7,
OutputPricePerToken: 4.5e-6,
@@ -667,6 +672,12 @@ func (s *BillingService) getFallbackPricing(model string) *ModelPricing {
// OpenAI(GPT-5 / Codex 族):仅匹配已知型号,避免未知 OpenAI 型号误计价。
if normalized := normalizeKnownOpenAICodexModel(modelLower); normalized != "" {
switch normalized {
case "gpt-5.6-sol":
return s.fallbackPrices["gpt-5.6-sol"]
case "gpt-5.6-terra":
return s.fallbackPrices["gpt-5.6-terra"]
case "gpt-5.6-luna":
return s.fallbackPrices["gpt-5.6-luna"]
case "gpt-5.5-pro":
return s.fallbackPrices["gpt-5.5-pro"]
case "gpt-5.5":
@@ -1060,7 +1071,8 @@ func isOpenAIGPT54Model(model string) bool {
// normalizeCodexModel 的默认兜底把非 OpenAI 模型(claude-*、gemini-*、gpt-4o)
// 误识别为 gpt-5.4。
normalized := normalizeKnownOpenAICodexModel(model)
return normalized == "gpt-5.4" || normalized == "gpt-5.5" || normalized == "gpt-5.5-pro"
return normalized == "gpt-5.4" || normalized == "gpt-5.5" || normalized == "gpt-5.5-pro" ||
normalized == "gpt-5.6-sol" || normalized == "gpt-5.6-terra" || normalized == "gpt-5.6-luna"
}
// CalculateCostWithConfig 使用配置中的默认倍率计算费用
@@ -4,6 +4,13 @@ import "strings"
const featureKeyCodexImageGenerationBridge = "codex_image_generation_bridge"
const (
featureKeyCodexImageGenerationExplicitToolPolicy = "codex_image_generation_explicit_tool_policy"
codexImageGenerationExplicitToolPolicyAllow = "allow"
codexImageGenerationExplicitToolPolicyStrip = "strip"
)
func boolOverridePtr(v bool) *bool {
return &v
}
@@ -20,6 +27,27 @@ func boolOverrideFromMap(values map[string]any, keys ...string) *bool {
return nil
}
func stringOverrideFromMap(values map[string]any, keys ...string) (string, bool) {
if values == nil {
return "", false
}
for _, key := range keys {
if v, ok := values[key].(string); ok {
return v, true
}
}
return "", false
}
func normalizeCodexImageGenerationExplicitToolPolicy(value string) string {
switch strings.ToLower(strings.TrimSpace(value)) {
case codexImageGenerationExplicitToolPolicyStrip, "remove", "drop":
return codexImageGenerationExplicitToolPolicyStrip
default:
return codexImageGenerationExplicitToolPolicyAllow
}
}
func platformBoolOverride(values map[string]any, key string, platform string) *bool {
if values == nil {
return nil
@@ -62,3 +90,20 @@ func (a *Account) CodexImageGenerationBridgeOverride() *bool {
openaiConfig, _ := a.Extra[PlatformOpenAI].(map[string]any)
return boolOverrideFromMap(openaiConfig, featureKeyCodexImageGenerationBridge, "codex_image_generation_bridge_enabled")
}
// CodexImageGenerationExplicitToolPolicy returns the account-level policy for
// client-provided Codex /responses image_generation tools. Unknown or unset
// values default to allow to preserve existing behavior.
func (a *Account) CodexImageGenerationExplicitToolPolicy() string {
if a == nil || a.Platform != PlatformOpenAI || a.Extra == nil {
return codexImageGenerationExplicitToolPolicyAllow
}
if policy, ok := stringOverrideFromMap(a.Extra, featureKeyCodexImageGenerationExplicitToolPolicy); ok {
return normalizeCodexImageGenerationExplicitToolPolicy(policy)
}
openaiConfig, _ := a.Extra[PlatformOpenAI].(map[string]any)
if policy, ok := stringOverrideFromMap(openaiConfig, featureKeyCodexImageGenerationExplicitToolPolicy); ok {
return normalizeCodexImageGenerationExplicitToolPolicy(policy)
}
return codexImageGenerationExplicitToolPolicyAllow
}
@@ -53,6 +53,12 @@ type ConcurrencyCache interface {
CleanupStaleProcessSlots(ctx context.Context, activeRequestPrefix string) error
}
type APIKeyConcurrencyCache interface {
TrackAPIKeySlot(ctx context.Context, apiKeyID int64, requestID string) error
ReleaseAPIKeySlot(ctx context.Context, apiKeyID int64, requestID string) error
GetAPIKeyConcurrencyBatch(ctx context.Context, apiKeyIDs []int64) (map[int64]int, error)
}
var (
requestIDPrefix = initRequestIDPrefix()
requestIDCounter atomic.Uint64
@@ -90,6 +96,8 @@ const (
defaultAccountLoadBatchCacheTTL = 200 * time.Millisecond
accountLoadBatchFetchTimeout = 3 * time.Second
maxAccountLoadBatchCacheEntries = 256
apiKeyConcurrencyFetchTimeout = 3 * time.Second
apiKeySlotTrackTimeout = 2 * time.Second
)
// ConcurrencyService 管理账号和用户的并发限制。
@@ -238,6 +246,77 @@ func (s *ConcurrencyService) AcquireUserSlot(ctx context.Context, userID int64,
}, nil
}
// TrackAPIKeySlot records one active request slot for an API key without
// applying key-level concurrency limits. It is fail-open: Redis errors are
// logged and return a no-op release function.
func (s *ConcurrencyService) TrackAPIKeySlot(ctx context.Context, apiKeyID int64) func() {
if s == nil || s.cache == nil || apiKeyID <= 0 {
return func() {}
}
cache, ok := s.cache.(APIKeyConcurrencyCache)
if !ok {
return func() {}
}
requestID := generateRequestID()
baseCtx := context.Background()
if ctx != nil {
baseCtx = context.WithoutCancel(ctx)
}
trackCtx, cancel := context.WithTimeout(baseCtx, apiKeySlotTrackTimeout)
err := cache.TrackAPIKeySlot(trackCtx, apiKeyID, requestID)
cancel()
if err != nil {
logger.LegacyPrintf("service.concurrency", "Warning: failed to track api key slot for %d (req=%s): %v", apiKeyID, requestID, err)
return func() {}
}
return func() {
bgCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
if err := cache.ReleaseAPIKeySlot(bgCtx, apiKeyID, requestID); err != nil {
logger.LegacyPrintf("service.concurrency", "Warning: failed to release api key slot for %d (req=%s): %v", apiKeyID, requestID, err)
}
}
}
// GetAPIKeyConcurrencyBatch gets real-time active request counts for API keys.
// Stats are best-effort: missing Redis support or Redis errors return zeroes.
func (s *ConcurrencyService) GetAPIKeyConcurrencyBatch(ctx context.Context, apiKeyIDs []int64) (map[int64]int, error) {
result := zeroAPIKeyConcurrencyMap(apiKeyIDs)
if len(apiKeyIDs) == 0 {
return result, nil
}
if s == nil || s.cache == nil {
return result, nil
}
cache, ok := s.cache.(APIKeyConcurrencyCache)
if !ok {
return result, nil
}
redisCtx, cancel := context.WithTimeout(context.Background(), apiKeyConcurrencyFetchTimeout)
defer cancel()
counts, err := cache.GetAPIKeyConcurrencyBatch(redisCtx, apiKeyIDs)
if err != nil {
logger.LegacyPrintf("service.concurrency", "Warning: get api key concurrency batch failed: %v", err)
return result, nil
}
for _, apiKeyID := range apiKeyIDs {
result[apiKeyID] = counts[apiKeyID]
}
return result, nil
}
func zeroAPIKeyConcurrencyMap(apiKeyIDs []int64) map[int64]int {
result := make(map[int64]int, len(apiKeyIDs))
for _, apiKeyID := range apiKeyIDs {
result[apiKeyID] = 0
}
return result
}
// ============================================
// Wait Queue Count Methods
// ============================================
@@ -16,25 +16,33 @@ import (
// stubConcurrencyCacheForTest 用于并发服务单元测试的缓存桩
type stubConcurrencyCacheForTest struct {
acquireResult bool
acquireErr error
releaseErr error
concurrency int
concurrencyErr error
waitAllowed bool
waitErr error
waitCount int
waitCountErr error
loadBatch map[int64]*AccountLoadInfo
loadBatchErr error
usersLoadBatch map[int64]*UserLoadInfo
usersLoadErr error
cleanupErr error
acquireResult bool
acquireErr error
releaseErr error
concurrency int
concurrencyErr error
waitAllowed bool
waitErr error
waitCount int
waitCountErr error
loadBatch map[int64]*AccountLoadInfo
loadBatchErr error
usersLoadBatch map[int64]*UserLoadInfo
usersLoadErr error
cleanupErr error
apiKeyTrackErr error
apiKeyReleaseErr error
apiKeyConcurrency map[int64]int
apiKeyConcurrencyErr error
// 记录调用
releasedAccountIDs []int64
releasedRequestIDs []string
loadBatchCalls atomic.Int64
releasedAccountIDs []int64
releasedRequestIDs []string
loadBatchCalls atomic.Int64
trackedAPIKeyIDs []int64
trackedAPIKeyRequestIDs []string
releasedAPIKeyIDs []int64
releasedAPIKeyRequestIDs []string
}
var _ ConcurrencyCache = (*stubConcurrencyCacheForTest)(nil)
@@ -78,6 +86,26 @@ func (c *stubConcurrencyCacheForTest) ReleaseUserSlot(_ context.Context, _ int64
func (c *stubConcurrencyCacheForTest) GetUserConcurrency(_ context.Context, _ int64) (int, error) {
return c.concurrency, c.concurrencyErr
}
func (c *stubConcurrencyCacheForTest) TrackAPIKeySlot(_ context.Context, apiKeyID int64, requestID string) error {
c.trackedAPIKeyIDs = append(c.trackedAPIKeyIDs, apiKeyID)
c.trackedAPIKeyRequestIDs = append(c.trackedAPIKeyRequestIDs, requestID)
return c.apiKeyTrackErr
}
func (c *stubConcurrencyCacheForTest) ReleaseAPIKeySlot(_ context.Context, apiKeyID int64, requestID string) error {
c.releasedAPIKeyIDs = append(c.releasedAPIKeyIDs, apiKeyID)
c.releasedAPIKeyRequestIDs = append(c.releasedAPIKeyRequestIDs, requestID)
return c.apiKeyReleaseErr
}
func (c *stubConcurrencyCacheForTest) GetAPIKeyConcurrencyBatch(_ context.Context, apiKeyIDs []int64) (map[int64]int, error) {
if c.apiKeyConcurrencyErr != nil {
return nil, c.apiKeyConcurrencyErr
}
result := make(map[int64]int, len(apiKeyIDs))
for _, apiKeyID := range apiKeyIDs {
result[apiKeyID] = c.apiKeyConcurrency[apiKeyID]
}
return result, nil
}
func (c *stubConcurrencyCacheForTest) IncrementWaitCount(_ context.Context, _ int64, _ int) (bool, error) {
return c.waitAllowed, c.waitErr
}
@@ -201,6 +229,62 @@ func TestAcquireUserSlot_UnlimitedConcurrency(t *testing.T) {
require.True(t, result.Acquired)
}
func TestTrackAPIKeySlot_ReleaseDecrements(t *testing.T) {
cache := &stubConcurrencyCacheForTest{}
svc := NewConcurrencyService(cache)
release := svc.TrackAPIKeySlot(context.Background(), 88)
require.NotNil(t, release)
require.Equal(t, []int64{88}, cache.trackedAPIKeyIDs)
require.Len(t, cache.trackedAPIKeyRequestIDs, 1)
require.NotEmpty(t, cache.trackedAPIKeyRequestIDs[0])
release()
require.Equal(t, []int64{88}, cache.releasedAPIKeyIDs)
require.Equal(t, cache.trackedAPIKeyRequestIDs, cache.releasedAPIKeyRequestIDs)
}
func TestTrackAPIKeySlot_FailOpen(t *testing.T) {
cache := &stubConcurrencyCacheForTest{apiKeyTrackErr: errors.New("redis down")}
svc := NewConcurrencyService(cache)
release := svc.TrackAPIKeySlot(context.Background(), 88)
require.NotNil(t, release)
require.Equal(t, []int64{88}, cache.trackedAPIKeyIDs)
require.NotPanics(t, release)
require.Empty(t, cache.releasedAPIKeyIDs)
}
func TestGetAPIKeyConcurrencyBatch_Fallbacks(t *testing.T) {
t.Run("nil cache returns zeroes", func(t *testing.T) {
svc := &ConcurrencyService{cache: nil}
counts, err := svc.GetAPIKeyConcurrencyBatch(context.Background(), []int64{1, 2})
require.NoError(t, err)
require.Equal(t, map[int64]int{1: 0, 2: 0}, counts)
})
t.Run("redis error returns zeroes", func(t *testing.T) {
cache := &stubConcurrencyCacheForTest{apiKeyConcurrencyErr: errors.New("redis down")}
svc := NewConcurrencyService(cache)
counts, err := svc.GetAPIKeyConcurrencyBatch(context.Background(), []int64{1, 2})
require.NoError(t, err)
require.Equal(t, map[int64]int{1: 0, 2: 0}, counts)
})
t.Run("success returns counts", func(t *testing.T) {
cache := &stubConcurrencyCacheForTest{apiKeyConcurrency: map[int64]int{1: 3, 2: 0}}
svc := NewConcurrencyService(cache)
counts, err := svc.GetAPIKeyConcurrencyBatch(context.Background(), []int64{1, 2})
require.NoError(t, err)
require.Equal(t, map[int64]int{1: 3, 2: 0}, counts)
})
}
func TestGenerateRequestID_UsesStablePrefixAndMonotonicCounter(t *testing.T) {
id1 := generateRequestID()
id2 := generateRequestID()
@@ -431,6 +431,20 @@ const (
// SettingKeyAllowUngroupedKeyScheduling 允许未分组 API Key 调度(默认 false:未分组 Key 返回 403)
SettingKeyAllowUngroupedKeyScheduling = "allow_ungrouped_key_scheduling"
// SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled OpenAI 高级调度下是否启用粘性加权。
SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled = "openai_advanced_scheduler_sticky_weighted_enabled"
// SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled OpenAI 高级调度下是否优先使用订阅账号池。
SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled = "openai_advanced_scheduler_subscription_priority_enabled"
SettingKeyOpenAIAdvancedSchedulerLBTopK = "openai_advanced_scheduler_lb_top_k"
SettingKeyOpenAIAdvancedSchedulerWeightPriority = "openai_advanced_scheduler_weight_priority"
SettingKeyOpenAIAdvancedSchedulerWeightLoad = "openai_advanced_scheduler_weight_load"
SettingKeyOpenAIAdvancedSchedulerWeightQueue = "openai_advanced_scheduler_weight_queue"
SettingKeyOpenAIAdvancedSchedulerWeightErrorRate = "openai_advanced_scheduler_weight_error_rate"
SettingKeyOpenAIAdvancedSchedulerWeightTTFT = "openai_advanced_scheduler_weight_ttft"
SettingKeyOpenAIAdvancedSchedulerWeightReset = "openai_advanced_scheduler_weight_reset"
SettingKeyOpenAIAdvancedSchedulerWeightQuotaHeadroom = "openai_advanced_scheduler_weight_quota_headroom"
SettingKeyOpenAIAdvancedSchedulerWeightPreviousResponse = "openai_advanced_scheduler_weight_previous_response"
SettingKeyOpenAIAdvancedSchedulerWeightSessionSticky = "openai_advanced_scheduler_weight_session_sticky"
// SettingKeyBackendModeEnabled Backend 模式:禁用用户注册和自助服务,仅管理员可登录
SettingKeyBackendModeEnabled = "backend_mode_enabled"
@@ -95,6 +95,9 @@ func (m *mockAccountRepoForPlatform) List(ctx context.Context, params pagination
func (m *mockAccountRepoForPlatform) ListWithFilters(ctx context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, *pagination.PaginationResult, error) {
return nil, nil, nil
}
func (m *mockAccountRepoForPlatform) ListAllWithFilters(ctx context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, error) {
return nil, nil
}
func (m *mockAccountRepoForPlatform) ListByGroup(ctx context.Context, groupID int64) ([]Account, error) {
return nil, nil
}
@@ -440,7 +440,9 @@ func TestGatewayServiceRecordUsage_GeneratesRequestIDWhenAllSourcesMissing(t *te
require.Equal(t, billingRepo.lastCmd.RequestID, usageRepo.lastLog.RequestID)
}
func TestGatewayServiceRecordUsage_DroppedUsageLogDoesNotSyncFallback(t *testing.T) {
func TestGatewayServiceRecordUsage_DroppedUsageLogFallsBackToSyncCreate(t *testing.T) {
// 计费成功后 best-effort 写入被丢弃(队列超时)时必须同步兜底,
// 否则出现“已扣费但无 usage_log”的对账缺口(issue #3656)。
usageRepo := &openAIRecordUsageBestEffortLogRepoStub{
bestEffortErr: MarkUsageLogCreateDropped(errors.New("usage log best-effort queue full")),
}
@@ -464,7 +466,9 @@ func TestGatewayServiceRecordUsage_DroppedUsageLogDoesNotSyncFallback(t *testing
require.NoError(t, err)
require.Equal(t, 1, usageRepo.bestEffortCalls)
require.Equal(t, 0, usageRepo.createCalls)
require.Equal(t, 1, usageRepo.createCalls)
// 兜底调用使用的 ctx 必须仍然存活,不能带着已死的 ctx 走过场。
require.NoError(t, usageRepo.lastCtxErr)
}
func TestGatewayServiceRecordUsage_BillingErrorSkipsUsageLogWrite(t *testing.T) {
+42 -3
View File
@@ -5920,6 +5920,10 @@ func (s *GatewayService) buildUpstreamRequestAnthropicAPIKeyPassthrough(
if c != nil && c.Request != nil {
clientBeta = getHeaderRaw(c.Request.Header, "anthropic-beta")
}
// 账号覆写了 anthropic-beta 时,覆写值即最终上游值:净化以覆写值为准
if beta, ok := account.HeaderOverrideValue("anthropic-beta"); ok {
clientBeta = beta
}
if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(body, clientBeta); changed {
body = sanitized
}
@@ -5956,6 +5960,9 @@ func (s *GatewayService) buildUpstreamRequestAnthropicAPIKeyPassthrough(
setHeaderRaw(req.Header, "anthropic-version", "2023-06-01")
}
// 账号级请求头覆写(最终生效,覆盖上面所有来源的同名头)
account.ApplyHeaderOverrides(req.Header)
return req, body, nil
}
@@ -6886,6 +6893,12 @@ func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Contex
tokenType, mimicClaudeCode, modelID, clientHeaders, body, effectiveDropSet,
)
// 账号覆写了 anthropic-beta 时,覆写值即最终上游值(由下方 ApplyHeaderOverrides 写入):
// body 能力净化必须以覆写值为准,否则 header/body 不对称会被上游 400。
if beta, ok := account.HeaderOverrideValue("anthropic-beta"); ok {
finalBetaHeader, finalBetaShouldSet = beta, true
}
// 能力维度 body sanitize:与最终 anthropic-beta header 对称
if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(body, finalBetaHeader); changed {
body = sanitized
@@ -6959,6 +6972,10 @@ func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Contex
}
}
// 账号级请求头覆写(仅 anthropic/openai api_key 账号启用时生效;OAuth 路径 no-op)。
// 放在所有 header 逻辑之后,确保配置值对同名头拥有最终决定权。
account.ApplyHeaderOverrides(req.Header)
// === DEBUG: 打印上游转发请求(headers + body 摘要),与 CLIENT_ORIGINAL 对比 ===
s.debugLogGatewaySnapshot("UPSTREAM_FORWARD", req.Header, body, map[string]string{
"url": req.URL.String(),
@@ -9473,10 +9490,17 @@ func writeUsageLogBestEffort(ctx context.Context, repo UsageLogRepository, usage
if writer, ok := repo.(usageLogBestEffortWriter); ok {
if err := writer.CreateBestEffort(usageCtx, usageLog); err != nil {
logger.LegacyPrintf(logKey, "Create usage log failed: %v", err)
if IsUsageLogCreateDropped(err) {
return
// 计费已在此前完成,日志必须落库:dropped(批处理队列超时)同样走同步兜底,
// 否则会出现“已扣费但无 usage_log”的对账缺口(issue #3656)。
// 重复写入由 usage_logs 的 ON CONFLICT (request_id, api_key_id) DO NOTHING 防护。
fallbackCtx := usageCtx
if usageCtx.Err() != nil {
// usageCtx 已耗尽(best-effort 入队阻塞到期限):换新的 detached 窗口,避免兜底必然失败。
var fallbackCancel context.CancelFunc
fallbackCtx, fallbackCancel = detachedBillingContext(context.Background())
defer fallbackCancel()
}
if _, syncErr := repo.Create(usageCtx, usageLog); syncErr != nil {
if _, syncErr := repo.Create(fallbackCtx, usageLog); syncErr != nil {
logger.LegacyPrintf(logKey, "Create usage log sync fallback failed: %v", syncErr)
}
}
@@ -10403,6 +10427,10 @@ func (s *GatewayService) buildCountTokensRequestAnthropicAPIKeyPassthrough(
if c != nil && c.Request != nil {
clientBeta = getHeaderRaw(c.Request.Header, "anthropic-beta")
}
// 账号覆写了 anthropic-beta 时,覆写值即最终上游值:净化以覆写值为准
if beta, ok := account.HeaderOverrideValue("anthropic-beta"); ok {
clientBeta = beta
}
if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(body, clientBeta); changed {
body = sanitized
}
@@ -10438,6 +10466,9 @@ func (s *GatewayService) buildCountTokensRequestAnthropicAPIKeyPassthrough(
req.Header.Set("anthropic-version", "2023-06-01")
}
// 账号级请求头覆写(最终生效,覆盖上面所有来源的同名头)
account.ApplyHeaderOverrides(req.Header)
return req, nil
}
@@ -10505,6 +10536,11 @@ func (s *GatewayService) buildCountTokensRequest(ctx context.Context, c *gin.Con
tokenType, mimicClaudeCode, modelID, clientHeaders, body, ctEffectiveDropSet,
)
// 账号覆写了 anthropic-beta 时,覆写值即最终上游值:净化以覆写值为准
if beta, ok := account.HeaderOverrideValue("anthropic-beta"); ok {
finalBetaHeader, finalBetaShouldSet = beta, true
}
// 能力维度 body sanitize:与最终 anthropic-beta header 对称
if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(body, finalBetaHeader); changed {
body = sanitized
@@ -10571,6 +10607,9 @@ func (s *GatewayService) buildCountTokensRequest(ctx context.Context, c *gin.Con
}
}
// 账号级请求头覆写(仅 anthropic/openai api_key 账号启用时生效;OAuth 路径 no-op)
account.ApplyHeaderOverrides(req.Header)
if c != nil && tokenType == "oauth" {
c.Set(claudeMimicDebugInfoKey, buildClaudeMimicDebugLine(req, body, account, tokenType, mimicClaudeCode))
}
@@ -82,6 +82,9 @@ func (m *mockAccountRepoForGemini) List(ctx context.Context, params pagination.P
func (m *mockAccountRepoForGemini) ListWithFilters(ctx context.Context, params pagination.PaginationParams, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, *pagination.PaginationResult, error) {
return nil, nil, nil
}
func (m *mockAccountRepoForGemini) ListAllWithFilters(ctx context.Context, platform, accountType, status, search string, groupID int64, privacyMode string) ([]Account, error) {
return nil, nil
}
func (m *mockAccountRepoForGemini) ListByGroup(ctx context.Context, groupID int64) ([]Account, error) {
return nil, nil
}
+1
View File
@@ -483,6 +483,7 @@ func grokMediaUsageFromResponse(endpoint GrokMediaEndpoint, requestInfo GrokMedi
meta.ImageOutputSizes = collectOpenAIResponseImageOutputSizesFromJSONBytes(responseBody)
case GrokMediaEndpointVideosGenerations:
meta.ResponseID = extractGrokMediaVideoRequestID(responseBody)
// Video generation is one billable media unit; the legacy usage schema stores it in ImageCount.
meta.ImageCount = 1
meta.ImageSize = requestInfo.SizeTier
meta.ImageInputSize = requestInfo.Size
@@ -16,6 +16,26 @@ type GroupCapacitySummary struct {
RPMMax int `json:"rpm_max"`
}
// GroupAccountCapacityRow is the lightweight account projection needed for
// capacity summary aggregation.
type GroupAccountCapacityRow struct {
GroupID int64
AccountID int64
Concurrency int
Extra map[string]any
SessionWindowStart *time.Time
SessionWindowEnd *time.Time
SessionWindowStatus string
}
type groupCapacityActiveGroupIDLister interface {
ListActiveIDs(ctx context.Context) ([]int64, error)
}
type groupCapacityAccountLister interface {
ListSchedulableCapacityByGroupIDs(ctx context.Context, groupIDs []int64) ([]GroupAccountCapacityRow, error)
}
// GroupCapacityService aggregates per-group capacity from runtime data.
type GroupCapacityService struct {
accountRepo AccountRepository
@@ -44,24 +64,176 @@ func NewGroupCapacityService(
// GetAllGroupCapacity returns capacity summary for all active groups.
func (s *GroupCapacityService) GetAllGroupCapacity(ctx context.Context) ([]GroupCapacitySummary, error) {
groups, err := s.groupRepo.ListActive(ctx)
groupIDs, err := s.listActiveGroupIDs(ctx)
if err != nil {
return nil, err
}
results := make([]GroupCapacitySummary, 0, len(groups))
if lister, ok := s.accountRepo.(groupCapacityAccountLister); ok {
return s.getGroupCapacitiesBatch(ctx, groupIDs, lister)
}
return s.getGroupCapacitiesSequential(ctx, groupIDs), nil
}
func (s *GroupCapacityService) listActiveGroupIDs(ctx context.Context) ([]int64, error) {
if lister, ok := s.groupRepo.(groupCapacityActiveGroupIDLister); ok {
return lister.ListActiveIDs(ctx)
}
groups, err := s.groupRepo.ListActive(ctx)
if err != nil {
return nil, err
}
groupIDs := make([]int64, 0, len(groups))
for i := range groups {
cap, err := s.getGroupCapacity(ctx, groups[i].ID)
groupIDs = append(groupIDs, groups[i].ID)
}
return groupIDs, nil
}
func (s *GroupCapacityService) getGroupCapacitiesSequential(ctx context.Context, groupIDs []int64) []GroupCapacitySummary {
results := make([]GroupCapacitySummary, 0, len(groupIDs))
for _, groupID := range groupIDs {
cap, err := s.getGroupCapacity(ctx, groupID)
if err != nil {
// Skip groups with errors, return partial results
continue
}
cap.GroupID = groups[i].ID
cap.GroupID = groupID
results = append(results, cap)
}
return results
}
type groupCapacityAccountRef struct {
groupID int64
accountID int64
}
func (s *GroupCapacityService) getGroupCapacitiesBatch(ctx context.Context, groupIDs []int64, lister groupCapacityAccountLister) ([]GroupCapacitySummary, error) {
results := make([]GroupCapacitySummary, len(groupIDs))
groupIndex := make(map[int64]int, len(groupIDs))
for i, groupID := range groupIDs {
results[i].GroupID = groupID
groupIndex[groupID] = i
}
if len(groupIDs) == 0 {
return results, nil
}
rows, err := lister.ListSchedulableCapacityByGroupIDs(ctx, groupIDs)
if err != nil {
return nil, err
}
if len(rows) == 0 {
return results, nil
}
refs := make([]groupCapacityAccountRef, 0, len(rows))
seenGroupAccount := make(map[groupCapacityAccountRef]struct{}, len(rows))
accountIDSet := make(map[int64]struct{}, len(rows))
accountIDs := make([]int64, 0, len(rows))
sessionTimeouts := make(map[int64]time.Duration)
for _, row := range rows {
idx, ok := groupIndex[row.GroupID]
if !ok || row.AccountID <= 0 {
continue
}
ref := groupCapacityAccountRef{groupID: row.GroupID, accountID: row.AccountID}
if _, ok := seenGroupAccount[ref]; ok {
continue
}
seenGroupAccount[ref] = struct{}{}
refs = append(refs, ref)
if _, ok := accountIDSet[row.AccountID]; !ok {
accountIDSet[row.AccountID] = struct{}{}
accountIDs = append(accountIDs, row.AccountID)
}
acc := Account{
ID: row.AccountID,
Concurrency: row.Concurrency,
Extra: row.Extra,
SessionWindowStart: row.SessionWindowStart,
SessionWindowEnd: row.SessionWindowEnd,
SessionWindowStatus: row.SessionWindowStatus,
}
results[idx].ConcurrencyMax += acc.Concurrency
if maxSessions := acc.GetMaxSessions(); maxSessions > 0 {
results[idx].SessionsMax += maxSessions
timeout := time.Duration(acc.GetSessionIdleTimeoutMinutes()) * time.Minute
if timeout <= 0 {
timeout = 5 * time.Minute
}
sessionTimeouts[acc.ID] = timeout
}
if rpm := acc.GetBaseRPM(); rpm > 0 {
results[idx].RPMMax += rpm
}
}
if len(accountIDs) == 0 {
return results, nil
}
concurrencyMap := map[int64]int{}
if s.concurrencyService != nil {
concurrencyMap, _ = s.concurrencyService.GetAccountConcurrencyBatch(ctx, accountIDs)
}
sessionAccountIDs := accountIDsForGroupsWithLimit(refs, groupIndex, results, func(summary GroupCapacitySummary) bool {
return summary.SessionsMax > 0
})
var sessionsMap map[int64]int
if len(sessionAccountIDs) > 0 && s.sessionLimitCache != nil {
sessionsMap, _ = s.sessionLimitCache.GetActiveSessionCountBatch(ctx, sessionAccountIDs, sessionTimeouts)
}
rpmAccountIDs := accountIDsForGroupsWithLimit(refs, groupIndex, results, func(summary GroupCapacitySummary) bool {
return summary.RPMMax > 0
})
var rpmMap map[int64]int
if len(rpmAccountIDs) > 0 && s.rpmCache != nil {
rpmMap, _ = s.rpmCache.GetRPMBatch(ctx, rpmAccountIDs)
}
for _, ref := range refs {
idx := groupIndex[ref.groupID]
results[idx].ConcurrencyUsed += concurrencyMap[ref.accountID]
if sessionsMap != nil && results[idx].SessionsMax > 0 {
results[idx].SessionsUsed += sessionsMap[ref.accountID]
}
if rpmMap != nil && results[idx].RPMMax > 0 {
results[idx].RPMUsed += rpmMap[ref.accountID]
}
}
return results, nil
}
func accountIDsForGroupsWithLimit(refs []groupCapacityAccountRef, groupIndex map[int64]int, summaries []GroupCapacitySummary, include func(GroupCapacitySummary) bool) []int64 {
seen := make(map[int64]struct{})
accountIDs := make([]int64, 0)
for _, ref := range refs {
idx, ok := groupIndex[ref.groupID]
if !ok || !include(summaries[idx]) {
continue
}
if _, ok := seen[ref.accountID]; ok {
continue
}
seen[ref.accountID] = struct{}{}
accountIDs = append(accountIDs, ref.accountID)
}
return accountIDs
}
func (s *GroupCapacityService) getGroupCapacity(ctx context.Context, groupID int64) (GroupCapacitySummary, error) {
accounts, err := s.accountRepo.ListSchedulableByGroupID(ctx, groupID)
if err != nil {
@@ -0,0 +1,179 @@
package service
import (
"context"
"testing"
"time"
"github.com/stretchr/testify/require"
)
type groupCapacityAccountRepoStub struct {
AccountRepository
rows []GroupAccountCapacityRow
requested []int64
}
func (s *groupCapacityAccountRepoStub) ListSchedulableCapacityByGroupIDs(_ context.Context, groupIDs []int64) ([]GroupAccountCapacityRow, error) {
s.requested = append([]int64(nil), groupIDs...)
return append([]GroupAccountCapacityRow(nil), s.rows...), nil
}
type groupCapacityGroupRepoStub struct {
GroupRepository
groupIDs []int64
listCalls int
}
func (s *groupCapacityGroupRepoStub) ListActiveIDs(context.Context) ([]int64, error) {
s.listCalls++
return append([]int64(nil), s.groupIDs...), nil
}
type groupCapacityConcurrencyCacheStub struct {
ConcurrencyCache
counts map[int64]int
requested []int64
}
func (s *groupCapacityConcurrencyCacheStub) GetAccountConcurrencyBatch(_ context.Context, accountIDs []int64) (map[int64]int, error) {
s.requested = append([]int64(nil), accountIDs...)
out := make(map[int64]int, len(accountIDs))
for _, id := range accountIDs {
out[id] = s.counts[id]
}
return out, nil
}
type groupCapacitySessionCacheStub struct {
SessionLimitCache
counts map[int64]int
requested []int64
idleTimeouts map[int64]time.Duration
}
func (s *groupCapacitySessionCacheStub) GetActiveSessionCountBatch(_ context.Context, accountIDs []int64, idleTimeouts map[int64]time.Duration) (map[int64]int, error) {
s.requested = append([]int64(nil), accountIDs...)
s.idleTimeouts = make(map[int64]time.Duration, len(idleTimeouts))
for id, timeout := range idleTimeouts {
s.idleTimeouts[id] = timeout
}
out := make(map[int64]int, len(accountIDs))
for _, id := range accountIDs {
out[id] = s.counts[id]
}
return out, nil
}
type groupCapacityRPMCacheStub struct {
RPMCache
counts map[int64]int
requested []int64
}
func (s *groupCapacityRPMCacheStub) GetRPMBatch(_ context.Context, accountIDs []int64) (map[int64]int, error) {
s.requested = append([]int64(nil), accountIDs...)
out := make(map[int64]int, len(accountIDs))
for _, id := range accountIDs {
out[id] = s.counts[id]
}
return out, nil
}
func TestGetAllGroupCapacityBatchAggregatesRuntimeAndLimits(t *testing.T) {
accountRepo := &groupCapacityAccountRepoStub{
rows: []GroupAccountCapacityRow{
{
GroupID: 10,
AccountID: 1,
Concurrency: 2,
Extra: map[string]any{
"max_sessions": 3,
"session_idle_timeout_minutes": 7,
"base_rpm": 11,
},
},
{
GroupID: 20,
AccountID: 1,
Concurrency: 2,
Extra: map[string]any{
"max_sessions": 3,
"session_idle_timeout_minutes": 7,
"base_rpm": 11,
},
},
{
GroupID: 20,
AccountID: 2,
Concurrency: 4,
Extra: map[string]any{
"max_sessions": 1,
"session_idle_timeout_minutes": 9,
"base_rpm": 13,
},
},
},
}
groupRepo := &groupCapacityGroupRepoStub{groupIDs: []int64{10, 20}}
concurrencyCache := &groupCapacityConcurrencyCacheStub{counts: map[int64]int{1: 1, 2: 2}}
sessionCache := &groupCapacitySessionCacheStub{counts: map[int64]int{1: 2, 2: 1}}
rpmCache := &groupCapacityRPMCacheStub{counts: map[int64]int{1: 5, 2: 7}}
svc := NewGroupCapacityService(
accountRepo,
groupRepo,
NewConcurrencyService(concurrencyCache),
sessionCache,
rpmCache,
)
results, err := svc.GetAllGroupCapacity(context.Background())
require.NoError(t, err)
require.Equal(t, 1, groupRepo.listCalls)
require.Equal(t, []int64{10, 20}, accountRepo.requested)
require.Equal(t, []int64{1, 2}, concurrencyCache.requested)
require.ElementsMatch(t, []int64{1, 2}, sessionCache.requested)
require.ElementsMatch(t, []int64{1, 2}, rpmCache.requested)
require.Equal(t, 7*time.Minute, sessionCache.idleTimeouts[1])
require.Equal(t, 9*time.Minute, sessionCache.idleTimeouts[2])
require.Equal(t, []GroupCapacitySummary{
{
GroupID: 10,
ConcurrencyUsed: 1,
ConcurrencyMax: 2,
SessionsUsed: 2,
SessionsMax: 3,
RPMUsed: 5,
RPMMax: 11,
},
{
GroupID: 20,
ConcurrencyUsed: 3,
ConcurrencyMax: 6,
SessionsUsed: 3,
SessionsMax: 4,
RPMUsed: 12,
RPMMax: 24,
},
}, results)
}
func TestGetAllGroupCapacityBatchKeepsEmptyGroupRows(t *testing.T) {
accountRepo := &groupCapacityAccountRepoStub{
rows: []GroupAccountCapacityRow{
{GroupID: 20, AccountID: 2, Concurrency: 4},
},
}
groupRepo := &groupCapacityGroupRepoStub{groupIDs: []int64{10, 20}}
svc := NewGroupCapacityService(accountRepo, groupRepo, nil, nil, nil)
results, err := svc.GetAllGroupCapacity(context.Background())
require.NoError(t, err)
require.Equal(t, []GroupCapacitySummary{
{GroupID: 10},
{GroupID: 20, ConcurrencyMax: 4},
}, results)
}
@@ -12,6 +12,9 @@ const (
modelRateLimitsKey = "model_rate_limits"
antigravityGeminiModelRateLimitKey = "antigravity:gemini"
openAIImageGenerationRateLimitKey = "openai:image_generation"
// anthropicFableRateLimitKey 是 Anthropic 7d_oi(Fable 专属 7d 窗口)限流的
// 家族级 scope:命中后所有 Fable 变体(含 [1m] 等后缀)都不再调度到该账号。
anthropicFableRateLimitKey = "claude-fable-5"
)
// isRateLimitActiveForKey 检查指定 key 的限流是否生效
@@ -82,10 +85,19 @@ func (a *Account) modelRateLimitKeysForRequest(ctx context.Context, requestedMod
if openAIImageGenerationRateLimitApplies(ctx, requestedModel, modelKey) && modelKey != openAIImageGenerationRateLimitKey {
keys = append(keys, openAIImageGenerationRateLimitKey)
}
case PlatformAnthropic:
if isAnthropicFableModel(modelKey) && modelKey != anthropicFableRateLimitKey {
keys = append(keys, anthropicFableRateLimitKey)
}
}
return keys
}
// isAnthropicFableModel 判断是否为 Fable 模型家族(claude-fable-5、claude-fable-5[1m] 等变体)
func isAnthropicFableModel(model string) bool {
return strings.Contains(strings.ToLower(model), "fable")
}
func openAIImageGenerationRateLimitApplies(ctx context.Context, requestedModel, modelKey string) bool {
if isOpenAIImageGenerationModel(requestedModel) || isOpenAIImageGenerationModel(modelKey) {
return true
@@ -499,3 +499,47 @@ func TestGetRateLimitRemainingTime(t *testing.T) {
})
}
}
func TestIsModelRateLimited_AnthropicFableFamilyKey(t *testing.T) {
now := time.Now()
future := now.Add(48 * time.Hour).Format(time.RFC3339)
account := &Account{
Platform: PlatformAnthropic,
Extra: map[string]any{
modelRateLimitsKey: map[string]any{
anthropicFableRateLimitKey: map[string]any{
"rate_limit_reset_at": future,
},
},
},
}
tests := []struct {
requestedModel string
expected bool
}{
{"claude-fable-5", true},
{"claude-fable-5[1m]", true}, // 家族 key 覆盖变体
{"Claude-Fable-5-20260601", true}, // 大小写不敏感
{"claude-sonnet-4-6", false}, // 其他模型不受影响
{"claude-opus-4-8", false},
}
for _, tc := range tests {
t.Run(tc.requestedModel, func(t *testing.T) {
got := account.isModelRateLimitedWithContext(context.Background(), tc.requestedModel)
require.Equal(t, tc.expected, got)
remaining := account.GetModelRateLimitRemainingTimeWithContext(context.Background(), tc.requestedModel)
require.Equal(t, tc.expected, remaining > 0)
})
}
}
func TestIsAnthropicFableModel(t *testing.T) {
require.True(t, isAnthropicFableModel("claude-fable-5"))
require.True(t, isAnthropicFableModel("claude-fable-5[1m]"))
require.True(t, isAnthropicFableModel("Claude-Fable-5"))
require.False(t, isAnthropicFableModel("claude-sonnet-4-6"))
require.False(t, isAnthropicFableModel(""))
}
@@ -36,8 +36,20 @@ const (
)
type cachedOpenAIAdvancedSchedulerSetting struct {
enabled bool
expiresAt int64
enabled bool
stickyWeightedEnabled bool
subscriptionPriorityEnabled bool
lbTopKOverride int
weightOverrides map[string]float64
expiresAt int64
}
type openAIAdvancedSchedulerRuntimeSettings struct {
enabled bool
stickyWeightedEnabled bool
subscriptionPriorityEnabled bool
lbTopKOverride int
weightOverrides map[string]float64
}
var openAIAdvancedSchedulerSettingCache atomic.Value // *cachedOpenAIAdvancedSchedulerSetting
@@ -48,8 +60,12 @@ type OpenAIAccountScheduleRequest struct {
Platform string
SessionHash string
StickyAccountID int64
StickyPreviousAccountID int64
StickyWeighted bool
SubscriptionPriority bool
PreserveStickyBinding bool
PreviousResponseID string
PreviousResponseCanMove bool
RequestedModel string
RequiredTransport OpenAIUpstreamTransport
RequiredCapability OpenAIEndpointCapability
@@ -111,6 +127,17 @@ type openAIAccountLoadPlan struct {
loadSkew float64
}
type openAIAccountLoadSelectionAttempt struct {
result *AccountSelectionResult
selectionOrder []openAIAccountCandidateScore
candidateCount int
topK int
loadSkew float64
compactBlocked bool
noCompactCandidates bool
err error
}
func (m *openAIAccountSchedulerMetrics) recordSelect(decision OpenAIAccountScheduleDecision) {
if m == nil {
return
@@ -277,7 +304,8 @@ func (s *defaultOpenAIAccountScheduler) Select(
}()
previousResponseID := strings.TrimSpace(req.PreviousResponseID)
if previousResponseID != "" && normalizeOpenAICompatiblePlatform(req.Platform) == PlatformOpenAI {
if previousResponseID != "" && normalizeOpenAICompatiblePlatform(req.Platform) == PlatformOpenAI &&
(!req.StickyWeighted || !req.PreviousResponseCanMove) {
selection, err := s.service.selectAccountByPreviousResponseIDForCapability(
ctx,
req.GroupID,
@@ -310,19 +338,21 @@ func (s *defaultOpenAIAccountScheduler) Select(
}
}
selection, escapedSticky, err := s.selectBySessionHash(ctx, req)
if err != nil {
return nil, decision, err
}
if selection != nil && selection.Account != nil {
decision.Layer = openAIAccountScheduleLayerSessionSticky
decision.StickySessionHit = true
decision.SelectedAccountID = selection.Account.ID
decision.SelectedAccountType = selection.Account.Type
return selection, decision, nil
}
if escapedSticky {
req.PreserveStickyBinding = true
if !req.StickyWeighted {
selection, escapedSticky, err := s.selectBySessionHash(ctx, req)
if err != nil {
return nil, decision, err
}
if selection != nil && selection.Account != nil {
decision.Layer = openAIAccountScheduleLayerSessionSticky
decision.StickySessionHit = true
decision.SelectedAccountID = selection.Account.ID
decision.SelectedAccountType = selection.Account.Type
return selection, decision, nil
}
if escapedSticky {
req.PreserveStickyBinding = true
}
}
selection, candidateCount, topK, loadSkew, err := s.selectByLoadBalance(ctx, req)
@@ -336,6 +366,14 @@ func (s *defaultOpenAIAccountScheduler) Select(
if selection != nil && selection.Account != nil {
decision.SelectedAccountID = selection.Account.ID
decision.SelectedAccountType = selection.Account.Type
if req.StickyWeighted {
if req.StickyPreviousAccountID > 0 && selection.Account.ID == req.StickyPreviousAccountID {
decision.StickyPreviousHit = true
}
if req.StickyAccountID > 0 && selection.Account.ID == req.StickyAccountID {
decision.StickySessionHit = true
}
}
}
return selection, decision, nil
}
@@ -453,6 +491,13 @@ func openAIStickyAccountMatchesGroup(account *Account, groupID *int64) bool {
return false
}
func openAIAccountSchedulingPriority(account *Account) int {
if account == nil {
return 0
}
return account.Priority
}
func (s *defaultOpenAIAccountScheduler) shouldEscapeStickyAccount(accountID int64, cfg openAIStickyEscapeConfig) (reason string, errorRate float64, ttft float64, shouldEscape bool) {
if !cfg.enabled || s == nil || s.stats == nil || accountID <= 0 {
return "", 0, 0, false
@@ -471,6 +516,7 @@ type openAIAccountCandidateScore struct {
account *Account
loadInfo *AccountLoadInfo
score float64
priority int
errorRate float64
ttft float64
hasTTFT bool
@@ -669,6 +715,7 @@ func buildOpenAIWeightedSelectionOrder(
}
func (s *defaultOpenAIAccountScheduler) buildOpenAIAccountLoadPlan(
ctx context.Context,
req OpenAIAccountScheduleRequest,
filtered []*Account,
loadMap map[int64]*AccountLoadInfo,
@@ -716,18 +763,20 @@ func (s *defaultOpenAIAccountScheduler) buildOpenAIAccountLoadPlan(
return plan
}
minPriority, maxPriority := candidates[0].account.Priority, candidates[0].account.Priority
minPriority, maxPriority := openAIAccountSchedulingPriority(candidates[0].account), openAIAccountSchedulingPriority(candidates[0].account)
maxWaiting := 1
loadRateSum := 0.0
loadRateSumSquares := 0.0
minTTFT, maxTTFT := 0.0, 0.0
hasTTFTSample := false
for _, candidate := range candidates {
if candidate.account.Priority < minPriority {
minPriority = candidate.account.Priority
for i := range candidates {
candidate := &candidates[i]
candidate.priority = openAIAccountSchedulingPriority(candidate.account)
if candidate.priority < minPriority {
minPriority = candidate.priority
}
if candidate.account.Priority > maxPriority {
maxPriority = candidate.account.Priority
if candidate.priority > maxPriority {
maxPriority = candidate.priority
}
if candidate.loadInfo.WaitingCount > maxWaiting {
maxWaiting = candidate.loadInfo.WaitingCount
@@ -751,7 +800,7 @@ func (s *defaultOpenAIAccountScheduler) buildOpenAIAccountLoadPlan(
}
plan.loadSkew = calcLoadSkewByMoments(loadRateSum, loadRateSumSquares, len(candidates))
weights := s.service.openAIWSSchedulerWeights()
weights := s.service.openAIWSSchedulerWeightsForRequest(ctx)
// Reset 因子(use-it-or-lose-it):在拥有「未来会话窗口结束时间」的账号中,
// 剩余时间越短 → 因子越接近 1(越早重置越优先用尽)。无活跃窗口的账号因子为 0。
@@ -785,7 +834,7 @@ func (s *defaultOpenAIAccountScheduler) buildOpenAIAccountLoadPlan(
item := &candidates[i]
priorityFactor := 1.0
if maxPriority > minPriority {
priorityFactor = 1 - float64(item.account.Priority-minPriority)/float64(maxPriority-minPriority)
priorityFactor = 1 - float64(item.priority-minPriority)/float64(maxPriority-minPriority)
}
loadFactor := 1 - clamp01(float64(item.loadInfo.LoadRate)/100.0)
queueFactor := 1 - clamp01(float64(item.loadInfo.WaitingCount)/float64(maxWaiting))
@@ -817,10 +866,18 @@ func (s *defaultOpenAIAccountScheduler) buildOpenAIAccountLoadPlan(
weights.TTFT*ttftFactor +
weights.Reset*resetFactor +
weights.QuotaHeadroom*quotaHeadroomFactor
if req.StickyWeighted {
if req.PreviousResponseCanMove && req.StickyPreviousAccountID > 0 && item.account.ID == req.StickyPreviousAccountID {
item.score += weights.Previous
}
if req.StickyAccountID > 0 && item.account.ID == req.StickyAccountID {
item.score += weights.SessionSticky
}
}
}
plan.candidates = candidates
plan.topK = s.service.openAIWSLBTopK()
plan.topK = s.service.openAIWSLBTopKForRequest(ctx)
if plan.topK > len(candidates) {
plan.topK = len(candidates)
}
@@ -845,6 +902,20 @@ func (s *defaultOpenAIAccountScheduler) buildOpenAISelectionOrder(
groupTopK = len(pool)
}
ranked := selectTopKOpenAICandidates(pool, groupTopK)
if req.StickyWeighted {
for _, stickyID := range []int64{req.StickyPreviousAccountID, req.StickyAccountID} {
if stickyID <= 0 {
continue
}
for i, candidate := range ranked {
if candidate.account != nil && candidate.account.ID == stickyID {
ordered := append([]openAIAccountCandidateScore{candidate}, ranked[:i]...)
ordered = append(ordered, ranked[i+1:]...)
return ordered
}
}
}
}
return buildOpenAIWeightedSelectionOrder(ranked, req)
}
@@ -939,6 +1010,74 @@ func (s *defaultOpenAIAccountScheduler) tryAcquireOpenAISelectionOrder(
return nil, compactBlocked, nil
}
func (s *defaultOpenAIAccountScheduler) tryFallbackToWeightedSticky(
ctx context.Context,
req OpenAIAccountScheduleRequest,
) (*AccountSelectionResult, error) {
if !req.StickyWeighted {
return nil, nil
}
for _, accountID := range []int64{req.StickyPreviousAccountID, req.StickyAccountID} {
if accountID <= 0 {
continue
}
if req.ExcludedIDs != nil {
if _, excluded := req.ExcludedIDs[accountID]; excluded {
continue
}
}
account, err := s.service.getSchedulableAccount(ctx, accountID)
if err != nil || account == nil {
continue
}
if !s.isAccountRequestCompatible(ctx, account, req) || !s.isAccountTransportCompatible(account, req.RequiredTransport) {
continue
}
account = s.service.recheckSelectedOpenAIAccountFromDB(ctx, account, req.Platform, req.RequestedModel, req.RequireCompact, req.RequiredCapability)
if account == nil || !s.isAccountRequestCompatible(ctx, account, req) || !s.isAccountTransportCompatible(account, req.RequiredTransport) {
continue
}
// 粘性绑定只证明绑定时账号在分组内;账号被移出分组后绑定仍会在 TTL 内存活,
// 必须与 selectBySessionHash 一样重验分组归属,否则会把分组流量泄漏到组外账号。
if !openAIStickyAccountMatchesGroup(account, req.GroupID) {
if accountID == req.StickyAccountID && strings.TrimSpace(req.SessionHash) != "" {
_ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, req.SessionHash)
}
continue
}
if req.RequireCompact && openAICompactSupportTier(account) == 0 {
continue
}
result, acquireErr := s.service.tryAcquireAccountSlot(ctx, account.ID, account.Concurrency)
if acquireErr != nil {
return nil, acquireErr
}
if result != nil && result.Acquired {
if req.SessionHash != "" && !req.PreserveStickyBinding {
_ = s.service.BindStickySession(ctx, req.GroupID, req.SessionHash, account.ID)
}
return &AccountSelectionResult{
Account: account,
Acquired: true,
ReleaseFunc: result.ReleaseFunc,
}, nil
}
if s.service.concurrencyService != nil {
cfg := s.service.schedulingConfig()
return &AccountSelectionResult{
Account: account,
WaitPlan: &AccountWaitPlan{
AccountID: account.ID,
MaxConcurrency: account.Concurrency,
Timeout: cfg.StickySessionWaitTimeout,
MaxWaiting: cfg.StickySessionMaxWaiting,
},
}, nil
}
}
return nil, nil
}
func (s *defaultOpenAIAccountScheduler) selectByLoadBalance(
ctx context.Context,
req OpenAIAccountScheduleRequest,
@@ -1002,52 +1141,175 @@ func (s *defaultOpenAIAccountScheduler) selectByLoadBalance(
}
}
plan := s.buildOpenAIAccountLoadPlan(req, filtered, loadMap)
candidateCount := plan.candidateCount
topK := plan.topK
loadSkew := plan.loadSkew
selectionOrder := plan.selectionOrder
if req.RequireCompact && len(plan.candidates) == 0 && len(plan.staleSnapshotCompactRetry) == 0 {
return nil, 0, 0, 0, ErrNoAvailableCompactAccounts
}
if req.RequireCompact && len(selectionOrder) == 0 && s.service.schedulerSnapshot == nil {
return nil, candidateCount, topK, loadSkew, ErrNoAvailableCompactAccounts
}
if len(selectionOrder) == 0 {
return nil, candidateCount, topK, loadSkew, noAvailableOpenAISelectionError(req.RequestedModel, req.RequireCompact && len(plan.allCandidates) > 0)
if req.SubscriptionPriority {
subscriptionAccounts, regularAccounts := partitionOpenAIChatGPTSubscriptionAccounts(filtered)
if len(subscriptionAccounts) > 0 {
attempt := s.trySelectByLoadBalancePool(ctx, req, subscriptionAccounts, loadMap)
if attempt.err != nil && (!attempt.noCompactCandidates || len(regularAccounts) <= 0) {
return nil, attempt.candidateCount, attempt.topK, attempt.loadSkew, attempt.err
}
if attempt.result != nil {
return attempt.result, attempt.candidateCount, attempt.topK, attempt.loadSkew, nil
}
if len(regularAccounts) > 0 {
regularAttempt := s.trySelectByLoadBalancePool(ctx, req, regularAccounts, loadMap)
if regularAttempt.err != nil && !regularAttempt.noCompactCandidates {
return nil, regularAttempt.candidateCount, regularAttempt.topK, regularAttempt.loadSkew, regularAttempt.err
}
if regularAttempt.result != nil {
return regularAttempt.result, regularAttempt.candidateCount, regularAttempt.topK, regularAttempt.loadSkew, nil
}
var result *AccountSelectionResult
candidateCount, topK, loadSkew := regularAttempt.candidateCount, regularAttempt.topK, regularAttempt.loadSkew
fallbackErr := regularAttempt.err
if regularAttempt.err == nil {
result, candidateCount, topK, loadSkew, fallbackErr = s.finishLoadBalanceSelectionFallback(ctx, req, regularAttempt)
if fallbackErr == nil && result != nil {
return result, candidateCount, topK, loadSkew, nil
}
}
// 常规池既无法获取也无法排队(含仅剩不支持 compact 的候选)时,
// 回退到订阅池的等待计划:busy-but-waitable 的订阅账号不应因常规池存在
// 而被丢弃,否则开启订阅优先反而让本可排队成功的请求硬失败。
subResult, subCandidateCount, subTopK, subLoadSkew, subErr := s.finishLoadBalanceSelectionFallback(ctx, req, attempt)
if subErr == nil && subResult != nil {
return subResult, subCandidateCount, subTopK, subLoadSkew, nil
}
return result, candidateCount, topK, loadSkew, fallbackErr
}
return s.finishLoadBalanceSelectionFallback(ctx, req, attempt)
}
}
result, compactBlocked, acquireErr := s.tryAcquireOpenAISelectionOrder(ctx, req, selectionOrder)
attempt := s.trySelectByLoadBalancePool(ctx, req, filtered, loadMap)
if attempt.err != nil {
return nil, attempt.candidateCount, attempt.topK, attempt.loadSkew, attempt.err
}
if attempt.result != nil {
return attempt.result, attempt.candidateCount, attempt.topK, attempt.loadSkew, nil
}
return s.finishLoadBalanceSelectionFallback(ctx, req, attempt)
}
func partitionOpenAIChatGPTSubscriptionAccounts(accounts []*Account) ([]*Account, []*Account) {
subscriptionAccounts := make([]*Account, 0, len(accounts))
regularAccounts := make([]*Account, 0, len(accounts))
for _, account := range accounts {
if account != nil && account.IsOpenAIChatGPTSubscription() {
subscriptionAccounts = append(subscriptionAccounts, account)
continue
}
regularAccounts = append(regularAccounts, account)
}
return subscriptionAccounts, regularAccounts
}
func (s *defaultOpenAIAccountScheduler) trySelectByLoadBalancePool(
ctx context.Context,
req OpenAIAccountScheduleRequest,
filtered []*Account,
loadMap map[int64]*AccountLoadInfo,
) openAIAccountLoadSelectionAttempt {
plan := s.buildOpenAIAccountLoadPlan(ctx, req, filtered, loadMap)
attempt := openAIAccountLoadSelectionAttempt{
selectionOrder: plan.selectionOrder,
candidateCount: plan.candidateCount,
topK: plan.topK,
loadSkew: plan.loadSkew,
}
if req.RequireCompact && len(plan.candidates) == 0 && len(plan.staleSnapshotCompactRetry) == 0 {
attempt.noCompactCandidates = true
attempt.err = ErrNoAvailableCompactAccounts
return attempt
}
if req.RequireCompact && len(attempt.selectionOrder) == 0 && s.service.schedulerSnapshot == nil {
attempt.noCompactCandidates = true
attempt.err = ErrNoAvailableCompactAccounts
return attempt
}
if len(attempt.selectionOrder) == 0 {
attempt.compactBlocked = req.RequireCompact && len(plan.allCandidates) > 0
return attempt
}
result, compactBlocked, acquireErr := s.tryAcquireOpenAISelectionOrder(ctx, req, attempt.selectionOrder)
attempt.compactBlocked = compactBlocked
if acquireErr != nil {
return nil, candidateCount, topK, loadSkew, acquireErr
attempt.err = acquireErr
return attempt
}
if result != nil {
return result, candidateCount, topK, loadSkew, nil
attempt.result = result
return attempt
}
if s.service.concurrencyService != nil {
loadReq := buildOpenAIAccountLoadRequest(filtered)
if freshLoadMap, loadErr := s.service.concurrencyService.GetAccountsLoadBatchFresh(ctx, loadReq); loadErr == nil {
freshPlan := s.buildOpenAIAccountLoadPlan(req, filtered, freshLoadMap)
freshPlan := s.buildOpenAIAccountLoadPlan(ctx, req, filtered, freshLoadMap)
if len(freshPlan.selectionOrder) > 0 {
freshResult, freshCompactBlocked, freshAcquireErr := s.tryAcquireOpenAISelectionOrder(ctx, req, freshPlan.selectionOrder)
if freshAcquireErr != nil {
return nil, candidateCount, topK, loadSkew, freshAcquireErr
attempt.err = freshAcquireErr
return attempt
}
if freshResult != nil {
return freshResult, freshPlan.candidateCount, freshPlan.topK, freshPlan.loadSkew, nil
attempt.result = freshResult
attempt.selectionOrder = freshPlan.selectionOrder
attempt.candidateCount = freshPlan.candidateCount
attempt.topK = freshPlan.topK
attempt.loadSkew = freshPlan.loadSkew
return attempt
}
compactBlocked = compactBlocked || freshCompactBlocked
selectionOrder = freshPlan.selectionOrder
candidateCount = freshPlan.candidateCount
topK = freshPlan.topK
loadSkew = freshPlan.loadSkew
attempt.compactBlocked = attempt.compactBlocked || freshCompactBlocked
attempt.selectionOrder = freshPlan.selectionOrder
attempt.candidateCount = freshPlan.candidateCount
attempt.topK = freshPlan.topK
attempt.loadSkew = freshPlan.loadSkew
}
}
}
return attempt
}
func buildOpenAIAccountLoadRequest(accounts []*Account) []AccountWithConcurrency {
loadReq := make([]AccountWithConcurrency, 0, len(accounts))
for _, account := range accounts {
if account == nil {
continue
}
loadReq = append(loadReq, AccountWithConcurrency{
ID: account.ID,
MaxConcurrency: account.EffectiveLoadFactor(),
})
}
return loadReq
}
func (s *defaultOpenAIAccountScheduler) finishLoadBalanceSelectionFallback(
ctx context.Context,
req OpenAIAccountScheduleRequest,
attempt openAIAccountLoadSelectionAttempt,
) (*AccountSelectionResult, int, int, float64, error) {
candidateCount := attempt.candidateCount
topK := attempt.topK
loadSkew := attempt.loadSkew
if len(attempt.selectionOrder) == 0 {
return nil, candidateCount, topK, loadSkew, noAvailableOpenAISelectionError(req.RequestedModel, attempt.compactBlocked)
}
if stickyFallback, stickyErr := s.tryFallbackToWeightedSticky(ctx, req); stickyErr != nil {
return nil, candidateCount, topK, loadSkew, stickyErr
} else if stickyFallback != nil {
return stickyFallback, candidateCount, topK, loadSkew, nil
}
cfg := s.service.schedulingConfig()
compactBlocked := attempt.compactBlocked
// WaitPlan.MaxConcurrency 使用 Concurrency(非 EffectiveLoadFactor),因为 WaitPlan 控制的是 Redis 实际并发槽位等待。
for _, candidate := range selectionOrder {
for _, candidate := range attempt.selectionOrder {
fresh := s.service.resolveFreshSchedulableOpenAIAccount(ctx, candidate.account, req.Platform, req.RequestedModel, false, req.RequiredCapability)
if fresh == nil || !s.isAccountTransportCompatible(fresh, req.RequiredTransport) || !s.isAccountRequestCompatible(ctx, fresh, req) {
continue
@@ -1184,40 +1446,169 @@ func (s *OpenAIGatewayService) openAIAdvancedSchedulerSettingRepo() SettingRepos
return s.rateLimitService.settingService.settingRepo
}
func (s *OpenAIGatewayService) isOpenAIAdvancedSchedulerEnabled(ctx context.Context) bool {
func (s *OpenAIGatewayService) openAIAdvancedSchedulerRuntimeSettings(ctx context.Context) openAIAdvancedSchedulerRuntimeSettings {
if cached, ok := openAIAdvancedSchedulerSettingCache.Load().(*cachedOpenAIAdvancedSchedulerSetting); ok && cached != nil {
if time.Now().UnixNano() < cached.expiresAt {
return cached.enabled
return openAIAdvancedSchedulerRuntimeSettings{
enabled: cached.enabled,
stickyWeightedEnabled: cached.stickyWeightedEnabled,
subscriptionPriorityEnabled: cached.subscriptionPriorityEnabled,
lbTopKOverride: cached.lbTopKOverride,
weightOverrides: cloneOpenAIAdvancedSchedulerWeightOverrides(cached.weightOverrides),
}
}
}
result, _, _ := openAIAdvancedSchedulerSettingSF.Do(openAIAdvancedSchedulerSettingKey, func() (any, error) {
if cached, ok := openAIAdvancedSchedulerSettingCache.Load().(*cachedOpenAIAdvancedSchedulerSetting); ok && cached != nil {
if time.Now().UnixNano() < cached.expiresAt {
return cached.enabled, nil
return openAIAdvancedSchedulerRuntimeSettings{
enabled: cached.enabled,
stickyWeightedEnabled: cached.stickyWeightedEnabled,
subscriptionPriorityEnabled: cached.subscriptionPriorityEnabled,
lbTopKOverride: cached.lbTopKOverride,
weightOverrides: cloneOpenAIAdvancedSchedulerWeightOverrides(cached.weightOverrides),
}, nil
}
}
enabled := false
stickyWeightedEnabled := false
subscriptionPriorityEnabled := false
lbTopKOverride := 0
weightOverrides := map[string]float64{}
if repo := s.openAIAdvancedSchedulerSettingRepo(); repo != nil {
dbCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), openAIAdvancedSchedulerSettingDBTimeout)
defer cancel()
value, err := repo.GetValue(dbCtx, openAIAdvancedSchedulerSettingKey)
if err == nil {
enabled = strings.EqualFold(strings.TrimSpace(value), "true")
if values, err := repo.GetMultiple(dbCtx, openAIAdvancedSchedulerRuntimeSettingKeys()); err == nil {
enabled = strings.EqualFold(strings.TrimSpace(values[openAIAdvancedSchedulerSettingKey]), "true")
stickyWeightedEnabled = strings.EqualFold(strings.TrimSpace(values[SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled]), "true")
subscriptionPriorityEnabled = strings.EqualFold(strings.TrimSpace(values[SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled]), "true")
lbTopKOverride = parsePositiveIntOverride(values[SettingKeyOpenAIAdvancedSchedulerLBTopK])
weightOverrides = parseOpenAIAdvancedSchedulerWeightOverrides(values)
} else {
// 批量读取失败时逐键降级,覆盖全部键(含 TopK/权重),避免只加载布尔开关
// 而静默丢弃管理员配置的覆盖值;降级状态会被缓存一个 TTL,必须留痕。
slog.Warn("openai_advanced_scheduler_settings_batch_load_failed", "error", err)
fallbackValues := make(map[string]string)
for _, key := range openAIAdvancedSchedulerRuntimeSettingKeys() {
if value, valueErr := repo.GetValue(dbCtx, key); valueErr == nil {
fallbackValues[key] = value
}
}
enabled = strings.EqualFold(strings.TrimSpace(fallbackValues[openAIAdvancedSchedulerSettingKey]), "true")
stickyWeightedEnabled = strings.EqualFold(strings.TrimSpace(fallbackValues[SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled]), "true")
subscriptionPriorityEnabled = strings.EqualFold(strings.TrimSpace(fallbackValues[SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled]), "true")
lbTopKOverride = parsePositiveIntOverride(fallbackValues[SettingKeyOpenAIAdvancedSchedulerLBTopK])
weightOverrides = parseOpenAIAdvancedSchedulerWeightOverrides(fallbackValues)
}
}
openAIAdvancedSchedulerSettingCache.Store(&cachedOpenAIAdvancedSchedulerSetting{
enabled: enabled,
expiresAt: time.Now().Add(openAIAdvancedSchedulerSettingCacheTTL).UnixNano(),
enabled: enabled,
stickyWeightedEnabled: stickyWeightedEnabled,
subscriptionPriorityEnabled: subscriptionPriorityEnabled,
lbTopKOverride: lbTopKOverride,
weightOverrides: cloneOpenAIAdvancedSchedulerWeightOverrides(weightOverrides),
expiresAt: time.Now().Add(openAIAdvancedSchedulerSettingCacheTTL).UnixNano(),
})
return enabled, nil
return openAIAdvancedSchedulerRuntimeSettings{
enabled: enabled,
stickyWeightedEnabled: stickyWeightedEnabled,
subscriptionPriorityEnabled: subscriptionPriorityEnabled,
lbTopKOverride: lbTopKOverride,
weightOverrides: weightOverrides,
}, nil
})
enabled, _ := result.(bool)
return enabled
settings, _ := result.(openAIAdvancedSchedulerRuntimeSettings)
return settings
}
func (s *OpenAIGatewayService) isOpenAIAdvancedSchedulerEnabled(ctx context.Context) bool {
return s.openAIAdvancedSchedulerRuntimeSettings(ctx).enabled
}
func (s *OpenAIGatewayService) isOpenAIAdvancedSchedulerStickyWeightedEnabled(ctx context.Context) bool {
settings := s.openAIAdvancedSchedulerRuntimeSettings(ctx)
return settings.enabled && settings.stickyWeightedEnabled
}
func (s *OpenAIGatewayService) isOpenAIAdvancedSchedulerSubscriptionPriorityEnabled(ctx context.Context) bool {
settings := s.openAIAdvancedSchedulerRuntimeSettings(ctx)
return settings.enabled && settings.subscriptionPriorityEnabled
}
func openAIAdvancedSchedulerRuntimeSettingKeys() []string {
keys := []string{
openAIAdvancedSchedulerSettingKey,
SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled,
SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled,
SettingKeyOpenAIAdvancedSchedulerLBTopK,
}
for _, spec := range openAIAdvancedSchedulerWeightOverrideSpecs() {
keys = append(keys, spec.key)
}
return keys
}
type openAIAdvancedSchedulerWeightOverrideSpec struct {
key string
name string
}
func openAIAdvancedSchedulerWeightOverrideSpecs() []openAIAdvancedSchedulerWeightOverrideSpec {
return []openAIAdvancedSchedulerWeightOverrideSpec{
{key: SettingKeyOpenAIAdvancedSchedulerWeightPriority, name: "priority"},
{key: SettingKeyOpenAIAdvancedSchedulerWeightLoad, name: "load"},
{key: SettingKeyOpenAIAdvancedSchedulerWeightQueue, name: "queue"},
{key: SettingKeyOpenAIAdvancedSchedulerWeightErrorRate, name: "error_rate"},
{key: SettingKeyOpenAIAdvancedSchedulerWeightTTFT, name: "ttft"},
{key: SettingKeyOpenAIAdvancedSchedulerWeightReset, name: "reset"},
{key: SettingKeyOpenAIAdvancedSchedulerWeightQuotaHeadroom, name: "quota_headroom"},
{key: SettingKeyOpenAIAdvancedSchedulerWeightPreviousResponse, name: "previous_response"},
{key: SettingKeyOpenAIAdvancedSchedulerWeightSessionSticky, name: "session_sticky"},
}
}
func parsePositiveIntOverride(raw string) int {
raw = strings.TrimSpace(raw)
if raw == "" {
return 0
}
value, err := strconv.Atoi(raw)
if err != nil || value <= 0 {
return 0
}
return value
}
func parseOpenAIAdvancedSchedulerWeightOverrides(values map[string]string) map[string]float64 {
overrides := map[string]float64{}
for _, spec := range openAIAdvancedSchedulerWeightOverrideSpecs() {
raw := strings.TrimSpace(values[spec.key])
if raw == "" {
continue
}
value, err := strconv.ParseFloat(raw, 64)
if err != nil || value < 0 || math.IsNaN(value) || math.IsInf(value, 0) {
continue
}
overrides[spec.name] = value
}
return overrides
}
func cloneOpenAIAdvancedSchedulerWeightOverrides(in map[string]float64) map[string]float64 {
if len(in) == 0 {
return nil
}
out := make(map[string]float64, len(in))
for key, value := range in {
out[key] = value
}
return out
}
func (s *OpenAIGatewayService) getOpenAIAccountScheduler(ctx context.Context) OpenAIAccountScheduler {
@@ -1253,9 +1644,12 @@ func (s *OpenAIGatewayService) SelectAccountWithScheduler(
requiredTransport OpenAIUpstreamTransport,
requireCompact bool,
) (*AccountSelectionResult, OpenAIAccountScheduleDecision, error) {
return s.selectAccountWithScheduler(ctx, groupID, previousResponseID, sessionHash, requestedModel, excludedIDs, requiredTransport, "", "", requireCompact, PlatformOpenAI)
return s.selectAccountWithScheduler(ctx, groupID, previousResponseID, sessionHash, requestedModel, excludedIDs, requiredTransport, "", "", requireCompact, PlatformOpenAI, false)
}
// SelectAccountWithSchedulerForCapability 按能力要求调度账号。
// previousResponseCanMove 表示首包 input 可自行重建工具续链,previous_response_id 允许跨账号迁移
// (粘性加权模式下改为加权偏好而非硬粘连)。
func (s *OpenAIGatewayService) SelectAccountWithSchedulerForCapability(
ctx context.Context,
groupID *int64,
@@ -1266,13 +1660,14 @@ func (s *OpenAIGatewayService) SelectAccountWithSchedulerForCapability(
requiredTransport OpenAIUpstreamTransport,
requiredCapability OpenAIEndpointCapability,
requireCompact bool,
previousResponseCanMove bool,
platformOverride ...string,
) (*AccountSelectionResult, OpenAIAccountScheduleDecision, error) {
platform := PlatformOpenAI
if len(platformOverride) > 0 {
platform = platformOverride[0]
}
return s.selectAccountWithScheduler(ctx, groupID, previousResponseID, sessionHash, requestedModel, excludedIDs, requiredTransport, requiredCapability, "", requireCompact, platform)
return s.selectAccountWithScheduler(ctx, groupID, previousResponseID, sessionHash, requestedModel, excludedIDs, requiredTransport, requiredCapability, "", requireCompact, platform, previousResponseCanMove)
}
func (s *OpenAIGatewayService) SelectAccountWithSchedulerForImages(
@@ -1283,13 +1678,13 @@ func (s *OpenAIGatewayService) SelectAccountWithSchedulerForImages(
excludedIDs map[int64]struct{},
requiredCapability OpenAIImagesCapability,
) (*AccountSelectionResult, OpenAIAccountScheduleDecision, error) {
selection, decision, err := s.selectAccountWithScheduler(ctx, groupID, "", sessionHash, requestedModel, excludedIDs, OpenAIUpstreamTransportHTTPSSE, "", requiredCapability, false, PlatformOpenAI)
selection, decision, err := s.selectAccountWithScheduler(ctx, groupID, "", sessionHash, requestedModel, excludedIDs, OpenAIUpstreamTransportHTTPSSE, "", requiredCapability, false, PlatformOpenAI, false)
if err == nil && selection != nil && selection.Account != nil {
return selection, decision, nil
}
// 如果要求 native 能力(如指定了模型)但没有可用的 APIKey 账号,回退到 basic(OAuth 账号)
if requiredCapability == OpenAIImagesCapabilityNative {
return s.selectAccountWithScheduler(ctx, groupID, "", sessionHash, requestedModel, excludedIDs, OpenAIUpstreamTransportHTTPSSE, "", OpenAIImagesCapabilityBasic, false, PlatformOpenAI)
return s.selectAccountWithScheduler(ctx, groupID, "", sessionHash, requestedModel, excludedIDs, OpenAIUpstreamTransportHTTPSSE, "", OpenAIImagesCapabilityBasic, false, PlatformOpenAI, false)
}
return selection, decision, err
}
@@ -1306,6 +1701,7 @@ func (s *OpenAIGatewayService) selectAccountWithScheduler(
requiredImageCapability OpenAIImagesCapability,
requireCompact bool,
platform string,
previousResponseCanMove bool,
) (*AccountSelectionResult, OpenAIAccountScheduleDecision, error) {
ctx = s.withOpenAIQuotaAutoPauseContext(ctx)
platform = normalizeOpenAICompatiblePlatform(platform)
@@ -1378,13 +1774,23 @@ func (s *OpenAIGatewayService) selectAccountWithScheduler(
stickyAccountID = accountID
}
}
stickyWeighted := s.isOpenAIAdvancedSchedulerStickyWeightedEnabled(ctx)
subscriptionPriority := s.isOpenAIAdvancedSchedulerSubscriptionPriorityEnabled(ctx)
stickyPreviousAccountID := int64(0)
if stickyWeighted && previousResponseCanMove && strings.TrimSpace(previousResponseID) != "" && platform == PlatformOpenAI {
stickyPreviousAccountID = s.ResolveAccountIDByPreviousResponseIDForScheduler(ctx, groupID, previousResponseID, requestedModel, excludedIDs, requiredCapability, requireCompact)
}
return scheduler.Select(ctx, OpenAIAccountScheduleRequest{
GroupID: groupID,
Platform: platform,
SessionHash: sessionHash,
StickyAccountID: stickyAccountID,
StickyPreviousAccountID: stickyPreviousAccountID,
StickyWeighted: stickyWeighted,
SubscriptionPriority: subscriptionPriority,
PreviousResponseID: previousResponseID,
PreviousResponseCanMove: previousResponseCanMove,
RequestedModel: requestedModel,
RequiredTransport: requiredTransport,
RequiredCapability: requiredCapability,
@@ -1473,6 +1879,20 @@ func (s *OpenAIGatewayService) openAIWSLBTopK() int {
return 7
}
func (s *OpenAIGatewayService) openAIWSLBTopKForRequest(ctx context.Context) int {
base := s.openAIWSLBTopK()
settings := s.openAIAdvancedSchedulerRuntimeSettings(ctx)
// DB 覆盖值与 stickyWeighted/subscriptionPriority 一样受总开关门控:
// 关闭高级调度器后所有调用方(含管理页分数快照)都应回到配置/默认行为。
if !settings.enabled {
return base
}
if settings.lbTopKOverride > 0 {
return settings.lbTopKOverride
}
return base
}
func (s *OpenAIGatewayService) openAIStickyEscapeConfig() openAIStickyEscapeConfig {
if s != nil && s.cfg != nil {
cfg := s.cfg.Gateway.OpenAIScheduler
@@ -1514,6 +1934,8 @@ func (s *OpenAIGatewayService) openAIWSSchedulerWeights() GatewayOpenAIWSSchedul
TTFT: s.cfg.Gateway.OpenAIWS.SchedulerScoreWeights.TTFT,
Reset: s.cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Reset,
QuotaHeadroom: s.cfg.Gateway.OpenAIWS.SchedulerScoreWeights.QuotaHeadroom,
Previous: s.cfg.Gateway.OpenAIWS.SchedulerScoreWeights.PreviousResponse,
SessionSticky: s.cfg.Gateway.OpenAIWS.SchedulerScoreWeights.SessionSticky,
}
}
return GatewayOpenAIWSSchedulerScoreWeightsView{
@@ -1524,9 +1946,50 @@ func (s *OpenAIGatewayService) openAIWSSchedulerWeights() GatewayOpenAIWSSchedul
TTFT: 0.5,
Reset: 0.0,
QuotaHeadroom: 0.0,
Previous: 5.0,
SessionSticky: 3.0,
}
}
func (s *OpenAIGatewayService) openAIWSSchedulerWeightsForRequest(ctx context.Context) GatewayOpenAIWSSchedulerScoreWeightsView {
weights := s.openAIWSSchedulerWeights()
settings := s.openAIAdvancedSchedulerRuntimeSettings(ctx)
// 同 openAIWSLBTopKForRequest:总开关关闭时不应用 DB 覆盖值。
if !settings.enabled {
return weights
}
return applyOpenAIAdvancedSchedulerWeightOverrides(weights, settings.weightOverrides)
}
func applyOpenAIAdvancedSchedulerWeightOverrides(
weights GatewayOpenAIWSSchedulerScoreWeightsView,
overrides map[string]float64,
) GatewayOpenAIWSSchedulerScoreWeightsView {
for key, value := range overrides {
switch key {
case "priority":
weights.Priority = value
case "load":
weights.Load = value
case "queue":
weights.Queue = value
case "error_rate":
weights.ErrorRate = value
case "ttft":
weights.TTFT = value
case "reset":
weights.Reset = value
case "quota_headroom":
weights.QuotaHeadroom = value
case "previous_response":
weights.Previous = value
case "session_sticky":
weights.SessionSticky = value
}
}
return weights
}
type GatewayOpenAIWSSchedulerScoreWeightsView struct {
Priority float64
Load float64
@@ -1536,6 +1999,149 @@ type GatewayOpenAIWSSchedulerScoreWeightsView struct {
// Reset 倾向「会话窗口最早重置」的账号;0 表示关闭(默认)。
Reset float64
QuotaHeadroom float64
Previous float64
SessionSticky float64
}
type OpenAIAccountSchedulerScoreSnapshot struct {
BaseScore float64
StickyScore float64
StickyScoreInfinity bool
StickyWeightedEnabled bool
}
func (s *RateLimitService) BuildOpenAIAccountSchedulerScoreSnapshot(
ctx context.Context,
accounts []*Account,
loadMap map[int64]*AccountLoadInfo,
) map[int64]OpenAIAccountSchedulerScoreSnapshot {
gateway := &OpenAIGatewayService{cfg: nil, rateLimitService: s}
if s != nil {
gateway.cfg = s.cfg
}
return buildOpenAIAccountSchedulerScoreSnapshot(accounts, loadMap, gateway.openAIWSSchedulerWeightsForRequest(ctx), gateway.isOpenAIAdvancedSchedulerStickyWeightedEnabled(ctx))
}
func BuildOpenAIAccountSchedulerScoreSnapshot(
accounts []*Account,
loadMap map[int64]*AccountLoadInfo,
) map[int64]OpenAIAccountSchedulerScoreSnapshot {
gateway := &OpenAIGatewayService{}
return buildOpenAIAccountSchedulerScoreSnapshot(accounts, loadMap, gateway.openAIWSSchedulerWeights(), false)
}
func buildOpenAIAccountSchedulerScoreSnapshot(
accounts []*Account,
loadMap map[int64]*AccountLoadInfo,
weights GatewayOpenAIWSSchedulerScoreWeightsView,
stickyWeightedEnabled bool,
) map[int64]OpenAIAccountSchedulerScoreSnapshot {
if len(accounts) == 0 {
return nil
}
candidates := make([]openAIAccountCandidateScore, 0, len(accounts))
for _, account := range accounts {
if account == nil {
continue
}
loadInfo := loadMap[account.ID]
if loadInfo == nil {
loadInfo = &AccountLoadInfo{AccountID: account.ID}
}
candidates = append(candidates, openAIAccountCandidateScore{
account: account,
loadInfo: loadInfo,
errorRate: 0,
ttft: 0,
hasTTFT: false,
})
}
if len(candidates) == 0 {
return nil
}
minPriority, maxPriority := openAIAccountSchedulingPriority(candidates[0].account), openAIAccountSchedulingPriority(candidates[0].account)
maxWaiting := 1
for i := range candidates {
candidate := &candidates[i]
candidate.priority = openAIAccountSchedulingPriority(candidate.account)
if candidate.priority < minPriority {
minPriority = candidate.priority
}
if candidate.priority > maxPriority {
maxPriority = candidate.priority
}
if candidate.loadInfo.WaitingCount > maxWaiting {
maxWaiting = candidate.loadInfo.WaitingCount
}
}
minResetRemaining, maxResetRemaining := 0.0, 0.0
hasResetSample := false
now := time.Now()
if weights.Reset > 0 {
for _, candidate := range candidates {
end := candidate.account.SessionWindowEnd
if end == nil || !now.Before(*end) {
continue
}
remaining := end.Sub(now).Seconds()
if !hasResetSample {
minResetRemaining, maxResetRemaining = remaining, remaining
hasResetSample = true
continue
}
if remaining < minResetRemaining {
minResetRemaining = remaining
}
if remaining > maxResetRemaining {
maxResetRemaining = remaining
}
}
}
result := make(map[int64]OpenAIAccountSchedulerScoreSnapshot, len(candidates))
for _, candidate := range candidates {
priorityFactor := 1.0
if maxPriority > minPriority {
priorityFactor = 1 - float64(candidate.priority-minPriority)/float64(maxPriority-minPriority)
}
loadFactor := 1 - clamp01(float64(candidate.loadInfo.LoadRate)/100.0)
queueFactor := 1 - clamp01(float64(candidate.loadInfo.WaitingCount)/float64(maxWaiting))
errorFactor := 1.0
ttftFactor := 0.5
resetFactor := 0.0
if weights.Reset > 0 && hasResetSample {
if end := candidate.account.SessionWindowEnd; end != nil && now.Before(*end) {
if maxResetRemaining > minResetRemaining {
resetFactor = 1 - clamp01((end.Sub(now).Seconds()-minResetRemaining)/(maxResetRemaining-minResetRemaining))
} else {
resetFactor = 1
}
}
}
quotaHeadroomFactor := 0.0
if weights.QuotaHeadroom > 0 {
quotaHeadroomFactor = openAIQuotaHeadroomFactor(candidate.account, now)
}
baseScore := weights.Priority*priorityFactor +
weights.Load*loadFactor +
weights.Queue*queueFactor +
weights.ErrorRate*errorFactor +
weights.TTFT*ttftFactor +
weights.Reset*resetFactor +
weights.QuotaHeadroom*quotaHeadroomFactor
score := OpenAIAccountSchedulerScoreSnapshot{
BaseScore: baseScore,
StickyWeightedEnabled: stickyWeightedEnabled,
StickyScoreInfinity: !stickyWeightedEnabled,
}
if stickyWeightedEnabled {
score.StickyScore = baseScore + weights.Previous + weights.SessionSticky
}
result[candidate.account.ID] = score
}
return result
}
func openAIQuotaHeadroomFactor(account *Account, now time.Time) float64 {
@@ -1,6 +1,7 @@
package service
import (
"context"
"testing"
"time"
@@ -49,7 +50,7 @@ func TestBuildOpenAIAccountLoadPlan_ResetWeightPrefersSoonestReset(t *testing.T)
}
sched := openAIResetTestScheduler(5.0)
plan := sched.buildOpenAIAccountLoadPlan(OpenAIAccountScheduleRequest{}, filtered, map[int64]*AccountLoadInfo{})
plan := sched.buildOpenAIAccountLoadPlan(context.Background(), OpenAIAccountScheduleRequest{}, filtered, map[int64]*AccountLoadInfo{})
scores := openAIPlanScores(plan)
require.Greater(t, scores[2], scores[1], "重置时间最早的账号(ID=2)得分更高")
}
@@ -65,7 +66,7 @@ func TestBuildOpenAIAccountLoadPlan_ResetWeightZeroNoEffect(t *testing.T) {
}
sched := openAIResetTestScheduler(0.0)
plan := sched.buildOpenAIAccountLoadPlan(OpenAIAccountScheduleRequest{}, filtered, map[int64]*AccountLoadInfo{})
plan := sched.buildOpenAIAccountLoadPlan(context.Background(), OpenAIAccountScheduleRequest{}, filtered, map[int64]*AccountLoadInfo{})
scores := openAIPlanScores(plan)
require.Equal(t, scores[1], scores[2], "Reset 权重为 0 时两账号得分相同")
}
@@ -80,7 +81,7 @@ func TestBuildOpenAIAccountLoadPlan_ResetWeightIgnoresNilWindow(t *testing.T) {
}
sched := openAIResetTestScheduler(5.0)
plan := sched.buildOpenAIAccountLoadPlan(OpenAIAccountScheduleRequest{}, filtered, map[int64]*AccountLoadInfo{})
plan := sched.buildOpenAIAccountLoadPlan(context.Background(), OpenAIAccountScheduleRequest{}, filtered, map[int64]*AccountLoadInfo{})
scores := openAIPlanScores(plan)
require.Greater(t, scores[2], scores[1], "拥有活跃窗口的账号得分高于无窗口账号")
}
@@ -161,7 +162,7 @@ func TestBuildOpenAIAccountLoadPlan_QuotaHeadroomPrefersHigher7dRemaining(t *tes
}
sched := openAIQuotaHeadroomTestScheduler(1.0)
plan := sched.buildOpenAIAccountLoadPlan(OpenAIAccountScheduleRequest{}, filtered, map[int64]*AccountLoadInfo{})
plan := sched.buildOpenAIAccountLoadPlan(context.Background(), OpenAIAccountScheduleRequest{}, filtered, map[int64]*AccountLoadInfo{})
scores := openAIPlanScores(plan)
require.Greater(t, scores[2], scores[1], "7d 剩余额度更高的账号得分应更高")
}
@@ -190,7 +191,7 @@ func TestBuildOpenAIAccountLoadPlan_QuotaHeadroomZeroNoEffect(t *testing.T) {
}
sched := openAIResetTestScheduler(0)
plan := sched.buildOpenAIAccountLoadPlan(OpenAIAccountScheduleRequest{}, filtered, map[int64]*AccountLoadInfo{})
plan := sched.buildOpenAIAccountLoadPlan(context.Background(), OpenAIAccountScheduleRequest{}, filtered, map[int64]*AccountLoadInfo{})
scores := openAIPlanScores(plan)
require.Equal(t, scores[1], scores[2], "quota_headroom 权重为 0 时不应影响打分")
}
@@ -183,6 +183,17 @@ func newSchedulerTestOpenAIWSV2Config() *config.Config {
return cfg
}
func newSchedulerTestSubscriptionPriorityConfig() *config.Config {
cfg := &config.Config{}
cfg.Gateway.OpenAIWS.LBTopK = 1
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Priority = 1
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Load = 1
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Queue = 1
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.ErrorRate = 0
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.TTFT = 0
return cfg
}
type openAIAdvancedSchedulerSettingRepoStub struct {
values map[string]string
}
@@ -210,8 +221,14 @@ func (s *openAIAdvancedSchedulerSettingRepoStub) Set(context.Context, string, st
panic("unexpected call to Set")
}
func (s *openAIAdvancedSchedulerSettingRepoStub) GetMultiple(context.Context, []string) (map[string]string, error) {
panic("unexpected call to GetMultiple")
func (s *openAIAdvancedSchedulerSettingRepoStub) GetMultiple(_ context.Context, keys []string) (map[string]string, error) {
result := make(map[string]string, len(keys))
for _, key := range keys {
if value, err := s.GetValue(context.Background(), key); err == nil {
result[key] = value
}
}
return result, nil
}
func (s *openAIAdvancedSchedulerSettingRepoStub) SetMultiple(context.Context, map[string]string) error {
@@ -226,7 +243,7 @@ func (s *openAIAdvancedSchedulerSettingRepoStub) Delete(context.Context, string)
panic("unexpected call to Delete")
}
func newOpenAIAdvancedSchedulerRateLimitService(enabled string) *RateLimitService {
func newOpenAIAdvancedSchedulerRateLimitService(enabled string, values ...string) *RateLimitService {
resetOpenAIAdvancedSchedulerSettingCacheForTest()
repo := &openAIAdvancedSchedulerSettingRepoStub{
values: map[string]string{},
@@ -234,6 +251,12 @@ func newOpenAIAdvancedSchedulerRateLimitService(enabled string) *RateLimitServic
if enabled != "" {
repo.values[openAIAdvancedSchedulerSettingKey] = enabled
}
if len(values) > 0 && values[0] != "" {
repo.values[SettingKeyOpenAIAdvancedSchedulerStickyWeightedEnabled] = values[0]
}
if len(values) > 1 && values[1] != "" {
repo.values[SettingKeyOpenAIAdvancedSchedulerSubscriptionPriorityEnabled] = values[1]
}
return &RateLimitService{
settingService: NewSettingService(repo, &config.Config{}),
}
@@ -266,6 +289,45 @@ func (s *openAISnapshotCacheStub) GetAccount(ctx context.Context, accountID int6
return &cloned, nil
}
func TestOpenAIGatewayService_OpenAIAdvancedSchedulerRuntimeSettings_DBOverridesConfig(t *testing.T) {
resetOpenAIAdvancedSchedulerSettingCacheForTest()
defer resetOpenAIAdvancedSchedulerSettingCacheForTest()
cfg := &config.Config{}
cfg.Gateway.OpenAIWS.LBTopK = 11
cfg.Gateway.OpenAIWS.SchedulerScoreWeights = config.GatewayOpenAIWSSchedulerScoreWeights{
Priority: 1,
Load: 2,
Queue: 3,
ErrorRate: 4,
TTFT: 5,
Reset: 6,
QuotaHeadroom: 7,
PreviousResponse: 8,
SessionSticky: 9,
}
repo := &openAIAdvancedSchedulerSettingRepoStub{
values: map[string]string{
openAIAdvancedSchedulerSettingKey: "true",
SettingKeyOpenAIAdvancedSchedulerLBTopK: "3",
SettingKeyOpenAIAdvancedSchedulerWeightPriority: "2.5",
SettingKeyOpenAIAdvancedSchedulerWeightPreviousResponse: "12",
},
}
svc := &OpenAIGatewayService{
cfg: cfg,
rateLimitService: &RateLimitService{settingService: NewSettingService(repo, cfg)},
}
ctx := context.Background()
require.Equal(t, 3, svc.openAIWSLBTopKForRequest(ctx))
weights := svc.openAIWSSchedulerWeightsForRequest(ctx)
require.Equal(t, 2.5, weights.Priority)
require.Equal(t, 2.0, weights.Load)
require.Equal(t, 12.0, weights.Previous)
require.Equal(t, 9.0, weights.SessionSticky)
}
func TestOpenAIGatewayService_SelectAccountWithScheduler_DefaultDisabledUsesLegacyLoadAwareness(t *testing.T) {
resetOpenAIAdvancedSchedulerSettingCacheForTest()
@@ -467,6 +529,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_DefaultDisabled_Embeddi
OpenAIUpstreamTransportHTTPSSE,
OpenAIEndpointCapabilityEmbeddings,
false,
false,
)
require.NoError(t, err)
require.NotNil(t, selection)
@@ -510,6 +573,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_DefaultDisabled_AllowsG
OpenAIUpstreamTransportAny,
OpenAIEndpointCapabilityChatCompletions,
false,
false,
PlatformGrok,
)
require.NoError(t, err)
@@ -584,6 +648,248 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_EnabledUsesAdvancedPrev
require.True(t, decision.StickyPreviousHit)
}
func TestOpenAIGatewayService_SelectAccountWithScheduler_StickyWeightedSessionInTopKUsesStickyFirst(t *testing.T) {
resetOpenAIAdvancedSchedulerSettingCacheForTest()
ctx := context.Background()
groupID := int64(101071)
accounts := []Account{
{
ID: 37101,
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Status: StatusActive,
Schedulable: true,
Concurrency: 1,
Priority: 100,
GroupIDs: []int64{groupID},
},
{
ID: 37102,
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Status: StatusActive,
Schedulable: true,
Concurrency: 1,
Priority: 0,
GroupIDs: []int64{groupID},
},
}
cfg := &config.Config{}
cfg.Gateway.OpenAIWS.LBTopK = 2
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Priority = 1
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Load = 1
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Queue = 0.7
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.ErrorRate = 0.8
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.TTFT = 0.5
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.SessionSticky = 3
cache := &schedulerTestGatewayCache{sessionBindings: map[string]int64{
"openai:session_hash_weighted_topk": 37101,
}}
svc := &OpenAIGatewayService{
accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts},
cache: cache,
cfg: cfg,
rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true", "true"),
concurrencyService: NewConcurrencyService(schedulerTestConcurrencyCache{}),
}
selection, decision, err := svc.SelectAccountWithScheduler(
ctx,
&groupID,
"",
"session_hash_weighted_topk",
"gpt-5.1",
nil,
OpenAIUpstreamTransportAny,
false,
)
require.NoError(t, err)
require.NotNil(t, selection)
require.NotNil(t, selection.Account)
require.Equal(t, int64(37101), selection.Account.ID)
require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer)
require.True(t, decision.StickySessionHit)
require.Equal(t, 2, decision.TopK)
if selection.ReleaseFunc != nil {
selection.ReleaseFunc()
}
}
func TestOpenAIGatewayService_SelectAccountWithScheduler_StickyWeightedPreviousRequiresMovableContext(t *testing.T) {
resetOpenAIAdvancedSchedulerSettingCacheForTest()
ctx := context.Background()
groupID := int64(101072)
accounts := []Account{
{
ID: 37111,
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Status: StatusActive,
Schedulable: true,
Concurrency: 1,
Priority: 100,
GroupIDs: []int64{groupID},
Extra: map[string]any{
"openai_apikey_responses_websockets_v2_enabled": true,
},
},
{
ID: 37112,
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Status: StatusActive,
Schedulable: true,
Concurrency: 1,
Priority: 0,
GroupIDs: []int64{groupID},
Extra: map[string]any{
"openai_apikey_responses_websockets_v2_enabled": true,
},
},
}
cfg := newSchedulerTestOpenAIWSV2Config()
cfg.Gateway.OpenAIWS.LBTopK = 1
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Priority = 1
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Load = 1
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Queue = 0.7
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.ErrorRate = 0.8
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.TTFT = 0.5
svc := &OpenAIGatewayService{
accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts},
cache: &schedulerTestGatewayCache{},
cfg: cfg,
rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true", "true"),
concurrencyService: NewConcurrencyService(schedulerTestConcurrencyCache{}),
}
store := svc.getOpenAIWSStateStore()
require.NoError(t, store.BindResponseAccount(ctx, groupID, "resp_weighted_unmovable", 37111, time.Hour))
selection, decision, err := svc.SelectAccountWithSchedulerForCapability(
ctx,
&groupID,
"resp_weighted_unmovable",
"",
"gpt-5.1",
nil,
OpenAIUpstreamTransportAny,
OpenAIEndpointCapabilityChatCompletions,
false,
false,
PlatformOpenAI,
)
require.NoError(t, err)
require.NotNil(t, selection)
require.NotNil(t, selection.Account)
require.Equal(t, int64(37111), selection.Account.ID)
require.Equal(t, openAIAccountScheduleLayerPreviousResponse, decision.Layer)
require.True(t, decision.StickyPreviousHit)
if selection.ReleaseFunc != nil {
selection.ReleaseFunc()
}
selection, decision, err = svc.SelectAccountWithSchedulerForCapability(
ctx,
&groupID,
"resp_weighted_unmovable",
"",
"gpt-5.1",
nil,
OpenAIUpstreamTransportAny,
OpenAIEndpointCapabilityChatCompletions,
false,
true,
PlatformOpenAI,
)
require.NoError(t, err)
require.NotNil(t, selection)
require.NotNil(t, selection.Account)
require.Equal(t, int64(37112), selection.Account.ID)
require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer)
require.False(t, decision.StickyPreviousHit)
if selection.ReleaseFunc != nil {
selection.ReleaseFunc()
}
}
func TestOpenAIGatewayService_SelectAccountWithScheduler_PreviousResponseCompactUnsupportedDeletesBinding(t *testing.T) {
resetOpenAIAdvancedSchedulerSettingCacheForTest()
ctx := context.Background()
groupID := int64(101073)
accounts := []Account{
{
ID: 37121,
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Status: StatusActive,
Schedulable: true,
Concurrency: 1,
Priority: 0,
GroupIDs: []int64{groupID},
Extra: map[string]any{
"openai_apikey_responses_websockets_v2_enabled": true,
"openai_compact_mode": OpenAICompactModeForceOff,
},
},
{
ID: 37122,
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Status: StatusActive,
Schedulable: true,
Concurrency: 1,
Priority: 10,
GroupIDs: []int64{groupID},
Extra: map[string]any{
"openai_apikey_responses_websockets_v2_enabled": true,
"openai_compact_mode": OpenAICompactModeForceOn,
},
},
}
cfg := newSchedulerTestOpenAIWSV2Config()
cfg.Gateway.OpenAIWS.LBTopK = 2
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Priority = 1
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Load = 1
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Queue = 0.7
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.ErrorRate = 0.8
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.TTFT = 0.5
svc := &OpenAIGatewayService{
accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts},
cache: &schedulerTestGatewayCache{},
cfg: cfg,
rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true"),
concurrencyService: NewConcurrencyService(schedulerTestConcurrencyCache{}),
}
store := svc.getOpenAIWSStateStore()
require.NoError(t, store.BindResponseAccount(ctx, groupID, "resp_compact_unsupported", 37121, time.Hour))
selection, decision, err := svc.SelectAccountWithScheduler(
ctx,
&groupID,
"resp_compact_unsupported",
"",
"gpt-5.1",
nil,
OpenAIUpstreamTransportAny,
true,
)
require.NoError(t, err)
require.NotNil(t, selection)
require.NotNil(t, selection.Account)
require.Equal(t, int64(37122), selection.Account.ID)
require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer)
require.False(t, decision.StickyPreviousHit)
if selection.ReleaseFunc != nil {
selection.ReleaseFunc()
}
accountID, err := store.GetResponseAccount(ctx, groupID, "resp_compact_unsupported")
require.NoError(t, err)
require.Zero(t, accountID)
}
func TestOpenAIGatewayService_SelectAccountWithScheduler_Enabled_EmbeddingsSkipsChatOnlyAccount(t *testing.T) {
resetOpenAIAdvancedSchedulerSettingCacheForTest()
@@ -635,6 +941,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_Enabled_EmbeddingsSkips
OpenAIUpstreamTransportHTTPSSE,
OpenAIEndpointCapabilityEmbeddings,
false,
false,
)
require.NoError(t, err)
require.NotNil(t, selection)
@@ -708,6 +1015,7 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_Enabled_EmbeddingsSkips
OpenAIUpstreamTransportHTTPSSE,
OpenAIEndpointCapabilityEmbeddings,
false,
false,
)
require.NoError(t, err)
require.NotNil(t, selection)
@@ -1560,6 +1868,217 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_SessionStickyEscapeDisa
require.True(t, decision.StickySessionHit)
}
func TestOpenAIGatewayService_SelectAccountWithScheduler_SubscriptionPriorityChoosesSubscriptionPoolFirst(t *testing.T) {
ctx := context.Background()
groupID := int64(10120)
accounts := []Account{
{
ID: 21601,
Platform: PlatformOpenAI,
Type: AccountTypeOAuth,
Status: StatusActive,
Schedulable: true,
Concurrency: 1,
Priority: 10,
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},
},
}
concurrencyCache := schedulerTestConcurrencyCache{
acquireResults: map[int64]bool{21601: true, 21602: true},
loadMap: map[int64]*AccountLoadInfo{
21601: {AccountID: 21601, LoadRate: 90, WaitingCount: 1},
21602: {AccountID: 21602, LoadRate: 0, WaitingCount: 0},
},
}
svc := &OpenAIGatewayService{
accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts},
cache: &schedulerTestGatewayCache{},
cfg: newSchedulerTestSubscriptionPriorityConfig(),
rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true", "", "true"),
concurrencyService: NewConcurrencyService(concurrencyCache),
}
selection, decision, err := svc.SelectAccountWithScheduler(ctx, &groupID, "", "session_subscription_first", "gpt-5.1", nil, OpenAIUpstreamTransportAny, false)
require.NoError(t, err)
require.NotNil(t, selection)
require.NotNil(t, selection.Account)
require.Equal(t, int64(21601), selection.Account.ID)
require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer)
require.Equal(t, 1, decision.TopK)
if selection.ReleaseFunc != nil {
selection.ReleaseFunc()
}
}
func TestOpenAIGatewayService_SelectAccountWithScheduler_SubscriptionPriorityFallsBackWhenSubscriptionFull(t *testing.T) {
ctx := context.Background()
groupID := int64(10121)
accounts := []Account{
{
ID: 21611,
Platform: PlatformOpenAI,
Type: AccountTypeOAuth,
Status: StatusActive,
Schedulable: true,
Concurrency: 1,
Priority: 0,
GroupIDs: []int64{groupID},
Credentials: map[string]any{"plan_type": "team"},
},
{
ID: 21612,
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Status: StatusActive,
Schedulable: true,
Concurrency: 1,
Priority: 9,
GroupIDs: []int64{groupID},
},
}
concurrencyCache := schedulerTestConcurrencyCache{
acquireResults: map[int64]bool{21611: false, 21612: true},
loadMap: map[int64]*AccountLoadInfo{
21611: {AccountID: 21611, LoadRate: 0, WaitingCount: 0},
21612: {AccountID: 21612, LoadRate: 90, WaitingCount: 1},
},
}
svc := &OpenAIGatewayService{
accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts},
cache: &schedulerTestGatewayCache{},
cfg: newSchedulerTestSubscriptionPriorityConfig(),
rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true", "", "true"),
concurrencyService: NewConcurrencyService(concurrencyCache),
}
selection, decision, err := svc.SelectAccountWithScheduler(ctx, &groupID, "", "session_subscription_fallback", "gpt-5.1", nil, OpenAIUpstreamTransportAny, false)
require.NoError(t, err)
require.NotNil(t, selection)
require.NotNil(t, selection.Account)
require.Equal(t, int64(21612), selection.Account.ID)
require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer)
require.True(t, selection.Acquired)
if selection.ReleaseFunc != nil {
selection.ReleaseFunc()
}
}
func TestOpenAIGatewayService_SelectAccountWithScheduler_SubscriptionPriorityDisabledUsesScore(t *testing.T) {
ctx := context.Background()
groupID := int64(10122)
accounts := []Account{
{
ID: 21621,
Platform: PlatformOpenAI,
Type: AccountTypeOAuth,
Status: StatusActive,
Schedulable: true,
Concurrency: 1,
Priority: 10,
GroupIDs: []int64{groupID},
Credentials: map[string]any{"plan_type": "pro"},
},
{
ID: 21622,
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Status: StatusActive,
Schedulable: true,
Concurrency: 1,
Priority: 0,
GroupIDs: []int64{groupID},
},
}
concurrencyCache := schedulerTestConcurrencyCache{
acquireResults: map[int64]bool{21621: true, 21622: true},
loadMap: map[int64]*AccountLoadInfo{
21621: {AccountID: 21621, LoadRate: 90, WaitingCount: 1},
21622: {AccountID: 21622, LoadRate: 0, WaitingCount: 0},
},
}
svc := &OpenAIGatewayService{
accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts},
cache: &schedulerTestGatewayCache{},
cfg: newSchedulerTestSubscriptionPriorityConfig(),
rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true", "", "false"),
concurrencyService: NewConcurrencyService(concurrencyCache),
}
selection, decision, err := svc.SelectAccountWithScheduler(ctx, &groupID, "", "session_subscription_disabled", "gpt-5.1", nil, OpenAIUpstreamTransportAny, false)
require.NoError(t, err)
require.NotNil(t, selection)
require.NotNil(t, selection.Account)
require.Equal(t, int64(21622), selection.Account.ID)
require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer)
if selection.ReleaseFunc != nil {
selection.ReleaseFunc()
}
}
func TestOpenAIGatewayService_SelectAccountWithScheduler_UsesAccountPriorityWithinGroupPool(t *testing.T) {
ctx := context.Background()
groupID := int64(10123)
accounts := []Account{
{
ID: 21631,
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Status: StatusActive,
Schedulable: true,
Concurrency: 1,
Priority: 1,
AccountGroups: []AccountGroup{
{AccountID: 21631, GroupID: groupID, Priority: 100},
},
GroupIDs: []int64{groupID},
},
{
ID: 21632,
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Status: StatusActive,
Schedulable: true,
Concurrency: 1,
Priority: 100000,
AccountGroups: []AccountGroup{
{AccountID: 21632, GroupID: groupID, Priority: 1},
},
GroupIDs: []int64{groupID},
},
}
cfg := newSchedulerTestSubscriptionPriorityConfig()
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Load = 0
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Queue = 0
svc := &OpenAIGatewayService{
accountRepo: schedulerGroupAwareOpenAIAccountRepo{schedulerTestOpenAIAccountRepo{accounts: accounts}},
cache: &schedulerTestGatewayCache{},
cfg: cfg,
rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true"),
concurrencyService: NewConcurrencyService(schedulerTestConcurrencyCache{}),
}
selection, decision, err := svc.SelectAccountWithScheduler(ctx, &groupID, "", "session_group_priority", "gpt-5.1", nil, OpenAIUpstreamTransportAny, false)
require.NoError(t, err)
require.NotNil(t, selection)
require.NotNil(t, selection.Account)
require.Equal(t, int64(21631), selection.Account.ID)
require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer)
if selection.ReleaseFunc != nil {
selection.ReleaseFunc()
}
}
func TestDefaultOpenAIAccountScheduler_ShouldEscapeStickyAccount_ThresholdBoundary(t *testing.T) {
stats := newOpenAIAccountRuntimeStats()
accountID := int64(21501)
@@ -2421,3 +2940,141 @@ func TestDefaultOpenAIAccountScheduler_IsAccountTransportCompatible_Branches(t *
func int64PtrForTest(v int64) *int64 {
return &v
}
func TestOpenAIGatewayService_SelectAccountWithScheduler_StickyWeightedFallbackSkipsOutOfGroupStickyAccount(t *testing.T) {
resetOpenAIAdvancedSchedulerSettingCacheForTest()
ctx := context.Background()
groupID := int64(101081)
otherGroupID := int64(101082)
accounts := []Account{
{
ID: 38001,
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Status: StatusActive,
Schedulable: true,
Concurrency: 1,
Priority: 10,
GroupIDs: []int64{groupID},
},
{
// 会话粘连绑定指向的账号已被移出请求分组(绑定 TTL 内账号改组的场景)。
ID: 38002,
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Status: StatusActive,
Schedulable: true,
Concurrency: 1,
Priority: 0,
GroupIDs: []int64{otherGroupID},
},
}
cfg := &config.Config{}
cfg.Gateway.OpenAIWS.LBTopK = 2
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Priority = 1
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Load = 1
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Queue = 0.7
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.ErrorRate = 0.8
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.TTFT = 0.5
cfg.Gateway.OpenAIWS.SchedulerScoreWeights.SessionSticky = 3
cache := &schedulerTestGatewayCache{sessionBindings: map[string]int64{
"openai:session_weighted_out_of_group": 38002,
}}
concurrencyCache := schedulerTestConcurrencyCache{
acquireResults: map[int64]bool{38001: false, 38002: true},
}
svc := &OpenAIGatewayService{
accountRepo: schedulerGroupAwareOpenAIAccountRepo{schedulerTestOpenAIAccountRepo{accounts: accounts}},
cache: cache,
cfg: cfg,
rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true", "true"),
concurrencyService: NewConcurrencyService(concurrencyCache),
}
selection, decision, err := svc.SelectAccountWithScheduler(
ctx,
&groupID,
"",
"session_weighted_out_of_group",
"gpt-5.1",
nil,
OpenAIUpstreamTransportAny,
false,
)
require.NoError(t, err)
require.NotNil(t, selection)
require.NotNil(t, selection.Account)
// 组内唯一候选 38001 满并发:必须返回其等待计划,绝不能把请求泄漏到组外的粘连账号 38002。
require.Equal(t, int64(38001), selection.Account.ID)
require.False(t, selection.Acquired)
require.NotNil(t, selection.WaitPlan)
require.Equal(t, int64(38001), selection.WaitPlan.AccountID)
require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer)
// 失效的粘连绑定应被清理,避免后续请求反复走同一条泄漏路径。
require.Positive(t, cache.deletedSessions["openai:session_weighted_out_of_group"])
}
func TestOpenAIGatewayService_SelectAccountWithScheduler_SubscriptionPriorityWaitsOnBusySubscriptionWhenRegularUnusable(t *testing.T) {
resetOpenAIAdvancedSchedulerSettingCacheForTest()
ctx := context.Background()
groupID := int64(101091)
accounts := []Account{
{
// 订阅账号:支持 compact,但并发已满(busy-but-waitable)。
ID: 38011,
Platform: PlatformOpenAI,
Type: AccountTypeOAuth,
Status: StatusActive,
Schedulable: true,
Concurrency: 1,
Priority: 0,
GroupIDs: []int64{groupID},
Credentials: map[string]any{"plan_type": "team"},
Extra: map[string]any{"openai_compact_supported": true},
},
{
// 常规账号:明确不支持 compact,无法服务本次请求。
ID: 38012,
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Status: StatusActive,
Schedulable: true,
Concurrency: 1,
Priority: 9,
GroupIDs: []int64{groupID},
Extra: map[string]any{"openai_compact_supported": false},
},
}
concurrencyCache := schedulerTestConcurrencyCache{
acquireResults: map[int64]bool{38011: false, 38012: true},
}
svc := &OpenAIGatewayService{
accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts},
cache: &schedulerTestGatewayCache{},
cfg: newSchedulerTestSubscriptionPriorityConfig(),
rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true", "", "true"),
concurrencyService: NewConcurrencyService(concurrencyCache),
}
selection, decision, err := svc.SelectAccountWithScheduler(
ctx,
&groupID,
"",
"session_subscription_wait",
"gpt-5.1",
nil,
OpenAIUpstreamTransportAny,
true,
)
// 常规池无可用候选时,忙碌的订阅账号应产生等待计划,而不是直接返回 no available accounts。
require.NoError(t, err)
require.NotNil(t, selection)
require.NotNil(t, selection.Account)
require.Equal(t, int64(38011), selection.Account.ID)
require.False(t, selection.Acquired)
require.NotNil(t, selection.WaitPlan)
require.Equal(t, int64(38011), selection.WaitPlan.AccountID)
require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer)
}
@@ -149,6 +149,9 @@ func (s *AccountTestService) ProbeOpenAIAPIKeyResponsesSupport(ctx context.Conte
req.Header.Set("Authorization", "Bearer "+apiKey)
req.Header.Set("Accept", "application/json")
// 账号级请求头覆写:能力探测与真实转发保持一致的最终头
account.ApplyHeaderOverrides(req.Header)
proxyURL := ""
if account.ProxyID != nil && account.Proxy != nil {
proxyURL = account.Proxy.URL()
@@ -1,6 +1,7 @@
package service
import (
"fmt"
"net/http"
"github.com/Wei-Shaw/sub2api/internal/config"
@@ -8,6 +9,11 @@ import (
"github.com/gin-gonic/gin"
)
// CodexOfficialClientsOnlyMessage 是 codex_cli_only 拒绝时面向客户端的通用兜底文案。
// 仅当拒绝原因不是「可解析版本但越界」(VersionTooLow/VersionTooHigh)时使用:
// 未命中官方/黑名单/缺指纹/版本无法识别都沿用这句(避免向伪装客户端泄露门控细节)。
const CodexOfficialClientsOnlyMessage = "This account only allows Codex official clients"
const (
// CodexClientRestrictionReasonDisabled 表示账号未开启 codex_cli_only。
CodexClientRestrictionReasonDisabled = "codex_cli_only_disabled"
@@ -51,6 +57,13 @@ type CodexClientRestrictionDetectionResult struct {
Enabled bool
Matched bool
Reason string
// DetectedVersion 是从官方 UA 解析出的 Codex 引擎版本;仅在版本门拒绝
// (VersionTooLow / VersionTooHigh) 时填充,供面向客户端的差异化文案使用。
DetectedVersion string
// MinCodexVersion 是触发 VersionTooLow 时的最低要求版本(来自策略快照)。
MinCodexVersion string
// MaxCodexVersion 是触发 VersionTooHigh 时的最高允许版本(来自策略快照)。
MaxCodexVersion string
}
// CodexClientRestrictionDetector 定义 codex_cli_only 统一检测入口。
@@ -127,10 +140,22 @@ func (d *OpenAICodexClientRestrictionDetector) Detect(c *gin.Context, account *A
return CodexClientRestrictionDetectionResult{Enabled: true, Matched: false, Reason: CodexClientRestrictionReasonVersionUndetectable}
}
if policy.MinCodexVersion != "" && CompareVersions(ver, policy.MinCodexVersion) < 0 {
return CodexClientRestrictionDetectionResult{Enabled: true, Matched: false, Reason: CodexClientRestrictionReasonVersionTooLow}
return CodexClientRestrictionDetectionResult{
Enabled: true,
Matched: false,
Reason: CodexClientRestrictionReasonVersionTooLow,
DetectedVersion: ver,
MinCodexVersion: policy.MinCodexVersion,
}
}
if policy.MaxCodexVersion != "" && CompareVersions(ver, policy.MaxCodexVersion) > 0 {
return CodexClientRestrictionDetectionResult{Enabled: true, Matched: false, Reason: CodexClientRestrictionReasonVersionTooHigh}
return CodexClientRestrictionDetectionResult{
Enabled: true,
Matched: false,
Reason: CodexClientRestrictionReasonVersionTooHigh,
DetectedVersion: ver,
MaxCodexVersion: policy.MaxCodexVersion,
}
}
}
@@ -145,3 +170,22 @@ func (d *OpenAICodexClientRestrictionDetector) Detect(c *gin.Context, account *A
return CodexClientRestrictionDetectionResult{Enabled: true, Matched: true, Reason: reason}
}
// CodexClientRestrictionMessage 把检测结果映射为面向客户端的 403 文案。
// 仅版本越界(VersionTooLow/VersionTooHigh)给出带实际版本号与边界的差异化提示——
// 这类请求其实已被识别为官方 Codex(命中官方 UA/originator),再回「只允许官方客户端」会误导;
// 其余拒绝原因统一沿用通用兜底句,不暴露门控细节。
func CodexClientRestrictionMessage(r CodexClientRestrictionDetectionResult) string {
switch r.Reason {
case CodexClientRestrictionReasonVersionTooLow:
return fmt.Sprintf(
"Your Codex version (%s) is below the minimum required version (%s). Please update Codex.",
r.DetectedVersion, r.MinCodexVersion)
case CodexClientRestrictionReasonVersionTooHigh:
return fmt.Sprintf(
"Your Codex version (%s) exceeds the maximum allowed version (%s). Please downgrade Codex to %s or lower.",
r.DetectedVersion, r.MaxCodexVersion, r.MaxCodexVersion)
default:
return CodexOfficialClientsOnlyMessage
}
}
@@ -284,6 +284,66 @@ func TestDetect_V3_AppServerAndSkipAndVersionScope(t *testing.T) {
})
}
func TestDetect_VersionGateCarriesVersionFields(t *testing.T) {
gin.SetMode(gin.TestMode)
d := NewOpenAICodexClientRestrictionDetector(nil)
acc := func() *Account {
return &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth, Extra: map[string]any{"codex_cli_only": true}}
}
t.Run("版本太低:携带 DetectedVersion + MinCodexVersion", func(t *testing.T) {
c := newCodexDetectorTestContext("codex_cli_rs/0.39.0 (x)", "")
r := d.Detect(c, acc(), CodexRestrictionPolicy{MinCodexVersion: "0.42.0"}, nil)
require.False(t, r.Matched)
require.Equal(t, CodexClientRestrictionReasonVersionTooLow, r.Reason)
require.Equal(t, "0.39.0", r.DetectedVersion)
require.Equal(t, "0.42.0", r.MinCodexVersion)
})
t.Run("版本太高:携带 DetectedVersion + MaxCodexVersion", func(t *testing.T) {
c := newCodexDetectorTestContext("codex_cli_rs/0.45.0 (x)", "")
r := d.Detect(c, acc(), CodexRestrictionPolicy{MaxCodexVersion: "0.42.0"}, nil)
require.False(t, r.Matched)
require.Equal(t, CodexClientRestrictionReasonVersionTooHigh, r.Reason)
require.Equal(t, "0.45.0", r.DetectedVersion)
require.Equal(t, "0.42.0", r.MaxCodexVersion)
})
}
func TestCodexClientRestrictionMessage(t *testing.T) {
t.Run("版本太低:带实际版本与最低要求", func(t *testing.T) {
msg := CodexClientRestrictionMessage(CodexClientRestrictionDetectionResult{
Reason: CodexClientRestrictionReasonVersionTooLow,
DetectedVersion: "0.39.0",
MinCodexVersion: "0.42.0",
})
require.Equal(t, "Your Codex version (0.39.0) is below the minimum required version (0.42.0). Please update Codex.", msg)
})
t.Run("版本太高:带实际版本与最高允许", func(t *testing.T) {
msg := CodexClientRestrictionMessage(CodexClientRestrictionDetectionResult{
Reason: CodexClientRestrictionReasonVersionTooHigh,
DetectedVersion: "0.45.0",
MaxCodexVersion: "0.42.0",
})
require.Equal(t, "Your Codex version (0.45.0) exceeds the maximum allowed version (0.42.0). Please downgrade Codex to 0.42.0 or lower.", msg)
})
t.Run("无法识别版本:保持原通用句", func(t *testing.T) {
msg := CodexClientRestrictionMessage(CodexClientRestrictionDetectionResult{
Reason: CodexClientRestrictionReasonVersionUndetectable,
})
require.Equal(t, "This account only allows Codex official clients", msg)
})
t.Run("未命中官方:保持原通用句", func(t *testing.T) {
msg := CodexClientRestrictionMessage(CodexClientRestrictionDetectionResult{
Reason: CodexClientRestrictionReasonNotMatchedUA,
})
require.Equal(t, "This account only allows Codex official clients", msg)
})
}
func TestDetect_EngineFingerprintSignals(t *testing.T) {
gin.SetMode(gin.TestMode)
det := NewOpenAICodexClientRestrictionDetector(&config.Config{})
@@ -9,6 +9,9 @@ import (
)
var codexModelMap = map[string]string{
"gpt-5.6-sol": "gpt-5.6-sol",
"gpt-5.6-terra": "gpt-5.6-terra",
"gpt-5.6-luna": "gpt-5.6-luna",
"gpt-5.5": "gpt-5.5",
"gpt-5.5-pro": "gpt-5.5-pro",
"codex-auto-review": "codex-auto-review",
@@ -54,6 +57,9 @@ var codexVersionModelPrefixes = []struct {
prefix string
target string
}{
{prefix: "gpt-5.6-sol", target: "gpt-5.6-sol"},
{prefix: "gpt-5.6-terra", target: "gpt-5.6-terra"},
{prefix: "gpt-5.6-luna", target: "gpt-5.6-luna"},
{prefix: "gpt-5.3-codex-spark", target: "gpt-5.3-codex-spark"},
{prefix: "gpt-5.3-codex", target: "gpt-5.3-codex"},
{prefix: "gpt-5.4-mini", target: "gpt-5.4-mini"},
@@ -607,18 +613,21 @@ func hasOpenAIImageGenerationTool(reqBody map[string]any) bool {
return false
}
// stripCodexSparkImageGenerationTools removes image_generation tool entries from
// reqBody["tools"]. gpt-5.3-codex-spark rejects that tool upstream with HTTP 400
// (invalid_request_error, param=tools), and Codex CLI advertises it by default, so
// it must be dropped for spark. When the tools list becomes empty the key is removed.
// Returns true when the body was modified.
func stripCodexSparkImageGenerationTools(reqBody map[string]any) bool {
func stripOpenAIImageGenerationTools(reqBody map[string]any) bool {
rawTools, ok := reqBody["tools"]
if !ok || rawTools == nil {
if openAIAnyToolChoiceSelectsImageGeneration(reqBody["tool_choice"]) {
delete(reqBody, "tool_choice")
return true
}
return false
}
tools, ok := rawTools.([]any)
if !ok {
if openAIAnyToolChoiceSelectsImageGeneration(reqBody["tool_choice"]) {
delete(reqBody, "tool_choice")
return true
}
return false
}
filtered := make([]any, 0, len(tools))
@@ -631,17 +640,31 @@ func stripCodexSparkImageGenerationTools(reqBody map[string]any) bool {
}
filtered = append(filtered, rawTool)
}
if !removed {
if !removed && !openAIAnyToolChoiceSelectsImageGeneration(reqBody["tool_choice"]) {
return false
}
if len(filtered) == 0 {
delete(reqBody, "tools")
} else {
reqBody["tools"] = filtered
if removed {
if len(filtered) == 0 {
delete(reqBody, "tools")
} else {
reqBody["tools"] = filtered
}
}
if openAIAnyToolChoiceSelectsImageGeneration(reqBody["tool_choice"]) {
delete(reqBody, "tool_choice")
}
return true
}
// stripCodexSparkImageGenerationTools removes image_generation tool entries from
// reqBody["tools"]. gpt-5.3-codex-spark rejects that tool upstream with HTTP 400
// (invalid_request_error, param=tools), and Codex CLI advertises it by default, so
// it must be dropped for spark. When the tools list becomes empty the key is removed.
// Returns true when the body was modified.
func stripCodexSparkImageGenerationTools(reqBody map[string]any) bool {
return stripOpenAIImageGenerationTools(reqBody)
}
func hasOpenAIInputImage(reqBody map[string]any) bool {
if reqBody == nil {
return false
@@ -82,6 +82,9 @@ func (s *OpenAIGatewayService) ForwardEmbeddings(
upstreamReq.Header.Set("user-agent", customUA)
}
// 账号级请求头覆写(仅 openai api_key 账号启用时生效)
account.ApplyHeaderOverrides(upstreamReq.Header)
proxyURL := ""
if account.Proxy != nil {
proxyURL = account.Proxy.URL()
@@ -166,6 +166,9 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions(
upstreamReq.Header.Set("user-agent", "sub2api-grok/1.0")
}
// 账号级请求头覆写(仅 openai api_key 账号启用时生效)
account.ApplyHeaderOverrides(upstreamReq.Header)
// 6. Send request
proxyURL := ""
if account.Proxy != nil {
@@ -231,6 +231,9 @@ func (s *OpenAIGatewayService) buildInputTokensUpstreamRequest(
}
}
// 账号级请求头覆写(仅 openai api_key 账号启用时生效;OAuth 路径 no-op)
account.ApplyHeaderOverrides(req.Header)
return req, nil
}
@@ -1306,6 +1306,66 @@ func TestOpenAIGatewayServiceRecordUsage_ChannelMappedOverridesBillingModelWhenM
require.True(t, usageRepo.lastLog.ActualCost > 0, "cost must not be zero")
}
func TestOpenAIGatewayServiceRecordUsage_ResponsesMappedBillingModelHonorsBillingModelSource(t *testing.T) {
usage := OpenAIUsage{InputTokens: 20, OutputTokens: 10}
tokens := UsageTokens{InputTokens: 20, OutputTokens: 10}
tests := []struct {
name string
billingModelSource string
wantBillingModel string
}{
{
name: "upstream uses mapped billing model",
billingModelSource: BillingModelSourceUpstream,
wantBillingModel: "gpt-5.5",
},
{
name: "requested overrides mapped billing model",
billingModelSource: BillingModelSourceRequested,
wantBillingModel: "gpt-5.4",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
userRepo := &openAIRecordUsageUserRepoStub{}
subRepo := &openAIRecordUsageSubRepoStub{}
svc := newOpenAIRecordUsageServiceForTest(usageRepo, userRepo, subRepo, nil)
expectedCost, err := svc.billingService.CalculateCost(tt.wantBillingModel, tokens, 1.1)
require.NoError(t, err)
err = svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{
Result: &OpenAIForwardResult{
RequestID: "resp_mapped_billing_model_source",
Model: "gpt-5.4",
BillingModel: "gpt-5.5",
UpstreamModel: "gpt-5.5",
Usage: usage,
Duration: time.Second,
},
APIKey: &APIKey{ID: 10},
User: &User{ID: 20},
Account: &Account{ID: 30},
ChannelUsageFields: ChannelUsageFields{
OriginalModel: "gpt-5.4",
ChannelMappedModel: "gpt-5.4",
BillingModelSource: tt.billingModelSource,
},
})
require.NoError(t, err)
require.NotNil(t, usageRepo.lastLog)
require.Equal(t, "gpt-5.4", usageRepo.lastLog.Model)
require.InDelta(t, expectedCost.ActualCost, usageRepo.lastLog.ActualCost, 1e-12)
require.InDelta(t, expectedCost.ActualCost, userRepo.lastAmount, 1e-12)
require.True(t, usageRepo.lastLog.ActualCost > 0, "cost must not be zero")
})
}
}
func TestOpenAIGatewayServiceRecordUsage_BillsCompactOpenAIModelAlias(t *testing.T) {
usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
userRepo := &openAIRecordUsageUserRepoStub{}
@@ -1743,6 +1803,52 @@ func TestOpenAIGatewayServiceRecordUsage_ImageIndependentMultiplierUsesImageRate
require.Equal(t, string(BillingModeImage), *usageRepo.lastLog.BillingMode)
}
func TestGrokVideoMediaBillingUsesImageRateMultiplier(t *testing.T) {
mediaPrice2K := 0.4
groupID := int64(126)
usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
svc := newOpenAIRecordUsageServiceForTest(usageRepo, &openAIRecordUsageUserRepoStub{}, &openAIRecordUsageSubRepoStub{}, nil)
err := svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{
Result: &OpenAIForwardResult{
RequestID: "video-request-123",
ResponseID: "video-request-123",
Model: "grok-imagine-video-1.5",
BillingModel: "grok-imagine-video-1.5",
// The usage schema has no separate video count; video generation is billed as one media unit.
ImageCount: 1,
ImageSize: ImageBillingSize2K,
Duration: time.Second,
},
APIKey: &APIKey{
ID: 10126,
GroupID: i64p(groupID),
Group: &Group{
ID: groupID,
Platform: PlatformGrok,
RateMultiplier: 0.15,
ImageRateIndependent: true,
ImageRateMultiplier: 0.5,
ImagePrice2K: &mediaPrice2K,
},
},
User: &User{ID: 20126},
Account: &Account{ID: 30126, Platform: PlatformGrok},
})
require.NoError(t, err)
require.NotNil(t, usageRepo.lastLog)
require.Equal(t, "grok-imagine-video-1.5", usageRepo.lastLog.Model)
require.Equal(t, 1, usageRepo.lastLog.ImageCount)
require.Equal(t, ImageBillingSize2K, *usageRepo.lastLog.ImageSize)
require.InDelta(t, 0.4, usageRepo.lastLog.TotalCost, 1e-12)
require.InDelta(t, 0.2, usageRepo.lastLog.ActualCost, 1e-12)
require.InDelta(t, 0.5, usageRepo.lastLog.RateMultiplier, 1e-12)
require.NotNil(t, usageRepo.lastLog.BillingMode)
require.Equal(t, string(BillingModeImage), *usageRepo.lastLog.BillingMode)
}
func TestOpenAIGatewayServiceRecordUsage_ChannelImageBillingUsesImageCountAndSharedMultiplier(t *testing.T) {
groupID := int64(123)
usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
@@ -138,6 +138,9 @@ func (s *OpenAIGatewayService) forwardResponsesViaRawChatCompletions(
upstreamReq.Header.Set("user-agent", customUA)
}
// 账号级请求头覆写(仅 openai api_key 账号启用时生效)
account.ApplyHeaderOverrides(upstreamReq.Header)
proxyURL := ""
if account.Proxy != nil {
proxyURL = account.Proxy.URL()
@@ -2617,7 +2617,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
c.JSON(http.StatusForbidden, gin.H{
"error": gin.H{
"type": "forbidden_error",
"message": "This account only allows Codex official clients",
"message": CodexClientRestrictionMessage(restrictionResult),
},
})
return nil, errors.New("codex_cli_only restriction: only codex official clients are allowed")
@@ -2738,8 +2738,25 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
if apiKey != nil {
imageGenerationAllowed = GroupAllowsImageGeneration(apiKey.Group)
}
codexImageGenerationBridgeEnabled := isCodexCLI && imageGenerationAllowed && s.isCodexImageGenerationBridgeEnabled(ctx, account, apiKey)
imageIntent := IsImageGenerationIntent(openAIResponsesEndpoint, reqModel, body)
codexImageGenerationExplicitToolPolicy := codexImageGenerationExplicitToolPolicyAllow
if isCodexCLI {
codexImageGenerationExplicitToolPolicy = account.CodexImageGenerationExplicitToolPolicy()
}
codexImageGenerationBridgeEnabled := isCodexCLI && imageGenerationAllowed && codexImageGenerationExplicitToolPolicy != codexImageGenerationExplicitToolPolicyStrip && s.isCodexImageGenerationBridgeEnabled(ctx, account, apiKey)
var imageIntent bool
if isCodexCLI && codexImageGenerationExplicitToolPolicy == codexImageGenerationExplicitToolPolicyStrip {
decoded, decodeErr := ensureReqBody()
if decodeErr != nil {
return nil, decodeErr
}
if stripOpenAIImageGenerationTools(decoded) {
markDecodedModified()
logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Stripped /responses image_generation tool for Codex client by account policy")
}
imageIntent = IsImageGenerationIntentMap(openAIResponsesEndpoint, reqModel, decoded)
} else {
imageIntent = IsImageGenerationIntent(openAIResponsesEndpoint, reqModel, body)
}
if imageIntent && !imageGenerationAllowed {
MarkOpsClientBusinessLimited(c, OpsClientBusinessLimitedReasonLocalFeatureGate)
c.JSON(http.StatusForbidden, gin.H{"error": gin.H{"type": "permission_error", "message": ImageGenerationPermissionMessage()}})
@@ -3216,6 +3233,9 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
wsAttempts,
)
wsResult.UpstreamModel = upstreamModel
if wsResult.BillingModel == "" {
wsResult.BillingModel = billingModel
}
if wsResult.ImageCount > 0 {
wsResult.ImageSize = imageSizeTier
wsResult.ImageInputSize = imageInputSize
@@ -3363,6 +3383,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
ResponseID: responseID,
Usage: *usage,
Model: originalModel,
BillingModel: billingModel,
UpstreamModel: upstreamModel,
ServiceTier: serviceTier,
ReasoningEffort: reasoningEffort,
@@ -3762,6 +3783,9 @@ func (s *OpenAIGatewayService) buildUpstreamRequestOpenAIPassthrough(
req.Header.Set("content-type", "application/json")
}
// 账号级请求头覆写(仅 openai api_key 账号启用时生效;OAuth 路径 no-op)
account.ApplyHeaderOverrides(req.Header)
return req, nil
}
@@ -4547,6 +4571,9 @@ func (s *OpenAIGatewayService) buildUpstreamRequest(ctx context.Context, c *gin.
req.Header.Set("content-type", "application/json")
}
// 账号级请求头覆写(仅 openai api_key 账号启用时生效;OAuth 路径 no-op)
account.ApplyHeaderOverrides(req.Header)
return req, nil
}
@@ -59,6 +59,52 @@ func TestOpenAIGatewayService_GetCodexClientRestrictionDetector(t *testing.T) {
})
}
func TestOpenAIGatewayService_Forward_VersionGateMessage(t *testing.T) {
gin.SetMode(gin.TestMode)
newCtx := func() (*httptest.ResponseRecorder, *gin.Context) {
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(nil))
return rec, c
}
account := func() *Account {
return &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth, Extra: map[string]any{"codex_cli_only": true}}
}
body := []byte(`{"model":"gpt-5.1-codex"}`)
t.Run("版本太低:返回带版本号的差异化文案", func(t *testing.T) {
rec, c := newCtx()
svc := &OpenAIGatewayService{codexDetector: &stubCodexRestrictionDetector{result: CodexClientRestrictionDetectionResult{
Enabled: true,
Matched: false,
Reason: CodexClientRestrictionReasonVersionTooLow,
DetectedVersion: "0.39.0",
MinCodexVersion: "0.42.0",
}}}
_, err := svc.Forward(context.Background(), c, account(), body)
require.Error(t, err)
require.Equal(t, http.StatusForbidden, rec.Code)
require.Contains(t, rec.Body.String(), "Your Codex version (0.39.0) is below the minimum required version (0.42.0)")
require.NotContains(t, rec.Body.String(), "This account only allows Codex official clients")
})
t.Run("未命中官方:仍返回通用兜底文案", func(t *testing.T) {
rec, c := newCtx()
svc := &OpenAIGatewayService{codexDetector: &stubCodexRestrictionDetector{result: CodexClientRestrictionDetectionResult{
Enabled: true,
Matched: false,
Reason: CodexClientRestrictionReasonNotMatchedUA,
}}}
_, err := svc.Forward(context.Background(), c, account(), body)
require.Error(t, err)
require.Equal(t, http.StatusForbidden, rec.Code)
require.Contains(t, rec.Body.String(), "This account only allows Codex official clients")
})
}
func TestGetAPIKeyIDFromContext(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -196,6 +196,144 @@ func TestOpenAIGatewayService_Forward_MappedImageModelUsesImageGate(t *testing.T
require.Equal(t, http.StatusForbidden, rec.Code)
}
func TestOpenAIGatewayService_Forward_TextResponsesSetsBillingModelToMappedModel(t *testing.T) {
gin.SetMode(gin.TestMode)
upstream := &httpUpstreamRecorder{
resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}, "x-request-id": []string{"rid_text_mapped_billing"}},
Body: io.NopCloser(strings.NewReader(
`{"id":"resp_text_mapped","object":"response","model":"gpt-5.5","status":"completed","output":[{"type":"message","role":"assistant","content":[{"type":"output_text","text":"ok"}]}],"usage":{"input_tokens":20,"output_tokens":10,"total_tokens":30}}`,
)),
},
}
cfg := &config.Config{}
cfg.Security.URLAllowlist.Enabled = false
svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream}
account := &Account{
ID: 4,
Name: "openai-apikey",
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Concurrency: 1,
Credentials: map[string]any{
"api_key": "sk-test",
"base_url": "https://example.com",
"model_mapping": map[string]any{"gpt-5.4": "gpt-5.5"},
},
Extra: map[string]any{"use_responses_api": true},
}
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
body := []byte(`{"model":"gpt-5.4","stream":false,"input":"hello"}`)
result, err := svc.Forward(context.Background(), c, account, body)
require.NoError(t, err)
require.NotNil(t, result)
require.Equal(t, "gpt-5.4", result.Model)
require.Equal(t, "gpt-5.5", result.BillingModel)
require.Equal(t, "gpt-5.5", result.UpstreamModel)
require.Equal(t, "gpt-5.5", gjson.GetBytes(upstream.lastBody, "model").String())
require.Equal(t, 0, result.ImageCount)
}
func TestOpenAIGatewayService_Forward_TextResponsesWithoutMappingKeepsRequestedBillingModel(t *testing.T) {
gin.SetMode(gin.TestMode)
upstream := &httpUpstreamRecorder{
resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}, "x-request-id": []string{"rid_text_unmapped_billing"}},
Body: io.NopCloser(strings.NewReader(`{"id":"resp_text_unmapped","object":"response","model":"gpt-5.4","status":"completed","usage":{"input_tokens":20,"output_tokens":10,"total_tokens":30}}`)),
},
}
cfg := &config.Config{}
cfg.Security.URLAllowlist.Enabled = false
svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream}
account := &Account{
ID: 4,
Name: "openai-apikey",
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Concurrency: 1,
Credentials: map[string]any{
"api_key": "sk-test",
"base_url": "https://example.com",
},
Extra: map[string]any{"use_responses_api": true},
}
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
result, err := svc.Forward(context.Background(), c, account, []byte(`{"model":"gpt-5.4","stream":false,"input":"hello"}`))
require.NoError(t, err)
require.NotNil(t, result)
require.Equal(t, "gpt-5.4", result.Model)
require.Equal(t, "gpt-5.4", result.BillingModel)
require.Equal(t, "gpt-5.4", result.UpstreamModel)
}
func TestOpenAIGatewayService_Forward_TextResponsesBillingModelMatchesChatCompletions(t *testing.T) {
gin.SetMode(gin.TestMode)
cfg := &config.Config{}
cfg.Security.URLAllowlist.Enabled = false
account := &Account{
ID: 5,
Name: "openai-apikey",
Platform: PlatformOpenAI,
Type: AccountTypeAPIKey,
Concurrency: 1,
Credentials: map[string]any{
"api_key": "sk-test",
"base_url": "https://example.com",
"model_mapping": map[string]any{"gpt-5.4": "gpt-5.5"},
},
Extra: map[string]any{"use_responses_api": true},
}
responsesUpstream := &httpUpstreamRecorder{
resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}, "x-request-id": []string{"rid_responses_mapped_billing"}},
Body: io.NopCloser(strings.NewReader(
`{"id":"resp_native","object":"response","model":"gpt-5.5","status":"completed","output":[{"type":"message","role":"assistant","content":[{"type":"output_text","text":"ok"}]}],"usage":{"input_tokens":20,"output_tokens":10,"total_tokens":30}}`,
)),
},
}
responsesSvc := &OpenAIGatewayService{cfg: cfg, httpUpstream: responsesUpstream}
responsesRecorder := httptest.NewRecorder()
responsesCtx, _ := gin.CreateTestContext(responsesRecorder)
responsesCtx.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
SetOpenAIClientTransport(responsesCtx, OpenAIClientTransportHTTP)
responsesResult, err := responsesSvc.Forward(context.Background(), responsesCtx, account, []byte(`{"model":"gpt-5.4","stream":false,"input":"hello"}`))
require.NoError(t, err)
require.NotNil(t, responsesResult)
chatUpstream := &httpUpstreamRecorder{
resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"rid_chat_mapped_billing"}},
Body: io.NopCloser(strings.NewReader(
`data: {"type":"response.completed","response":{"id":"resp_chat","object":"response","model":"gpt-5.5","status":"completed","output":[{"type":"message","role":"assistant","content":[{"type":"output_text","text":"ok"}]}],"usage":{"input_tokens":20,"output_tokens":10,"total_tokens":30}}}` + "\n\n",
)),
},
}
chatSvc := &OpenAIGatewayService{cfg: cfg, httpUpstream: chatUpstream}
chatRecorder := httptest.NewRecorder()
chatCtx, _ := gin.CreateTestContext(chatRecorder)
chatCtx.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/chat/completions", nil)
chatResult, err := chatSvc.ForwardAsChatCompletions(context.Background(), chatCtx, account, []byte(`{"model":"gpt-5.4","stream":false,"messages":[{"role":"user","content":"hello"}]}`), "", "")
require.NoError(t, err)
require.NotNil(t, chatResult)
require.Equal(t, chatResult.BillingModel, responsesResult.BillingModel)
require.Equal(t, "gpt-5.5", responsesResult.BillingModel)
require.Equal(t, "gpt-5.5", chatResult.BillingModel)
}
func TestOpenAIGatewayService_Forward_TextDataImageDoesNotForceMapMarshal(t *testing.T) {
gin.SetMode(gin.TestMode)
upstream := &httpUpstreamRecorder{
@@ -152,6 +152,45 @@ func TestOpenAIGatewayServiceForward_ExplicitImageToolWorksWithBridgeDisabled(t
require.NotContains(t, instructions, "image_generation")
}
func TestOpenAIGatewayServiceForward_AccountPolicyStripsExplicitImageTool(t *testing.T) {
gin.SetMode(gin.TestMode)
upstream := &httpUpstreamRecorder{
resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(`{"id":"resp_stripped_image","model":"gpt-5.4","usage":{"input_tokens":2,"output_tokens":1}}`)),
},
}
svc := newOpenAIImageGenerationControlTestService(upstream)
c, _ := newOpenAIImageGenerationControlTestContext(true, "codex_cli_rs/0.98.0")
account := newOpenAIImageGenerationControlTestAccount()
account.Extra = map[string]any{
featureKeyCodexImageGenerationExplicitToolPolicy: codexImageGenerationExplicitToolPolicyStrip,
}
body := []byte(`{
"model":"gpt-5.4",
"input":"draw",
"stream":false,
"tools":[
{"type":"function","name":"shell","parameters":{"type":"object"}},
{"type":"image_generation","format":"jpeg"}
],
"tool_choice":{"type":"image_generation"}
}`)
result, err := svc.Forward(context.Background(), c, account, body)
require.NoError(t, err)
require.NotNil(t, result)
require.NotNil(t, upstream.lastReq)
require.False(t, gjson.GetBytes(upstream.lastBody, `tools.#(type=="image_generation")`).Exists())
require.True(t, gjson.GetBytes(upstream.lastBody, `tools.#(type=="function")`).Exists())
require.False(t, gjson.GetBytes(upstream.lastBody, "tool_choice").Exists())
instructions := gjson.GetBytes(upstream.lastBody, "instructions").String()
require.NotContains(t, instructions, "image_generation")
}
func TestOpenAIGatewayServiceForward_ChannelBridgeOverrideEnablesCodexInjection(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -760,6 +760,8 @@ func (s *OpenAIGatewayService) buildOpenAIImagesRequest(
if strings.TrimSpace(contentType) != "" {
req.Header.Set("Content-Type", contentType)
}
// 账号级请求头覆写(仅 openai api_key 账号启用时生效;OAuth 路径 no-op)
account.ApplyHeaderOverrides(req.Header)
return req, nil
}
@@ -65,6 +65,12 @@ func normalizeKnownOpenAICodexModel(model string) string {
}
switch {
case strings.Contains(normalized, "gpt-5.6-sol"):
return "gpt-5.6-sol"
case strings.Contains(normalized, "gpt-5.6-terra"):
return "gpt-5.6-terra"
case strings.Contains(normalized, "gpt-5.6-luna"):
return "gpt-5.6-luna"
case strings.Contains(normalized, "gpt-5.5-pro"):
return "gpt-5.5-pro"
case strings.Contains(normalized, "gpt-5.5"):
@@ -215,6 +215,84 @@ func ValidateFunctionCallOutputContextBytes(body []byte) FunctionCallOutputValid
return result
}
// ToolCallOutputContextCoverage 描述 input 中工具输出与可重建上下文的覆盖关系,
// 用于判断剥离 previous_response_id 后上游能否仅凭 input 重建工具续链。
type ToolCallOutputContextCoverage struct {
HasFunctionCallOutput bool
// ContextCoversAllCallIDs 表示每个工具输出的 call_id 都能在 input 内找到
// 同 call_id 的工具调用上下文项或同 id 的 item_reference,且不存在缺失 call_id 的输出。
// 任一输出无法由 input 自身重建时为 false,此时剥离 previous_response_id 会导致
// 上游以 "No tool call found for function call output" 拒绝请求。
ContextCoversAllCallIDs bool
}
// AnalyzeToolCallOutputContextCoverageBytes 全量扫描 input,按 call_id 精确匹配工具输出
// 与可重建上下文。不能复用 ValidateFunctionCallOutputContextBytes 的 HasToolCallContext:
// 该标志只代表"存在某一个上下文项",部分覆盖的续链仍会被上游拒绝。
func AnalyzeToolCallOutputContextCoverageBytes(body []byte) ToolCallOutputContextCoverage {
coverage := ToolCallOutputContextCoverage{}
if len(body) == 0 {
return coverage
}
input := parseRawJSONView(body).Get("input")
if !input.IsArray() {
return coverage
}
missingCallID := false
var outputCallIDs map[string]struct{}
var contextIDs map[string]struct{}
input.ForEach(func(_, item gjson.Result) bool {
if !item.IsObject() {
return true
}
itemType := item.Get("type").String()
switch {
case isCodexToolCallOutputItemType(itemType):
coverage.HasFunctionCallOutput = true
callID := strings.TrimSpace(item.Get("call_id").String())
if callID == "" {
missingCallID = true
return true
}
if outputCallIDs == nil {
outputCallIDs = make(map[string]struct{})
}
outputCallIDs[callID] = struct{}{}
case isCodexToolCallContextItemType(itemType):
callID := strings.TrimSpace(item.Get("call_id").String())
if callID == "" {
return true
}
if contextIDs == nil {
contextIDs = make(map[string]struct{})
}
contextIDs[callID] = struct{}{}
case itemType == "item_reference":
idValue := strings.TrimSpace(item.Get("id").String())
if idValue == "" {
return true
}
if contextIDs == nil {
contextIDs = make(map[string]struct{})
}
contextIDs[idValue] = struct{}{}
}
return true
})
if !coverage.HasFunctionCallOutput || missingCallID {
return coverage
}
for callID := range outputCallIDs {
if _, ok := contextIDs[callID]; !ok {
return coverage
}
}
coverage.ContextCoversAllCallIDs = true
return coverage
}
// ValidateFunctionCallOutputContext 为 handler 提供低开销校验结果:
// 1) 无工具输出直接返回
// 2) 若已存在工具调用上下文则提前返回
@@ -184,3 +184,109 @@ func TestValidateFunctionCallOutputContextBytesMatchesMapValidation(t *testing.T
})
}
}
func TestAnalyzeToolCallOutputContextCoverageBytes(t *testing.T) {
cases := []struct {
name string
body map[string]any
hasOutput bool
coversAllIDs bool
}{
{
name: "no_input",
body: map[string]any{"model": "gpt-5.1"},
hasOutput: false,
coversAllIDs: false,
},
{
name: "no_tool_output",
body: map[string]any{"input": []any{
map[string]any{"type": "message", "content": "hi"},
}},
hasOutput: false,
coversAllIDs: false,
},
{
name: "all_outputs_covered_by_context",
body: map[string]any{"input": []any{
map[string]any{"type": "function_call", "call_id": "call_a"},
map[string]any{"type": "function_call_output", "call_id": "call_a"},
}},
hasOutput: true,
coversAllIDs: true,
},
{
name: "all_outputs_covered_by_item_reference",
body: map[string]any{"input": []any{
map[string]any{"type": "function_call_output", "call_id": "call_a"},
map[string]any{"type": "item_reference", "id": "call_a"},
}},
hasOutput: true,
coversAllIDs: true,
},
{
// 关键回归用例:input 内存在某一个上下文项,但另一个输出的 call_id
// 只能由上游会话链(previous_response_id)解析——不可剥离。
name: "partial_coverage_not_movable",
body: map[string]any{"input": []any{
map[string]any{"type": "function_call", "call_id": "call_a"},
map[string]any{"type": "function_call_output", "call_id": "call_a"},
map[string]any{"type": "function_call_output", "call_id": "call_b"},
}},
hasOutput: true,
coversAllIDs: false,
},
{
name: "unrelated_context_does_not_cover",
body: map[string]any{"input": []any{
map[string]any{"type": "function_call", "call_id": "call_x"},
map[string]any{"type": "function_call_output", "call_id": "call_b"},
}},
hasOutput: true,
coversAllIDs: false,
},
{
name: "output_missing_call_id_not_movable",
body: map[string]any{"input": []any{
map[string]any{"type": "function_call", "call_id": "call_a"},
map[string]any{"type": "function_call_output"},
map[string]any{"type": "function_call_output", "call_id": "call_a"},
}},
hasOutput: true,
coversAllIDs: false,
},
{
name: "mixed_context_and_reference_cover_all",
body: map[string]any{"input": []any{
map[string]any{"type": "function_call", "call_id": "call_a"},
map[string]any{"type": "function_call_output", "call_id": "call_a"},
map[string]any{"type": "function_call_output", "call_id": "call_b"},
map[string]any{"type": "item_reference", "id": "call_b"},
}},
hasOutput: true,
coversAllIDs: true,
},
{
name: "all_codex_output_types_covered",
body: map[string]any{"input": []any{
map[string]any{"type": "tool_search_output", "call_id": "call_s"},
map[string]any{"type": "tool_search_call", "call_id": "call_s"},
map[string]any{"type": "mcp_tool_call_output", "call_id": "call_m"},
map[string]any{"type": "mcp_tool_call", "call_id": "call_m"},
}},
hasOutput: true,
coversAllIDs: true,
},
}
for _, tt := range cases {
t.Run(tt.name, func(t *testing.T) {
bodyBytes, err := json.Marshal(tt.body)
require.NoError(t, err)
coverage := AnalyzeToolCallOutputContextCoverageBytes(bodyBytes)
require.Equal(t, tt.hasOutput, coverage.HasFunctionCallOutput, "HasFunctionCallOutput")
require.Equal(t, tt.coversAllIDs, coverage.ContextCoversAllCallIDs, "ContextCoversAllCallIDs")
})
}
}
+135 -83
View File
@@ -1183,6 +1183,10 @@ func (s *OpenAIGatewayService) buildOpenAIWSHeaders(
headers.Set("user-agent", codexCLIUserAgent)
}
// 账号级请求头覆写(仅 openai api_key 账号启用时生效;OAuth 路径 no-op)。
// 覆盖所有 WS 模式(ctx_pool/dedicated/passthrough)的握手头。
account.ApplyHeaderOverrides(headers)
return headers, sessionResolution, nil
}
@@ -2448,11 +2452,15 @@ func stripCodexSparkImageGenerationToolFromRawPayload(payload []byte, model stri
if !isCodexSparkModel(model) || !openAIRequestBodyHasImageGenerationTool(payload) {
return payload, false, nil
}
return stripOpenAIImageGenerationToolFromRawPayload(payload)
}
func stripOpenAIImageGenerationToolFromRawPayload(payload []byte) ([]byte, bool, error) {
payloadMap := make(map[string]any)
if err := json.Unmarshal(payload, &payloadMap); err != nil {
return payload, false, err
}
if !stripCodexSparkImageGenerationTools(payloadMap) {
if !stripOpenAIImageGenerationTools(payloadMap) {
return payload, false, nil
}
rebuilt, err := json.Marshal(payloadMap)
@@ -2671,7 +2679,11 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
}
apiKey := getAPIKeyFromContext(c)
imageGenerationAllowed := GroupAllowsImageGeneration(apiKeyGroup(apiKey))
codexBridgeEnabled := isCodexCLI && imageGenerationAllowed && s.isCodexImageGenerationBridgeEnabled(ctx, account, apiKey)
codexImageGenerationExplicitToolPolicy := codexImageGenerationExplicitToolPolicyAllow
if isCodexCLI {
codexImageGenerationExplicitToolPolicy = account.CodexImageGenerationExplicitToolPolicy()
}
codexBridgeEnabled := isCodexCLI && imageGenerationAllowed && codexImageGenerationExplicitToolPolicy != codexImageGenerationExplicitToolPolicyStrip && s.isCodexImageGenerationBridgeEnabled(ctx, account, apiKey)
if codexBridgeEnabled {
payloadMap := make(map[string]any)
if err := json.Unmarshal(normalized, &payloadMap); err != nil {
@@ -2709,6 +2721,14 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
}
normalized = next
}
if isCodexCLI && codexImageGenerationExplicitToolPolicy == codexImageGenerationExplicitToolPolicyStrip {
if stripped, changed, stripErr := stripOpenAIImageGenerationToolFromRawPayload(normalized); stripErr != nil {
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", stripErr)
} else if changed {
normalized = stripped
logOpenAIWSModeInfo("ingress_ws_codex_image_tool_stripped_by_policy account_id=%d", account.ID)
}
}
if stripped, changed, stripErr := stripCodexSparkImageGenerationToolFromRawPayload(normalized, upstreamModel); stripErr != nil {
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", stripErr)
} else if changed {
@@ -4309,87 +4329,8 @@ func (s *OpenAIGatewayService) selectAccountByPreviousResponseIDForCapability(
if s == nil {
return nil, nil
}
responseID := strings.TrimSpace(previousResponseID)
if responseID == "" {
return nil, nil
}
store := s.getOpenAIWSStateStore()
if store == nil {
return nil, nil
}
accountID, err := store.GetResponseAccount(ctx, derefGroupID(groupID), responseID)
if err != nil || accountID <= 0 {
return nil, nil
}
if excludedIDs != nil {
if _, excluded := excludedIDs[accountID]; excluded {
return nil, nil
}
}
account, err := s.getSchedulableAccount(ctx, accountID)
if err != nil || account == nil {
_ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
return nil, nil
}
// 非 WSv2 场景(如 force_http/全局关闭)不应使用 previous_response_id 粘连,
// 以保持“回滚到 HTTP”后的历史行为一致性。
if s.getOpenAIWSProtocolResolver().Resolve(account).Transport != OpenAIUpstreamTransportResponsesWebsocketV2 {
return nil, nil
}
if shouldClearStickySession(account, requestedModel) || !account.IsOpenAI() || !account.IsSchedulable() {
_ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
return nil, nil
}
if !parentHealthyForShadow(account, s.parentAccountLookup(ctx)) {
_ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
return nil, nil
}
if requestedModel != "" && !account.IsModelSupported(requestedModel) {
return nil, nil
}
if !account.SupportsOpenAIEndpointCapability(requiredCapability) {
return nil, nil
}
// Quota auto-pause must also gate the previous_response_id sticky path; otherwise an
// account over its 5h/7d threshold keeps serving the same response chain even though
// normal scheduling skips it. Pause is transient, so fall through to normal scheduling
// without deleting the binding (the window may reset before the next turn).
if paused, _ := shouldAutoPauseOpenAIAccountByQuota(ctx, account); paused {
return nil, nil
}
if s.schedulerSnapshot != nil && s.accountRepo != nil {
latest, latestErr := s.accountRepo.GetByID(ctx, account.ID)
if latestErr != nil || latest == nil {
_ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
return nil, nil
}
if shouldClearStickySession(latest, requestedModel) || !latest.IsOpenAI() || !latest.IsSchedulable() {
_ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
return nil, nil
}
if !parentHealthyForShadow(latest, s.parentAccountLookup(ctx)) {
_ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
return nil, nil
}
if requestedModel != "" && !latest.IsModelSupported(requestedModel) {
return nil, nil
}
if !latest.SupportsOpenAIEndpointCapability(requiredCapability) {
return nil, nil
}
if paused, _ := shouldAutoPauseOpenAIAccountByQuota(ctx, latest); paused {
return nil, nil
}
if s.isOpenAIAccountRuntimeBlocked(latest) {
_ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
return nil, nil
}
account = latest
}
if requireCompact && openAICompactSupportTier(account) == 0 {
_ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
accountID, account, responseID, store := s.resolveAccountByPreviousResponseIDForCapability(ctx, groupID, previousResponseID, requestedModel, excludedIDs, requiredCapability, requireCompact)
if accountID <= 0 || account == nil || store == nil {
return nil, nil
}
@@ -4423,6 +4364,117 @@ func (s *OpenAIGatewayService) selectAccountByPreviousResponseIDForCapability(
return nil, nil
}
func (s *OpenAIGatewayService) ResolveAccountIDByPreviousResponseIDForScheduler(
ctx context.Context,
groupID *int64,
previousResponseID string,
requestedModel string,
excludedIDs map[int64]struct{},
requiredCapability OpenAIEndpointCapability,
requireCompact bool,
) int64 {
accountID, _, _, _ := s.resolveAccountByPreviousResponseIDForCapability(ctx, groupID, previousResponseID, requestedModel, excludedIDs, requiredCapability, requireCompact)
return accountID
}
func (s *OpenAIGatewayService) resolveAccountByPreviousResponseIDForCapability(
ctx context.Context,
groupID *int64,
previousResponseID string,
requestedModel string,
excludedIDs map[int64]struct{},
requiredCapability OpenAIEndpointCapability,
requireCompact bool,
) (int64, *Account, string, OpenAIWSStateStore) {
if s == nil {
return 0, nil, "", nil
}
responseID := strings.TrimSpace(previousResponseID)
if responseID == "" {
return 0, nil, "", nil
}
store := s.getOpenAIWSStateStore()
if store == nil {
return 0, nil, "", nil
}
accountID, err := store.GetResponseAccount(ctx, derefGroupID(groupID), responseID)
if err != nil || accountID <= 0 {
return 0, nil, "", nil
}
if excludedIDs != nil {
if _, excluded := excludedIDs[accountID]; excluded {
return 0, nil, "", nil
}
}
account, err := s.getSchedulableAccount(ctx, accountID)
if err != nil || account == nil {
_ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
return 0, nil, "", nil
}
// 非 WSv2 场景(如 force_http/全局关闭)不应使用 previous_response_id 粘连,
// 以保持“回滚到 HTTP”后的历史行为一致性。
if s.getOpenAIWSProtocolResolver().Resolve(account).Transport != OpenAIUpstreamTransportResponsesWebsocketV2 {
return 0, nil, "", nil
}
if shouldClearStickySession(account, requestedModel) || !account.IsOpenAI() || !account.IsSchedulable() {
_ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
return 0, nil, "", nil
}
if !parentHealthyForShadow(account, s.parentAccountLookup(ctx)) {
_ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
return 0, nil, "", nil
}
if requestedModel != "" && !account.IsModelSupported(requestedModel) {
return 0, nil, "", nil
}
if !account.SupportsOpenAIEndpointCapability(requiredCapability) {
return 0, nil, "", nil
}
// Quota auto-pause must also gate the previous_response_id sticky path; otherwise an
// account over its 5h/7d threshold keeps serving the same response chain even though
// normal scheduling skips it. Pause is transient, so fall through to normal scheduling
// without deleting the binding (the window may reset before the next turn).
if paused, _ := shouldAutoPauseOpenAIAccountByQuota(ctx, account); paused {
return 0, nil, "", nil
}
if s.schedulerSnapshot != nil && s.accountRepo != nil {
latest, latestErr := s.accountRepo.GetByID(ctx, account.ID)
if latestErr != nil || latest == nil {
_ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
return 0, nil, "", nil
}
if shouldClearStickySession(latest, requestedModel) || !latest.IsOpenAI() || !latest.IsSchedulable() {
_ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
return 0, nil, "", nil
}
if !parentHealthyForShadow(latest, s.parentAccountLookup(ctx)) {
_ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
return 0, nil, "", nil
}
if requestedModel != "" && !latest.IsModelSupported(requestedModel) {
return 0, nil, "", nil
}
if !latest.SupportsOpenAIEndpointCapability(requiredCapability) {
return 0, nil, "", nil
}
if paused, _ := shouldAutoPauseOpenAIAccountByQuota(ctx, latest); paused {
return 0, nil, "", nil
}
if s.isOpenAIAccountRuntimeBlocked(latest) {
_ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
return 0, nil, "", nil
}
account = latest
}
if requireCompact && openAICompactSupportTier(account) == 0 {
_ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
return 0, nil, "", nil
}
return accountID, account, responseID, store
}
func classifyOpenAIWSAcquireError(err error) string {
if err == nil {
return "acquire_conn"
@@ -169,6 +169,26 @@ func TestStripCodexSparkImageGenerationToolFromRawPayload(t *testing.T) {
})
}
func TestStripOpenAIImageGenerationToolFromRawPayload(t *testing.T) {
payload := []byte(`{
"type":"response.create",
"model":"gpt-5.4",
"tools":[
{"type":"function","name":"shell"},
{"type":"image_generation","output_format":"png"}
],
"tool_choice":{"type":"image_generation"}
}`)
updated, changed, err := stripOpenAIImageGenerationToolFromRawPayload(payload)
require.NoError(t, err)
require.True(t, changed)
require.False(t, gjson.GetBytes(updated, `tools.#(type=="image_generation")`).Exists())
require.True(t, gjson.GetBytes(updated, `tools.#(type=="function")`).Exists())
require.False(t, gjson.GetBytes(updated, "tool_choice").Exists())
}
func TestAlignStoreDisabledPreviousResponseID(t *testing.T) {
t.Parallel()
+34 -8
View File
@@ -1,6 +1,9 @@
package service
import "time"
import (
"strings"
"time"
)
type OpsSystemLog struct {
ID int64 `json:"id"`
@@ -65,17 +68,22 @@ type OpsErrorLog struct {
RequestedModel string `json:"requested_model"`
UpstreamModel string `json:"upstream_model"`
RequestType *int16 `json:"request_type"`
UserAgent string `json:"user_agent"`
// 关联 api_key 名称(LEFT JOIN api_keys 取得;软删只覆盖 key 列,name 保留,故已删 key 仍有原名)。
APIKeyName string `json:"api_key_name,omitempty"`
APIKeyDeleted bool `json:"api_key_deleted,omitempty"`
// 已删除 KEY 所有者(INVALID_API_KEY 且该 key 曾存在时的归因快照)。
// 认证失败行 user_id 为空,列表用户列以此回退显示所有者。
DeletedKeyOwnerUserID *int64 `json:"deleted_key_owner_user_id,omitempty"`
DeletedKeyOwnerEmail string `json:"deleted_key_owner_email,omitempty"`
}
type OpsErrorLogDetail struct {
OpsErrorLog
ErrorBody string `json:"error_body"`
UserAgent string `json:"user_agent"`
// Upstream context (optional)
UpstreamStatusCode *int `json:"upstream_status_code,omitempty"`
@@ -93,11 +101,10 @@ type OpsErrorLogDetail struct {
// vNext metric semantics
IsBusinessLimited bool `json:"is_business_limited"`
// Deleted key owner info (populated when INVALID_API_KEY and key was previously deleted)
AttemptedKeyPrefix string `json:"attempted_key_prefix,omitempty"`
DeletedKeyOwnerUserID *int64 `json:"deleted_key_owner_user_id,omitempty"`
DeletedKeyOwnerEmail string `json:"deleted_key_owner_email,omitempty"`
DeletedKeyName string `json:"deleted_key_name,omitempty"`
// Deleted key owner info (populated when INVALID_API_KEY and key was previously deleted).
// OwnerUserID/OwnerEmail 已上移到 OpsErrorLog(列表用户列回退需要)。
AttemptedKeyPrefix string `json:"attempted_key_prefix,omitempty"`
DeletedKeyName string `json:"deleted_key_name,omitempty"`
// Bound (non-deleted) key prefix, snapshotted at error time; mutually exclusive with AttemptedKeyPrefix.
APIKeyPrefix string `json:"api_key_prefix,omitempty"`
@@ -142,8 +149,14 @@ type OpsErrorLogFilter struct {
// ExcludeCountTokens drops count_tokens probe errors (is_count_tokens=true).
ExcludeCountTokens bool
// IncludeRecoveredUpstream 显式豁免 status>=400 守卫(仅在 Phase=="upstream" 时生效):
// ops 专用上游错误列表需要看到 status<400 的 recovered upstream 行。
// 请求错误语义的端点不设此开关,phase=upstream 过滤照常生效且守卫保留。
IncludeRecoveredUpstream bool
// ErrorPhasesAny / ErrorTypesAny add plain ANY() filters WITHOUT touching the
// special-cased single `Phase` field (only Phase=="upstream" bypasses the status>=400 clause).
// special-cased single `Phase` field (only Phase=="upstream" with
// IncludeRecoveredUpstream bypasses the status>=400 clause).
// NOTE: these ANY filters do NOT bypass status>=400; records with error_phase='upstream'
// but status_code<400 (recovered upstream errors) remain excluded.
// Used to map user-facing coarse categories to backend conditions.
@@ -158,6 +171,19 @@ type OpsErrorLogFilter struct {
Page int
PageSize int
// SortBy/SortOrder: server-side sorting aligned with the usage-log list.
// Repo whitelists columns (created_at/model/status_code); anything else
// falls back to created_at. SortOrder is "asc"/"desc" (default desc).
SortBy string
SortOrder string
}
// SetSort normalizes raw sort_by/sort_order query values into the filter.
// Shared by the admin and user-facing error list handlers.
func (f *OpsErrorLogFilter) SetSort(sortBy, sortOrder string) {
f.SortBy = strings.TrimSpace(sortBy)
f.SortOrder = strings.TrimSpace(sortOrder)
}
type OpsErrorLogList struct {
+5 -3
View File
@@ -359,10 +359,12 @@ func (s *OpsService) ListUserErrorRequests(ctx context.Context, userID int64, fi
filter.UserQuery = ""
filter.Owner = ""
filter.Source = ""
// 清空 Phase 是防御:Phase 是单值特殊字段,仅当其 == "upstream" 时 buildOpsErrorLogsWhere 才跳过 status>=400 子句。
// 用户端一律改走 category→ErrorPhasesAny/ErrorTypesAny(纯 ANY 过滤,不影响 status>=400 子句),
// 因此 recovered upstream(error_phase='upstream' 但 status<400,最终成功返回)记录对用户不可见——符合预期。
// 清空 Phase 是防御:用户端一律改走 category→ErrorPhasesAny/ErrorTypesAny
//(纯 ANY 过滤,不影响 status>=400 子句)。守卫豁免现在还需要
// IncludeRecoveredUpstream(用户端永不设置),recovered upstream
//(error_phase='upstream' 但 status<400,最终成功返回)记录对用户不可见——符合预期。
filter.Phase = ""
filter.IncludeRecoveredUpstream = false
list, err := s.opsRepo.ListErrorLogs(ctx, filter)
if err != nil {
@@ -184,16 +184,16 @@ func TestGetUserErrorRequestDetail_DeletedKeyOwnerAccess(t *testing.T) {
mk := func() *OpsErrorLogDetail {
return &OpsErrorLogDetail{
OpsErrorLog: OpsErrorLog{
ID: 55,
Phase: "auth",
Type: "api_error",
StatusCode: 401,
Message: "Invalid API key",
UserID: nil,
APIKeyName: "my-old-key",
APIKeyDeleted: true,
ID: 55,
Phase: "auth",
Type: "api_error",
StatusCode: 401,
Message: "Invalid API key",
UserID: nil,
APIKeyName: "my-old-key",
APIKeyDeleted: true,
DeletedKeyOwnerUserID: &ownerUID,
},
DeletedKeyOwnerUserID: &ownerUID,
}
}
+19 -2
View File
@@ -3,9 +3,12 @@ package service
import "time"
// UserErrorRequest 是面向终端用户的"错误请求"精简脱敏视图(白名单)。
// 严禁包含 client_ip / user_agent / account / api_key_prefix / upstream_endpoint /
// user_email 等敏感或内部字段。注:message(网关标准化错误描述)与 key_name
// 严禁包含 account / api_key_prefix / upstream_endpoint / user_email 等
// 敏感或内部字段。注:message(网关标准化错误描述)与 key_name
// (用户自有 API Key 名称,KeysView 中本就可见)经产品决策对该用户开放;
// client_ip / user_agent / group_name / request_type / stream 均为该用户
// 自己请求的属性,经产品决策(2026-07-03)开放,
// 与用量明细已向用户展示自身 ip_address/user_agent/分组/类型 的口径对齐;
// error_body 仅在详情接口(GetUserErrorRequestDetail)按归属校验后返回。
type UserErrorRequest struct {
ID int64 `json:"id"`
@@ -18,6 +21,11 @@ type UserErrorRequest struct {
Message string `json:"message"`
KeyName string `json:"key_name"`
KeyDeleted bool `json:"key_deleted"`
ClientIP string `json:"client_ip,omitempty"`
GroupName string `json:"group_name,omitempty"`
RequestType *int16 `json:"request_type,omitempty"`
Stream bool `json:"stream"`
UserAgent string `json:"user_agent,omitempty"`
}
// UserErrorRequestList 是用户错误请求分页结果。
@@ -90,6 +98,10 @@ func ToUserErrorRequest(e *OpsErrorLog) *UserErrorRequest {
if model == "" {
model = e.Model
}
clientIP := ""
if e.ClientIP != nil {
clientIP = *e.ClientIP
}
return &UserErrorRequest{
ID: e.ID,
CreatedAt: e.CreatedAt,
@@ -101,6 +113,11 @@ func ToUserErrorRequest(e *OpsErrorLog) *UserErrorRequest {
Message: e.Message,
KeyName: e.APIKeyName,
KeyDeleted: e.APIKeyDeleted,
ClientIP: clientIP,
GroupName: e.GroupName,
RequestType: e.RequestType,
Stream: e.Stream,
UserAgent: e.UserAgent,
}
}

Some files were not shown because too many files have changed in this diff Show More