mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
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:
@@ -1 +1 @@
|
||||
0.1.143
|
||||
0.1.145
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 +
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 映射
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"`
|
||||
|
||||
@@ -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"`
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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"},
|
||||
|
||||
@@ -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))
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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"))
|
||||
}
|
||||
@@ -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) {
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user