diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index 75786b9c70..2b339f3d96 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -96,9 +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) - batchImageRepository := repository.NewBatchImageRepository(db) - batchImageQueue := repository.NewBatchImageQueue(redisClient, configConfig) - batchImageDownloadLimiter := repository.NewBatchImageDownloadLimiter(redisClient, configConfig) usageBillingRepository := repository.NewUsageBillingRepository(client, db) gatewayCache := repository.NewGatewayCache(redisClient) schedulerOutboxRepository := repository.NewSchedulerOutboxRepository(db) @@ -137,11 +134,6 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { channelRepository := repository.NewChannelRepository(db) channelService := service.NewChannelService(channelRepository, groupRepository, apiKeyAuthCacheInvalidator, pricingService) modelPricingResolver := service.NewModelPricingResolver(channelService, billingService) - batchImageModelPricingResolver := service.ProvideBatchImageModelPricingResolver(modelPricingResolver) - batchImagePublicService := service.NewBatchImagePublicService(batchImageRepository, accountRepository, groupRepository, userGroupRateRepository, batchImageQueue, batchImageModelPricingResolver, usageBillingRepository, apiKeyAuthCacheInvalidator, configConfig) - batchImageDownloadService := service.NewBatchImageDownloadService(batchImageRepository, accountRepository, batchImageDownloadLimiter, configConfig) - batchImageCleanupService := service.ProvideBatchImageCleanupService(batchImageRepository, accountRepository, configConfig) - batchImageWorkerRuntime := service.ProvideBatchImageWorkerRuntime(batchImageRepository, accountRepository, batchImageQueue, usageBillingRepository, usageLogRepository, batchImageModelPricingResolver, apiKeyAuthCacheInvalidator, configConfig) notificationEmailService := service.NewNotificationEmailService(settingRepository, emailService) balanceNotifyService := service.ProvideBalanceNotifyService(emailService, settingRepository, accountRepository, notificationEmailService) gatewayService := service.NewGatewayService(accountRepository, groupRepository, usageLogRepository, usageBillingRepository, userRepository, userSubscriptionRepository, userGroupRateRepository, gatewayCache, configConfig, schedulerSnapshotService, concurrencyService, billingService, rateLimitService, billingCacheService, identityService, httpUpstream, deferredService, claudeTokenProvider, sessionLimitCache, rpmCache, digestSessionStore, settingService, tlsFingerprintProfileService, channelService, modelPricingResolver, balanceNotifyService, serviceUserPlatformQuotaRepository) @@ -191,10 +183,10 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { antigravityQuotaFetcher := service.NewAntigravityQuotaFetcher(proxyRepository) grokQuotaFetcher := service.NewGrokQuotaFetcher() grokQuotaService := service.ProvideGrokQuotaService(accountRepository, proxyRepository, grokTokenProvider, httpUpstream, usageLogRepository) - openAIQuotaService := service.ProvideOpenAIQuotaService(accountRepository, proxyRepository, openAITokenProvider, privacyClientFactory) + openAIQuotaService := service.ProvideOpenAIQuotaService(accountRepository, proxyRepository, openAITokenProvider, privacyClientFactory, openAIGatewayService) usageCache := service.NewUsageCache() - accountUsageService := service.NewAccountUsageService(accountRepository, usageLogRepository, claudeUsageFetcher, geminiQuotaService, antigravityQuotaFetcher, grokQuotaFetcher, grokQuotaService, openAIQuotaService, usageCache, identityCache, tlsFingerprintProfileService) - accountTestService := service.NewAccountTestService(accountRepository, geminiTokenProvider, claudeTokenProvider, grokTokenProvider, antigravityGatewayService, httpUpstream, configConfig, tlsFingerprintProfileService) + accountUsageService := service.ProvideAccountUsageService(accountRepository, usageLogRepository, claudeUsageFetcher, geminiQuotaService, antigravityQuotaFetcher, grokQuotaFetcher, grokQuotaService, openAIQuotaService, usageCache, identityCache, tlsFingerprintProfileService, openAIGatewayService) + accountTestService := service.ProvideAccountTestService(accountRepository, geminiTokenProvider, claudeTokenProvider, grokTokenProvider, antigravityGatewayService, httpUpstream, configConfig, tlsFingerprintProfileService, openAIGatewayService) crsSyncService := service.NewCRSSyncService(accountRepository, proxyRepository, oAuthService, openAIOAuthService, geminiOAuthService, configConfig) accountHandler := admin.ProvideAccountHandler(adminService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, rateLimitService, accountUsageService, accountTestService, concurrencyService, crsSyncService, sessionLimitCache, rpmCache, compositeTokenCacheInvalidator, grokQuotaService) adminAnnouncementHandler := admin.NewAnnouncementHandler(announcementService) @@ -267,6 +259,13 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { handlerPaymentHandler := handler.NewPaymentHandler(paymentService, paymentConfigService) paymentWebhookHandler := handler.NewPaymentWebhookHandler(paymentService, registry) availableChannelHandler := handler.NewAvailableChannelHandler(channelService, apiKeyService, settingService) + batchImageRepository := repository.NewBatchImageRepository(db) + batchImageQueue := repository.NewBatchImageQueue(redisClient, configConfig) + batchImageModelPricingResolver := service.ProvideBatchImageModelPricingResolver(modelPricingResolver) + batchImagePublicService := service.NewBatchImagePublicService(batchImageRepository, accountRepository, groupRepository, userGroupRateRepository, batchImageQueue, batchImageModelPricingResolver, usageBillingRepository, apiKeyAuthCacheInvalidator, configConfig) + batchImageDownloadLimiter := repository.NewBatchImageDownloadLimiter(redisClient, configConfig) + batchImageDownloadService := service.NewBatchImageDownloadService(batchImageRepository, accountRepository, batchImageDownloadLimiter, configConfig) + batchImageCleanupService := service.ProvideBatchImageCleanupService(batchImageRepository, accountRepository, configConfig) batchImageHandler := handler.NewBatchImageHandler(batchImagePublicService, batchImageDownloadService, batchImageCleanupService) idempotencyCoordinator := service.ProvideIdempotencyCoordinator(idempotencyRepository, configConfig) idempotencyCleanupService := service.ProvideIdempotencyCleanupService(idempotencyRepository, configConfig) @@ -285,6 +284,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { accountExpiryService := service.ProvideAccountExpiryService(accountRepository) proxyExpiryService := service.ProvideProxyExpiryService(proxyRepository) subscriptionExpiryService := service.ProvideSubscriptionExpiryService(userSubscriptionRepository, settingRepository, notificationEmailService, leaderLockCache, db) + batchImageWorkerRuntime := service.ProvideBatchImageWorkerRuntime(batchImageRepository, accountRepository, batchImageQueue, usageBillingRepository, usageLogRepository, batchImageModelPricingResolver, apiKeyAuthCacheInvalidator, configConfig) scheduledTestRunnerService := service.ProvideScheduledTestRunnerService(scheduledTestPlanRepository, scheduledTestService, accountTestService, rateLimitService, configConfig) paymentOrderExpiryService := service.ProvidePaymentOrderExpiryService(paymentService, leaderLockCache, db) channelMonitorRunner := service.ProvideChannelMonitorRunner(channelMonitorService, settingService) diff --git a/backend/internal/handler/admin/account_codex_agent_identity_import_test.go b/backend/internal/handler/admin/account_codex_agent_identity_import_test.go new file mode 100644 index 0000000000..db55182b6b --- /dev/null +++ b/backend/internal/handler/admin/account_codex_agent_identity_import_test.go @@ -0,0 +1,47 @@ +package admin + +import ( + "crypto/ed25519" + "crypto/rand" + "crypto/x509" + "encoding/base64" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" +) + +func TestNormalizeCodexImportEntryAcceptsAgentIdentityAuthJSON(t *testing.T) { + _, privateKey, err := ed25519.GenerateKey(rand.Reader) + require.NoError(t, err) + der, err := x509.MarshalPKCS8PrivateKey(privateKey) + require.NoError(t, err) + privateKeyBase64 := base64.StdEncoding.EncodeToString(der) + + item, err := normalizeCodexImportEntry(codexImportEntry{ + Index: 1, + Value: map[string]any{ + "auth_mode": "agentIdentity", + "agent_identity": map[string]any{ + "agent_runtime_id": "runtime-import", + "agent_private_key": privateKeyBase64, + "account_id": "account-import", + "chatgpt_user_id": "user-import", + "email": "agent@example.invalid", + "plan_type": "pro", + "chatgpt_account_is_fedramp": false, + }, + }, + }) + require.NoError(t, err) + require.NotNil(t, item) + require.True(t, item.IsAgentIdentity) + require.Equal(t, service.OpenAIAuthModeAgentIdentity, item.Credentials["auth_mode"]) + require.Equal(t, "runtime-import", item.Credentials["agent_runtime_id"]) + require.Equal(t, privateKeyBase64, item.Credentials["agent_private_key"]) + require.Equal(t, "account-import", item.Credentials["chatgpt_account_id"]) + require.Equal(t, "user-import", item.Credentials["chatgpt_user_id"]) + require.NotContains(t, item.Credentials, "access_token") + require.NotContains(t, item.Credentials, "refresh_token") + require.NotEmpty(t, item.WarningTexts) +} diff --git a/backend/internal/handler/admin/account_codex_import.go b/backend/internal/handler/admin/account_codex_import.go index 271bd19470..57b97f8476 100644 --- a/backend/internal/handler/admin/account_codex_import.go +++ b/backend/internal/handler/admin/account_codex_import.go @@ -72,20 +72,25 @@ type codexImportEntry struct { } type codexImportAccount struct { - Name string - AccessToken string - RefreshToken string - IDToken string - Email string - AccountID string - UserID string - PlanType string - Organization string - Credentials map[string]any - Extra map[string]any - TokenExpiresAt *time.Time - IdentityKeys []string - WarningTexts []string + Name string + AccessToken string + RefreshToken string + IDToken string + Email string + AccountID string + UserID string + PlanType string + Organization string + AgentRuntimeID string + AgentPrivateKey string + AgentTaskID string + AgentFedRAMP bool + IsAgentIdentity bool + Credentials map[string]any + Extra map[string]any + TokenExpiresAt *time.Time + IdentityKeys []string + WarningTexts []string } type codexJWTClaims struct { @@ -496,6 +501,41 @@ func normalizeCodexImportEntry(entry codexImportEntry) (*codexImportAccount, err case string: item.AccessToken = strings.TrimSpace(raw) case map[string]any: + if agentIdentity, ok := firstCodexMap(raw, []string{"agent_identity"}, []string{"agentIdentity"}); ok || strings.EqualFold(firstCodexString(raw, []string{"auth_mode"}, []string{"authMode"}), service.OpenAIAuthModeAgentIdentity) { + if !ok { + agentIdentity = raw + } + item.IsAgentIdentity = true + item.AgentRuntimeID = firstCodexString(agentIdentity, []string{"agent_runtime_id"}, []string{"agentRuntimeId"}) + item.AgentPrivateKey = firstCodexString(agentIdentity, []string{"agent_private_key"}, []string{"agentPrivateKey"}) + item.AgentTaskID = firstCodexString(agentIdentity, []string{"task_id"}, []string{"taskId"}) + item.AccountID = firstCodexString(agentIdentity, []string{"account_id"}, []string{"accountId"}) + item.UserID = firstCodexString(agentIdentity, []string{"chatgpt_user_id"}, []string{"chatgptUserId"}) + item.Email = firstCodexString(agentIdentity, []string{"email"}) + item.PlanType = firstCodexString(agentIdentity, []string{"plan_type"}, []string{"planType"}) + item.AgentFedRAMP = firstCodexBool(agentIdentity, []string{"chatgpt_account_is_fedramp"}, []string{"chatgptAccountIsFedramp"}) + if item.AgentRuntimeID == "" || item.AgentPrivateKey == "" || item.AccountID == "" || item.UserID == "" { + return nil, errors.New("agent identity 缺少必要字段") + } + if err := service.ValidateOpenAIAgentIdentityPrivateKey(item.AgentPrivateKey); err != nil { + return nil, errors.New("agent identity private key 格式无效") + } + item.Credentials["auth_mode"] = service.OpenAIAuthModeAgentIdentity + item.Credentials["agent_runtime_id"] = item.AgentRuntimeID + item.Credentials["agent_private_key"] = item.AgentPrivateKey + item.Credentials["chatgpt_account_id"] = item.AccountID + item.Credentials["chatgpt_user_id"] = item.UserID + item.Credentials["chatgpt_account_is_fedramp"] = item.AgentFedRAMP + setCodexCredentialIfNotEmpty(item.Credentials, "task_id", item.AgentTaskID) + setCodexCredentialIfNotEmpty(item.Credentials, "email", item.Email) + setCodexCredentialIfNotEmpty(item.Credentials, "plan_type", item.PlanType) + if item.AgentTaskID == "" { + item.WarningTexts = append(item.WarningTexts, "未包含 task_id,首次请求会使用现有 runtime 注册新 task") + } + item.IdentityKeys = buildCodexAgentIdentityKeys(item.AccountID, item.UserID, item.Email, item.AgentRuntimeID) + item.Name = buildCodexImportAccountName(item, entry.Index) + return item, nil + } item.AccessToken = firstCodexString(raw, []string{"tokens", "access_token"}, []string{"tokens", "accessToken"}, @@ -577,6 +617,9 @@ func normalizeCodexImportEntry(entry codexImportEntry) (*codexImportAccount, err return nil, fmt.Errorf("第 %d 条格式不支持", entry.Index) } + if item.IsAgentIdentity { + return item, nil + } if item.AccessToken == "" { return nil, errors.New("缺少 accessToken/access_token") } @@ -812,6 +855,9 @@ func sanitizeCodexImportCredentialExtras(input map[string]any) map[string]any { "openai_auth_mode": {}, "token_type": {}, "chatgpt_account_is_fedramp": {}, + "agent_runtime_id": {}, + "agent_private_key": {}, + "task_id": {}, } out := make(map[string]any, len(input)) for key, value := range input { @@ -842,6 +888,14 @@ func buildCodexImportIdentityKeys(accountID, userID, email, accessToken, refresh return buildCodexStoredIdentityKeys(accountID, userID, email, accessToken) } +func buildCodexAgentIdentityKeys(accountID, userID, email, runtimeID string) []string { + keys := buildCodexStoredIdentityKeys(accountID, userID, email, "") + if runtimeID = strings.TrimSpace(runtimeID); runtimeID != "" { + keys = append([]string{"agent:" + runtimeID}, keys...) + } + return keys +} + // buildCodexStoredIdentityKeys 生成存量账号索引键,保留 user/account 维度, // 让 accessToken-only 账号后续升级为完整 OAuth 时仍能命中并更新原账号。 func buildCodexStoredIdentityKeys(accountID, userID, email, accessToken string) []string { @@ -891,6 +945,10 @@ func (i *codexAccountIndex) Add(account service.Account) { for _, key := range keys { i.accountsByKey[key] = upsertCodexAccount(i.accountsByKey[key], account) } + if runtimeID := codexCredentialString(account.Credentials, "agent_runtime_id"); runtimeID != "" { + key := "agent:" + runtimeID + i.accountsByKey[key] = upsertCodexAccount(i.accountsByKey[key], account) + } } func (i *codexAccountIndex) remove(accountID int64) { @@ -1043,6 +1101,38 @@ func firstCodexString(obj map[string]any, paths ...[]string) string { return "" } +func firstCodexMap(obj map[string]any, paths ...[]string) (map[string]any, bool) { + for _, path := range paths { + value, ok := codexPathValue(obj, path) + if !ok || value == nil { + continue + } + if mapped, ok := value.(map[string]any); ok { + return mapped, true + } + } + return nil, false +} + +func firstCodexBool(obj map[string]any, paths ...[]string) bool { + for _, path := range paths { + value, ok := codexPathValue(obj, path) + if !ok { + continue + } + switch value := value.(type) { + case bool: + return value + case string: + parsed, err := strconv.ParseBool(strings.TrimSpace(value)) + if err == nil { + return parsed + } + } + } + return false +} + func copyCodexExtraString(obj map[string]any, extra map[string]any, key string, path []string) { value := firstCodexString(obj, path) if value != "" { diff --git a/backend/internal/handler/admin/grok_oauth_handler_test.go b/backend/internal/handler/admin/grok_oauth_handler_test.go index 64ea044aa3..9065e3321b 100644 --- a/backend/internal/handler/admin/grok_oauth_handler_test.go +++ b/backend/internal/handler/admin/grok_oauth_handler_test.go @@ -90,10 +90,13 @@ func TestGrokOAuthHandlerQueryQuotaProbesUpstream(t *testing.T) { ID: 42, Platform: service.PlatformGrok, Type: service.AccountTypeOAuth, + Status: service.StatusActive, + Schedulable: true, Concurrency: 1, Credentials: map[string]any{ - "access_token": "access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + "access_token": "access-token", + "refresh_token": "refresh-token", + "expires_at": time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339), }, }} upstream := &grokQuotaHandlerUpstream{} diff --git a/backend/internal/handler/admin/ops_handler.go b/backend/internal/handler/admin/ops_handler.go index e820aef0c6..0284dbfd0b 100644 --- a/backend/internal/handler/admin/ops_handler.go +++ b/backend/internal/handler/admin/ops_handler.go @@ -386,8 +386,8 @@ func (h *OpsHandler) ListRequestErrorUpstreamErrors(c *gin.Context) { filter.EndTime = &endTime } filter.View = "all" - filter.Phase = "upstream" - // 上游错误列表需含 status<400 的 recovered 行,显式豁免客户端可见守卫。 + filter.ErrorPhasesAny = []string{"upstream", "account_auth"} + // Provider-health list includes recovered inference and credential rows. filter.IncludeRecoveredUpstream = true filter.Owner = "provider" filter.Source = strings.TrimSpace(c.Query("error_source")) @@ -470,8 +470,8 @@ func (h *OpsHandler) ListUpstreamErrors(c *gin.Context) { } filter.View = parseOpsViewParam(c) - filter.Phase = "upstream" - // 上游错误列表需含 status<400 的 recovered 行,显式豁免客户端可见守卫。 + filter.ErrorPhasesAny = []string{"upstream", "account_auth"} + // Provider-health list includes recovered inference and credential rows. filter.IncludeRecoveredUpstream = true filter.Owner = "provider" filter.Source = strings.TrimSpace(c.Query("error_source")) diff --git a/backend/internal/handler/dto/credentials_redact_test.go b/backend/internal/handler/dto/credentials_redact_test.go index 431078fafd..9bf4f78039 100644 --- a/backend/internal/handler/dto/credentials_redact_test.go +++ b/backend/internal/handler/dto/credentials_redact_test.go @@ -20,6 +20,7 @@ func TestRedactCredentials_StripsSensitiveKeysAndReportsStatus(t *testing.T) { "aws_secret_access_key": "aws-secret", "service_account_json": map[string]any{"private_key": "..."}, "private_key": "raw-key", + "agent_private_key": "agent-key-secret", // 非敏感 "base_url": "https://api.example.com", "model_mapping": map[string]any{"foo": "bar"}, @@ -35,6 +36,7 @@ func TestRedactCredentials_StripsSensitiveKeysAndReportsStatus(t *testing.T) { require.NotContains(t, out, "aws_secret_access_key") require.NotContains(t, out, "service_account_json") require.NotContains(t, out, "private_key") + require.NotContains(t, out, "agent_private_key") require.Equal(t, "https://api.example.com", out["base_url"]) require.Equal(t, map[string]any{"foo": "bar"}, out["model_mapping"]) @@ -47,6 +49,7 @@ func TestRedactCredentials_StripsSensitiveKeysAndReportsStatus(t *testing.T) { require.True(t, status["has_aws_secret_access_key"]) require.True(t, status["has_service_account_json"]) require.True(t, status["has_private_key"]) + require.True(t, status["has_agent_private_key"]) // 状态 map 不应携带非敏感键的 has_* require.NotContains(t, status, "has_base_url") @@ -84,6 +87,7 @@ func TestRedactCredentials_AllKnownSensitiveKeys(t *testing.T) { "api_key", "session_key", "cookie", "aws_secret_access_key", "aws_session_token", "service_account_json", "service_account", "private_key", + "agent_private_key", } in := make(map[string]any, len(keys)) for _, k := range keys { diff --git a/backend/internal/handler/failover_loop.go b/backend/internal/handler/failover_loop.go index 5838e58f48..7ce1d03a8c 100644 --- a/backend/internal/handler/failover_loop.go +++ b/backend/internal/handler/failover_loop.go @@ -5,6 +5,8 @@ import ( "net/http" "time" + "github.com/gin-gonic/gin" + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" "github.com/Wei-Shaw/sub2api/internal/service" "go.uber.org/zap" @@ -71,7 +73,15 @@ func (s *FailoverState) HandleFailoverError( retryLimit int, failoverErr *service.UpstreamFailoverError, ) FailoverAction { + // 客户端已断开:failover 只会用已取消的 context 重新选号并必然失败, + // 不应再被当成账号耗尽处理(误报 502)。 + if ctx != nil && ctx.Err() != nil { + return FailoverCanceled + } s.LastFailoverErr = failoverErr + if failoverErr == nil || !failoverErr.ShouldRetryNextAccount() { + return FailoverExhausted + } // 缓存计费判断 if needForceCacheBilling(s.hasBoundSession, failoverErr) { @@ -135,6 +145,12 @@ func (s *FailoverState) HandleFailoverError( // 返回 FailoverExhausted 时,调用方应返回错误响应。 // 返回 FailoverCanceled 时,调用方应直接 return。 func (s *FailoverState) HandleSelectionExhausted(ctx context.Context) FailoverAction { + // 客户端已断开时选号失败是 context canceled 的必然结果, + // 不代表账号耗尽,直接按取消终止。 + if ctx.Err() != nil { + return FailoverCanceled + } + if s.LastFailoverErr != nil && s.LastFailoverErr.StatusCode == http.StatusServiceUnavailable && s.SwitchCount <= s.MaxSwitches { @@ -163,6 +179,28 @@ func needForceCacheBilling(hasBoundSession bool, failoverErr *service.UpstreamFa return hasBoundSession || (failoverErr != nil && failoverErr.ForceCacheBilling) } +// failoverClientGone 判断下游客户端是否已断开(请求 context 已取消)。 +// 客户端断开后 failover 必须静默终止:用已取消的 context 重新选号只会得到 +// context.Canceled,并被误报成账号耗尽(通用 502);上游 detach 的在途请求 +// 照常完成计费,但不再为无人接收的响应启动新的上游尝试。 +// 响应尚未提交时把状态码标记为 499(client closed request),供访问日志归类。 +func failoverClientGone(c *gin.Context) bool { + if c == nil || c.Request == nil || c.Request.Context().Err() == nil { + return false + } + // 先停 compact 心跳(接管 ResponseWriter,建立 happens-before),与 + // handleStreamingAwareError/errorResponse 等终结路径对齐,避免心跳 + // goroutine 与下面的状态标记并发触碰同一 writer。心跳已提交 200 时 + // 状态码已固化,不再标 499。 + if service.StopOpenAICompactSSEKeepaliveCommitted(c) { + return true + } + if !c.Writer.Written() { + c.Status(statusClientClosedRequest) + } + return true +} + // sleepWithContext 等待指定时长,返回 false 表示 context 已取消。 func sleepWithContext(ctx context.Context, d time.Duration) bool { if d <= 0 { diff --git a/backend/internal/handler/failover_loop_test.go b/backend/internal/handler/failover_loop_test.go index 9fabe75f25..8946eb5236 100644 --- a/backend/internal/handler/failover_loop_test.go +++ b/backend/internal/handler/failover_loop_test.go @@ -2,11 +2,14 @@ package handler import ( "context" + "net/http" + "net/http/httptest" "testing" "time" "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" "github.com/stretchr/testify/require" ) @@ -128,6 +131,50 @@ func TestSleepWithContext(t *testing.T) { // --------------------------------------------------------------------------- func TestHandleFailoverError_BasicSwitch(t *testing.T) { + t.Run("显式停止不切换账号且旧错误默认仍切换", func(t *testing.T) { + mock := &mockTempUnscheduler{} + fs := NewFailoverState(3, false) + stopErr := &service.UpstreamFailoverError{ + Stage: service.GatewayFailureStageAccountAuth, + Scope: service.GatewayFailureScopeProvider, + NextAccountAction: service.NextAccountStop, + } + + action := fs.HandleFailoverError(context.Background(), mock, 100, service.PlatformGrok, maxSameAccountRetries, stopErr) + + require.Equal(t, FailoverExhausted, action) + require.Zero(t, fs.SwitchCount) + require.Empty(t, fs.FailedAccountIDs) + require.Equal(t, stopErr, fs.LastFailoverErr) + + legacyErr := newTestFailoverErr(http.StatusTooManyRequests, false, false) + action = fs.HandleFailoverError(context.Background(), mock, 100, service.PlatformGrok, maxSameAccountRetries, legacyErr) + + require.Equal(t, FailoverContinue, action) + require.Equal(t, 1, fs.SwitchCount) + require.Contains(t, fs.FailedAccountIDs, int64(100)) + }) + + t.Run("已取消的认证失败不改变切换状态", func(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + mock := &mockTempUnscheduler{} + fs := NewFailoverState(3, false) + err := &service.UpstreamFailoverError{ + Stage: service.GatewayFailureStageAccountAuth, + Scope: service.GatewayFailureScopeAccount, + NextAccountAction: service.NextAccountRetry, + } + + action := fs.HandleFailoverError(ctx, mock, 101, service.PlatformGrok, maxSameAccountRetries, err) + + require.Equal(t, FailoverCanceled, action) + require.Zero(t, fs.SwitchCount) + require.Empty(t, fs.FailedAccountIDs) + require.Nil(t, fs.LastFailoverErr) + require.Empty(t, mock.calls) + }) + t.Run("非重试错误_非Antigravity_直接切换", func(t *testing.T) { mock := &mockTempUnscheduler{} fs := NewFailoverState(3, false) @@ -444,7 +491,28 @@ func TestHandleFailoverError_ContextCanceled(t *testing.T) { err := newTestFailoverErr(400, true, false) ctx, cancel := context.WithCancel(context.Background()) - cancel() // 立即取消 + go func() { + time.Sleep(30 * time.Millisecond) + cancel() // 通过入口检查后、sleep 期间取消 + }() + + start := time.Now() + action := fs.HandleFailoverError(ctx, mock, 100, "openai", maxSameAccountRetries, err) + elapsed := time.Since(start) + + require.Equal(t, FailoverCanceled, action) + require.Less(t, elapsed, 400*time.Millisecond, "sleep 应被取消打断") + // 进入重试分支后才取消:重试计数已递增 + require.Equal(t, 1, fs.SameAccountRetryCount[100]) + }) + + t.Run("入口即已取消_不改动任何failover状态", func(t *testing.T) { + mock := &mockTempUnscheduler{} + fs := NewFailoverState(3, false) + err := newTestFailoverErr(520, false, false) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() // 调用前客户端已断开 start := time.Now() action := fs.HandleFailoverError(ctx, mock, 100, "openai", maxSameAccountRetries, err) @@ -452,8 +520,12 @@ func TestHandleFailoverError_ContextCanceled(t *testing.T) { require.Equal(t, FailoverCanceled, action) require.Less(t, elapsed, 100*time.Millisecond, "应立即返回") - // 重试计数仍应递增 - require.Equal(t, 1, fs.SameAccountRetryCount[100]) + // 入口已取消时不得改变任何 failover 状态。 + require.Equal(t, 0, fs.SwitchCount, "取消的请求不应计入切换") + require.Equal(t, 0, fs.SameAccountRetryCount[100], "取消的请求不应改动重试计数") + require.NotContains(t, fs.FailedAccountIDs, int64(100)) + require.Nil(t, fs.LastFailoverErr) + require.Empty(t, mock.calls, "不应触发 TempUnschedule") }) t.Run("Antigravity延迟期间context取消", func(t *testing.T) { @@ -755,6 +827,29 @@ func TestHandleSelectionExhausted(t *testing.T) { require.Less(t, elapsed, 100*time.Millisecond, "应立即返回") }) + t.Run("context已取消_非503也返回Canceled而非Exhausted", func(t *testing.T) { + // #4257 核心场景:客户端断开后选号失败源于 context canceled, + // 不应被当成账号耗尽转成 502。 + fs := NewFailoverState(3, false) + fs.LastFailoverErr = newTestFailoverErr(520, false, false) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + action := fs.HandleSelectionExhausted(ctx) + require.Equal(t, FailoverCanceled, action) + }) + + t.Run("context已取消_无LastFailoverErr也返回Canceled", func(t *testing.T) { + fs := NewFailoverState(3, false) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + action := fs.HandleSelectionExhausted(ctx) + require.Equal(t, FailoverCanceled, action) + }) + t.Run("503且SwitchCount等于MaxSwitches_仍可重试", func(t *testing.T) { fs := NewFailoverState(2, false) fs.LastFailoverErr = newTestFailoverErr(503, false, false) @@ -764,3 +859,47 @@ func TestHandleSelectionExhausted(t *testing.T) { require.Equal(t, FailoverContinue, action) }) } + +// --------------------------------------------------------------------------- +// failoverClientGone 测试 +// --------------------------------------------------------------------------- + +func TestFailoverClientGone(t *testing.T) { + gin.SetMode(gin.TestMode) + + t.Run("活跃请求返回false", func(t *testing.T) { + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + + require.False(t, failoverClientGone(c)) + require.Equal(t, http.StatusOK, c.Writer.Status(), "不应改动状态码") + }) + + t.Run("客户端已断开_返回true并标记499", func(t *testing.T) { + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil).WithContext(ctx) + + require.True(t, failoverClientGone(c)) + require.Equal(t, statusClientClosedRequest, c.Writer.Status()) + }) + + t.Run("响应已提交_不改状态码", func(t *testing.T) { + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil).WithContext(ctx) + c.String(http.StatusOK, "partial") + + require.True(t, failoverClientGone(c)) + require.Equal(t, http.StatusOK, c.Writer.Status(), "已提交的状态码不应被覆盖") + }) + + t.Run("nil安全", func(t *testing.T) { + require.False(t, failoverClientGone(nil)) + }) +} diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go index fb07b098d1..1c51713656 100644 --- a/backend/internal/handler/gateway_handler.go +++ b/backend/internal/handler/gateway_handler.go @@ -327,6 +327,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) { c.Request = c.Request.WithContext(ctx) continue case FailoverCanceled: + failoverClientGone(c) return default: // FailoverExhausted if fs.LastFailoverErr != nil { @@ -456,6 +457,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) { h.handleFailoverExhausted(c, fs.LastFailoverErr, service.PlatformGemini, streamStarted) return case FailoverCanceled: + failoverClientGone(c) return } } @@ -613,6 +615,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) { c.Request = c.Request.WithContext(ctx) continue case FailoverCanceled: + failoverClientGone(c) return default: // FailoverExhausted if fs.LastFailoverErr != nil { @@ -876,6 +879,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) { h.handleFailoverExhausted(c, fs.LastFailoverErr, account.Platform, streamStarted) return case FailoverCanceled: + failoverClientGone(c) return } } diff --git a/backend/internal/handler/gateway_handler_cancellation_test.go b/backend/internal/handler/gateway_handler_cancellation_test.go new file mode 100644 index 0000000000..d53961174a --- /dev/null +++ b/backend/internal/handler/gateway_handler_cancellation_test.go @@ -0,0 +1,98 @@ +//go:build unit + +package handler + +import ( + "bytes" + "context" + "net/http" + "net/http/httptest" + "sync/atomic" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey" + middleware "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" +) + +type countingGatewaySchedulerCache struct { + *fakeSchedulerCache + snapshotCalls atomic.Int64 +} + +func (c *countingGatewaySchedulerCache) GetSnapshot(ctx context.Context, bucket service.SchedulerBucket) ([]*service.Account, bool, error) { + c.snapshotCalls.Add(1) + return c.fakeSchedulerCache.GetSnapshot(ctx, bucket) +} + +func TestGatewayHandlerPreCancelledCompatibleRequestsDoNotSelectAccount(t *testing.T) { + gin.SetMode(gin.TestMode) + groupID := int64(9100) + group := &service.Group{ID: groupID, Hydrated: true, Platform: service.PlatformAnthropic, Status: service.StatusActive} + account := &service.Account{ + ID: 9101, Platform: service.PlatformAnthropic, Type: service.AccountTypeAPIKey, + Status: service.StatusActive, Schedulable: true, Concurrency: 1, + AccountGroups: []service.AccountGroup{{AccountID: 9101, GroupID: groupID}}, + } + schedulerCache := &countingGatewaySchedulerCache{fakeSchedulerCache: &fakeSchedulerCache{accounts: []*service.Account{account}}} + schedulerSnapshot := service.NewSchedulerSnapshotService(schedulerCache, nil, nil, nil, nil) + gatewayService := service.NewGatewayService( + nil, &fakeGroupRepo{group: group}, nil, nil, nil, nil, nil, nil, nil, + schedulerSnapshot, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, + ) + cfg := &config.Config{RunMode: config.RunModeSimple} + billingCacheService := service.NewBillingCacheService(nil, nil, nil, nil, nil, nil, cfg, nil) + t.Cleanup(billingCacheService.Stop) + h := &GatewayHandler{ + gatewayService: gatewayService, + billingCacheService: billingCacheService, + concurrencyHelper: NewConcurrencyHelper(service.NewConcurrencyService(&fakeConcurrencyCache{}), SSEPingFormatClaude, 0), + maxAccountSwitches: 1, + cfg: cfg, + } + apiKey := &service.APIKey{ + ID: 9102, UserID: 9103, GroupID: &groupID, Group: group, Status: service.StatusActive, + User: &service.User{ID: 9103, Concurrency: 10, Balance: 100}, + } + + tests := []struct { + name string + path string + body string + call func(*gin.Context) + }{ + { + name: "responses", path: "/v1/responses", body: `{"model":"claude-test","input":"hello","stream":false}`, + call: h.Responses, + }, + { + name: "chat completions", path: "/v1/chat/completions", body: `{"model":"claude-test","messages":[{"role":"user","content":"hello"}],"stream":false}`, + call: h.ChatCompletions, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + schedulerCache.snapshotCalls.Store(0) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + ctx = context.WithValue(ctx, ctxkey.Group, group) + req := httptest.NewRequest(http.MethodPost, tt.path, bytes.NewBufferString(tt.body)).WithContext(ctx) + req.Header.Set("Content-Type", "application/json") + c.Request = req + c.Set(string(middleware.ContextKeyAPIKey), apiKey) + c.Set(string(middleware.ContextKeyUser), middleware.AuthSubject{UserID: apiKey.UserID, Concurrency: 10}) + + tt.call(c) + + require.Zero(t, schedulerCache.snapshotCalls.Load(), "a cancelled request must stop before the account selector") + _, selected := c.Get(opsAccountIDKey) + require.False(t, selected) + }) + } +} diff --git a/backend/internal/handler/gateway_handler_chat_completions.go b/backend/internal/handler/gateway_handler_chat_completions.go index af9bcdb344..cc9ff14282 100644 --- a/backend/internal/handler/gateway_handler_chat_completions.go +++ b/backend/internal/handler/gateway_handler_chat_completions.go @@ -159,6 +159,9 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) { } for { + if c.Request.Context().Err() != nil { + return + } selection, err := h.gatewayService.SelectAccountWithLoadAwareness(c.Request.Context(), apiKey.GroupID, selectionSessionHash, reqModel, fs.FailedAccountIDs, "", int64(0)) if err != nil { if len(fs.FailedAccountIDs) == 0 { @@ -178,6 +181,7 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) { case FailoverContinue: continue case FailoverCanceled: + failoverClientGone(c) return default: if fs.LastFailoverErr != nil { @@ -262,6 +266,7 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) { h.handleCCFailoverExhausted(c, fs.LastFailoverErr, streamStarted) return case FailoverCanceled: + failoverClientGone(c) return } } @@ -328,6 +333,14 @@ func (h *GatewayHandler) handleCCFailoverExhausted(c *gin.Context, lastErr *serv if streamStarted { return } + if lastErr != nil { + copyFailoverRetryAfter(c, lastErr.ResponseHeaders) + } + if lastErr != nil && lastErr.IsCredentialFailure() { + status, message := credentialFailoverClientResponse(lastErr) + h.chatCompletionsErrorResponse(c, status, "server_error", message) + return + } statusCode := http.StatusBadGateway if lastErr != nil && lastErr.StatusCode > 0 { statusCode = lastErr.StatusCode diff --git a/backend/internal/handler/gateway_handler_responses.go b/backend/internal/handler/gateway_handler_responses.go index 8a88d5fa57..a5060b7263 100644 --- a/backend/internal/handler/gateway_handler_responses.go +++ b/backend/internal/handler/gateway_handler_responses.go @@ -157,6 +157,9 @@ func (h *GatewayHandler) Responses(c *gin.Context) { fs := NewFailoverState(h.maxAccountSwitches, false) for { + if requestCtx.Err() != nil { + return + } selection, err := h.gatewayService.SelectAccountWithLoadAwareness(requestCtx, apiKey.GroupID, sessionHash, reqModel, fs.FailedAccountIDs, "", int64(0)) if err != nil { if len(fs.FailedAccountIDs) == 0 { @@ -176,6 +179,7 @@ func (h *GatewayHandler) Responses(c *gin.Context) { case FailoverContinue: continue case FailoverCanceled: + failoverClientGone(c) return default: if fs.LastFailoverErr != nil { @@ -241,6 +245,7 @@ func (h *GatewayHandler) Responses(c *gin.Context) { h.handleResponsesFailoverExhausted(c, fs.LastFailoverErr, streamStarted) return case FailoverCanceled: + failoverClientGone(c) return } } @@ -307,6 +312,14 @@ func (h *GatewayHandler) handleResponsesFailoverExhausted(c *gin.Context, lastEr if streamStarted { return // Can't write error after stream started } + if lastErr != nil { + copyFailoverRetryAfter(c, lastErr.ResponseHeaders) + } + if lastErr != nil && lastErr.IsCredentialFailure() { + status, message := credentialFailoverClientResponse(lastErr) + h.responsesErrorResponse(c, status, "server_error", message) + return + } statusCode := http.StatusBadGateway if lastErr != nil && lastErr.StatusCode > 0 { statusCode = lastErr.StatusCode diff --git a/backend/internal/handler/gateway_helper.go b/backend/internal/handler/gateway_helper.go index 8489b17f6d..1f6daf5291 100644 --- a/backend/internal/handler/gateway_helper.go +++ b/backend/internal/handler/gateway_helper.go @@ -163,20 +163,15 @@ func wrapReleaseOnDone(ctx context.Context, releaseFunc func()) func() { return nil } var once sync.Once - var stop func() bool - - release := func() { - once.Do(func() { - if stop != nil { - _ = stop() - } - releaseFunc() - }) + releaseOnce := func() { + once.Do(releaseFunc) } + stop := context.AfterFunc(ctx, releaseOnce) - stop = context.AfterFunc(ctx, release) - - return release + return func() { + _ = stop() + releaseOnce() + } } // IncrementWaitCount increments the wait count for a user diff --git a/backend/internal/handler/gateway_helper_test.go b/backend/internal/handler/gateway_helper_test.go index 664258f8c9..bc0f50f0b6 100644 --- a/backend/internal/handler/gateway_helper_test.go +++ b/backend/internal/handler/gateway_helper_test.go @@ -6,6 +6,8 @@ import ( "sync/atomic" "testing" "time" + + "github.com/stretchr/testify/require" ) // TestWrapReleaseOnDone_NoGoroutineLeak 验证 wrapReleaseOnDone 修复后不会泄露 goroutine @@ -67,6 +69,21 @@ func TestWrapReleaseOnDone_ContextCancellation(t *testing.T) { } } +func TestWrapReleaseOnDone_AlreadyCancelledReleasesExactlyOnce(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + var releaseCount int32 + release := wrapReleaseOnDone(ctx, func() { + atomic.AddInt32(&releaseCount, 1) + }) + release() + + require.Eventually(t, func() bool { + return atomic.LoadInt32(&releaseCount) == 1 + }, time.Second, time.Millisecond) +} + // TestWrapReleaseOnDone_MultipleCallsOnlyReleaseOnce 验证多次调用 release 只释放一次 func TestWrapReleaseOnDone_MultipleCallsOnlyReleaseOnce(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) diff --git a/backend/internal/handler/gemini_v1beta_handler.go b/backend/internal/handler/gemini_v1beta_handler.go index b1653c6b3f..c595abfe55 100644 --- a/backend/internal/handler/gemini_v1beta_handler.go +++ b/backend/internal/handler/gemini_v1beta_handler.go @@ -368,6 +368,7 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) { c.Request = c.Request.WithContext(ctx) continue case FailoverCanceled: + failoverClientGone(c) return default: // FailoverExhausted h.handleGeminiFailoverExhausted(c, fs.LastFailoverErr) @@ -490,6 +491,7 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) { h.handleGeminiFailoverExhausted(c, fs.LastFailoverErr) return case FailoverCanceled: + failoverClientGone(c) return } } diff --git a/backend/internal/handler/grok_media.go b/backend/internal/handler/grok_media.go index b7092fcc42..01928028e6 100644 --- a/backend/internal/handler/grok_media.go +++ b/backend/internal/handler/grok_media.go @@ -166,6 +166,7 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service. failedAccountIDs := make(map[int64]struct{}) sameAccountRetryCount := make(map[int64]int) var lastFailoverErr *service.UpstreamFailoverError + var oauth429FailoverState service.OpenAIOAuth429FailoverState switchCount := 0 maxAccountSwitches := h.maxAccountSwitches if maxAccountSwitches <= 0 { @@ -174,6 +175,9 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service. routingStart := time.Now() for { + if failoverClientGone(c) { + return + } selection, scheduleDecision, err := h.gatewayService.SelectAccountWithSchedulerForCapability( requestCtx, apiKey.GroupID, @@ -188,6 +192,10 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service. service.PlatformGrok, ) if err != nil { + if failoverClientGone(c) { + reqLog.Info("grok_media.account_select_aborted_client_disconnected", zap.Error(err)) + return + } reqLog.Warn("grok_media.account_select_failed", zap.Error(err), zap.Int("excluded_account_count", len(failedAccountIDs)), @@ -257,11 +265,24 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service. if err != nil { var failoverErr *service.UpstreamFailoverError if errors.As(err, &failoverErr) { - h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil) + if failoverClientGone(c) { + reqLog.Info("grok_media.failover_aborted_client_disconnected", + zap.Int64("account_id", account.ID), + zap.Int("upstream_status", failoverErr.StatusCode), + ) + return + } + if failoverErr.ShouldReportAccountScheduleFailure() { + h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil) + } if c.Writer.Size() != writerSizeBeforeForward { h.handleFailoverExhausted(c, failoverErr, true) return } + if !failoverErr.ShouldRetryNextAccount() { + h.handleFailoverExhausted(c, failoverErr, false) + return + } if failoverErr.RetryableOnSameAccount { retryLimit := account.GetPoolModeRetryCount() if sameAccountRetryCount[account.ID] < retryLimit { @@ -288,6 +309,10 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service. return } switchCount++ + if h.gatewayService.ShouldStopOpenAIOAuth429Failover(account, failoverErr.StatusCode, switchCount, &oauth429FailoverState) { + h.handleFailoverExhausted(c, failoverErr, false) + return + } reqLog.Warn("grok_media.upstream_failover_switching", zap.Int64("account_id", account.ID), zap.Int("upstream_status", failoverErr.StatusCode), diff --git a/backend/internal/handler/openai_alpha_search.go b/backend/internal/handler/openai_alpha_search.go index a808145235..a1d185a1cb 100644 --- a/backend/internal/handler/openai_alpha_search.go +++ b/backend/internal/handler/openai_alpha_search.go @@ -106,6 +106,7 @@ func (h *OpenAIGatewayHandler) AlphaSearch(c *gin.Context) { failedAccountIDs := make(map[int64]struct{}) var lastFailoverErr *service.UpstreamFailoverError switchCount := 0 + var oauth429FailoverState service.OpenAIOAuth429FailoverState routingStart := time.Now() for { @@ -123,6 +124,10 @@ func (h *OpenAIGatewayHandler) AlphaSearch(c *gin.Context) { service.PlatformOpenAI, ) if err != nil || selection == nil || selection.Account == nil { + if failoverClientGone(c) { + reqLog.Info("openai_alpha_search.account_select_aborted_client_disconnected", zap.Error(err)) + return + } if len(failedAccountIDs) == 0 { cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, requestedModel, requestedModel, service.PlatformOpenAI) if !cls.ModelNotFound { @@ -180,6 +185,13 @@ func (h *OpenAIGatewayHandler) AlphaSearch(c *gin.Context) { h.handleFailoverExhausted(c, failoverErr, true) return } + if failoverClientGone(c) { + reqLog.Info("openai_alpha_search.failover_aborted_client_disconnected", + zap.Int64("account_id", account.ID), + zap.Int("upstream_status", failoverErr.StatusCode), + ) + return + } h.gatewayService.RecordOpenAIAccountSwitch() failedAccountIDs[account.ID] = struct{}{} lastFailoverErr = failoverErr @@ -188,7 +200,7 @@ func (h *OpenAIGatewayHandler) AlphaSearch(c *gin.Context) { return } switchCount++ - if h.gatewayService.ShouldStopOpenAIOAuth429Failover(account, failoverErr.StatusCode, switchCount) { + if h.gatewayService.ShouldStopOpenAIOAuth429Failover(account, failoverErr.StatusCode, switchCount, &oauth429FailoverState) { h.handleFailoverExhausted(c, failoverErr, false) return } diff --git a/backend/internal/handler/openai_chat_completions.go b/backend/internal/handler/openai_chat_completions.go index 636e143740..62a2b70b59 100644 --- a/backend/internal/handler/openai_chat_completions.go +++ b/backend/internal/handler/openai_chat_completions.go @@ -133,8 +133,12 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { failedAccountIDs := make(map[int64]struct{}) sameAccountRetryCount := make(map[int64]int) var lastFailoverErr *service.UpstreamFailoverError + var oauth429FailoverState service.OpenAIOAuth429FailoverState for { + if failoverClientGone(c) { + return + } reqLog.Debug("openai_chat_completions.account_selecting", zap.Int("excluded_account_count", len(failedAccountIDs))) selection, scheduleDecision, err := h.gatewayService.SelectAccountWithSchedulerForCapability( c.Request.Context(), @@ -150,6 +154,10 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { requestPlatform, ) if err != nil { + if failoverClientGone(c) { + reqLog.Info("openai_chat_completions.account_select_aborted_client_disconnected", zap.Error(err)) + return + } reqLog.Warn("openai_chat_completions.account_select_failed", zap.Error(openAICompatibleSelectionErrorForLog(err, requestPlatform)), zap.Int("excluded_account_count", len(failedAccountIDs)), @@ -231,11 +239,24 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { } else { var failoverErr *service.UpstreamFailoverError if errors.As(err, &failoverErr) { + if failoverClientGone(c) { + reqLog.Info("openai_chat_completions.failover_aborted_client_disconnected", + zap.Int64("account_id", account.ID), + zap.Int("upstream_status", failoverErr.StatusCode), + ) + return + } if c.Writer.Size() != writerSizeBeforeForward { h.handleFailoverExhausted(c, failoverErr, true) return } - h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil) + if failoverErr.ShouldReportAccountScheduleFailure() { + h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil) + } + if !failoverErr.ShouldRetryNextAccount() { + h.handleFailoverExhausted(c, failoverErr, streamStarted) + return + } // Pool mode: retry on the same account if failoverErr.RetryableOnSameAccount { retryLimit := account.GetPoolModeRetryCount() @@ -263,7 +284,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { return } switchCount++ - if h.gatewayService.ShouldStopOpenAIOAuth429Failover(account, failoverErr.StatusCode, switchCount) { + if h.gatewayService.ShouldStopOpenAIOAuth429Failover(account, failoverErr.StatusCode, switchCount, &oauth429FailoverState) { h.handleFailoverExhausted(c, failoverErr, streamStarted) return } diff --git a/backend/internal/handler/openai_embeddings.go b/backend/internal/handler/openai_embeddings.go index 56d775eb7c..99f869f14c 100644 --- a/backend/internal/handler/openai_embeddings.go +++ b/backend/internal/handler/openai_embeddings.go @@ -121,6 +121,10 @@ func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) { false, ) if err != nil { + if failoverClientGone(c) { + reqLog.Info("openai_embeddings.account_select_aborted_client_disconnected", zap.Error(err)) + return + } reqLog.Warn("openai_embeddings.account_select_failed", zap.Error(err), zap.Int("excluded_account_count", len(failedAccountIDs)), @@ -189,6 +193,13 @@ func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) { return } h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil) + if failoverClientGone(c) { + reqLog.Info("openai_embeddings.failover_aborted_client_disconnected", + zap.Int64("account_id", account.ID), + zap.Int("upstream_status", failoverErr.StatusCode), + ) + return + } h.gatewayService.RecordOpenAIAccountSwitch() failedAccountIDs[account.ID] = struct{}{} lastFailoverErr = failoverErr diff --git a/backend/internal/handler/openai_gateway_credential_failover_loop_test.go b/backend/internal/handler/openai_gateway_credential_failover_loop_test.go new file mode 100644 index 0000000000..42a6bf6cd3 --- /dev/null +++ b/backend/internal/handler/openai_gateway_credential_failover_loop_test.go @@ -0,0 +1,926 @@ +//go:build unit + +package handler + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "io" + "net/http" + "net/http/httptest" + "strings" + "sync" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/Wei-Shaw/sub2api/internal/server/middleware" + "github.com/Wei-Shaw/sub2api/internal/service" + coderws "github.com/coder/websocket" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +type grokCredentialHandlerRepo struct { + service.AccountRepository + mu sync.Mutex + accounts []service.Account + setErrorIDs []int64 + setTempIDs []int64 + rateLimitIDs []int64 + updateExtraIDs []int64 + selectionCalls int + setErrorErr error + setTempErr error + missingOnGet map[int64]bool +} + +func (r *grokCredentialHandlerRepo) ListSchedulableByPlatform(_ context.Context, platform string) ([]service.Account, error) { + r.mu.Lock() + defer r.mu.Unlock() + r.selectionCalls++ + out := make([]service.Account, 0, len(r.accounts)) + for _, account := range r.accounts { + if account.Platform == platform && account.IsSchedulable() { + out = append(out, account) + } + } + return out, nil +} + +func (r *grokCredentialHandlerRepo) ListSchedulableByGroupIDAndPlatform(ctx context.Context, _ int64, platform string) ([]service.Account, error) { + return r.ListSchedulableByPlatform(ctx, platform) +} + +func (r *grokCredentialHandlerRepo) ListSchedulableUngroupedByPlatform(ctx context.Context, platform string) ([]service.Account, error) { + return r.ListSchedulableByPlatform(ctx, platform) +} + +func (r *grokCredentialHandlerRepo) GetByID(_ context.Context, id int64) (*service.Account, error) { + r.mu.Lock() + defer r.mu.Unlock() + if r.missingOnGet[id] { + return nil, nil + } + for _, account := range r.accounts { + if account.ID == id { + copy := account + copy.Credentials = cloneCredentialMap(account.Credentials) + return ©, nil + } + } + return nil, nil +} + +func (r *grokCredentialHandlerRepo) SetError(_ context.Context, id int64, message string) error { + r.mu.Lock() + defer r.mu.Unlock() + r.setErrorIDs = append(r.setErrorIDs, id) + if r.setErrorErr != nil { + return r.setErrorErr + } + for i := range r.accounts { + if r.accounts[i].ID == id { + r.accounts[i].Status = service.StatusError + r.accounts[i].Schedulable = false + r.accounts[i].ErrorMessage = message + } + } + return nil +} + +func (r *grokCredentialHandlerRepo) SetTempUnschedulable(_ context.Context, id int64, until time.Time, _ string) error { + r.mu.Lock() + defer r.mu.Unlock() + r.setTempIDs = append(r.setTempIDs, id) + if r.setTempErr != nil { + return r.setTempErr + } + for i := range r.accounts { + if r.accounts[i].ID == id { + value := until + r.accounts[i].TempUnschedulableUntil = &value + } + } + return nil +} + +func (r *grokCredentialHandlerRepo) SetRateLimited(_ context.Context, id int64, resetAt time.Time) error { + r.mu.Lock() + defer r.mu.Unlock() + r.rateLimitIDs = append(r.rateLimitIDs, id) + for i := range r.accounts { + if r.accounts[i].ID != id { + continue + } + now := time.Now() + r.accounts[i].RateLimitedAt = &now + value := resetAt + r.accounts[i].RateLimitResetAt = &value + } + return nil +} + +func (r *grokCredentialHandlerRepo) SetRateLimitedIfLater(ctx context.Context, id int64, resetAt time.Time) error { + r.mu.Lock() + for i := range r.accounts { + if r.accounts[i].ID == id && r.accounts[i].RateLimitResetAt != nil && !resetAt.After(*r.accounts[i].RateLimitResetAt) { + r.mu.Unlock() + return nil + } + } + r.mu.Unlock() + return r.SetRateLimited(ctx, id, resetAt) +} + +func (r *grokCredentialHandlerRepo) SetGrokCredentialErrorIfMatch( + _ context.Context, + id int64, + snapshot service.GrokCredentialMutationSnapshot, + message string, +) (bool, error) { + r.mu.Lock() + defer r.mu.Unlock() + for i := range r.accounts { + account := &r.accounts[i] + if account.ID != id || !handlerGrokCredentialSnapshotMatches(account, snapshot) { + continue + } + r.setErrorIDs = append(r.setErrorIDs, id) + if r.setErrorErr != nil { + return false, r.setErrorErr + } + account.Status = service.StatusError + account.Schedulable = false + account.ErrorMessage = message + return true, nil + } + return false, nil +} + +func (r *grokCredentialHandlerRepo) SetGrokCredentialTempUnschedulableIfMatch( + _ context.Context, + id int64, + snapshot service.GrokCredentialMutationSnapshot, + until time.Time, + _ string, +) (bool, error) { + r.mu.Lock() + defer r.mu.Unlock() + for i := range r.accounts { + account := &r.accounts[i] + if account.ID != id || !handlerGrokCredentialSnapshotMatches(account, snapshot) { + continue + } + r.setTempIDs = append(r.setTempIDs, id) + if r.setTempErr != nil { + return false, r.setTempErr + } + value := until + account.TempUnschedulableUntil = &value + return true, nil + } + return false, nil +} + +func handlerGrokCredentialSnapshotMatches(account *service.Account, snapshot service.GrokCredentialMutationSnapshot) bool { + if account == nil { + return false + } + credentialsJSON, err := json.Marshal(account.Credentials) + return err == nil && account.IsGrokOAuth() && account.IsSchedulable() && string(credentialsJSON) == snapshot.CredentialsJSON && + handlerGrokCredentialProxyIDsEqual(account.ProxyID, snapshot.ProxyID) +} + +func handlerGrokCredentialProxyIDsEqual(left, right *int64) bool { + if left == nil || right == nil { + return left == nil && right == nil + } + return *left == *right +} + +func (r *grokCredentialHandlerRepo) UpdateExtra(_ context.Context, id int64, updates map[string]any) error { + r.mu.Lock() + defer r.mu.Unlock() + r.updateExtraIDs = append(r.updateExtraIDs, id) + for i := range r.accounts { + if r.accounts[i].ID != id { + continue + } + if r.accounts[i].Extra == nil { + r.accounts[i].Extra = map[string]any{} + } + for key, value := range updates { + r.accounts[i].Extra[key] = value + } + } + return nil +} + +func (r *grokCredentialHandlerRepo) errorIDs() []int64 { + r.mu.Lock() + defer r.mu.Unlock() + return append([]int64(nil), r.setErrorIDs...) +} + +func (r *grokCredentialHandlerRepo) selectorCalls() int { + r.mu.Lock() + defer r.mu.Unlock() + return r.selectionCalls +} + +func (r *grokCredentialHandlerRepo) rateLimitedAccountIDs() []int64 { + r.mu.Lock() + defer r.mu.Unlock() + return append([]int64(nil), r.rateLimitIDs...) +} + +type grokCredentialHandlerTokenCache struct { + service.GrokTokenCache + mu sync.Mutex + deleteErr error +} + +func (c *grokCredentialHandlerTokenCache) GetAccessToken(context.Context, string) (string, error) { + return "", errors.New("not cached") +} + +func (c *grokCredentialHandlerTokenCache) SetAccessToken(context.Context, string, string, time.Duration) error { + return nil +} + +func (c *grokCredentialHandlerTokenCache) DeleteAccessToken(context.Context, string) error { + c.mu.Lock() + defer c.mu.Unlock() + return c.deleteErr +} + +func (c *grokCredentialHandlerTokenCache) AcquireRefreshLock(context.Context, string, time.Duration) (bool, error) { + return true, nil +} + +func (c *grokCredentialHandlerTokenCache) ReleaseRefreshLock(context.Context, string) error { + return nil +} + +func cloneCredentialMap(source map[string]any) map[string]any { + cloned := make(map[string]any, len(source)) + for key, value := range source { + cloned[key] = value + } + return cloned +} + +type grokCredentialHandlerRefresher struct { + mode string + started chan struct{} + once sync.Once +} + +func (r *grokCredentialHandlerRefresher) CacheKey(account *service.Account) string { + return service.GrokTokenCacheKey(account) +} + +func (r *grokCredentialHandlerRefresher) CanRefresh(account *service.Account) bool { + return account != nil && account.IsGrokOAuth() +} + +func (r *grokCredentialHandlerRefresher) NeedsRefresh(account *service.Account, _ time.Duration) bool { + return account != nil && (account.ID == 801 || r.mode == "all_revoked") +} + +func (r *grokCredentialHandlerRefresher) Refresh(ctx context.Context, _ *service.Account) (map[string]any, error) { + switch r.mode { + case "revoked", "all_revoked", "mutation_set_error", "mutation_cache": + return nil, infraerrors.New(http.StatusBadGateway, "GROK_OAUTH_TOKEN_REFRESH_FAILED", "invalid_grant") + case "provider": + return nil, infraerrors.New(http.StatusBadGateway, "GROK_OAUTH_TOKEN_REFRESH_FAILED", "invalid_client") + case "cancel": + r.once.Do(func() { close(r.started) }) + <-ctx.Done() + return nil, ctx.Err() + case "transient", "mutation_temp": + return nil, errors.New("temporary refresh transport failure") + default: + return nil, nil + } +} + +type grokCredentialHandlerUpstream struct { + service.HTTPUpstream + mu sync.Mutex + hits []int64 + requestURLs []string + authorization []string + failAccountID int64 + rateLimitIDs map[int64]bool + failureStatus map[int64]int + cancelRequest context.CancelFunc +} + +func (u *grokCredentialHandlerUpstream) Do(req *http.Request, _ string, accountID int64, _ int) (*http.Response, error) { + var requestBody []byte + if req.Body != nil { + requestBody, _ = io.ReadAll(req.Body) + } + u.mu.Lock() + u.hits = append(u.hits, accountID) + u.requestURLs = append(u.requestURLs, req.URL.String()) + u.authorization = append(u.authorization, req.Header.Get("Authorization")) + failAccountID := u.failAccountID + rateLimited := u.rateLimitIDs[accountID] + failureStatus := u.failureStatus[accountID] + cancelRequest := u.cancelRequest + u.mu.Unlock() + if rateLimited { + return &http.Response{ + StatusCode: http.StatusTooManyRequests, + Header: http.Header{ + "Content-Type": []string{"application/json"}, + "Retry-After": []string{"60"}, + }, + Body: io.NopCloser(bytes.NewBufferString(`{"error":{"message":"rate limited"}}`)), + }, nil + } + if failureStatus > 0 { + return &http.Response{ + StatusCode: failureStatus, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(bytes.NewBufferString(`{"error":{"message":"upstream unavailable"}}`)), + }, nil + } + if accountID == failAccountID { + if cancelRequest != nil { + cancelRequest() + } + return &http.Response{ + StatusCode: http.StatusPaymentRequired, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(bytes.NewBufferString(`{"error":{"message":"payment required"}}`)), + }, nil + } + if bytes.Contains(requestBody, []byte(`"stream":true`)) { + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(bytes.NewBufferString( + "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_healthy\",\"model\":\"grok-4.5\",\"usage\":{\"input_tokens\":1,\"output_tokens\":1}}}\n\n", + )), + }, nil + } + if strings.Contains(req.URL.Path, "/chat/completions") { + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(bytes.NewBufferString( + `{"id":"chatcmpl_healthy","object":"chat.completion","model":"grok-4.5","choices":[{"index":0,"message":{"role":"assistant","content":"ok"},"finish_reason":"stop"}],"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}`, + )), + }, nil + } + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(bytes.NewBufferString( + `{"id":"resp_healthy","object":"response","model":"grok-4.5","status":"completed","output":[],"usage":{"input_tokens":1,"output_tokens":1}}`, + )), + }, nil +} + +func (u *grokCredentialHandlerUpstream) accountHits() []int64 { + u.mu.Lock() + defer u.mu.Unlock() + return append([]int64(nil), u.hits...) +} + +func (u *grokCredentialHandlerUpstream) requests() ([]string, []string) { + u.mu.Lock() + defer u.mu.Unlock() + return append([]string(nil), u.requestURLs...), append([]string(nil), u.authorization...) +} + +func TestResponsesCredentialFailoverLoop(t *testing.T) { + gin.SetMode(gin.TestMode) + + t.Run("revoked account selects healthy account", func(t *testing.T) { + h, repo, upstream, router, cleanup := newGrokCredentialFailoverHandler(t, "revoked") + defer cleanup() + _ = h + + recorder := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/openai/v1/responses", bytes.NewBufferString(`{"model":"grok","input":"hello","stream":false}`)) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(recorder, req) + + require.Equal(t, http.StatusOK, recorder.Code, recorder.Body.String()) + require.Contains(t, recorder.Body.String(), "resp_healthy") + require.Equal(t, []int64{801}, repo.errorIDs()) + require.Equal(t, []int64{802}, upstream.accountHits()) + requestURLs, authorization := upstream.requests() + require.Equal(t, []string{xai.DefaultCLIBaseURL + "/responses"}, requestURLs) + require.Equal(t, []string{"Bearer healthy-access"}, authorization) + }) + + t.Run("provider configuration stops before healthy account", func(t *testing.T) { + h, repo, upstream, router, cleanup := newGrokCredentialFailoverHandler(t, "provider") + defer cleanup() + + recorder := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/openai/v1/responses", bytes.NewBufferString(`{"model":"grok","input":"hello","stream":false}`)) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(recorder, req) + + require.Equal(t, http.StatusServiceUnavailable, recorder.Code) + require.Contains(t, recorder.Body.String(), service.GrokCredentialUnavailableClientMessage) + require.Empty(t, repo.errorIDs()) + require.Empty(t, upstream.accountHits()) + require.Equal(t, 1, repo.selectorCalls()) + require.Zero(t, h.gatewayService.SnapshotOpenAIAccountSchedulerMetrics().RuntimeStatsAccountCount, + "provider-scoped auth failure must not penalize the selected account") + }) + + t.Run("parent cancellation stops before healthy account", func(t *testing.T) { + _, repo, upstream, router, cleanup := newGrokCredentialFailoverHandler(t, "cancel") + defer cleanup() + + ctx, cancel := context.WithCancel(context.Background()) + req := httptest.NewRequest(http.MethodPost, "/openai/v1/responses", bytes.NewBufferString(`{"model":"grok","input":"hello","stream":false}`)).WithContext(ctx) + req.Header.Set("Content-Type", "application/json") + recorder := httptest.NewRecorder() + done := make(chan struct{}) + go func() { + defer close(done) + router.ServeHTTP(recorder, req) + }() + + select { + case <-time.After(2 * time.Second): + t.Fatal("credential refresh did not start") + case <-findHandlerRefresherStarted(router): + cancel() + } + select { + case <-time.After(2 * time.Second): + t.Fatal("handler did not stop after cancellation") + case <-done: + } + + require.Empty(t, repo.errorIDs()) + require.Empty(t, upstream.accountHits()) + }) + + t.Run("post-mapping cancellation stops before scheduler mutation or reselection", func(t *testing.T) { + h, repo, upstream, router, cleanup := newGrokCredentialFailoverHandler(t, "postmap_cancel") + defer cleanup() + ctx, cancel := context.WithCancel(context.Background()) + upstream.mu.Lock() + upstream.failAccountID = 801 + upstream.cancelRequest = cancel + upstream.mu.Unlock() + + recorder := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/openai/v1/responses", bytes.NewBufferString(`{"model":"grok","input":"hello","stream":false}`)).WithContext(ctx) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(recorder, req) + + require.Equal(t, []int64{801}, upstream.accountHits()) + require.Empty(t, repo.errorIDs()) + require.Equal(t, 1, repo.selectorCalls()) + require.Zero(t, h.gatewayService.SnapshotOpenAIAccountSchedulerMetrics().RuntimeStatsAccountCount) + }) + + t.Run("pre-cancelled request never invokes an account selector", func(t *testing.T) { + tests := []struct { + name string + method string + path string + body string + }{ + {name: "responses", method: http.MethodPost, path: "/openai/v1/responses", body: `{"model":"grok","input":"hello","stream":false}`}, + {name: "messages", method: http.MethodPost, path: "/openai/v1/messages", body: `{"model":"grok","max_tokens":16,"messages":[{"role":"user","content":"hello"}]}`}, + {name: "chat completions", method: http.MethodPost, path: "/openai/v1/chat/completions", body: `{"model":"grok","messages":[{"role":"user","content":"hello"}],"stream":false}`}, + {name: "grok media", method: http.MethodGet, path: "/openai/v1/videos/request-1"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + _, repo, upstream, router, cleanup := newGrokCredentialFailoverHandler(t, "revoked") + defer cleanup() + ctx, cancel := context.WithCancel(context.Background()) + cancel() + recorder := httptest.NewRecorder() + req := httptest.NewRequest(tt.method, tt.path, bytes.NewBufferString(tt.body)).WithContext(ctx) + req.Header.Set("Content-Type", "application/json") + + router.ServeHTTP(recorder, req) + + require.Zero(t, repo.selectorCalls()) + require.Empty(t, upstream.accountHits()) + }) + } + }) + + t.Run("credential state mutation failures stop before reselection", func(t *testing.T) { + for _, mode := range []string{"mutation_set_error", "mutation_temp", "mutation_cache"} { + t.Run(mode, func(t *testing.T) { + _, repo, upstream, router, cleanup := newGrokCredentialFailoverHandler(t, mode) + defer cleanup() + + recorder := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/openai/v1/responses", bytes.NewBufferString(`{"model":"grok","input":"hello","stream":false}`)) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(recorder, req) + + require.Equal(t, http.StatusServiceUnavailable, recorder.Code, recorder.Body.String()) + require.Contains(t, recorder.Body.String(), service.GrokCredentialUnavailableClientMessage) + require.Empty(t, upstream.accountHits()) + require.Equal(t, 1, repo.selectorCalls()) + }) + } + }) + + t.Run("missing credential provider stops before upstream or reselection", func(t *testing.T) { + _, repo, upstream, router, cleanup := newGrokCredentialFailoverHandler(t, "nil_provider") + defer cleanup() + recorder := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/openai/v1/responses", bytes.NewBufferString(`{"model":"grok","input":"hello","stream":false}`)) + req.Header.Set("Content-Type", "application/json") + + router.ServeHTTP(recorder, req) + + require.Equal(t, http.StatusServiceUnavailable, recorder.Code, recorder.Body.String()) + require.Contains(t, recorder.Body.String(), service.GrokCredentialUnavailableClientMessage) + require.Equal(t, 1, repo.selectorCalls()) + require.Empty(t, upstream.accountHits()) + require.Empty(t, repo.errorIDs()) + }) +} + +func TestResponsesGrok429FailoverIsBounded(t *testing.T) { + gin.SetMode(gin.TestMode) + + t.Run("first rate limited account selects healthy account", func(t *testing.T) { + _, repo, upstream, router, cleanup := newGrokCredentialFailoverHandler(t, "first_429") + defer cleanup() + recorder := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/openai/v1/responses", bytes.NewBufferString(`{"model":"grok","input":"hello","stream":false}`)) + req.Header.Set("Content-Type", "application/json") + + router.ServeHTTP(recorder, req) + + require.Equal(t, http.StatusOK, recorder.Code, recorder.Body.String()) + require.Contains(t, recorder.Body.String(), "resp_healthy") + require.Equal(t, []int64{801, 802}, upstream.accountHits()) + require.Equal(t, []int64{801}, repo.rateLimitedAccountIDs()) + }) + + t.Run("two rate limited accounts stop without sweeping the pool", func(t *testing.T) { + _, repo, upstream, router, cleanup := newGrokCredentialFailoverHandler(t, "all_429") + defer cleanup() + recorder := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/openai/v1/responses", bytes.NewBufferString(`{"model":"grok","input":"hello","stream":false}`)) + req.Header.Set("Content-Type", "application/json") + + router.ServeHTTP(recorder, req) + + require.Equal(t, http.StatusTooManyRequests, recorder.Code, recorder.Body.String()) + require.Equal(t, []int64{801, 802}, upstream.accountHits()) + require.Equal(t, []int64{801, 802}, repo.rateLimitedAccountIDs()) + require.NotContains(t, recorder.Body.String(), "expired") + require.NotContains(t, recorder.Body.String(), "healthy-access") + require.NotContains(t, recorder.Body.String(), "rate limited") + }) +} + +func TestResponsesGrok429FailoverHandlesMixedStatuses(t *testing.T) { + gin.SetMode(gin.TestMode) + + t.Run("429 then 500 stops after the bounded followup", func(t *testing.T) { + _, _, upstream, router, cleanup := newGrokCredentialFailoverHandler(t, "mixed_429_500") + defer cleanup() + recorder := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/openai/v1/responses", bytes.NewBufferString(`{"model":"grok","input":"hello","stream":false}`)) + req.Header.Set("Content-Type", "application/json") + + router.ServeHTTP(recorder, req) + + require.Equal(t, http.StatusBadGateway, recorder.Code, recorder.Body.String()) + require.Equal(t, []int64{801, 802}, upstream.accountHits()) + require.NotContains(t, recorder.Body.String(), "upstream unavailable") + }) + + t.Run("500 then 429 permits one healthy followup", func(t *testing.T) { + _, _, upstream, router, cleanup := newGrokCredentialFailoverHandler(t, "mixed_500_429") + defer cleanup() + recorder := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/openai/v1/responses", bytes.NewBufferString(`{"model":"grok","input":"hello","stream":false}`)) + req.Header.Set("Content-Type", "application/json") + + router.ServeHTTP(recorder, req) + + require.Equal(t, http.StatusOK, recorder.Code, recorder.Body.String()) + require.Equal(t, []int64{801, 802, 803}, upstream.accountHits()) + }) + + t.Run("OAuth 429 then API-key failure cannot bypass the bound", func(t *testing.T) { + _, _, upstream, router, cleanup := newGrokCredentialFailoverHandler(t, "oauth_429_apikey_500") + defer cleanup() + recorder := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/openai/v1/responses", bytes.NewBufferString(`{"model":"grok","input":"hello","stream":false}`)) + req.Header.Set("Content-Type", "application/json") + + router.ServeHTTP(recorder, req) + + require.Equal(t, http.StatusBadGateway, recorder.Code, recorder.Body.String()) + require.Equal(t, []int64{801, 802}, upstream.accountHits()) + }) +} + +func TestGrokMedia429FailoverIsBounded(t *testing.T) { + gin.SetMode(gin.TestMode) + + t.Run("first 429 selects one healthy followup", func(t *testing.T) { + _, _, upstream, router, cleanup := newGrokCredentialFailoverHandler(t, "first_429") + defer cleanup() + recorder := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/openai/v1/videos/request-1", nil) + + router.ServeHTTP(recorder, req) + + require.Equal(t, http.StatusOK, recorder.Code, recorder.Body.String()) + require.Equal(t, []int64{801, 802}, upstream.accountHits()) + }) + + t.Run("second 429 stops without sweeping a third account", func(t *testing.T) { + _, _, upstream, router, cleanup := newGrokCredentialFailoverHandler(t, "all_429") + defer cleanup() + recorder := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/openai/v1/videos/request-1", nil) + + router.ServeHTTP(recorder, req) + + require.Equal(t, http.StatusTooManyRequests, recorder.Code, recorder.Body.String()) + require.Equal(t, []int64{801, 802}, upstream.accountHits()) + require.NotContains(t, recorder.Body.String(), "rate limited") + }) +} + +func TestGrokOAuthCredentialFailoverAcrossHTTPHandlers(t *testing.T) { + gin.SetMode(gin.TestMode) + endpoints := []struct { + name string + method string + path string + body string + }{ + {name: "messages", method: http.MethodPost, path: "/openai/v1/messages", body: `{"model":"grok","max_tokens":16,"messages":[{"role":"user","content":"hello"}]}`}, + {name: "chat completions", method: http.MethodPost, path: "/openai/v1/chat/completions", body: `{"model":"grok","messages":[{"role":"user","content":"hello"}],"stream":false}`}, + {name: "chat completions raw fallback", method: http.MethodPost, path: "/openai/v1/chat/completions", body: `{"model":"grok","messages":[{"role":"user","content":"hello"}],"stop":["END"],"stream":false}`}, + {name: "grok media", method: http.MethodGet, path: "/openai/v1/videos/request-1"}, + } + + for _, endpoint := range endpoints { + t.Run(endpoint.name+" revoked selects healthy", func(t *testing.T) { + _, repo, upstream, router, cleanup := newGrokCredentialFailoverHandler(t, "revoked") + defer cleanup() + recorder := httptest.NewRecorder() + req := httptest.NewRequest(endpoint.method, endpoint.path, bytes.NewBufferString(endpoint.body)) + req.Header.Set("Content-Type", "application/json") + + router.ServeHTTP(recorder, req) + + require.Equal(t, http.StatusOK, recorder.Code, recorder.Body.String()) + require.Equal(t, []int64{801}, repo.errorIDs()) + require.Equal(t, []int64{802}, upstream.accountHits()) + }) + + t.Run(endpoint.name+" all accounts exhausted safely", func(t *testing.T) { + _, repo, upstream, router, cleanup := newGrokCredentialFailoverHandler(t, "all_revoked") + defer cleanup() + recorder := httptest.NewRecorder() + req := httptest.NewRequest(endpoint.method, endpoint.path, bytes.NewBufferString(endpoint.body)) + req.Header.Set("Content-Type", "application/json") + + router.ServeHTTP(recorder, req) + + require.Equal(t, http.StatusServiceUnavailable, recorder.Code, recorder.Body.String()) + require.Contains(t, recorder.Body.String(), service.GrokCredentialUnavailableClientMessage) + require.NotContains(t, recorder.Body.String(), "revoked-refresh") + require.NotContains(t, recorder.Body.String(), "healthy-refresh") + require.Equal(t, []int64{801, 802}, repo.errorIDs()) + require.Empty(t, upstream.accountHits()) + }) + } +} + +func TestGrokOAuthMissingSelectedRowRetriesHealthyAccountWithoutMutation(t *testing.T) { + gin.SetMode(gin.TestMode) + _, repo, upstream, router, cleanup := newGrokCredentialFailoverHandler(t, "missing_row") + defer cleanup() + recorder := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/openai/v1/responses", bytes.NewBufferString(`{"model":"grok","input":"hello","stream":false}`)) + req.Header.Set("Content-Type", "application/json") + + router.ServeHTTP(recorder, req) + + require.Equal(t, http.StatusOK, recorder.Code, recorder.Body.String()) + require.Equal(t, []int64{802}, upstream.accountHits()) + require.Empty(t, repo.errorIDs()) + require.Empty(t, repo.setTempIDs) +} + +func TestResponsesWebSocketCredentialFailoverLoop(t *testing.T) { + gin.SetMode(gin.TestMode) + dial := func(t *testing.T, router *gin.Engine) (*coderws.Conn, func()) { + t.Helper() + server := httptest.NewServer(router) + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + conn, _, err := coderws.Dial(ctx, "ws"+strings.TrimPrefix(server.URL, "http")+"/openai/v1/responses", nil) + cancel() + require.NoError(t, err) + return conn, func() { + _ = conn.CloseNow() + server.Close() + } + } + writeFirst := func(t *testing.T, conn *coderws.Conn) { + t.Helper() + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + require.NoError(t, conn.Write(ctx, coderws.MessageText, []byte(`{"type":"response.create","model":"grok","input":"hello","stream":false}`))) + } + + t.Run("revoked account selects healthy account", func(t *testing.T) { + _, repo, upstream, router, cleanup := newGrokCredentialFailoverHandler(t, "revoked") + defer cleanup() + conn, closeConn := dial(t, router) + defer closeConn() + writeFirst(t, conn) + + readCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + _, payload, err := conn.Read(readCtx) + cancel() + require.NoError(t, err) + require.Contains(t, string(payload), "resp_healthy") + require.Equal(t, []int64{801}, repo.errorIDs()) + require.Equal(t, 2, repo.selectorCalls()) + require.Equal(t, []int64{802}, upstream.accountHits()) + }) + + t.Run("provider configuration stops", func(t *testing.T) { + _, repo, upstream, router, cleanup := newGrokCredentialFailoverHandler(t, "provider") + defer cleanup() + conn, closeConn := dial(t, router) + defer closeConn() + writeFirst(t, conn) + + readCtx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + _, _, err := conn.Read(readCtx) + cancel() + var closeErr coderws.CloseError + require.ErrorAs(t, err, &closeErr) + require.Contains(t, closeErr.Reason, service.GrokCredentialUnavailableClientMessage) + require.Equal(t, 1, repo.selectorCalls()) + require.Empty(t, upstream.accountHits()) + }) + + t.Run("parent cancellation prevents reselection", func(t *testing.T) { + _, repo, upstream, router, cleanup := newGrokCredentialFailoverHandler(t, "cancel") + defer cleanup() + conn, closeConn := dial(t, router) + writeFirst(t, conn) + select { + case <-findHandlerRefresherStarted(router): + case <-time.After(2 * time.Second): + t.Fatal("credential refresh did not start") + } + closeConn() + + require.Eventually(t, func() bool { return repo.selectorCalls() == 1 }, 2*time.Second, 20*time.Millisecond) + require.Empty(t, repo.errorIDs()) + require.Empty(t, upstream.accountHits()) + }) +} + +var handlerRefresherStarted sync.Map + +func findHandlerRefresherStarted(router *gin.Engine) <-chan struct{} { + value, _ := handlerRefresherStarted.Load(router) + return value.(chan struct{}) +} + +func newGrokCredentialFailoverHandler(t *testing.T, mode string) (*OpenAIGatewayHandler, *grokCredentialHandlerRepo, *grokCredentialHandlerUpstream, *gin.Engine, func()) { + t.Helper() + groupID := int64(901) + accounts := []service.Account{ + { + ID: 801, Name: "revoked", Platform: service.PlatformGrok, Type: service.AccountTypeOAuth, + Status: service.StatusActive, Schedulable: true, Concurrency: 1, Priority: 1, + Credentials: map[string]any{ + "access_token": "expired", "refresh_token": "revoked-refresh", + "expires_at": time.Now().Add(-time.Minute).UTC().Format(time.RFC3339), + }, + }, + { + ID: 802, Name: "healthy", Platform: service.PlatformGrok, Type: service.AccountTypeOAuth, + Status: service.StatusActive, Schedulable: true, Concurrency: 1, Priority: 2, + Credentials: map[string]any{ + "access_token": "healthy-access", "refresh_token": "healthy-refresh", + "expires_at": time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339), + }, + }, + } + if mode == "postmap_cancel" || mode == "first_429" || mode == "all_429" || mode == "mixed_429_500" || mode == "mixed_500_429" || mode == "oauth_429_apikey_500" { + accounts[0].Credentials["expires_at"] = time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339) + } + if mode == "all_429" || mode == "mixed_429_500" || mode == "mixed_500_429" || mode == "oauth_429_apikey_500" { + accounts = append(accounts, service.Account{ + ID: 803, Name: "untried-healthy", Platform: service.PlatformGrok, Type: service.AccountTypeOAuth, + Status: service.StatusActive, Schedulable: true, Concurrency: 1, Priority: 3, + Credentials: map[string]any{ + "access_token": "untried-healthy-access", "refresh_token": "untried-healthy-refresh", + "expires_at": time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339), + }, + }) + } + if mode == "oauth_429_apikey_500" { + accounts[1].Type = service.AccountTypeAPIKey + accounts[1].Credentials = map[string]any{"api_key": "third-party-key"} + } + if mode == "all_revoked" { + accounts[1].Credentials["expires_at"] = time.Now().Add(-time.Minute).UTC().Format(time.RFC3339) + } + repo := &grokCredentialHandlerRepo{accounts: accounts, missingOnGet: map[int64]bool{}} + if mode == "missing_row" { + repo.missingOnGet[801] = true + } + if mode == "mutation_set_error" { + repo.setErrorErr = errors.New("database write failed") + } + if mode == "mutation_temp" { + repo.setTempErr = errors.New("database write failed") + } + refresher := &grokCredentialHandlerRefresher{mode: mode, started: make(chan struct{})} + tokenCache := &grokCredentialHandlerTokenCache{} + if mode == "mutation_cache" { + tokenCache.deleteErr = errors.New("cache delete failed") + } + var provider *service.GrokTokenProvider + if mode != "nil_provider" { + provider = service.NewGrokTokenProvider(repo, tokenCache) + provider.SetRefreshAPI(service.NewOAuthRefreshAPI(repo, tokenCache), refresher) + } + upstream := &grokCredentialHandlerUpstream{} + switch mode { + case "first_429": + upstream.rateLimitIDs = map[int64]bool{801: true} + case "all_429": + upstream.rateLimitIDs = map[int64]bool{801: true, 802: true} + case "mixed_429_500": + upstream.rateLimitIDs = map[int64]bool{801: true} + upstream.failureStatus = map[int64]int{802: http.StatusInternalServerError} + case "mixed_500_429": + upstream.failureStatus = map[int64]int{801: http.StatusInternalServerError} + upstream.rateLimitIDs = map[int64]bool{802: true} + case "oauth_429_apikey_500": + upstream.rateLimitIDs = map[int64]bool{801: true} + upstream.failureStatus = map[int64]int{802: http.StatusInternalServerError} + } + cfg := &config.Config{RunMode: config.RunModeSimple} + cfg.Gateway.MaxAccountSwitches = 3 + billingCache := service.NewBillingCacheService(nil, nil, nil, nil, nil, nil, cfg, nil) + gateway := service.NewOpenAIGatewayService( + repo, nil, nil, nil, nil, nil, nil, cfg, nil, nil, + service.NewBillingService(cfg, nil), nil, billingCache, upstream, + &service.DeferredService{}, nil, provider, nil, nil, nil, nil, nil, + ) + cache := &concurrencyCacheMock{ + acquireUserSlotFn: func(context.Context, int64, int, string) (bool, error) { return true, nil }, + acquireAccountSlotFn: func(context.Context, int64, int, string) (bool, error) { return true, nil }, + } + h := NewOpenAIGatewayHandler(gateway, service.NewConcurrencyService(cache), billingCache, &service.APIKeyService{}, nil, nil, nil, nil, cfg) + apiKey := &service.APIKey{ + ID: 902, GroupID: &groupID, + User: &service.User{ID: 903, Status: service.StatusActive}, + Group: &service.Group{ID: groupID, Platform: service.PlatformGrok, Status: service.StatusActive}, + } + router := gin.New() + router.Use(func(c *gin.Context) { + c.Set(string(middleware.ContextKeyAPIKey), apiKey) + c.Set(string(middleware.ContextKeyUser), middleware.AuthSubject{UserID: apiKey.User.ID, Concurrency: 1}) + c.Next() + }) + router.POST("/openai/v1/responses", h.Responses) + router.GET("/openai/v1/responses", h.ResponsesWebSocket) + router.POST("/openai/v1/messages", h.Messages) + router.POST("/openai/v1/chat/completions", h.ChatCompletions) + router.GET("/openai/v1/videos/:request_id", h.GrokVideoStatus) + handlerRefresherStarted.Store(router, refresher.started) + cleanup := func() { + handlerRefresherStarted.Delete(router) + billingCache.Stop() + } + return h, repo, upstream, router, cleanup +} diff --git a/backend/internal/handler/openai_gateway_credential_failover_test.go b/backend/internal/handler/openai_gateway_credential_failover_test.go new file mode 100644 index 0000000000..07a5e4999f --- /dev/null +++ b/backend/internal/handler/openai_gateway_credential_failover_test.go @@ -0,0 +1,211 @@ +//go:build unit + +package handler + +import ( + "context" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestGatewayChatCredentialStopDoesNotSelectAnotherAccountAndReturnsSafe503(t *testing.T) { + gin.SetMode(gin.TestMode) + stopErr := &service.UpstreamFailoverError{ + Stage: service.GatewayFailureStageAccountAuth, + Scope: service.GatewayFailureScopeProvider, + Reason: service.GrokCredentialReasonProviderConfig, + NextAccountAction: service.NextAccountStop, + ClientStatusCode: http.StatusTeapot, + ClientMessage: "invalid_client client_secret=must-not-leak", + } + state := NewFailoverState(3, false) + action := state.HandleFailoverError(context.Background(), &mockTempUnscheduler{}, 71, service.PlatformGrok, 0, stopErr) + + require.Equal(t, FailoverExhausted, action) + require.Zero(t, state.SwitchCount) + require.Empty(t, state.FailedAccountIDs) + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + (&GatewayHandler{}).handleCCFailoverExhausted(c, state.LastFailoverErr, false) + + require.Equal(t, http.StatusServiceUnavailable, recorder.Code) + require.Contains(t, recorder.Body.String(), service.GrokCredentialUnavailableClientMessage) + require.NotContains(t, recorder.Body.String(), "invalid_client") + require.NotContains(t, recorder.Body.String(), "client_secret") +} + +func TestGatewayChatInferenceExhaustionRestoresRetryAfter(t *testing.T) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + + (&GatewayHandler{}).handleCCFailoverExhausted(c, &service.UpstreamFailoverError{ + StatusCode: http.StatusTooManyRequests, + ResponseHeaders: http.Header{"Retry-After": []string{"45"}}, + }, false) + + require.Equal(t, http.StatusTooManyRequests, recorder.Code) + require.Equal(t, "45", recorder.Header().Get("Retry-After")) +} + +func TestCredentialFailoverExhaustionReturnsFixedSafe503(t *testing.T) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + h := &OpenAIGatewayHandler{} + + h.handleFailoverExhausted(c, &service.UpstreamFailoverError{ + Stage: service.GatewayFailureStageAccountAuth, + Scope: service.GatewayFailureScopeAccount, + Reason: service.GrokCredentialReasonRevoked, + NextAccountAction: service.NextAccountRetry, + ClientStatusCode: http.StatusTeapot, + ClientMessage: "invalid_grant refresh_token=must-not-leak", + }, false) + + require.Equal(t, http.StatusServiceUnavailable, recorder.Code) + require.Contains(t, recorder.Body.String(), service.GrokCredentialUnavailableClientMessage) + require.NotContains(t, strings.ToLower(recorder.Body.String()), "invalid_grant") + require.NotContains(t, strings.ToLower(recorder.Body.String()), "refresh_token") + require.NotContains(t, recorder.Body.String(), "must-not-leak") +} + +func TestInferenceFailoverExhaustionRestoresRetryAfter(t *testing.T) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + h := &OpenAIGatewayHandler{} + + h.handleFailoverExhausted(c, &service.UpstreamFailoverError{ + StatusCode: http.StatusTooManyRequests, + ResponseHeaders: http.Header{"Retry-After": []string{"17"}}, + }, false) + + require.Equal(t, http.StatusTooManyRequests, recorder.Code) + require.Equal(t, "17", recorder.Header().Get("Retry-After")) +} + +func TestFailoverExhaustionRejectsSecretBearingRetryAfter(t *testing.T) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + h := &OpenAIGatewayHandler{} + + h.handleFailoverExhausted(c, &service.UpstreamFailoverError{ + StatusCode: http.StatusTooManyRequests, + ResponseHeaders: http.Header{"Retry-After": []string{"refresh_token=must-not-leak"}}, + }, false) + + require.Equal(t, http.StatusTooManyRequests, recorder.Code) + require.Empty(t, recorder.Header().Get("Retry-After")) + require.NotContains(t, recorder.Body.String(), "must-not-leak") +} + +func TestFailoverExhaustionRejectsFarFutureRetryAfterDate(t *testing.T) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + h := &OpenAIGatewayHandler{} + + h.handleFailoverExhausted(c, &service.UpstreamFailoverError{ + StatusCode: http.StatusTooManyRequests, + ResponseHeaders: http.Header{ + "Retry-After": []string{time.Now().Add(30 * 24 * time.Hour).UTC().Format(http.TimeFormat)}, + }, + }, false) + + require.Equal(t, http.StatusTooManyRequests, recorder.Code) + require.Empty(t, recorder.Header().Get("Retry-After")) +} + +func TestFailoverExhaustionAllowsBoundedRetryAfterDate(t *testing.T) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + h := &OpenAIGatewayHandler{} + retryAfter := time.Now().Add(time.Hour).UTC().Format(http.TimeFormat) + + h.handleFailoverExhausted(c, &service.UpstreamFailoverError{ + StatusCode: http.StatusTooManyRequests, + ResponseHeaders: http.Header{"Retry-After": []string{retryAfter}}, + }, false) + + require.Equal(t, http.StatusTooManyRequests, recorder.Code) + require.Equal(t, retryAfter, recorder.Header().Get("Retry-After")) +} + +func TestOpsClassificationTreatsCredentialFailureAsAuthNotInference(t *testing.T) { + gin.SetMode(gin.TestMode) + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + c.Set(service.OpsUpstreamStatusCodeKey, http.StatusForbidden) + c.Set(service.OpsUpstreamErrorMessageKey, "stale inference message") + c.Set(service.OpsUpstreamErrorDetailKey, "stale inference detail") + c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{ + {Stage: string(service.GatewayFailureStageInference), UpstreamStatusCode: http.StatusForbidden, Message: "stale inference message", Detail: "stale inference detail"}, + { + Stage: string(service.GatewayFailureStageAccountAuth), + Scope: string(service.GatewayFailureScopeAccount), + Reason: string(service.GrokCredentialReasonRevoked), + UpstreamStatusCode: 0, + Message: "Grok OAuth credentials require account action", + }, + }) + + phase, _, owner, source := classifyOpsErrorLog(c, "upstream_error", service.GrokCredentialUnavailableClientMessage, "", http.StatusServiceUnavailable) + require.Equal(t, "account_auth", phase) + require.Equal(t, "provider", owner) + require.Equal(t, "gateway", source) + + entry := &service.OpsInsertErrorLogInput{} + applyOpsUpstreamFieldsFromContext(c, entry) + require.NotNil(t, entry.UpstreamStatusCode) + require.Zero(t, *entry.UpstreamStatusCode) + require.NotNil(t, entry.UpstreamErrorMessage) + require.Equal(t, "Grok OAuth credentials require account action", *entry.UpstreamErrorMessage) + require.Nil(t, entry.UpstreamErrorDetail) + require.Len(t, entry.UpstreamErrors, 2) + require.Equal(t, http.StatusForbidden, entry.UpstreamErrors[0].UpstreamStatusCode) +} + +func TestOpsRecoveredCredentialFailoverUsesAccountAuthAttribution(t *testing.T) { + setupOpsErrorLogTestQueue(t, 2) + gin.SetMode(gin.TestMode) + ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + router := gin.New() + router.Use(OpsErrorLoggerMiddleware(ops)) + router.GET("/openai/v1/responses", func(c *gin.Context) { + c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{ + {Stage: string(service.GatewayFailureStageInference), UpstreamStatusCode: http.StatusForbidden, Message: "earlier inference failure"}, + { + Stage: string(service.GatewayFailureStageAccountAuth), Scope: string(service.GatewayFailureScopeAccount), + Reason: string(service.GrokCredentialReasonRevoked), Message: "Grok OAuth credentials require account action", + }, + }) + c.JSON(http.StatusOK, gin.H{"ok": true}) + }) + + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, httptest.NewRequest(http.MethodGet, "/openai/v1/responses", nil)) + + require.Equal(t, http.StatusOK, recorder.Code) + require.Equal(t, int64(1), OpsErrorLogQueueLength()) + job := <-opsErrorLogQueue + require.Equal(t, "account_auth", job.entry.ErrorPhase) + require.Equal(t, "provider", job.entry.ErrorOwner) + require.Equal(t, "gateway", job.entry.ErrorSource) + require.Contains(t, job.entry.ErrorMessage, "Recovered account authentication failure") + require.NotContains(t, job.entry.ErrorMessage, "403") + require.NotContains(t, job.entry.ErrorMessage, "earlier inference failure") + require.NotNil(t, job.entry.UpstreamStatusCode) + require.Zero(t, *job.entry.UpstreamStatusCode) + require.Len(t, job.entry.UpstreamErrors, 2) + require.Equal(t, http.StatusForbidden, job.entry.UpstreamErrors[0].UpstreamStatusCode) +} diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index 231363a8a3..eea8836ecf 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -332,8 +332,12 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { failedAccountIDs := make(map[int64]struct{}) sameAccountRetryCount := make(map[int64]int) var lastFailoverErr *service.UpstreamFailoverError + var oauth429FailoverState service.OpenAIOAuth429FailoverState for { + if failoverClientGone(c) { + return + } // Select account supporting the requested model reqLog.Debug("openai.account_selecting", zap.Int("excluded_account_count", len(failedAccountIDs))) selection, scheduleDecision, err := h.gatewayService.SelectAccountWithSchedulerForCapability( @@ -350,6 +354,10 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { requestPlatform, ) if err != nil { + if failoverClientGone(c) { + reqLog.Info("openai.account_select_aborted_client_disconnected", zap.Error(err)) + return + } reqLog.Warn("openai.account_select_failed", zap.Error(openAICompatibleSelectionErrorForLog(err, requestPlatform)), zap.Int("excluded_account_count", len(failedAccountIDs)), @@ -443,11 +451,24 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { } else { var failoverErr *service.UpstreamFailoverError if errors.As(err, &failoverErr) { + if failoverClientGone(c) { + reqLog.Info("openai.failover_aborted_client_disconnected", + zap.Int64("account_id", account.ID), + zap.Int("upstream_status", failoverErr.StatusCode), + ) + return + } if service.OpenAICompactKeepaliveAdjustedWrittenSize(c) != writerSizeBeforeForward { h.handleFailoverExhausted(c, failoverErr, true) return } - h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil) + if failoverErr.ShouldReportAccountScheduleFailure() { + h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil) + } + if !failoverErr.ShouldRetryNextAccount() { + h.handleFailoverExhausted(c, failoverErr, streamStarted) + return + } // 池模式:同账号重试 if failoverErr.RetryableOnSameAccount { retryLimit := account.GetPoolModeRetryCount() @@ -475,16 +496,27 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { return } switchCount++ - if h.gatewayService.ShouldStopOpenAIOAuth429Failover(account, failoverErr.StatusCode, switchCount) { + if h.gatewayService.ShouldStopOpenAIOAuth429Failover(account, failoverErr.StatusCode, switchCount, &oauth429FailoverState) { h.handleFailoverExhausted(c, failoverErr, streamStarted) return } - reqLog.Warn("openai.upstream_failover_switching", + failoverSwitchFields := []zap.Field{ zap.Int64("account_id", account.ID), zap.Int("upstream_status", failoverErr.StatusCode), zap.Int("switch_count", switchCount), zap.Int("max_switches", maxAccountSwitches), - ) + } + if account.Proxy != nil { + failoverSwitchFields = append(failoverSwitchFields, + zap.Int64("proxy_id", account.Proxy.ID), + zap.String("proxy_name", account.Proxy.Name), + zap.String("proxy_host", account.Proxy.Host), + zap.Int("proxy_port", account.Proxy.Port), + ) + } else if account.ProxyID != nil { + failoverSwitchFields = append(failoverSwitchFields, zap.Int64p("proxy_id", account.ProxyID)) + } + reqLog.Warn("openai.upstream_failover_switching", failoverSwitchFields...) continue } h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil) @@ -832,9 +864,13 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { failedAccountIDs := make(map[int64]struct{}) sameAccountRetryCount := make(map[int64]int) var lastFailoverErr *service.UpstreamFailoverError + var oauth429FailoverState service.OpenAIOAuth429FailoverState effectiveMappedModel := preferredMappedModel for { + if failoverClientGone(c) { + return + } currentRoutingModel := routingModel if effectiveMappedModel != "" { currentRoutingModel = effectiveMappedModel @@ -854,6 +890,10 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { requestPlatform, ) if err != nil { + if failoverClientGone(c) { + reqLog.Info("openai_messages.account_select_aborted_client_disconnected", zap.Error(err)) + return + } reqLog.Warn("openai_messages.account_select_failed", zap.Error(openAICompatibleSelectionErrorForLog(err, requestPlatform)), zap.Int("excluded_account_count", len(failedAccountIDs)), @@ -935,11 +975,24 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { } else { var failoverErr *service.UpstreamFailoverError if errors.As(err, &failoverErr) { + if failoverClientGone(c) { + reqLog.Info("openai_messages.failover_aborted_client_disconnected", + zap.Int64("account_id", account.ID), + zap.Int("upstream_status", failoverErr.StatusCode), + ) + return + } if c.Writer.Size() != writerSizeBeforeForward { h.handleAnthropicFailoverExhausted(c, failoverErr, true) return } - h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil) + if failoverErr.ShouldReportAccountScheduleFailure() { + h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil) + } + if !failoverErr.ShouldRetryNextAccount() { + h.handleAnthropicFailoverExhausted(c, failoverErr, streamStarted) + return + } // 池模式:同账号重试 if failoverErr.RetryableOnSameAccount { retryLimit := account.GetPoolModeRetryCount() @@ -967,7 +1020,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { return } switchCount++ - if h.gatewayService.ShouldStopOpenAIOAuth429Failover(account, failoverErr.StatusCode, switchCount) { + if h.gatewayService.ShouldStopOpenAIOAuth429Failover(account, failoverErr.StatusCode, switchCount, &oauth429FailoverState) { h.handleAnthropicFailoverExhausted(c, failoverErr, streamStarted) return } @@ -1093,6 +1146,14 @@ func (h *OpenAIGatewayHandler) anthropicStreamingAwareError(c *gin.Context, stat // handleAnthropicFailoverExhausted maps upstream failover errors to Anthropic format. func (h *OpenAIGatewayHandler) handleAnthropicFailoverExhausted(c *gin.Context, failoverErr *service.UpstreamFailoverError, streamStarted bool) { + if failoverErr != nil { + copyFailoverRetryAfter(c, failoverErr.ResponseHeaders) + } + if failoverErr != nil && failoverErr.IsCredentialFailure() { + status, message := credentialFailoverClientResponse(failoverErr) + h.anthropicStreamingAwareError(c, status, "api_error", message, streamStarted) + return + } status, errType, errMsg := h.mapUpstreamError(failoverErr.StatusCode) h.anthropicStreamingAwareError(c, status, errType, errMsg, streamStarted) } @@ -1464,8 +1525,50 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { switchCount := 0 failedAccountIDs := make(map[int64]struct{}) var lastFailoverErr *service.UpstreamFailoverError + var oauth429FailoverState service.OpenAIOAuth429FailoverState + handleWSFailover := func(account *service.Account, failoverErr *service.UpstreamFailoverError) bool { + if ctx.Err() != nil { + return false + } + if failoverErr.ShouldReportAccountScheduleFailure() { + h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil) + } + releaseAccountSlot() + if !failoverErr.ShouldRetryNextAccount() { + closeOpenAIWSFailoverExhausted(wsConn, failoverErr) + return false + } + if ctx.Err() != nil { + return false + } + h.gatewayService.RecordOpenAIAccountSwitch() + failedAccountIDs[account.ID] = struct{}{} + lastFailoverErr = failoverErr + if switchCount >= maxAccountSwitches { + closeOpenAIWSFailoverExhausted(wsConn, failoverErr) + return false + } + switchCount++ + if h.gatewayService.ShouldStopOpenAIOAuth429Failover(account, failoverErr.StatusCode, switchCount, &oauth429FailoverState) { + closeOpenAIWSFailoverExhausted(wsConn, failoverErr) + return false + } + reqLog.Warn("openai.websocket_upstream_failover_switching", + zap.Int64("account_id", account.ID), + zap.Int("upstream_status", failoverErr.StatusCode), + zap.Int("switch_count", switchCount), + zap.Int("max_switches", maxAccountSwitches), + ) + if ctx.Err() != nil { + return false + } + return ensureUserSlotHeld() + } for { + if ctx.Err() != nil { + return + } reqLog.Debug("openai.websocket_account_selecting", zap.Int("excluded_account_count", len(failedAccountIDs))) selection, scheduleDecision, err := h.gatewayService.SelectAccountWithSchedulerForCapability( ctx, @@ -1533,9 +1636,19 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { reqLog.Warn("openai.websocket_bind_sticky_session_failed", zap.Int64("account_id", account.ID), zap.Error(err)) } - token, _, err := h.gatewayService.GetAccessToken(ctx, account) + token, _, err := h.gatewayService.GetRequestCredential(ctx, c, account) if err != nil { reqLog.Warn("openai.websocket_get_access_token_failed", zap.Int64("account_id", account.ID), zap.Error(err)) + if ctx.Err() != nil { + return + } + var failoverErr *service.UpstreamFailoverError + if errors.As(err, &failoverErr) { + if handleWSFailover(account, failoverErr) { + continue + } + return + } closeOpenAIClientWS(wsConn, coderws.StatusInternalError, "failed to get access token") return } @@ -1692,30 +1805,10 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { if err := h.gatewayService.ProxyResponsesWebSocketFromClient(ctx, c, wsConn, account, token, wsFirstMessage, hooks); err != nil { var failoverErr *service.UpstreamFailoverError if errors.As(err, &failoverErr) { - h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil) - releaseAccountSlot() - failedAccountIDs[account.ID] = struct{}{} - lastFailoverErr = failoverErr - if switchCount >= maxAccountSwitches { - closeOpenAIWSFailoverExhausted(wsConn, failoverErr) - return + if handleWSFailover(account, failoverErr) { + continue } - switchCount++ - if h.gatewayService.ShouldStopOpenAIOAuth429Failover(account, failoverErr.StatusCode, switchCount) { - closeOpenAIWSFailoverExhausted(wsConn, failoverErr) - return - } - h.gatewayService.RecordOpenAIAccountSwitch() - reqLog.Warn("openai.websocket_upstream_failover_switching", - zap.Int64("account_id", account.ID), - zap.Int("upstream_status", failoverErr.StatusCode), - zap.Int("switch_count", switchCount), - zap.Int("max_switches", maxAccountSwitches), - ) - if !ensureUserSlotHeld() { - return - } - continue + return } if errors.Is(context.Cause(ctx), service.ErrOpenAIWSIngressLeaseLost) { @@ -1739,12 +1832,23 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil) closeStatus, closeReason := summarizeWSCloseErrorForLog(err) - reqLog.Warn("openai.websocket_proxy_failed", + proxyFailedFields := []zap.Field{ zap.Int64("account_id", account.ID), zap.Error(err), zap.String("close_status", closeStatus), zap.String("close_reason", closeReason), - ) + } + if account.Proxy != nil { + proxyFailedFields = append(proxyFailedFields, + zap.Int64("proxy_id", account.Proxy.ID), + zap.String("proxy_name", account.Proxy.Name), + zap.String("proxy_host", account.Proxy.Host), + zap.Int("proxy_port", account.Proxy.Port), + ) + } else if account.ProxyID != nil { + proxyFailedFields = append(proxyFailedFields, zap.Int64p("proxy_id", account.ProxyID)) + } + reqLog.Warn("openai.websocket_proxy_failed", proxyFailedFields...) if errors.As(err, &closeErr) { closeOpenAIClientWS(wsConn, closeErr.StatusCode(), closeErr.Reason()) return @@ -1946,6 +2050,16 @@ func (h *OpenAIGatewayHandler) handleConcurrencyError(c *gin.Context, err error, } func (h *OpenAIGatewayHandler) handleFailoverExhausted(c *gin.Context, failoverErr *service.UpstreamFailoverError, streamStarted bool) { + if failoverErr == nil { + h.handleFailoverExhaustedSimple(c, http.StatusBadGateway, streamStarted) + return + } + copyFailoverRetryAfter(c, failoverErr.ResponseHeaders) + if failoverErr.IsCredentialFailure() { + status, message := credentialFailoverClientResponse(failoverErr) + h.handleStreamingAwareError(c, status, "upstream_error", message, streamStarted) + return + } statusCode := failoverErr.StatusCode responseBody := failoverErr.ResponseBody if service.IsOpenAISilentRefusalErrorBody(responseBody) { @@ -1987,6 +2101,41 @@ func (h *OpenAIGatewayHandler) handleFailoverExhausted(c *gin.Context, failoverE h.handleStreamingAwareError(c, status, errType, errMsg, streamStarted) } +func credentialFailoverClientResponse(failoverErr *service.UpstreamFailoverError) (int, string) { + _ = failoverErr + return http.StatusServiceUnavailable, service.GrokCredentialUnavailableClientMessage +} + +func copyFailoverRetryAfter(c *gin.Context, headers http.Header) { + if c == nil || headers == nil { + return + } + retryAfter := strings.TrimSpace(headers.Get("Retry-After")) + if retryAfter == "" || len(retryAfter) > 128 || strings.ContainsAny(retryAfter, "\r\n") || !isSafeRetryAfter(retryAfter) { + return + } + c.Header("Retry-After", retryAfter) +} + +func isSafeRetryAfter(value string) bool { + digitsOnly := true + for _, char := range value { + if char < '0' || char > '9' { + digitsOnly = false + break + } + } + if digitsOnly { + seconds, err := strconv.ParseUint(value, 10, 32) + return err == nil && seconds <= uint64((7*24*time.Hour)/time.Second) + } + retryAt, err := http.ParseTime(value) + if err != nil { + return false + } + return !retryAt.After(time.Now().Add(7 * 24 * time.Hour)) +} + // handleFailoverExhaustedSimple 简化版本,用于没有响应体的情况 func (h *OpenAIGatewayHandler) handleFailoverExhaustedSimple(c *gin.Context, statusCode int, streamStarted bool) { status, errType, errMsg := h.mapUpstreamError(statusCode) @@ -2194,6 +2343,10 @@ func closeOpenAIWSFailoverExhausted(conn *coderws.Conn, failoverErr *service.Ups closeOpenAIClientWS(conn, coderws.StatusInternalError, "upstream websocket proxy failed") return } + if failoverErr.Stage == service.GatewayFailureStageAccountAuth { + closeOpenAIClientWS(conn, coderws.StatusTryAgainLater, service.GrokCredentialUnavailableClientMessage) + return + } switch failoverErr.StatusCode { case http.StatusTooManyRequests: closeOpenAIClientWS(conn, coderws.StatusTryAgainLater, "upstream rate limit exceeded, please retry later") diff --git a/backend/internal/handler/openai_images.go b/backend/internal/handler/openai_images.go index c5982fb7d1..f5709225d1 100644 --- a/backend/internal/handler/openai_images.go +++ b/backend/internal/handler/openai_images.go @@ -145,6 +145,7 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) { stopJSONKeepalive := func() {} jsonKeepaliveStarted := false defer func() { stopJSONKeepalive() }() + var oauth429FailoverState service.OpenAIOAuth429FailoverState for { reqLog.Debug("openai.images.account_selecting", zap.Int("excluded_account_count", len(failedAccountIDs))) @@ -157,6 +158,10 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) { parsed.RequiredCapability, ) if err != nil { + if failoverClientGone(c) { + reqLog.Info("openai.images.account_select_aborted_client_disconnected", zap.Error(err)) + return + } reqLog.Warn("openai.images.account_select_failed", zap.Error(err), zap.Int("excluded_account_count", len(failedAccountIDs)), @@ -273,6 +278,13 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) { h.handleFailoverExhausted(c, failoverErr, true) return } + if failoverClientGone(c) { + reqLog.Info("openai.images.failover_aborted_client_disconnected", + zap.Int64("account_id", account.ID), + zap.Int("upstream_status", failoverErr.StatusCode), + ) + return + } if failoverErr.RetryableOnSameAccount { retryLimit := account.GetPoolModeRetryCount() if sameAccountRetryCount[account.ID] < retryLimit { @@ -299,7 +311,7 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) { return } switchCount++ - if h.gatewayService.ShouldStopOpenAIOAuth429Failover(account, failoverErr.StatusCode, switchCount) { + if h.gatewayService.ShouldStopOpenAIOAuth429Failover(account, failoverErr.StatusCode, switchCount, &oauth429FailoverState) { h.handleFailoverExhausted(c, failoverErr, streamStarted) return } diff --git a/backend/internal/handler/openai_responses_failover_cancel_test.go b/backend/internal/handler/openai_responses_failover_cancel_test.go new file mode 100644 index 0000000000..0d806e19d0 --- /dev/null +++ b/backend/internal/handler/openai_responses_failover_cancel_test.go @@ -0,0 +1,194 @@ +//go:build unit + +package handler + +import ( + "bytes" + "context" + "io" + "net/http" + "net/http/httptest" + "sync" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + 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" + "github.com/tidwall/gjson" +) + +// openAIResponsesFailoverCancelUpstream 固定返回 HTTP 520,可在首次上游调用时 +// 触发回调(用于模拟“上游在途期间客户端断开”)。 +type openAIResponsesFailoverCancelUpstream struct { + service.HTTPUpstream + mu sync.Mutex + accountIDs []int64 + onFirstDo func() +} + +func (u *openAIResponsesFailoverCancelUpstream) Do(_ *http.Request, _ string, accountID int64, _ int) (*http.Response, error) { + u.mu.Lock() + u.accountIDs = append(u.accountIDs, accountID) + first := len(u.accountIDs) == 1 + u.mu.Unlock() + if first && u.onFirstDo != nil { + u.onFirstDo() + } + return &http.Response{ + StatusCode: 520, + Header: http.Header{"Content-Type": []string{"text/html"}}, + Body: io.NopCloser(bytes.NewBufferString("520: unknown error")), + }, nil +} + +func (u *openAIResponsesFailoverCancelUpstream) calls() []int64 { + u.mu.Lock() + defer u.mu.Unlock() + return append([]int64(nil), u.accountIDs...) +} + +func newOpenAIResponsesFailoverTestHandler(t *testing.T, upstream service.HTTPUpstream) *OpenAIGatewayHandler { + t.Helper() + accounts := []service.Account{ + { + ID: 1, + Name: "responses-account-1", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeOAuth, + Status: service.StatusActive, + Schedulable: true, + Concurrency: 0, + Priority: 0, + Credentials: map[string]any{"access_token": "token-1"}, + }, + { + ID: 2, + Name: "responses-account-2", + Platform: service.PlatformOpenAI, + Type: service.AccountTypeOAuth, + Status: service.StatusActive, + Schedulable: true, + Concurrency: 0, + Priority: 1, + Credentials: map[string]any{"access_token": "token-2"}, + }, + } + accountRepo := openAIImagesFailoverAccountRepo{accounts: accounts} + cfg := &config.Config{RunMode: config.RunModeSimple} + gatewayService := service.NewOpenAIGatewayService( + accountRepo, + nil, + nil, + nil, + nil, + nil, + nil, + cfg, + nil, + nil, + nil, + nil, + nil, + upstream, + nil, + nil, + nil, + nil, + nil, + nil, + nil, + nil, + ) + billingService := service.NewBillingCacheService(nil, nil, nil, nil, nil, nil, cfg, nil) + t.Cleanup(billingService.Stop) + concurrencyService := service.NewConcurrencyService(nil) + handler := NewOpenAIGatewayHandler( + gatewayService, + concurrencyService, + billingService, + service.NewAPIKeyService(nil, nil, nil, nil, nil, nil, cfg), + nil, + nil, + nil, + nil, + cfg, + ) + handler.maxAccountSwitches = 10 + return handler +} + +func newOpenAIResponsesFailoverTestContext(t *testing.T, ctx context.Context) (*gin.Context, *httptest.ResponseRecorder) { + t.Helper() + groupID := int64(3131) + body := []byte(`{"model":"gpt-5.1","stream":false,"input":"hello"}`) + req := httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body)) + if ctx != nil { + req = req.WithContext(ctx) + } + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = req + c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{ + ID: 99, + GroupID: &groupID, + Group: &service.Group{ + ID: groupID, + Platform: service.PlatformOpenAI, + }, + User: &service.User{ID: 100}, + }) + c.Set(string(middleware2.ContextKeyUser), middleware2.AuthSubject{UserID: 100, Concurrency: 0}) + return c, rec +} + +// TestOpenAIGatewayHandlerResponses_FailoverAbortsWhenClientDisconnected 复现 +// #4257:客户端在上游请求在途期间断开,上游随后返回可 failover 的 520。 +// 期望:不再用已取消的 context 重新选号(不触达账号 2)、不把取消误报成 +// 502 账号耗尽、请求按 499 归类。 +func TestOpenAIGatewayHandlerResponses_FailoverAbortsWhenClientDisconnected(t *testing.T) { + gin.SetMode(gin.TestMode) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + upstream := &openAIResponsesFailoverCancelUpstream{onFirstDo: cancel} + handler := newOpenAIResponsesFailoverTestHandler(t, upstream) + c, rec := newOpenAIResponsesFailoverTestContext(t, ctx) + + handler.Responses(c) + + require.Equal(t, []int64{1}, upstream.calls(), "客户端断开后不应再切换到账号 2") + require.Equal(t, statusClientClosedRequest, c.Writer.Status(), "应按 499 归类") + require.Zero(t, rec.Body.Len(), "不应写入 502 错误响应体") + + _, hasFinalUpstreamErr := c.Get(service.OpsUpstreamStatusCodeKey) + require.False(t, hasFinalUpstreamErr, "不应记录 failover 耗尽的上游错误终态") + + // 真实发生过的 520 应保留 failover 事件(service 层在返回 failover 错误前记录) + rawEvents, ok := c.Get(service.OpsUpstreamErrorsKey) + require.True(t, ok) + events, ok := rawEvents.([]*service.OpsUpstreamErrorEvent) + require.True(t, ok) + require.Len(t, events, 1) + require.Equal(t, "failover", events[0].Kind) + require.Equal(t, 520, events[0].UpstreamStatusCode) +} + +// TestOpenAIGatewayHandlerResponses_FailoverContinuesForConnectedClient 回归 +// 守卫:客户端在线时 failover 行为不变——切换到账号 2,两个账号都 520 后按 +// 耗尽返回 502。 +func TestOpenAIGatewayHandlerResponses_FailoverContinuesForConnectedClient(t *testing.T) { + gin.SetMode(gin.TestMode) + + upstream := &openAIResponsesFailoverCancelUpstream{} + handler := newOpenAIResponsesFailoverTestHandler(t, upstream) + c, rec := newOpenAIResponsesFailoverTestContext(t, nil) + + handler.Responses(c) + + require.Equal(t, []int64{1, 2}, upstream.calls(), "在线客户端应正常切换账号") + require.Equal(t, http.StatusBadGateway, rec.Code) + require.Equal(t, "upstream_error", gjson.GetBytes(rec.Body.Bytes(), "error.type").String()) +} diff --git a/backend/internal/handler/ops_error_logger.go b/backend/internal/handler/ops_error_logger.go index 64aa3ba495..f32ca9c054 100644 --- a/backend/internal/handler/ops_error_logger.go +++ b/backend/internal/handler/ops_error_logger.go @@ -721,10 +721,15 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc { var upstreamStatusCode *int var upstreamErrorMessage *string var upstreamErrorDetail *string + finalAccountAuth := false if len(events) > 0 { last := events[len(events)-1] if last != nil { - if last.UpstreamStatusCode > 0 { + finalAccountAuth = last.Stage == string(service.GatewayFailureStageAccountAuth) + if finalAccountAuth { + code := 0 + upstreamStatusCode = &code + } else if last.UpstreamStatusCode > 0 { code := last.UpstreamStatusCode upstreamStatusCode = &code } @@ -737,7 +742,7 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc { } } - if upstreamStatusCode == nil { + if !finalAccountAuth && upstreamStatusCode == nil { if v, ok := c.Get(service.OpsUpstreamStatusCodeKey); ok { switch t := v.(type) { case int: @@ -753,7 +758,7 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc { } } } - if upstreamErrorMessage == nil { + if !finalAccountAuth && upstreamErrorMessage == nil { if v, ok := c.Get(service.OpsUpstreamErrorMessageKey); ok { if s, ok := v.(string); ok && strings.TrimSpace(s) != "" { msg := strings.TrimSpace(s) @@ -761,7 +766,7 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc { } } } - if upstreamErrorDetail == nil { + if !finalAccountAuth && upstreamErrorDetail == nil { if v, ok := c.Get(service.OpsUpstreamErrorDetailKey); ok { if s, ok := v.(string); ok && strings.TrimSpace(s) != "" { detail := strings.TrimSpace(s) @@ -781,13 +786,18 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc { } recoveredMsg := "Recovered upstream error" - if effectiveUpstreamStatus > 0 { + if finalAccountAuth { + recoveredMsg = "Recovered account authentication failure" + } else if effectiveUpstreamStatus > 0 { recoveredMsg += " " + strconvItoa(effectiveUpstreamStatus) } if upstreamErrorMessage != nil && strings.TrimSpace(*upstreamErrorMessage) != "" { recoveredMsg += ": " + strings.TrimSpace(*upstreamErrorMessage) } recoveredMsg = truncateString(recoveredMsg, 2048) + recoveredPhase, recoveredBusinessLimited, recoveredOwner, recoveredSource := classifyOpsErrorLog( + c, "upstream_error", recoveredMsg, "", effectiveUpstreamStatus, + ) entry := &service.OpsInsertErrorLogInput{ RequestID: requestID, @@ -828,19 +838,19 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc { }(), UserAgent: c.GetHeader("User-Agent"), - ErrorPhase: "upstream", + ErrorPhase: recoveredPhase, ErrorType: "upstream_error", // Severity should reflect the upstream failure, not the final client status (200). Severity: classifyOpsSeverity("upstream_error", effectiveUpstreamStatus), StatusCode: status, - IsBusinessLimited: false, + IsBusinessLimited: recoveredBusinessLimited, IsCountTokens: isCountTokensRequest(c), ErrorMessage: recoveredMsg, ErrorBody: "", - ErrorSource: "upstream_http", - ErrorOwner: "provider", + ErrorSource: recoveredSource, + ErrorOwner: recoveredOwner, UpstreamStatusCode: upstreamStatusCode, UpstreamErrorMessage: upstreamErrorMessage, @@ -850,6 +860,7 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc { CreatedAt: time.Now(), } applyOpsLatencyFieldsFromContext(c, entry) + applyOpsUpstreamFieldsFromContext(c, entry) if apiKey != nil { entry.APIKeyID = &apiKey.ID @@ -987,60 +998,7 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc { CreatedAt: time.Now(), } applyOpsLatencyFieldsFromContext(c, entry) - - // Capture upstream error context set by gateway services (if present). - // This does NOT affect the client response; it enriches Ops troubleshooting data. - { - if v, ok := c.Get(service.OpsUpstreamStatusCodeKey); ok { - switch t := v.(type) { - case int: - if t > 0 { - code := t - entry.UpstreamStatusCode = &code - } - case int64: - if t > 0 { - code := int(t) - entry.UpstreamStatusCode = &code - } - } - } - if v, ok := c.Get(service.OpsUpstreamErrorMessageKey); ok { - if s, ok := v.(string); ok { - if msg := strings.TrimSpace(s); msg != "" { - entry.UpstreamErrorMessage = &msg - } - } - } - if v, ok := c.Get(service.OpsUpstreamErrorDetailKey); ok { - if s, ok := v.(string); ok { - if detail := strings.TrimSpace(s); detail != "" { - entry.UpstreamErrorDetail = &detail - } - } - } - if v, ok := c.Get(service.OpsUpstreamErrorsKey); ok { - if events, ok := v.([]*service.OpsUpstreamErrorEvent); ok && len(events) > 0 { - entry.UpstreamErrors = events - // Best-effort backfill the single upstream fields from the last event when missing. - last := events[len(events)-1] - if last != nil { - if entry.UpstreamStatusCode == nil && last.UpstreamStatusCode > 0 { - code := last.UpstreamStatusCode - entry.UpstreamStatusCode = &code - } - if entry.UpstreamErrorMessage == nil && strings.TrimSpace(last.Message) != "" { - msg := strings.TrimSpace(last.Message) - entry.UpstreamErrorMessage = &msg - } - if entry.UpstreamErrorDetail == nil && strings.TrimSpace(last.Detail) != "" { - detail := strings.TrimSpace(last.Detail) - entry.UpstreamErrorDetail = &detail - } - } - } - } - } + applyOpsUpstreamFieldsFromContext(c, entry) if apiKey != nil { entry.APIKeyID = &apiKey.ID @@ -1235,6 +1193,77 @@ func applyOpsLatencyFieldsFromContext(c *gin.Context, entry *service.OpsInsertEr entry.TimeToFirstTokenMs = getContextLatencyMs(c, service.OpsTimeToFirstTokenMsKey) } +// applyOpsUpstreamFieldsFromContext captures attempt-level upstream context. +// A final account_auth event owns the top-level status and forces it to zero; +// prior inference statuses remain available in UpstreamErrors. +func applyOpsUpstreamFieldsFromContext(c *gin.Context, entry *service.OpsInsertErrorLogInput) { + if c == nil || entry == nil { + return + } + if v, ok := c.Get(service.OpsUpstreamStatusCodeKey); ok { + switch t := v.(type) { + case int: + if t > 0 { + code := t + entry.UpstreamStatusCode = &code + } + case int64: + if t > 0 { + code := int(t) + entry.UpstreamStatusCode = &code + } + } + } + if v, ok := c.Get(service.OpsUpstreamErrorMessageKey); ok { + if value, ok := v.(string); ok { + if message := strings.TrimSpace(value); message != "" { + entry.UpstreamErrorMessage = &message + } + } + } + if v, ok := c.Get(service.OpsUpstreamErrorDetailKey); ok { + if value, ok := v.(string); ok { + if detail := strings.TrimSpace(value); detail != "" { + entry.UpstreamErrorDetail = &detail + } + } + } + if v, ok := c.Get(service.OpsUpstreamErrorsKey); ok { + if events, ok := v.([]*service.OpsUpstreamErrorEvent); ok && len(events) > 0 { + entry.UpstreamErrors = events + last := events[len(events)-1] + if last == nil { + return + } + if last.Stage == string(service.GatewayFailureStageAccountAuth) { + code := 0 + entry.UpstreamStatusCode = &code + entry.UpstreamErrorMessage = nil + if message := strings.TrimSpace(last.Message); message != "" { + entry.UpstreamErrorMessage = &message + } + entry.UpstreamErrorDetail = nil + if detail := strings.TrimSpace(last.Detail); detail != "" { + entry.UpstreamErrorDetail = &detail + } + } else { + if entry.UpstreamStatusCode == nil && last.UpstreamStatusCode > 0 { + code := last.UpstreamStatusCode + entry.UpstreamStatusCode = &code + } + if entry.UpstreamErrorMessage == nil && strings.TrimSpace(last.Message) != "" { + message := strings.TrimSpace(last.Message) + entry.UpstreamErrorMessage = &message + } + if entry.UpstreamErrorDetail == nil && strings.TrimSpace(last.Detail) != "" { + detail := strings.TrimSpace(last.Detail) + entry.UpstreamErrorDetail = &detail + } + } + } + } +} + func getContextLatencyMs(c *gin.Context, key string) *int64 { if c == nil || strings.TrimSpace(key) == "" { return nil @@ -1390,7 +1419,7 @@ func normalizeOpsErrorType(errType string, code string) string { func classifyOpsPhase(errType, message, code string) string { msg := strings.ToLower(message) - // Standardized phases: request|auth|routing|upstream|network|internal + // Standardized phases: request|auth|account_auth|routing|upstream|network|internal // Map billing/concurrency/response => request; scheduling => routing. if isOpsClientAuthError(code, msg) { return "auth" @@ -1445,7 +1474,10 @@ func classifyOpsErrorLog(c *gin.Context, errType, message, code string, status i routingCapacityLimited := isOpsRoutingCapacityLimited(c) clientBusinessLimited := service.HasOpsClientBusinessLimited(c) upstreamError := hasOpsUpstreamErrorContext(c) - if upstreamError && !routingCapacityLimited { + accountAuthFailure := hasOpsAccountAuthFailure(c) + if accountAuthFailure && !routingCapacityLimited { + phase = "account_auth" + } else if upstreamError && !routingCapacityLimited { phase = "upstream" } if clientBusinessLimited && !upstreamError && !routingCapacityLimited { @@ -1570,6 +1602,22 @@ func hasOpsUpstreamErrorContext(c *gin.Context) bool { return false } +func hasOpsAccountAuthFailure(c *gin.Context) bool { + if c == nil { + return false + } + if v, ok := c.Get(service.OpsUpstreamErrorsKey); ok { + if events, ok := v.([]*service.OpsUpstreamErrorEvent); ok { + for i := len(events) - 1; i >= 0; i-- { + if events[i] != nil { + return events[i].Stage == string(service.GatewayFailureStageAccountAuth) + } + } + } + } + return false +} + func isOpsNoAvailableAccountMessage(message string) bool { msg := strings.ToLower(message) return strings.Contains(msg, opsErrNoAvailableAccounts) || @@ -1584,6 +1632,8 @@ func classifyOpsErrorOwner(phase string, message string) string { switch phase { case "upstream", "network": return "provider" + case "account_auth": + return "provider" case "request", "auth": return "client" case "routing", "internal": @@ -1601,6 +1651,8 @@ func classifyOpsErrorSource(phase string, message string) string { switch phase { case "upstream": return "upstream_http" + case "account_auth": + return "gateway" case "network": return "gateway" case "request", "auth": diff --git a/backend/internal/pkg/xai/oauth.go b/backend/internal/pkg/xai/oauth.go index a8a549c9dc..ea837fd78a 100644 --- a/backend/internal/pkg/xai/oauth.go +++ b/backend/internal/pkg/xai/oauth.go @@ -5,6 +5,7 @@ import ( "crypto/sha256" "encoding/base64" "encoding/hex" + "errors" "fmt" "net/url" "os" @@ -164,6 +165,24 @@ func ValidatedBaseURL(override string) (string, error) { return ValidateBaseURL(EffectiveBaseURL(override)) } +// BaseURLValidator applies the caller's outbound URL trust policy before xAI +// endpoint paths are appended. The service layer uses this for API-key accounts +// so the global security.url_allowlist policy remains the single source of +// truth; OAuth callers keep using the strict trusted-host validator. +type BaseURLValidator func(string) (string, error) + +func validatedBaseURLWithValidator(override string, validator BaseURLValidator) (string, error) { + if validator == nil { + return ValidatedBaseURL(override) + } + raw := EffectiveBaseURL(override) + validated, err := validator(raw) + if err != nil { + return "", err + } + return normalizeKnownBaseURLPath(validated) +} + type RuntimeSanityCheck struct { Value string `json:"value"` Valid bool `json:"valid"` @@ -282,7 +301,16 @@ func ValidateTrustedBaseURL(raw string) (string, error) { func normalizeKnownBaseURLPath(raw string) (string, error) { parsed, err := url.Parse(raw) if err != nil || parsed.Scheme == "" || parsed.Host == "" { - return "", fmt.Errorf("invalid url: %s", raw) + return "", errors.New("invalid base URL") + } + if parsed.User != nil { + return "", errors.New("base URL must not include userinfo") + } + if parsed.RawQuery != "" { + return "", errors.New("base URL must not include a query") + } + if parsed.Fragment != "" { + return "", errors.New("base URL must not include a fragment") } path := strings.TrimRight(parsed.Path, "/") if path == "" { @@ -435,7 +463,11 @@ func ParseAuthorizationInput(raw string) AuthorizationInput { } func BuildResponsesURL(baseURL string) (string, error) { - validatedBaseURL, err := ValidatedBaseURL(baseURL) + return BuildResponsesURLWithValidator(baseURL, nil) +} + +func BuildResponsesURLWithValidator(baseURL string, validator BaseURLValidator) (string, error) { + validatedBaseURL, err := validatedBaseURLWithValidator(baseURL, validator) if err != nil { return "", fmt.Errorf("invalid base url: %w", err) } @@ -443,7 +475,11 @@ func BuildResponsesURL(baseURL string) (string, error) { } func BuildChatCompletionsURL(baseURL string) (string, error) { - validatedBaseURL, err := ValidatedBaseURL(baseURL) + return BuildChatCompletionsURLWithValidator(baseURL, nil) +} + +func BuildChatCompletionsURLWithValidator(baseURL string, validator BaseURLValidator) (string, error) { + validatedBaseURL, err := validatedBaseURLWithValidator(baseURL, validator) if err != nil { return "", fmt.Errorf("invalid base url: %w", err) } @@ -451,7 +487,11 @@ func BuildChatCompletionsURL(baseURL string) (string, error) { } func BuildImagesGenerationsURL(baseURL string) (string, error) { - validatedBaseURL, err := ValidatedBaseURL(baseURL) + return BuildImagesGenerationsURLWithValidator(baseURL, nil) +} + +func BuildImagesGenerationsURLWithValidator(baseURL string, validator BaseURLValidator) (string, error) { + validatedBaseURL, err := validatedBaseURLWithValidator(baseURL, validator) if err != nil { return "", fmt.Errorf("invalid base url: %w", err) } @@ -459,7 +499,11 @@ func BuildImagesGenerationsURL(baseURL string) (string, error) { } func BuildImagesEditsURL(baseURL string) (string, error) { - validatedBaseURL, err := ValidatedBaseURL(baseURL) + return BuildImagesEditsURLWithValidator(baseURL, nil) +} + +func BuildImagesEditsURLWithValidator(baseURL string, validator BaseURLValidator) (string, error) { + validatedBaseURL, err := validatedBaseURLWithValidator(baseURL, validator) if err != nil { return "", fmt.Errorf("invalid base url: %w", err) } @@ -467,7 +511,11 @@ func BuildImagesEditsURL(baseURL string) (string, error) { } func BuildVideosGenerationsURL(baseURL string) (string, error) { - validatedBaseURL, err := ValidatedBaseURL(baseURL) + return BuildVideosGenerationsURLWithValidator(baseURL, nil) +} + +func BuildVideosGenerationsURLWithValidator(baseURL string, validator BaseURLValidator) (string, error) { + validatedBaseURL, err := validatedBaseURLWithValidator(baseURL, validator) if err != nil { return "", fmt.Errorf("invalid base url: %w", err) } @@ -475,7 +523,11 @@ func BuildVideosGenerationsURL(baseURL string) (string, error) { } func BuildVideosEditsURL(baseURL string) (string, error) { - validatedBaseURL, err := ValidatedBaseURL(baseURL) + return BuildVideosEditsURLWithValidator(baseURL, nil) +} + +func BuildVideosEditsURLWithValidator(baseURL string, validator BaseURLValidator) (string, error) { + validatedBaseURL, err := validatedBaseURLWithValidator(baseURL, validator) if err != nil { return "", fmt.Errorf("invalid base url: %w", err) } @@ -483,7 +535,11 @@ func BuildVideosEditsURL(baseURL string) (string, error) { } func BuildVideosExtensionsURL(baseURL string) (string, error) { - validatedBaseURL, err := ValidatedBaseURL(baseURL) + return BuildVideosExtensionsURLWithValidator(baseURL, nil) +} + +func BuildVideosExtensionsURLWithValidator(baseURL string, validator BaseURLValidator) (string, error) { + validatedBaseURL, err := validatedBaseURLWithValidator(baseURL, validator) if err != nil { return "", fmt.Errorf("invalid base url: %w", err) } @@ -491,7 +547,11 @@ func BuildVideosExtensionsURL(baseURL string) (string, error) { } func BuildVideoURL(baseURL, requestID string) (string, error) { - validatedBaseURL, err := ValidatedBaseURL(baseURL) + return BuildVideoURLWithValidator(baseURL, requestID, nil) +} + +func BuildVideoURLWithValidator(baseURL, requestID string, validator BaseURLValidator) (string, error) { + validatedBaseURL, err := validatedBaseURLWithValidator(baseURL, validator) if err != nil { return "", fmt.Errorf("invalid base url: %w", err) } diff --git a/backend/internal/pkg/xai/oauth_test.go b/backend/internal/pkg/xai/oauth_test.go index 39200ffc39..a82a26e299 100644 --- a/backend/internal/pkg/xai/oauth_test.go +++ b/backend/internal/pkg/xai/oauth_test.go @@ -6,6 +6,7 @@ import ( "net/url" "testing" + "github.com/Wei-Shaw/sub2api/internal/util/urlvalidator" "github.com/stretchr/testify/require" ) @@ -165,6 +166,44 @@ func TestValidateBaseURLAllowsPublicThirdPartyGrokAPI(t *testing.T) { require.Error(t, err) } +func TestBuildResponsesURLWithValidatorUsesCallerPolicy(t *testing.T) { + validator := func(raw string) (string, error) { + return urlvalidator.ValidateURLFormat(raw, true) + } + + target, err := BuildResponsesURLWithValidator("http://grok.example.test/v1/", validator) + require.NoError(t, err) + require.Equal(t, "http://grok.example.test/v1/responses", target) +} + +func TestBuildResponsesURLPreservesUnsafeOverrideCustomPath(t *testing.T) { + t.Setenv(EnvAllowUnsafeURLOverrides, "true") + + target, err := BuildResponsesURL("http://localhost:8080/custom") + require.NoError(t, err) + require.Equal(t, "http://localhost:8080/custom/responses", target) +} + +func TestBuildResponsesURLWithValidatorRejectsBaseURLComponents(t *testing.T) { + permissive := func(raw string) (string, error) { return raw, nil } + tests := []struct { + name string + raw string + }{ + {name: "userinfo", raw: "https://user:secret@grok.example.test/v1"}, + {name: "query", raw: "https://grok.example.test/v1?token=secret"}, + {name: "fragment", raw: "https://grok.example.test/v1#secret"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + _, err := BuildResponsesURLWithValidator(tt.raw, permissive) + require.Error(t, err) + require.NotContains(t, err.Error(), "secret") + }) + } +} + func TestValidateXAIURLsAllowUnsafeDevOverride(t *testing.T) { t.Setenv(EnvAllowUnsafeURLOverrides, "true") diff --git a/backend/internal/repository/account_repo.go b/backend/internal/repository/account_repo.go index 8eb819aeab..e85bf46fc1 100644 --- a/backend/internal/repository/account_repo.go +++ b/backend/internal/repository/account_repo.go @@ -859,6 +859,54 @@ func (r *accountRepository) SetError(ctx context.Context, id int64, errorMsg str return nil } +func (r *accountRepository) SetGrokCredentialErrorIfMatch( + ctx context.Context, + id int64, + snapshot service.GrokCredentialMutationSnapshot, + errorMsg string, +) (bool, error) { + result, err := r.sql.ExecContext(ctx, ` + WITH updated AS ( + UPDATE accounts AS a + SET status = $1, + error_message = $2, + schedulable = false, + updated_at = NOW() + WHERE a.id = $3 + AND a.deleted_at IS NULL + AND a.status = $4 + AND a.platform = $5 + AND a.type = $6 + AND a.schedulable IS TRUE + AND (a.temp_unschedulable_until IS NULL OR a.temp_unschedulable_until <= NOW()) + AND (a.rate_limit_reset_at IS NULL OR a.rate_limit_reset_at <= NOW()) + AND (a.overload_until IS NULL OR a.overload_until <= NOW()) + AND (a.auto_pause_on_expired IS NOT TRUE OR a.expires_at IS NULL OR a.expires_at > NOW()) + AND a.credentials = $7::jsonb + AND a.proxy_id IS NOT DISTINCT FROM $8 + AND ($2 <> $9 OR ( + a.proxy_id IS NOT NULL AND NOT EXISTS ( + SELECT 1 FROM proxies p WHERE p.id = a.proxy_id AND p.deleted_at IS NULL + ) + )) + RETURNING a.id + ) + INSERT INTO scheduler_outbox (event_type, account_id, group_id, payload) + SELECT $10, updated.id, NULL, NULL FROM updated + `, service.StatusError, errorMsg, id, service.StatusActive, service.PlatformGrok, service.AccountTypeOAuth, + snapshot.CredentialsJSON, snapshot.ProxyID, string(service.GrokCredentialReasonProxyInvalid), + service.SchedulerOutboxEventAccountChanged) + if err != nil { + return false, err + } + affected, err := result.RowsAffected() + if err != nil || affected == 0 { + return false, err + } + r.syncSchedulerAccountSnapshotDetached(ctx, id) + return true, nil +} + // syncSchedulerAccountSnapshot 在账号状态变更时主动同步快照到调度器缓存。 // 当账号被设置为错误、禁用、不可调度或临时不可调度时调用, // 确保调度器和粘性会话逻辑能及时感知账号的最新状态,避免继续使用不可用账号。 @@ -881,6 +929,16 @@ func (r *accountRepository) syncSchedulerAccountSnapshot(ctx context.Context, ac } } +func (r *accountRepository) syncSchedulerAccountSnapshotDetached(ctx context.Context, accountID int64) { + base := context.Background() + if ctx != nil { + base = context.WithoutCancel(ctx) + } + propagationCtx, cancel := context.WithTimeout(base, 2*time.Second) + defer cancel() + r.syncSchedulerAccountSnapshot(propagationCtx, accountID) +} + func (r *accountRepository) deleteSchedulerAccountSnapshot(ctx context.Context, accountID int64) { if r == nil || r.schedulerCache == nil || accountID <= 0 { return @@ -1050,8 +1108,42 @@ func (r *accountRepository) BindGroups(ctx context.Context, accountID int64, gro } func (r *accountRepository) ListSchedulable(ctx context.Context) ([]service.Account, error) { - now := time.Now() - accounts, err := r.client.Account.Query(). + accounts, err := r.schedulableAccountsQuery(time.Now()).All(ctx) + if err != nil { + return nil, err + } + return r.accountsToService(ctx, accounts) +} + +func (r *accountRepository) ListSchedulableAccountLoads(ctx context.Context) ([]service.AccountWithConcurrency, error) { + accounts, err := r.schedulableAccountsQuery(time.Now()). + Select( + dbaccount.FieldID, + dbaccount.FieldConcurrency, + dbaccount.FieldLoadFactor, + ). + All(ctx) + if err != nil { + return nil, err + } + + loads := make([]service.AccountWithConcurrency, 0, len(accounts)) + for _, account := range accounts { + projection := service.Account{ + ID: account.ID, + Concurrency: account.Concurrency, + LoadFactor: account.LoadFactor, + } + loads = append(loads, service.AccountWithConcurrency{ + ID: account.ID, + MaxConcurrency: projection.EffectiveLoadFactor(), + }) + } + return loads, nil +} + +func (r *accountRepository) schedulableAccountsQuery(now time.Time) *dbent.AccountQuery { + return r.client.Account.Query(). Where( dbaccount.StatusEQ(service.StatusActive), dbaccount.SchedulableEQ(true), @@ -1060,12 +1152,7 @@ func (r *accountRepository) ListSchedulable(ctx context.Context) ([]service.Acco dbaccount.Or(dbaccount.OverloadUntilIsNil(), dbaccount.OverloadUntilLTE(now)), dbaccount.Or(dbaccount.RateLimitResetAtIsNil(), dbaccount.RateLimitResetAtLTE(now)), ). - Order(dbent.Asc(dbaccount.FieldPriority)). - All(ctx) - if err != nil { - return nil, err - } - return r.accountsToService(ctx, accounts) + Order(dbent.Asc(dbaccount.FieldPriority)) } func (r *accountRepository) ListSchedulableByGroupID(ctx context.Context, groupID int64) ([]service.Account, error) { @@ -1319,6 +1406,35 @@ func (r *accountRepository) SetRateLimitedIfLater(ctx context.Context, id int64, return nil } +// ClearRateLimitIfObserved clears exactly the Grok rate-limit generation seen +// by a successful request. Matching both timestamps prevents a stale success +// from erasing a later clear/re-arm generation with an equal or shorter reset. +func (r *accountRepository) ClearRateLimitIfObserved(ctx context.Context, id int64, observedLimitedAt, observedResetAt time.Time) (bool, error) { + updated, err := r.client.Account.Update(). + Where( + dbaccount.IDEQ(id), + dbaccount.PlatformEQ(service.PlatformGrok), + dbaccount.TypeEQ(service.AccountTypeOAuth), + dbaccount.RateLimitedAtEQ(observedLimitedAt), + dbaccount.RateLimitResetAtEQ(observedResetAt), + ). + ClearRateLimitedAt(). + ClearRateLimitResetAt(). + Save(ctx) + if err != nil { + return false, err + } + if updated == 0 { + r.syncSchedulerAccountSnapshot(ctx, id) + return false, nil + } + if err := enqueueSchedulerOutbox(ctx, r.sql, service.SchedulerOutboxEventAccountChanged, &id, nil, nil); err != nil { + logger.LegacyPrintf("repository.account", "[SchedulerOutbox] enqueue observed rate-limit clear failed: account=%d err=%v", id, err) + } + r.syncSchedulerAccountSnapshot(ctx, id) + return true, nil +} + func (r *accountRepository) SetModelRateLimit(ctx context.Context, id int64, scope string, resetAt time.Time, reason ...string) error { if scope == "" { return nil @@ -1414,6 +1530,51 @@ func (r *accountRepository) SetTempUnschedulable(ctx context.Context, id int64, return nil } +func (r *accountRepository) SetGrokCredentialTempUnschedulableIfMatch( + ctx context.Context, + id int64, + snapshot service.GrokCredentialMutationSnapshot, + until time.Time, + reason string, +) (bool, error) { + result, err := r.sql.ExecContext(ctx, ` + WITH updated AS ( + UPDATE accounts AS a + SET temp_unschedulable_until = CASE + WHEN a.temp_unschedulable_until IS NULL OR a.temp_unschedulable_until < $1 THEN $1 + ELSE a.temp_unschedulable_until + END, + temp_unschedulable_reason = $2, + updated_at = NOW() + WHERE a.id = $3 + AND a.deleted_at IS NULL + AND a.status = $4 + AND a.platform = $5 + AND a.type = $6 + AND a.schedulable IS TRUE + AND (a.temp_unschedulable_until IS NULL OR a.temp_unschedulable_until <= NOW()) + AND (a.rate_limit_reset_at IS NULL OR a.rate_limit_reset_at <= NOW()) + AND (a.overload_until IS NULL OR a.overload_until <= NOW()) + AND (a.auto_pause_on_expired IS NOT TRUE OR a.expires_at IS NULL OR a.expires_at > NOW()) + AND a.credentials = $7::jsonb + AND a.proxy_id IS NOT DISTINCT FROM $8 + RETURNING a.id + ) + INSERT INTO scheduler_outbox (event_type, account_id, group_id, payload) + SELECT $9, updated.id, NULL, NULL FROM updated + `, until, reason, id, service.StatusActive, service.PlatformGrok, service.AccountTypeOAuth, + snapshot.CredentialsJSON, snapshot.ProxyID, service.SchedulerOutboxEventAccountChanged) + if err != nil { + return false, err + } + affected, err := result.RowsAffected() + if err != nil || affected == 0 { + return false, err + } + r.syncSchedulerAccountSnapshotDetached(ctx, id) + return true, nil +} + func (r *accountRepository) ClearTempUnschedulable(ctx context.Context, id int64) error { _, err := r.sql.ExecContext(ctx, ` UPDATE accounts diff --git a/backend/internal/repository/account_repo_integration_test.go b/backend/internal/repository/account_repo_integration_test.go index f82ad4ca28..ae519f3df4 100644 --- a/backend/internal/repository/account_repo_integration_test.go +++ b/backend/internal/repository/account_repo_integration_test.go @@ -722,6 +722,56 @@ func (s *AccountRepoSuite) TestSetRateLimitedIfLaterDoesNotShortenReset() { s.Require().WithinDuration(later, *cacheRecorder.setAccounts[1].RateLimitResetAt, time.Second) } +func (s *AccountRepoSuite) TestClearRateLimitIfObservedProtectsRearmed429Generation() { + account := mustCreateAccount(s.T(), s.client, &service.Account{ + Name: "acc-rl-conditional-clear", + Platform: service.PlatformGrok, + Type: service.AccountTypeOAuth, + }) + firstReset := time.Now().Add(30 * time.Minute).UTC().Truncate(time.Second) + rearmedReset := time.Now().Add(5 * time.Minute).UTC().Truncate(time.Second) + + s.Require().NoError(s.repo.SetRateLimitedIfLater(s.ctx, account.ID, firstReset)) + staleGeneration, err := s.repo.GetByID(s.ctx, account.ID) + s.Require().NoError(err) + s.Require().NotNil(staleGeneration.RateLimitedAt) + s.Require().NotNil(staleGeneration.RateLimitResetAt) + cleared, err := s.repo.ClearRateLimitIfObserved(s.ctx, account.ID, *staleGeneration.RateLimitedAt, *staleGeneration.RateLimitResetAt) + s.Require().NoError(err) + s.Require().True(cleared) + + // A newer generation may legitimately re-arm a shorter boundary after the + // first generation was cleared. The stale success must not erase it. + s.Require().NoError(s.repo.SetRateLimitedIfLater(s.ctx, account.ID, rearmedReset)) + cleared, err = s.repo.ClearRateLimitIfObserved(s.ctx, account.ID, *staleGeneration.RateLimitedAt, *staleGeneration.RateLimitResetAt) + s.Require().NoError(err) + s.Require().False(cleared) + + got, err := s.repo.GetByID(s.ctx, account.ID) + s.Require().NoError(err) + s.Require().NotNil(got.RateLimitedAt) + s.Require().NotNil(got.RateLimitResetAt) + s.Require().WithinDuration(rearmedReset, *got.RateLimitResetAt, time.Second) + + // An admin can retype the row while the successful OAuth request is still + // in flight. The stale OAuth recovery must not cross into API-key state even + // when both observed timestamps still match. + _, err = s.client.Account.UpdateOneID(account.ID). + SetType(service.AccountTypeAPIKey). + Save(s.ctx) + s.Require().NoError(err) + cleared, err = s.repo.ClearRateLimitIfObserved(s.ctx, account.ID, *got.RateLimitedAt, *got.RateLimitResetAt) + s.Require().NoError(err) + s.Require().False(cleared) + + retyped, err := s.repo.GetByID(s.ctx, account.ID) + s.Require().NoError(err) + s.Require().Equal(service.AccountTypeAPIKey, retyped.Type) + s.Require().NotNil(retyped.RateLimitedAt) + s.Require().NotNil(retyped.RateLimitResetAt) + s.Require().WithinDuration(rearmedReset, *retyped.RateLimitResetAt, time.Second) +} + func (s *AccountRepoSuite) TestClearRateLimit() { account := mustCreateAccount(s.T(), s.client, &service.Account{Name: "acc-clear"}) until := time.Now().Add(1 * time.Hour) diff --git a/backend/internal/repository/account_repo_schedulable_projection_integration_test.go b/backend/internal/repository/account_repo_schedulable_projection_integration_test.go new file mode 100644 index 0000000000..82d62a1448 --- /dev/null +++ b/backend/internal/repository/account_repo_schedulable_projection_integration_test.go @@ -0,0 +1,107 @@ +//go:build integration + +package repository + +import ( + "context" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" +) + +func TestListSchedulableAccountLoadsMatchesListSchedulable(t *testing.T) { + ctx := context.Background() + tx := testEntTx(t) + client := tx.Client() + repo := newAccountRepositoryWithSQL(client, tx, nil) + now := time.Now() + past := now.Add(-time.Hour) + future := now.Add(time.Hour) + + create := func(name string) *service.Account { + return mustCreateAccount(t, client, &service.Account{Name: name, Schedulable: true}) + } + + positiveLoad := create("projection-positive-load") + _, err := client.Account.UpdateOneID(positiveLoad.ID).SetConcurrency(2).SetLoadFactor(9).SetPriority(30).Save(ctx) + require.NoError(t, err) + concurrencyFallback := create("projection-concurrency-fallback") + _, err = client.Account.UpdateOneID(concurrencyFallback.ID).SetConcurrency(4).SetPriority(10).Save(ctx) + require.NoError(t, err) + zeroFallback := create("projection-zero-fallback") + _, err = client.Account.UpdateOneID(zeroFallback.ID).SetConcurrency(0).SetLoadFactor(0).SetPriority(20).Save(ctx) + require.NoError(t, err) + + disabled := create("projection-disabled") + _, err = client.Account.UpdateOneID(disabled.ID).SetStatus(service.StatusDisabled).Save(ctx) + require.NoError(t, err) + unschedulable := create("projection-unschedulable") + _, err = client.Account.UpdateOneID(unschedulable.ID).SetSchedulable(false).Save(ctx) + require.NoError(t, err) + expired := create("projection-expired") + _, err = client.Account.UpdateOneID(expired.ID).SetExpiresAt(past).SetAutoPauseOnExpired(true).Save(ctx) + require.NoError(t, err) + expiredAllowed := create("projection-expired-allowed") + _, err = client.Account.UpdateOneID(expiredAllowed.ID).SetExpiresAt(past).SetAutoPauseOnExpired(false).Save(ctx) + require.NoError(t, err) + overloaded := create("projection-overloaded") + _, err = client.Account.UpdateOneID(overloaded.ID).SetOverloadUntil(future).Save(ctx) + require.NoError(t, err) + overloadCleared := create("projection-overload-cleared") + _, err = client.Account.UpdateOneID(overloadCleared.ID).SetOverloadUntil(past).Save(ctx) + require.NoError(t, err) + rateLimited := create("projection-rate-limited") + _, err = client.Account.UpdateOneID(rateLimited.ID).SetRateLimitResetAt(future).Save(ctx) + require.NoError(t, err) + rateLimitCleared := create("projection-rate-limit-cleared") + _, err = client.Account.UpdateOneID(rateLimitCleared.ID).SetRateLimitResetAt(past).Save(ctx) + require.NoError(t, err) + tempBlocked := create("projection-temp-blocked") + _, err = client.Account.UpdateOneID(tempBlocked.ID).SetTempUnschedulableUntil(future).Save(ctx) + require.NoError(t, err) + tempCleared := create("projection-temp-cleared") + _, err = client.Account.UpdateOneID(tempCleared.ID).SetTempUnschedulableUntil(past).Save(ctx) + require.NoError(t, err) + + accounts, err := repo.ListSchedulable(ctx) + require.NoError(t, err) + loads, err := repo.ListSchedulableAccountLoads(ctx) + require.NoError(t, err) + + accountIDs := make([]int64, 0, len(accounts)) + wantByID := make(map[int64]int, len(accounts)) + for i := range accounts { + accountIDs = append(accountIDs, accounts[i].ID) + wantByID[accounts[i].ID] = accounts[i].EffectiveLoadFactor() + } + + loadIDs := make([]int64, 0, len(loads)) + byID := make(map[int64]int, len(loads)) + for _, load := range loads { + loadIDs = append(loadIDs, load.ID) + byID[load.ID] = load.MaxConcurrency + } + require.Equal(t, accountIDs, loadIDs) + targetIDs := map[int64]struct{}{ + positiveLoad.ID: {}, concurrencyFallback.ID: {}, zeroFallback.ID: {}, + } + targetOrder := make([]int64, 0, len(targetIDs)) + for _, id := range loadIDs { + if _, ok := targetIDs[id]; ok { + targetOrder = append(targetOrder, id) + } + } + require.Equal(t, []int64{concurrencyFallback.ID, zeroFallback.ID, positiveLoad.ID}, targetOrder) + require.Equal(t, wantByID, byID) + require.Equal(t, 9, byID[positiveLoad.ID]) + require.Equal(t, 4, byID[concurrencyFallback.ID]) + require.Equal(t, 1, byID[zeroFallback.ID]) + for _, included := range []*service.Account{expiredAllowed, overloadCleared, rateLimitCleared, tempCleared} { + require.Contains(t, byID, included.ID) + } + for _, excluded := range []*service.Account{disabled, unschedulable, expired, overloaded, rateLimited, tempBlocked} { + require.NotContains(t, byID, excluded.ID) + } +} diff --git a/backend/internal/repository/account_repo_schedulable_projection_test.go b/backend/internal/repository/account_repo_schedulable_projection_test.go new file mode 100644 index 0000000000..76d7e0bdc5 --- /dev/null +++ b/backend/internal/repository/account_repo_schedulable_projection_test.go @@ -0,0 +1,82 @@ +package repository + +import ( + "context" + "fmt" + "strings" + "testing" + + "github.com/DATA-DOG/go-sqlmock" + dbent "github.com/Wei-Shaw/sub2api/ent" + _ "github.com/Wei-Shaw/sub2api/ent/runtime" + "github.com/stretchr/testify/require" + + "entgo.io/ent/dialect" + entsql "entgo.io/ent/dialect/sql" +) + +type captureEntQueryMatcher struct { + actual *string +} + +func (m captureEntQueryMatcher) Match(_, actual string) error { + if m.actual == nil { + return fmt.Errorf("query capture target is nil") + } + *m.actual = actual + return nil +} + +func TestListSchedulableAccountLoadsUsesSingleProjectionQuery(t *testing.T) { + var capturedSQL string + db, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(captureEntQueryMatcher{actual: &capturedSQL})) + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + + driver := entsql.OpenDB(dialect.Postgres, db) + client := dbent.NewClient(dbent.Driver(driver)) + t.Cleanup(func() { _ = client.Close() }) + repo := newAccountRepositoryWithSQL(client, db, nil) + + mock.ExpectQuery("schedulable account load projection"). + WillReturnRows(sqlmock.NewRows([]string{"id", "concurrency", "load_factor"}). + AddRow(int64(11), 3, nil). + AddRow(int64(12), 2, 7)) + + loads, err := repo.ListSchedulableAccountLoads(context.Background()) + require.NoError(t, err) + require.Len(t, loads, 2) + require.Equal(t, int64(11), loads[0].ID) + require.Equal(t, 3, loads[0].MaxConcurrency) + require.Equal(t, int64(12), loads[1].ID) + require.Equal(t, 7, loads[1].MaxConcurrency) + require.NoError(t, mock.ExpectationsWereMet(), "projection path must execute exactly one query") + + normalized := normalizeSQLWhitespace(capturedSQL) + selectClause, _, found := strings.Cut(normalized, " FROM ") + require.True(t, found, "unexpected projection SQL: %s", normalized) + require.Equal(t, 2, strings.Count(selectClause, ","), "projection must select exactly three columns: %s", selectClause) + require.Contains(t, selectClause, `"id"`) + require.Contains(t, selectClause, `"concurrency"`) + require.Contains(t, selectClause, `"load_factor"`) + require.NotContains(t, selectClause, "credentials") + require.NotContains(t, selectClause, "extra") + require.NotContains(t, selectClause, "proxy_id") + require.NotContains(t, normalized, "account_groups") + require.NotContains(t, normalized, "proxies") + for _, predicateColumn := range []string{ + "status", + "schedulable", + "temp_unschedulable_until", + "expires_at", + "auto_pause_on_expired", + "overload_until", + "rate_limit_reset_at", + "deleted_at", + } { + require.Contains(t, normalized, predicateColumn) + } + _, orderClause, hasOrder := strings.Cut(normalized, " ORDER BY ") + require.True(t, hasOrder, "projection query must preserve schedulable account order: %s", normalized) + require.Contains(t, orderClause, `"priority" ASC`) +} diff --git a/backend/internal/repository/account_repo_temp_unsched_test.go b/backend/internal/repository/account_repo_temp_unsched_test.go index 9c3e4cbdf5..2ea18bd214 100644 --- a/backend/internal/repository/account_repo_temp_unsched_test.go +++ b/backend/internal/repository/account_repo_temp_unsched_test.go @@ -9,6 +9,7 @@ import ( "time" sqlmock "github.com/DATA-DOG/go-sqlmock" + "github.com/Wei-Shaw/sub2api/internal/service" "github.com/stretchr/testify/require" ) @@ -24,6 +25,103 @@ func TestAccountRepository_SetTempUnschedulable_NoRowsAffectedDoesNotWriteOutbox require.NotContains(t, strings.Join(exec.execQueries, "\n"), "scheduler_outbox") } +func TestAccountRepository_GrokCredentialConditionalMutationsAreEligibleAndAtomicallyPropagated(t *testing.T) { + proxyID := int64(77) + snapshot := service.GrokCredentialMutationSnapshot{ + CredentialsJSON: `{"access_token":"access","refresh_token":"refresh","_token_version":123}`, + ProxyID: &proxyID, + } + + t.Run("permanent", func(t *testing.T) { + exec := &recordingSQLExecutor{result: rowsAffectedResult(0)} + repo := newAccountRepositoryWithSQL(nil, exec, nil) + + updated, err := repo.SetGrokCredentialErrorIfMatch(context.Background(), 42, snapshot, "revoked") + + require.NoError(t, err) + require.False(t, updated) + require.Len(t, exec.execQueries, 1) + normalized := normalizeSQLWhitespace(exec.execQueries[0]) + require.Contains(t, normalized, "WITH updated AS ( UPDATE accounts AS a") + require.Contains(t, normalized, "a.schedulable IS TRUE") + require.Contains(t, normalized, "a.temp_unschedulable_until IS NULL OR a.temp_unschedulable_until <= NOW()") + require.Contains(t, normalized, "a.rate_limit_reset_at IS NULL OR a.rate_limit_reset_at <= NOW()") + require.Contains(t, normalized, "a.overload_until IS NULL OR a.overload_until <= NOW()") + require.Contains(t, normalized, "a.credentials = $7::jsonb") + require.Contains(t, normalized, "a.proxy_id IS NOT DISTINCT FROM $8") + require.Contains(t, normalized, "NOT EXISTS ( SELECT 1 FROM proxies p") + require.Contains(t, normalized, "INSERT INTO scheduler_outbox") + require.Len(t, exec.execArgs[0], 10) + require.Equal(t, snapshot.CredentialsJSON, exec.execArgs[0][6]) + require.Equal(t, &proxyID, exec.execArgs[0][7]) + require.Equal(t, string(service.GrokCredentialReasonProxyInvalid), exec.execArgs[0][8]) + require.Equal(t, service.SchedulerOutboxEventAccountChanged, exec.execArgs[0][9]) + }) + + t.Run("transient", func(t *testing.T) { + exec := &recordingSQLExecutor{result: rowsAffectedResult(0)} + repo := newAccountRepositoryWithSQL(nil, exec, nil) + + updated, err := repo.SetGrokCredentialTempUnschedulableIfMatch( + context.Background(), 42, snapshot, time.Now().Add(time.Minute), "temporary", + ) + + require.NoError(t, err) + require.False(t, updated) + require.Len(t, exec.execQueries, 1) + normalized := normalizeSQLWhitespace(exec.execQueries[0]) + require.Contains(t, normalized, "WITH updated AS ( UPDATE accounts AS a") + require.Contains(t, normalized, "a.schedulable IS TRUE") + require.Contains(t, normalized, "a.temp_unschedulable_until IS NULL OR a.temp_unschedulable_until <= NOW()") + require.Contains(t, normalized, "a.rate_limit_reset_at IS NULL OR a.rate_limit_reset_at <= NOW()") + require.Contains(t, normalized, "a.overload_until IS NULL OR a.overload_until <= NOW()") + require.Contains(t, normalized, "a.credentials = $7::jsonb") + require.Contains(t, normalized, "a.proxy_id IS NOT DISTINCT FROM $8") + require.Contains(t, normalized, "INSERT INTO scheduler_outbox") + require.Len(t, exec.execArgs[0], 9) + require.Equal(t, snapshot.CredentialsJSON, exec.execArgs[0][6]) + require.Equal(t, &proxyID, exec.execArgs[0][7]) + require.Equal(t, service.SchedulerOutboxEventAccountChanged, exec.execArgs[0][8]) + }) +} + +func TestAccountRepository_GrokCredentialCommitCarriesOutboxAcrossCallerCancellation(t *testing.T) { + snapshot := service.GrokCredentialMutationSnapshot{CredentialsJSON: `{"access_token":"access","refresh_token":"refresh"}`} + tests := []struct { + name string + mutate func(context.Context, *accountRepository) (bool, error) + }{ + { + name: "permanent", + mutate: func(ctx context.Context, repo *accountRepository) (bool, error) { + return repo.SetGrokCredentialErrorIfMatch(ctx, 42, snapshot, string(service.GrokCredentialReasonRevoked)) + }, + }, + { + name: "transient", + mutate: func(ctx context.Context, repo *accountRepository) (bool, error) { + return repo.SetGrokCredentialTempUnschedulableIfMatch(ctx, 42, snapshot, time.Now().Add(time.Minute), string(service.GrokCredentialReasonRefreshTransient)) + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + exec := &recordingSQLExecutor{result: rowsAffectedResult(1), afterExec: cancel} + repo := newAccountRepositoryWithSQL(nil, exec, nil) + + updated, err := tt.mutate(ctx, repo) + + require.NoError(t, err) + require.True(t, updated) + require.ErrorIs(t, ctx.Err(), context.Canceled) + require.Len(t, exec.execQueries, 1, "state update and scheduler outbox must share one atomic SQL statement") + require.Contains(t, normalizeSQLWhitespace(exec.execQueries[0]), "INSERT INTO scheduler_outbox") + }) + } +} + func TestAccountRepository_ListOAuthRefreshCandidates_SQLFilter(t *testing.T) { db, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(sqlmock.QueryMatcherRegexp)) require.NoError(t, err) @@ -87,14 +185,20 @@ func (r rowsAffectedResult) RowsAffected() (int64, error) { return int64(r), nil type recordingSQLExecutor struct { result sql.Result err error + afterExec func() execQueries []string + execArgs [][]any } func (e *recordingSQLExecutor) ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error) { e.execQueries = append(e.execQueries, query) + e.execArgs = append(e.execArgs, append([]any(nil), args...)) if e.err != nil { return nil, e.err } + if e.afterExec != nil { + e.afterExec() + } return e.result, nil } diff --git a/backend/internal/repository/grok_oauth_client.go b/backend/internal/repository/grok_oauth_client.go index 6c9c2c407f..38f6cfb96e 100644 --- a/backend/internal/repository/grok_oauth_client.go +++ b/backend/internal/repository/grok_oauth_client.go @@ -153,14 +153,29 @@ func grokOAuthStatusError(code, message string, resp *req.Response) error { statusCode := http.StatusBadGateway errorCode := code upstreamStatus := 0 - if resp != nil && resp.StatusCode == http.StatusForbidden { - statusCode = http.StatusForbidden - errorCode = "GROK_OAUTH_ENTITLEMENT_DENIED" - } body := "" if resp != nil { upstreamStatus = resp.StatusCode body = logredact.RedactText(resp.String()) + if resp.StatusCode == http.StatusForbidden && grokOAuthHasExplicitEntitlementDenial(body) { + statusCode = http.StatusForbidden + errorCode = "GROK_OAUTH_ENTITLEMENT_DENIED" + } } return infraerrors.Newf(statusCode, errorCode, "%s: status %d, body: %s", message, upstreamStatus, body) } + +func grokOAuthHasExplicitEntitlementDenial(body string) bool { + lower := strings.ToLower(body) + compact := strings.NewReplacer(" ", "", "\n", "", "\r", "", "\t", "").Replace(lower) + for _, field := range []string{"error", "code", "reason"} { + for _, value := range []string{"access_denied", "entitlement_denied", "subscription_required", "no_active_subscription"} { + if strings.Contains(compact, `"`+field+`":"`+value+`"`) { + return true + } + } + } + return strings.Contains(lower, "entitlement denied") || + strings.Contains(lower, "subscription required") || + strings.Contains(lower, "no active grok subscription") +} diff --git a/backend/internal/repository/grok_oauth_client_test.go b/backend/internal/repository/grok_oauth_client_test.go index eabeb641c7..690a8e2139 100644 --- a/backend/internal/repository/grok_oauth_client_test.go +++ b/backend/internal/repository/grok_oauth_client_test.go @@ -50,15 +50,7 @@ func TestGrokOAuthClientExchangeAndRefreshUseFormFields(t *testing.T) { t.Setenv(xai.EnvTokenURL, server.URL) client := NewGrokOAuthClient() - - exchanged, err := client.ExchangeCode( - context.Background(), - "auth-code", - "verifier", - "http://127.0.0.1:56121/callback", - "", - "client-id", - ) + exchanged, err := client.ExchangeCode(context.Background(), "auth-code", "verifier", "http://127.0.0.1:56121/callback", "", "client-id") require.NoError(t, err) require.Equal(t, "exchange-access", exchanged.AccessToken) require.Equal(t, "exchange-refresh", exchanged.RefreshToken) @@ -72,18 +64,30 @@ func TestGrokOAuthClientExchangeAndRefreshUseFormFields(t *testing.T) { require.Equal(t, int64(7200), refreshed.ExpiresIn) } -func TestGrokOAuthClientRefreshForbiddenClassifiesEntitlement(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.WriteHeader(http.StatusForbidden) - _, _ = w.Write([]byte(`{"error":"subscription required"}`)) - })) - defer server.Close() - t.Setenv(xai.EnvTokenURL, server.URL) +func TestGrokOAuthClientRefreshForbiddenClassifiesOnlyExplicitEntitlement(t *testing.T) { + tests := []struct { + name string + body string + wantReason string + }{ + {name: "explicit entitlement", body: `{"error":"access_denied"}`, wantReason: "GROK_OAUTH_ENTITLEMENT_DENIED"}, + {name: "generic forbidden", body: `{"error":"forbidden"}`, wantReason: "GROK_OAUTH_TOKEN_REFRESH_FAILED"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusForbidden) + _, _ = w.Write([]byte(tt.body)) + })) + defer server.Close() + t.Setenv(xai.EnvTokenURL, server.URL) - client := NewGrokOAuthClient() - _, err := client.RefreshToken(context.Background(), "refresh-token", "", "client-id") - require.Error(t, err) - require.Contains(t, strings.ToUpper(err.Error()), "GROK_OAUTH_ENTITLEMENT_DENIED") + client := NewGrokOAuthClient() + _, err := client.RefreshToken(context.Background(), "refresh-token", "", "client-id") + require.Error(t, err) + require.Contains(t, strings.ToUpper(err.Error()), tt.wantReason) + }) + } } func TestGrokOAuthClientStatusErrorRedactsSensitiveResponseBody(t *testing.T) { @@ -105,3 +109,13 @@ func TestGrokOAuthClientStatusErrorRedactsSensitiveResponseBody(t *testing.T) { require.NotContains(t, errText, "refresh-secret") require.NotContains(t, errText, "verifier-secret") } + +func TestGrokOAuthEntitlementDenialRequiresExplicitEvidence(t *testing.T) { + t.Parallel() + + require.True(t, grokOAuthHasExplicitEntitlementDenial(`{"error":"access_denied"}`)) + require.True(t, grokOAuthHasExplicitEntitlementDenial(`{"code":"entitlement_denied"}`)) + require.True(t, grokOAuthHasExplicitEntitlementDenial(`{"message":"no active Grok subscription"}`)) + require.False(t, grokOAuthHasExplicitEntitlementDenial(`{"error":"forbidden","message":"request forbidden"}`)) + require.False(t, grokOAuthHasExplicitEntitlementDenial(`403 Forbidden`)) +} diff --git a/backend/internal/repository/ops_error_where_test.go b/backend/internal/repository/ops_error_where_test.go index c997865ba4..c52461e8ff 100644 --- a/backend/internal/repository/ops_error_where_test.go +++ b/backend/internal/repository/ops_error_where_test.go @@ -101,6 +101,40 @@ func TestBuildOpsErrorLogsWhere_CyberPolicyStatusExemption(t *testing.T) { if strings.Contains(whereRecovered, "status_code") { t.Fatalf("upstream phase with IncludeRecoveredUpstream must not add any status_code clause\nfull: %s", whereRecovered) } + + // account_auth uses the same explicit provider-health opt-in but remains a + // distinct phase from inference upstream errors. + whereAccountAuth, _ := buildOpsErrorLogsWhere(&service.OpsErrorLogFilter{Phase: "account_auth", IncludeRecoveredUpstream: true}) + if strings.Contains(whereAccountAuth, "status_code") { + t.Fatalf("account_auth phase with IncludeRecoveredUpstream must expose recovered rows\nfull: %s", whereAccountAuth) + } + if !strings.Contains(whereAccountAuth, "e.error_phase = $") { + t.Fatalf("account_auth recovered filter must retain its explicit phase\nfull: %s", whereAccountAuth) + } + + whereProviderHealth, _ := buildOpsErrorLogsWhere(&service.OpsErrorLogFilter{ + ErrorPhasesAny: []string{"upstream", "account_auth"}, + IncludeRecoveredUpstream: true, + }) + if strings.Contains(whereProviderHealth, "status_code") { + t.Fatalf("provider-health ANY filter must expose recovered inference and credential rows\nfull: %s", whereProviderHealth) + } + if !strings.Contains(whereProviderHealth, "e.error_phase = ANY($") { + t.Fatalf("provider-health filter must preserve distinct phase values\nfull: %s", whereProviderHealth) + } + + whereUserAccountAuth, _ := buildOpsErrorLogsWhere(&service.OpsErrorLogFilter{ErrorPhasesAny: []string{"account_auth"}}) + if !strings.Contains(whereUserAccountAuth, "COALESCE(e.status_code, 0) >= 400") { + t.Fatalf("request-error account_auth filters must exclude recovered successes\nfull: %s", whereUserAccountAuth) + } + + whereMixed, _ := buildOpsErrorLogsWhere(&service.OpsErrorLogFilter{ + ErrorPhasesAny: []string{"account_auth", "request"}, + IncludeRecoveredUpstream: true, + }) + if !strings.Contains(whereMixed, "COALESCE(e.status_code, 0) >= 400") { + t.Fatalf("recovered opt-in must not bypass the guard for non-provider phases\nfull: %s", whereMixed) + } } func TestBuildOpsErrorLogsWhere_MatchDeletedKeyOwner(t *testing.T) { diff --git a/backend/internal/repository/ops_repo.go b/backend/internal/repository/ops_repo.go index 900abcf212..635bcd1a1f 100644 --- a/backend/internal/repository/ops_repo.go +++ b/backend/internal/repository/ops_repo.go @@ -160,7 +160,7 @@ func opsInsertErrorLogArgs(input *service.OpsInsertErrorLogInput) []any { opsNullString(input.ErrorBody), opsNullString(input.ErrorSource), opsNullString(input.ErrorOwner), - opsNullInt(input.UpstreamStatusCode), + opsNullableIntPointer(input.UpstreamStatusCode), opsNullString(input.UpstreamErrorMessage), opsNullString(input.UpstreamErrorDetail), opsNullString(input.UpstreamErrorsJSON), @@ -582,7 +582,7 @@ LIMIT 1` s := clientIP.String out.ClientIP = &s } - if upstreamStatusCode.Valid && upstreamStatusCode.Int64 > 0 { + if upstreamStatusCode.Valid { v := int(upstreamStatusCode.Int64) out.UpstreamStatusCode = &v } @@ -978,13 +978,13 @@ func buildOpsErrorLogsWhere(filter *service.OpsErrorLogFilter) (string, []any) { resolvedFilter = filter.Resolved } // Keep list endpoints scoped to client errors unless the caller explicitly opts - // into recovered upstream rows (Phase=="upstream" + IncludeRecoveredUpstream, - // ops 专用上游列表)。请求错误语义的端点即便过滤 phase=upstream 也保留该守卫。 + // into recovered provider-health rows (upstream/account_auth). Request-error + // endpoints never set the opt-in and retain this guard. // 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" || filter == nil || !filter.IncludeRecoveredUpstream { + if !opsFilterIncludesRecoveredProviderRows(filter, phaseFilter) { clauses = append(clauses, "(COALESCE(e.status_code, 0) >= 400 OR e.error_type = 'cyber_policy')") } @@ -1117,6 +1117,28 @@ func buildOpsErrorLogsWhere(filter *service.OpsErrorLogFilter) (string, []any) { return "WHERE " + strings.Join(clauses, " AND "), args } +func opsFilterIncludesRecoveredProviderRows(filter *service.OpsErrorLogFilter, phaseFilter string) bool { + if filter == nil || !filter.IncludeRecoveredUpstream { + return false + } + if phaseFilter != "" { + return phaseFilter == "upstream" || phaseFilter == "account_auth" + } + if len(filter.ErrorPhasesAny) == 0 { + return false + } + sawProviderPhase := false + for _, rawPhase := range filter.ErrorPhasesAny { + switch strings.TrimSpace(strings.ToLower(rawPhase)) { + case "upstream", "account_auth": + sawProviderPhase = true + default: + return false + } + } + return sawProviderPhase +} + func buildOpsSystemLogsWhere(filter *service.OpsSystemLogFilter) (string, []any, bool) { clauses := make([]string, 0, 10) args := make([]any, 0, 10) @@ -1269,6 +1291,16 @@ func opsNullInt(v any) any { } } +// opsNullableIntPointer distinguishes an absent value from an explicitly +// observed zero. Credential-stage failures intentionally persist upstream +// status 0 because no inference request was sent. +func opsNullableIntPointer(v *int) any { + if v == nil { + return sql.NullInt64{} + } + return sql.NullInt64{Int64: int64(*v), Valid: true} +} + func opsNullInt16(v *int16) any { if v == nil { return sql.NullInt64{} diff --git a/backend/internal/repository/ops_repo_args_test.go b/backend/internal/repository/ops_repo_args_test.go new file mode 100644 index 0000000000..5a88758f0a --- /dev/null +++ b/backend/internal/repository/ops_repo_args_test.go @@ -0,0 +1,37 @@ +//go:build unit + +package repository + +import ( + "database/sql" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" +) + +func TestOpsInsertErrorLogArgsPreservesExplicitZeroUpstreamStatus(t *testing.T) { + zero := 0 + args := opsInsertErrorLogArgs(&service.OpsInsertErrorLogInput{UpstreamStatusCode: &zero}) + + require.Len(t, args, 41) + encoded, ok := args[27].(sql.NullInt64) + require.True(t, ok) + require.True(t, encoded.Valid) + require.Zero(t, encoded.Int64) +} + +func TestOpsNullableIntPointerDistinguishesNilZeroAndStatus(t *testing.T) { + missing := opsNullableIntPointer(nil).(sql.NullInt64) + require.False(t, missing.Valid) + + zeroValue := 0 + zero := opsNullableIntPointer(&zeroValue).(sql.NullInt64) + require.True(t, zero.Valid) + require.Zero(t, zero.Int64) + + statusValue := 503 + status := opsNullableIntPointer(&statusValue).(sql.NullInt64) + require.True(t, status.Valid) + require.EqualValues(t, 503, status.Int64) +} diff --git a/backend/internal/repository/ops_repo_get_error_log_by_id_integration_test.go b/backend/internal/repository/ops_repo_get_error_log_by_id_integration_test.go index 470b1c0d0b..1411002be0 100644 --- a/backend/internal/repository/ops_repo_get_error_log_by_id_integration_test.go +++ b/backend/internal/repository/ops_repo_get_error_log_by_id_integration_test.go @@ -91,4 +91,21 @@ func TestGetErrorLogByID_DeletedKeyOwner(t *testing.T) { require.Equal(t, "sk-valid", valid.APIKeyPrefix) require.Empty(t, valid.AttemptedKeyPrefix, "attempted prefix and api key prefix are mutually exclusive") require.Nil(t, valid.DeletedKeyOwnerUserID, "valid key error has no deleted owner") + + // ── Case 4: account_auth with no inference attempt preserves explicit 0 ── + zero := 0 + credentialFailureID, err := repo.InsertErrorLog(ctx, &service.OpsInsertErrorLogInput{ + ErrorPhase: "account_auth", + ErrorType: "upstream_error", + Severity: "error", + StatusCode: 503, + UpstreamStatusCode: &zero, + CreatedAt: time.Now(), + }) + require.NoError(t, err) + + credentialFailure, err := repo.GetErrorLogByID(ctx, credentialFailureID) + require.NoError(t, err) + require.NotNil(t, credentialFailure.UpstreamStatusCode) + require.Zero(t, *credentialFailure.UpstreamStatusCode) } diff --git a/backend/internal/service/account.go b/backend/internal/service/account.go index 3e67fa5982..5214518fb5 100644 --- a/backend/internal/service/account.go +++ b/backend/internal/service/account.go @@ -6,7 +6,6 @@ import ( "errors" "hash/fnv" "log/slog" - "net/url" "reflect" "sort" "strconv" @@ -1268,16 +1267,13 @@ func (a *Account) GetGrokBaseURL() string { if !a.IsGrok() { return "" } - baseURL := a.GetCredential("base_url") if a.IsGrokOAuth() { - if strings.TrimSpace(baseURL) == "" || isOfficialGrokAPIBaseURL(baseURL) { - return xai.DefaultCLIBaseURL - } - if _, err := xai.ValidateTrustedBaseURL(baseURL); err == nil { - return baseURL - } + // OAuth bearer credentials are subscription credentials and may only be + // sent to the supported CLI gateway. Stored base_url values and unsafe + // development overrides apply exclusively to API-key accounts. return xai.DefaultCLIBaseURL } + baseURL := a.GetCredential("base_url") if baseURL != "" { return baseURL } @@ -1286,57 +1282,17 @@ func (a *Account) GetGrokBaseURL() string { // GetGrokMediaBaseURL selects the upstream used by Grok Imagine APIs. // -// OAuth text requests need the CLI subscription proxy, but that proxy has a -// smaller request-body limit than the official Imagine API. Media requests can -// contain large base64 inputs, so default OAuth accounts must use api.x.ai. -// API-key accounts and explicit unsafe development overrides retain their -// configured base URL. +// OAuth media credentials have the same trust boundary as OAuth text traffic: +// they are pinned to the supported CLI gateway even for large request bodies. +// API-key accounts retain their configured public/custom upstream behavior. func (a *Account) GetGrokMediaBaseURL() string { if !a.IsGrok() { return "" } - if !a.IsGrokOAuth() { - return a.GetGrokBaseURL() + if a.IsGrokOAuth() { + return xai.DefaultCLIBaseURL } - - baseURL := a.GetCredential("base_url") - if strings.TrimSpace(baseURL) == "" || isOfficialGrokAPIBaseURL(baseURL) || isOfficialGrokCLIBaseURL(baseURL) { - return xai.DefaultBaseURL - } - if _, err := xai.ValidateTrustedBaseURL(baseURL); err == nil { - return baseURL - } - return xai.DefaultBaseURL -} - -func isOfficialGrokAPIBaseURL(raw string) bool { - return isOfficialGrokBaseURL(raw, xai.DefaultBaseURL) -} - -func isOfficialGrokCLIBaseURL(raw string) bool { - return isOfficialGrokBaseURL(raw, xai.DefaultCLIBaseURL) -} - -func isOfficialGrokBaseURL(raw, expected string) bool { - parsed, err := url.Parse(strings.TrimSpace(raw)) - if err != nil || parsed == nil || parsed.Opaque != "" || parsed.User != nil || parsed.RawQuery != "" || parsed.Fragment != "" { - return false - } - defaultURL, err := url.Parse(expected) - if err != nil { - return false - } - if !strings.EqualFold(parsed.Scheme, defaultURL.Scheme) || !strings.EqualFold(parsed.Hostname(), defaultURL.Hostname()) { - return false - } - if port := parsed.Port(); port != "" { - portNumber, err := strconv.Atoi(port) - if err != nil || portNumber != 443 { - return false - } - } - path := strings.TrimRight(parsed.Path, "/") - return path == "" || path == strings.TrimRight(defaultURL.Path, "/") + return a.GetGrokBaseURL() } func (a *Account) GetGrokAccessToken() string { diff --git a/backend/internal/service/account_base_url_test.go b/backend/internal/service/account_base_url_test.go index 59f53db3db..701a37e131 100644 --- a/backend/internal/service/account_base_url_test.go +++ b/backend/internal/service/account_base_url_test.go @@ -255,7 +255,7 @@ func TestGetGrokBaseURLUsesSubscriptionProxyForOAuth(t *testing.T) { expected: xai.DefaultCLIBaseURL, }, { - name: "oauth non-default API port remains an explicit override", + name: "oauth non-default API port remains pinned to CLI proxy", account: Account{ Type: AccountTypeOAuth, Platform: PlatformGrok, @@ -263,7 +263,7 @@ func TestGetGrokBaseURLUsesSubscriptionProxyForOAuth(t *testing.T) { "base_url": "https://api.x.ai:8443/v1", }, }, - expected: "https://api.x.ai:8443/v1", + expected: xai.DefaultCLIBaseURL, }, { name: "oauth explicit custom base_url stays pinned to CLI proxy by default", @@ -294,7 +294,7 @@ func TestGetGrokBaseURLUsesSubscriptionProxyForOAuth(t *testing.T) { } } -func TestGetGrokBaseURLAllowsExplicitOAuthOverrideWhenUnsafeOverridesEnabled(t *testing.T) { +func TestGetGrokBaseURLPinsOAuthWhenUnsafeOverridesEnabled(t *testing.T) { t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true") account := Account{ Type: AccountTypeOAuth, @@ -304,26 +304,26 @@ func TestGetGrokBaseURLAllowsExplicitOAuthOverrideWhenUnsafeOverridesEnabled(t * }, } - require.Equal(t, "https://custom.example.com/v1", account.GetGrokBaseURL()) + require.Equal(t, xai.DefaultCLIBaseURL, account.GetGrokBaseURL()) } -func TestGetGrokMediaBaseURLSeparatesOAuthMediaFromCLIProxy(t *testing.T) { +func TestGetGrokMediaBaseURLPinsOAuthMediaToCLIProxy(t *testing.T) { tests := []struct { name string account Account expected string }{ { - name: "oauth without base_url uses official media API", + name: "oauth without base_url uses CLI subscription proxy", account: Account{ Type: AccountTypeOAuth, Platform: PlatformGrok, Credentials: map[string]any{}, }, - expected: xai.DefaultBaseURL, + expected: xai.DefaultCLIBaseURL, }, { - name: "oauth stored CLI proxy uses official media API", + name: "oauth stored CLI proxy stays on CLI subscription proxy", account: Account{ Type: AccountTypeOAuth, Platform: PlatformGrok, @@ -331,10 +331,10 @@ func TestGetGrokMediaBaseURLSeparatesOAuthMediaFromCLIProxy(t *testing.T) { "base_url": xai.DefaultCLIBaseURL, }, }, - expected: xai.DefaultBaseURL, + expected: xai.DefaultCLIBaseURL, }, { - name: "oauth stored CLI proxy variant uses official media API", + name: "oauth stored CLI proxy variant is canonicalized to CLI proxy", account: Account{ Type: AccountTypeOAuth, Platform: PlatformGrok, @@ -342,10 +342,10 @@ func TestGetGrokMediaBaseURLSeparatesOAuthMediaFromCLIProxy(t *testing.T) { "base_url": "HTTPS://CLI-CHAT-PROXY.GROK.COM:443/%76%31/", }, }, - expected: xai.DefaultBaseURL, + expected: xai.DefaultCLIBaseURL, }, { - name: "oauth legacy official API remains on official media API", + name: "oauth legacy official API is pinned to CLI proxy", account: Account{ Type: AccountTypeOAuth, Platform: PlatformGrok, @@ -353,10 +353,10 @@ func TestGetGrokMediaBaseURLSeparatesOAuthMediaFromCLIProxy(t *testing.T) { "base_url": xai.DefaultBaseURL, }, }, - expected: xai.DefaultBaseURL, + expected: xai.DefaultCLIBaseURL, }, { - name: "oauth untrusted custom base_url is pinned to official media API", + name: "oauth untrusted custom base_url is pinned to CLI proxy", account: Account{ Type: AccountTypeOAuth, Platform: PlatformGrok, @@ -364,7 +364,7 @@ func TestGetGrokMediaBaseURLSeparatesOAuthMediaFromCLIProxy(t *testing.T) { "base_url": "https://custom.example.com/v1", }, }, - expected: xai.DefaultBaseURL, + expected: xai.DefaultCLIBaseURL, }, { name: "API key retains its configured media API", @@ -395,7 +395,7 @@ func TestGetGrokMediaBaseURLSeparatesOAuthMediaFromCLIProxy(t *testing.T) { } } -func TestGetGrokMediaBaseURLAllowsExplicitOAuthOverrideWhenUnsafeOverridesEnabled(t *testing.T) { +func TestGetGrokMediaBaseURLPinsOAuthWhenUnsafeOverridesEnabled(t *testing.T) { t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true") account := Account{ Type: AccountTypeOAuth, @@ -405,5 +405,5 @@ func TestGetGrokMediaBaseURLAllowsExplicitOAuthOverrideWhenUnsafeOverridesEnable }, } - require.Equal(t, "https://custom.example.com/v1", account.GetGrokMediaBaseURL()) + require.Equal(t, xai.DefaultCLIBaseURL, account.GetGrokMediaBaseURL()) } diff --git a/backend/internal/service/account_credentials_redact.go b/backend/internal/service/account_credentials_redact.go index 76c2d1de5b..8eb70513ba 100644 --- a/backend/internal/service/account_credentials_redact.go +++ b/backend/internal/service/account_credentials_redact.go @@ -4,7 +4,7 @@ package service // dto 层做响应脱敏、service 层做更新合并都引用此清单——新增凭证类型时务必同步。 var SensitiveCredentialKeys = []string{ // OAuth - "access_token", "refresh_token", "id_token", + "access_token", "refresh_token", "id_token", "agent_private_key", // API Key 类 "api_key", "session_key", "cookie", // 云服务凭据 diff --git a/backend/internal/service/account_stats_pricing.go b/backend/internal/service/account_stats_pricing.go index 221021d85c..1098fff9e5 100644 --- a/backend/internal/service/account_stats_pricing.go +++ b/backend/internal/service/account_stats_pricing.go @@ -65,6 +65,13 @@ func tryModelFilePricing(billingService *BillingService, model string, tokens Us if err != nil || pricing == nil { return nil } + if billingService.shouldApplySessionLongContextPricing(tokens, pricing) { + breakdown, err := billingService.CalculateCost(model, tokens, 1) + if err != nil || breakdown == nil || breakdown.TotalCost <= 0 { + return nil + } + return &breakdown.TotalCost + } cost := float64(tokens.InputTokens)*pricing.InputPricePerToken + float64(tokens.OutputTokens)*pricing.OutputPricePerToken + float64(tokens.CacheCreationTokens)*pricing.CacheCreationPricePerToken + diff --git a/backend/internal/service/account_stats_pricing_test.go b/backend/internal/service/account_stats_pricing_test.go index 36e5eb7400..57ace01c22 100644 --- a/backend/internal/service/account_stats_pricing_test.go +++ b/backend/internal/service/account_stats_pricing_test.go @@ -459,6 +459,26 @@ func TestTryModelFilePricing_Success(t *testing.T) { require.InDelta(t, 0.2, *result, 1e-12) } +func TestTryModelFilePricing_AppliesLongContextPricing(t *testing.T) { + bs := newTestBillingServiceWithPrices(map[string]*ModelPricing{ + "gpt-5.6-sol": { + InputPricePerToken: 0.001, + OutputPricePerToken: 0.002, + CacheReadPricePerToken: 0.0001, + LongContextInputThreshold: 100, + LongContextInputMultiplier: 2, + LongContextOutputMultiplier: 1.5, + }, + }) + tokens := UsageTokens{InputTokens: 101, OutputTokens: 10, CacheReadTokens: 5} + + result := tryModelFilePricing(bs, "gpt-5.6-sol", tokens) + + require.NotNil(t, result) + // Input and cache-read use the 2x input tier; output uses the 1.5x tier. + require.InDelta(t, 0.233, *result, 1e-12) +} + func TestTryModelFilePricing_PricingNotFound(t *testing.T) { // "nonexistent-model" does not match any fallback pattern bs := newTestBillingServiceWithPrices(map[string]*ModelPricing{}) diff --git a/backend/internal/service/account_test_service.go b/backend/internal/service/account_test_service.go index 222b3f8a4d..f34cc909b3 100644 --- a/backend/internal/service/account_test_service.go +++ b/backend/internal/service/account_test_service.go @@ -15,6 +15,7 @@ import ( "net/http/httptest" "regexp" "strings" + "sync" "time" "github.com/Wei-Shaw/sub2api/internal/config" @@ -72,6 +73,8 @@ type AccountTestService struct { httpUpstream HTTPUpstream cfg *config.Config tlsFPProfileService *TLSFingerprintProfileService + agentIdentityTaskMu sync.Mutex + agentIdentityWS agentIdentityWSConnectionInvalidator } // NewAccountTestService creates a new AccountTestService @@ -544,9 +547,11 @@ func (s *AccountTestService) testOpenAIAccountConnection(c *gin.Context, account if credentialAccount.IsOAuth() { isOAuth = true - // OAuth - use Bearer token with ChatGPT internal API - authToken = credentialAccount.GetOpenAIAccessToken() - if authToken == "" { + // Agent Identity signs each request and does not retain the OAuth token. + if !credentialAccount.IsOpenAIAgentIdentity() { + authToken = credentialAccount.GetOpenAIAccessToken() + } + if authToken == "" && !credentialAccount.IsOpenAIAgentIdentity() { return s.sendErrorAndEnd(c, "No access token available") } @@ -591,8 +596,11 @@ func (s *AccountTestService) testOpenAIAccountConnection(c *gin.Context, account payload := createOpenAITestPayload(upstreamTestModelID, isOAuth) payloadBytes, _ := json.Marshal(payload) - // Send test_start event - s.sendEvent(c, TestEvent{Type: "test_start", Model: testModelID}) + // Send test_start event once. A task-invalid Agent Identity response may + // restart this probe after registering a replacement task. + if !agentIdentityTaskRecoveryWasTried(ctx) { + s.sendEvent(c, TestEvent{Type: "test_start", Model: testModelID}) + } req, err := http.NewRequestWithContext(ctx, "POST", apiURL, bytes.NewReader(payloadBytes)) if err != nil { @@ -602,7 +610,19 @@ func (s *AccountTestService) testOpenAIAccountConnection(c *gin.Context, account // Set common headers req.Header.Set("Content-Type", "application/json") - req.Header.Set("Authorization", "Bearer "+authToken) + if credentialAccount.IsOpenAIAgentIdentity() { + authHeaders, authErr := buildAgentIdentityAuthenticationHeaders(ctx, s.accountRepo, s.agentIdentityWS, &s.agentIdentityTaskMu, credentialAccount) + if authErr != nil { + return s.sendErrorAndEnd(c, "Failed to build Agent Identity authentication") + } + for key, values := range authHeaders { + for _, value := range values { + req.Header.Add(key, value) + } + } + } else { + req.Header.Set("Authorization", "Bearer "+authToken) + } // Set OAuth-specific headers for ChatGPT internal API if isOAuth { @@ -644,6 +664,15 @@ func (s *AccountTestService) testOpenAIAccountConnection(c *gin.Context, account if resp.StatusCode != http.StatusOK { body, _ := io.ReadAll(resp.Body) + body = redactAgentIdentitySensitiveBodyForAccount(ctx, s.accountRepo, credentialAccount, body) + if !agentIdentityTaskRecoveryWasTried(ctx) && credentialAccount.IsOpenAIAgentIdentity() && isAgentIdentityTaskInvalidHTTPResponse(resp.StatusCode, body) { + expectedTaskID := credentialAccount.GetCredential("task_id") + if err := ensureAgentIdentityTaskForAccount(ctx, s.accountRepo, s.agentIdentityWS, &s.agentIdentityTaskMu, credentialAccount, expectedTaskID); err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Agent Identity task recovery failed: %s", err.Error())) + } + c.Request = c.Request.WithContext(markAgentIdentityTaskRecoveryTried(ctx)) + return s.testOpenAIAccountConnection(c, account, modelID, prompt, mode) + } if resp.StatusCode == http.StatusTooManyRequests { s.reconcileOpenAI429State(ctx, account, resp.Header, body) } @@ -695,7 +724,7 @@ func (s *AccountTestService) testGrokAccountConnection(c *gin.Context, account * return s.sendErrorAndEnd(c, fmt.Sprintf("Unsupported Grok account type: %s", account.Type)) } - apiURL, err := xai.BuildResponsesURL(account.GetGrokBaseURL()) + apiURL, err := buildGrokResponsesURL(account, s.cfg) if err != nil { return s.sendErrorAndEnd(c, fmt.Sprintf("Invalid Grok base URL: %s", err.Error())) } @@ -715,7 +744,9 @@ func (s *AccountTestService) testGrokAccountConnection(c *gin.Context, account * return s.sendErrorAndEnd(c, "Failed to create Grok test payload") } - s.sendEvent(c, TestEvent{Type: "test_start", Model: testModelID}) + if !agentIdentityTaskRecoveryWasTried(ctx) { + s.sendEvent(c, TestEvent{Type: "test_start", Model: testModelID}) + } req, err := http.NewRequestWithContext(ctx, http.MethodPost, apiURL, bytes.NewReader(payloadBytes)) if err != nil { @@ -724,7 +755,9 @@ func (s *AccountTestService) testGrokAccountConnection(c *gin.Context, account * req.Header.Set("Content-Type", "application/json") req.Header.Set("Accept", "application/json, text/event-stream") req.Header.Set("Authorization", "Bearer "+authToken) - applyGrokCLIHeaders(req.Header) + if account.IsGrokOAuth() { + applyGrokCLIHeaders(req.Header) + } proxyURL := "" if account.ProxyID != nil && account.Proxy != nil { @@ -740,7 +773,7 @@ func (s *AccountTestService) testGrokAccountConnection(c *gin.Context, account * now := time.Now() snapshot := parseGrokQuotaSnapshot(resp.Header, resp.StatusCode, now) if snapshot != nil && s.accountRepo != nil { - resetAt, limited := grokRateLimitResetAt(snapshot, now) + resetAt, limited := grokRateLimitResetAtForAccount(account, snapshot, now) if limited { normalizeGrokExhaustedWindowResets(snapshot, resetAt, now) } @@ -749,7 +782,11 @@ func (s *AccountTestService) testGrokAccountConnection(c *gin.Context, account * }) if limited { persistGrokRateLimit(ctx, s.accountRepo, account, resetAt) + } else if isSuccessfulGrokRateLimitRecovery(account, snapshot) { + clearGrokRateLimitAfterRecovery(ctx, s.accountRepo, account) } + } else if s.accountRepo != nil && isSuccessfulGrokRateLimitRecovery(account, &xai.QuotaSnapshot{StatusCode: resp.StatusCode}) { + clearGrokRateLimitAfterRecovery(ctx, s.accountRepo, account) } if resp.StatusCode != http.StatusOK { @@ -827,16 +864,26 @@ func (s *AccountTestService) testOpenAIChatCompletionsConnection( // resulting capability state on the account. func (s *AccountTestService) testOpenAICompactConnection(c *gin.Context, account *Account, testModelID string) error { ctx := c.Request.Context() + credentialAccount := account + if account.IsShadow() { + resolved, err := resolveCredentialAccount(ctx, s.accountRepo, account) + if err != nil { + return s.sendErrorAndEnd(c, "Failed to resolve account credentials") + } + credentialAccount = resolved + } authToken := "" apiURL := "" isOAuth := false switch { - case account.IsOAuth(): + case credentialAccount.IsOAuth(): isOAuth = true - authToken = account.GetOpenAIAccessToken() - if authToken == "" { + if !credentialAccount.IsOpenAIAgentIdentity() { + authToken = credentialAccount.GetOpenAIAccessToken() + } + if authToken == "" && !credentialAccount.IsOpenAIAgentIdentity() { return s.sendErrorAndEnd(c, "No access token available") } apiURL = chatgptCodexAPIURL + "/compact" @@ -865,7 +912,9 @@ func (s *AccountTestService) testOpenAICompactConnection(c *gin.Context, account c.Writer.Flush() payloadBytes, _ := json.Marshal(createOpenAICompactProbePayload(testModelID)) - s.sendEvent(c, TestEvent{Type: "test_start", Model: testModelID}) + if !agentIdentityTaskRecoveryWasTried(ctx) { + s.sendEvent(c, TestEvent{Type: "test_start", Model: testModelID}) + } req, err := http.NewRequestWithContext(ctx, "POST", apiURL, bytes.NewReader(payloadBytes)) if err != nil { @@ -875,7 +924,19 @@ func (s *AccountTestService) testOpenAICompactConnection(c *gin.Context, account req.Header.Set("Content-Type", "application/json") req.Header.Set("Accept", "application/json") - req.Header.Set("Authorization", "Bearer "+authToken) + if credentialAccount.IsOpenAIAgentIdentity() { + authHeaders, authErr := buildAgentIdentityAuthenticationHeaders(ctx, s.accountRepo, s.agentIdentityWS, &s.agentIdentityTaskMu, credentialAccount) + if authErr != nil { + return s.sendErrorAndEnd(c, "Failed to build Agent Identity authentication") + } + for key, values := range authHeaders { + for _, value := range values { + req.Header.Add(key, value) + } + } + } else { + req.Header.Set("Authorization", "Bearer "+authToken) + } req.Header.Set("OpenAI-Beta", "responses=experimental") req.Header.Set("Originator", "codex_cli_rs") req.Header.Set("User-Agent", codexCLIUserAgent) @@ -886,7 +947,7 @@ func (s *AccountTestService) testOpenAICompactConnection(c *gin.Context, account if isOAuth { req.Host = "chatgpt.com" - setOpenAIChatGPTAccountHeaders(req.Header, account) + setOpenAIChatGPTAccountHeaders(req.Header, credentialAccount) } // 账号级请求头覆写:测试请求与真实转发保持一致的最终头 @@ -909,6 +970,15 @@ func (s *AccountTestService) testOpenAICompactConnection(c *gin.Context, account defer func() { _ = resp.Body.Close() }() body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20)) + body = redactAgentIdentitySensitiveBodyForAccount(ctx, s.accountRepo, credentialAccount, body) + if !agentIdentityTaskRecoveryWasTried(ctx) && credentialAccount.IsOpenAIAgentIdentity() && isAgentIdentityTaskInvalidHTTPResponse(resp.StatusCode, body) { + expectedTaskID := credentialAccount.GetCredential("task_id") + if err := ensureAgentIdentityTaskForAccount(ctx, s.accountRepo, s.agentIdentityWS, &s.agentIdentityTaskMu, credentialAccount, expectedTaskID); err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Agent Identity task recovery failed: %s", err.Error())) + } + c.Request = c.Request.WithContext(markAgentIdentityTaskRecoveryTried(ctx)) + return s.testOpenAICompactConnection(c, account, testModelID) + } if s.accountRepo != nil { updates := buildOpenAICompactProbeExtraUpdates(resp, body, nil, time.Now()) @@ -1700,8 +1770,19 @@ func (s *AccountTestService) testOpenAIImageAPIKey(c *gin.Context, ctx context.C // testOpenAIImageOAuth tests OpenAI image generation using an OAuth account via Codex /responses API. func (s *AccountTestService) testOpenAIImageOAuth(c *gin.Context, ctx context.Context, account *Account, modelID, prompt string) error { - authToken := account.GetOpenAIAccessToken() - if authToken == "" { + credentialAccount := account + if account.IsShadow() { + resolved, err := resolveCredentialAccount(ctx, s.accountRepo, account) + if err != nil { + return s.sendErrorAndEnd(c, "Failed to resolve account credentials") + } + credentialAccount = resolved + } + authToken := "" + if !credentialAccount.IsOpenAIAgentIdentity() { + authToken = credentialAccount.GetOpenAIAccessToken() + } + if authToken == "" && !credentialAccount.IsOpenAIAgentIdentity() { return s.sendErrorAndEnd(c, "No access token available") } @@ -1733,17 +1814,29 @@ func (s *AccountTestService) testOpenAIImageOAuth(c *gin.Context, ctx context.Co } req = req.WithContext(WithHTTPUpstreamProfile(req.Context(), HTTPUpstreamProfileOpenAI)) req.Host = "chatgpt.com" - req.Header.Set("Authorization", "Bearer "+authToken) + if credentialAccount.IsOpenAIAgentIdentity() { + authHeaders, authErr := buildAgentIdentityAuthenticationHeaders(ctx, s.accountRepo, s.agentIdentityWS, &s.agentIdentityTaskMu, credentialAccount) + if authErr != nil { + return s.sendErrorAndEnd(c, "Failed to build Agent Identity authentication") + } + for key, values := range authHeaders { + for _, value := range values { + req.Header.Add(key, value) + } + } + } else { + req.Header.Set("Authorization", "Bearer "+authToken) + } req.Header.Set("Content-Type", "application/json") 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(account.GetOpenAIUserAgent()); customUA != "" { + if customUA := strings.TrimSpace(credentialAccount.GetOpenAIUserAgent()); customUA != "" { req.Header.Set("User-Agent", customUA) } else { req.Header.Set("User-Agent", codexCLIUserAgent) } - setOpenAIChatGPTAccountHeaders(req.Header, account) + setOpenAIChatGPTAccountHeaders(req.Header, credentialAccount) // 与真实转发一致:originator 与最终 User-Agent 首段配套(原 opencode 与 Codex UA 错配会 404,issue #3901)。 enforceCodexIdentityHeaders(req.Header) @@ -1762,6 +1855,7 @@ func (s *AccountTestService) testOpenAIImageOAuth(c *gin.Context, ctx context.Co }() if resp.StatusCode >= 400 { body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20)) + body = redactAgentIdentitySensitiveBodyForAccount(ctx, s.accountRepo, credentialAccount, body) message := strings.TrimSpace(extractUpstreamErrorMessage(body)) if message == "" { message = fmt.Sprintf("Responses API returned %d", resp.StatusCode) @@ -1773,6 +1867,7 @@ func (s *AccountTestService) testOpenAIImageOAuth(c *gin.Context, ctx context.Co if err != nil { return s.sendErrorAndEnd(c, fmt.Sprintf("Failed to read image response: %s", err.Error())) } + body = redactAgentIdentitySensitiveBodyForAccount(ctx, s.accountRepo, credentialAccount, body) results, _, _, _, _, err := collectOpenAIImagesFromResponsesBody(body) if err != nil { diff --git a/backend/internal/service/account_test_service_grok_test.go b/backend/internal/service/account_test_service_grok_test.go index 4b0890ff44..72f995a0c7 100644 --- a/backend/internal/service/account_test_service_grok_test.go +++ b/backend/internal/service/account_test_service_grok_test.go @@ -40,8 +40,9 @@ func TestAccountTestService_TestAccountConnection_GrokUsesXAIResponses(t *testin Schedulable: true, Concurrency: 1, Credentials: map[string]any{ - "access_token": "grok-access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + "access_token": "grok-access-token", + "refresh_token": "grok-refresh-token", + "expires_at": time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339), "model_mapping": map[string]any{ "grok": "grok-4.3", }, @@ -92,8 +93,9 @@ func TestAccountTestService_TestAccountConnection_GrokDefaultsEmptyModelTo45(t * Schedulable: true, Concurrency: 1, Credentials: map[string]any{ - "access_token": "grok-access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + "access_token": "grok-access-token", + "refresh_token": "grok-refresh-token", + "expires_at": time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339), }, } repo := &mockAccountRepoForGemini{accountsByID: map[int64]*Account{account.ID: account}} @@ -133,8 +135,9 @@ func TestAccountTestService_Grok429PersistsRateLimitReset(t *testing.T) { Schedulable: true, Concurrency: 1, Credentials: map[string]any{ - "access_token": "grok-access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + "access_token": "grok-access-token", + "refresh_token": "grok-refresh-token", + "expires_at": time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339), }, } baseRepo := &mockAccountRepoForGemini{accountsByID: map[int64]*Account{account.ID: account}} @@ -166,8 +169,9 @@ func TestAccountTestService_Grok429WithoutQuotaHeadersUsesFallback(t *testing.T) ID: 15, Name: "grok-oauth-limited-no-headers", Platform: PlatformGrok, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Credentials: map[string]any{ - "access_token": "grok-access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + "access_token": "grok-access-token", + "refresh_token": "grok-refresh-token", + "expires_at": time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339), }, } baseRepo := &mockAccountRepoForGemini{accountsByID: map[int64]*Account{account.ID: account}} diff --git a/backend/internal/service/account_usage_service.go b/backend/internal/service/account_usage_service.go index 7b42a1d1c5..12dc25a673 100644 --- a/backend/internal/service/account_usage_service.go +++ b/backend/internal/service/account_usage_service.go @@ -299,6 +299,8 @@ type AccountUsageService struct { cache *UsageCache identityCache IdentityCache tlsFPProfileService *TLSFingerprintProfileService + agentIdentityTaskMu sync.Mutex + agentIdentityWS agentIdentityWSConnectionInvalidator } // NewAccountUsageService 创建AccountUsageService实例 @@ -692,8 +694,11 @@ func (s *AccountUsageService) probeOpenAICodexSnapshot(ctx context.Context, acco if account == nil || !account.IsOAuth() { return nil, nil } - accessToken := account.GetOpenAIAccessToken() - if accessToken == "" { + accessToken := "" + if !account.IsOpenAIAgentIdentity() { + accessToken = account.GetOpenAIAccessToken() + } + if accessToken == "" && !account.IsOpenAIAgentIdentity() { return nil, fmt.Errorf("no access token available") } modelID := openaipkg.DefaultTestModel @@ -711,7 +716,19 @@ func (s *AccountUsageService) probeOpenAICodexSnapshot(ctx context.Context, acco } req.Host = "chatgpt.com" req.Header.Set("Content-Type", "application/json") - req.Header.Set("Authorization", "Bearer "+accessToken) + if account.IsOpenAIAgentIdentity() { + authHeaders, authErr := buildAgentIdentityAuthenticationHeaders(ctx, s.accountRepo, s.agentIdentityWS, &s.agentIdentityTaskMu, account) + if authErr != nil { + return nil, fmt.Errorf("build Agent Identity authentication: %w", authErr) + } + for key, values := range authHeaders { + for _, value := range values { + req.Header.Add(key, value) + } + } + } else { + req.Header.Set("Authorization", "Bearer "+accessToken) + } req.Header.Set("Accept", "text/event-stream") req.Header.Set("OpenAI-Beta", "responses=experimental") req.Header.Set("Originator", "codex_cli_rs") diff --git a/backend/internal/service/content_moderation.go b/backend/internal/service/content_moderation.go index f633c8ad17..8102f34a90 100644 --- a/backend/internal/service/content_moderation.go +++ b/backend/internal/service/content_moderation.go @@ -92,6 +92,9 @@ const ( contentModerationCleanupInterval = 24 * time.Hour contentModerationCleanupTimeout = 30 * time.Minute contentModerationCleanupDelay = 5 * time.Minute + + contentModerationRuntimeCacheTTL = time.Second + contentModerationRuntimeRefreshTimeout = 5 * time.Second ) var contentModerationCategoryOrder = []string{ @@ -512,10 +515,22 @@ type ContentModerationService struct { lastCleanupUnix atomic.Int64 lastCleanupDeletedHit atomic.Int64 lastCleanupDeletedNonHit atomic.Int64 + runtimeSnapshot atomic.Pointer[contentModerationRuntimeSnapshot] + runtimeRefreshMu sync.Mutex + runtimeCacheTTL time.Duration + runtimeRefreshRetryAt atomic.Int64 keyHealthMu sync.Mutex keyHealth map[string]*contentModerationKeyHealth } +type contentModerationRuntimeSnapshot struct { + riskControlEnabled bool + config *ContentModerationConfig + keywordMatcher *contentModerationKeywordMatcher + configDigest [sha256.Size]byte + loadedAt time.Time +} + type contentModerationTask struct { input ContentModerationCheckInput content ContentModerationInput @@ -700,6 +715,7 @@ func (s *ContentModerationService) UpdateConfig(ctx context.Context, input Updat if err := s.settingRepo.Set(ctx, SettingKeyContentModerationConfig, string(raw)); err != nil { return nil, fmt.Errorf("save content moderation config: %w", err) } + s.replaceRuntimeConfig(cfg, raw) return s.configView(cfg), nil } @@ -776,16 +792,7 @@ func (s *ContentModerationService) Check(ctx context.Context, input ContentModer "protocol", input.Protocol) return allow, nil } - if !s.isRiskControlEnabled(ctx) { - slog.Info("content_moderation.skip_feature_disabled", - "user_id", input.UserID, - "api_key_id", input.APIKeyID, - "group_id", contentModerationLogGroupID(input.GroupID), - "endpoint", input.Endpoint, - "protocol", input.Protocol) - return allow, nil - } - cfg, err := s.loadConfig(ctx) + runtimeSnapshot, err := s.loadRuntimeSnapshot(ctx) if err != nil { slog.Warn("content_moderation.skip_config_load_failed", "user_id", input.UserID, @@ -796,6 +803,16 @@ func (s *ContentModerationService) Check(ctx context.Context, input ContentModer "error", err) return allow, nil } + if !runtimeSnapshot.riskControlEnabled { + slog.Info("content_moderation.skip_feature_disabled", + "user_id", input.UserID, + "api_key_id", input.APIKeyID, + "group_id", contentModerationLogGroupID(input.GroupID), + "endpoint", input.Endpoint, + "protocol", input.Protocol) + return allow, nil + } + cfg := runtimeSnapshot.config inGroupScope := cfg.includesGroup(input.GroupID) inModelScope := cfg.includesModel(input.Model) slog.Info("content_moderation.config_loaded", @@ -885,7 +902,7 @@ func (s *ContentModerationService) Check(ctx context.Context, input ContentModer hashText := content.Hash() if cfg.Mode == ContentModerationModePreBlock { if cfg.KeywordBlockingMode != ContentModerationKeywordModeAPIOnly && len(cfg.BlockedKeywords) > 0 { - if keyword, hit := matchBlockedKeyword(content.Text, cfg.BlockedKeywords); hit { + if keyword, hit := runtimeSnapshot.matchBlockedKeyword(content.Text); hit { s.recordPreBlockSyncMetric(0, ContentModerationActionKeywordBlock) slog.Info("content_moderation.keyword_block", "user_id", input.UserID, @@ -1178,12 +1195,13 @@ func (s *ContentModerationService) enqueueRecord(input ContentModerationCheckInp func (s *ContentModerationService) worker(id int) { for { ctx, cancel := context.WithTimeout(context.Background(), maxContentModerationTimeoutMS*time.Millisecond+10*time.Second) - cfg, err := s.loadConfig(ctx) - if err != nil || id >= cfg.WorkerCount { + runtimeSnapshot, err := s.loadRuntimeSnapshot(ctx) + if err != nil || runtimeSnapshot == nil || runtimeSnapshot.config == nil || id >= runtimeSnapshot.config.WorkerCount { cancel() time.Sleep(time.Second) continue } + cfg := runtimeSnapshot.config task, ok := s.dequeueAsyncTask(ctx, time.Second) if !ok { cancel() @@ -1438,15 +1456,18 @@ func (s *ContentModerationService) runCleanupOnce() { } func (s *ContentModerationService) loadConfig(ctx context.Context) (*ContentModerationConfig, error) { - cfg := defaultContentModerationConfig() raw, err := s.settingRepo.GetValue(ctx, SettingKeyContentModerationConfig) if err != nil { if errors.Is(err, ErrSettingNotFound) { - cfg.normalize() - return cfg, nil + return parseContentModerationConfig("") } return nil, fmt.Errorf("get content moderation config: %w", err) } + return parseContentModerationConfig(raw) +} + +func parseContentModerationConfig(raw string) (*ContentModerationConfig, error) { + cfg := defaultContentModerationConfig() if strings.TrimSpace(raw) == "" { cfg.normalize() return cfg, nil @@ -1458,6 +1479,137 @@ func (s *ContentModerationService) loadConfig(ctx context.Context) (*ContentMode return cfg, nil } +func (s *ContentModerationService) loadRuntimeSnapshot(ctx context.Context) (*contentModerationRuntimeSnapshot, error) { + if s == nil || s.settingRepo == nil { + return nil, errors.New("content moderation setting repository unavailable") + } + now := time.Now() + if snapshot := s.runtimeSnapshot.Load(); snapshot != nil { + if now.Sub(snapshot.loadedAt) < s.runtimeSnapshotTTL() { + return snapshot, nil + } + s.triggerRuntimeSnapshotRefresh() + return snapshot, nil + } + + s.runtimeRefreshMu.Lock() + defer s.runtimeRefreshMu.Unlock() + if snapshot := s.runtimeSnapshot.Load(); snapshot != nil { + return snapshot, nil + } + return s.refreshRuntimeSnapshot(ctx) +} + +func (s *ContentModerationService) runtimeSnapshotTTL() time.Duration { + if s != nil && s.runtimeCacheTTL > 0 { + return s.runtimeCacheTTL + } + return contentModerationRuntimeCacheTTL +} + +func (s *ContentModerationService) triggerRuntimeSnapshotRefresh() { + if s == nil || s.runtimeRefreshDeferred() || !s.runtimeRefreshMu.TryLock() { + return + } + if s.runtimeRefreshDeferred() { + s.runtimeRefreshMu.Unlock() + return + } + go func() { + defer s.runtimeRefreshMu.Unlock() + ctx, cancel := context.WithTimeout(context.Background(), contentModerationRuntimeRefreshTimeout) + defer cancel() + if _, err := s.refreshRuntimeSnapshot(ctx); err != nil { + s.runtimeRefreshRetryAt.Store(time.Now().Add(s.runtimeSnapshotTTL()).UnixNano()) + slog.Warn("content_moderation.runtime_snapshot_refresh_failed", "error", err) + } + }() +} + +func (s *ContentModerationService) runtimeRefreshDeferred() bool { + if s == nil { + return false + } + return time.Now().UnixNano() < s.runtimeRefreshRetryAt.Load() +} + +func (s *ContentModerationService) refreshRuntimeSnapshot(ctx context.Context) (*contentModerationRuntimeSnapshot, error) { + values, err := s.settingRepo.GetMultiple(ctx, []string{ + SettingKeyRiskControlEnabled, + SettingKeyContentModerationConfig, + }) + if err != nil { + return nil, fmt.Errorf("get content moderation runtime settings: %w", err) + } + rawConfig := values[SettingKeyContentModerationConfig] + configDigest := sha256.Sum256([]byte(rawConfig)) + if current := s.runtimeSnapshot.Load(); current != nil && current.configDigest == configDigest { + snapshot := &contentModerationRuntimeSnapshot{ + riskControlEnabled: values[SettingKeyRiskControlEnabled] == "true", + config: current.config, + keywordMatcher: current.keywordMatcher, + configDigest: configDigest, + loadedAt: time.Now(), + } + s.runtimeSnapshot.Store(snapshot) + s.runtimeRefreshRetryAt.Store(0) + return snapshot, nil + } + cfg, err := parseContentModerationConfig(rawConfig) + if err != nil { + return nil, err + } + snapshot := &contentModerationRuntimeSnapshot{ + riskControlEnabled: values[SettingKeyRiskControlEnabled] == "true", + config: cfg, + keywordMatcher: newContentModerationKeywordMatcher(cfg.BlockedKeywords), + configDigest: configDigest, + loadedAt: time.Now(), + } + s.runtimeSnapshot.Store(snapshot) + s.runtimeRefreshRetryAt.Store(0) + return snapshot, nil +} + +func (s *ContentModerationService) replaceRuntimeConfig(cfg *ContentModerationConfig, raw []byte) { + if s == nil || cfg == nil { + return + } + s.runtimeRefreshMu.Lock() + hasSnapshot := s.runtimeSnapshot.Load() != nil + s.runtimeRefreshMu.Unlock() + if !hasSnapshot { + return + } + config := cloneContentModerationConfig(cfg) + keywordMatcher := newContentModerationKeywordMatcher(cfg.BlockedKeywords) + configDigest := sha256.Sum256(raw) + + s.runtimeRefreshMu.Lock() + defer s.runtimeRefreshMu.Unlock() + current := s.runtimeSnapshot.Load() + if current == nil { + return + } + s.runtimeSnapshot.Store(&contentModerationRuntimeSnapshot{ + riskControlEnabled: current.riskControlEnabled, + config: config, + keywordMatcher: keywordMatcher, + configDigest: configDigest, + loadedAt: time.Now(), + }) +} + +func (s *contentModerationRuntimeSnapshot) matchBlockedKeyword(text string) (string, bool) { + if s == nil || s.config == nil { + return "", false + } + if s.keywordMatcher != nil { + return s.keywordMatcher.Match(text) + } + return matchBlockedKeyword(text, s.config.BlockedKeywords) +} + func (s *ContentModerationService) isRiskControlEnabled(ctx context.Context) bool { raw, err := s.settingRepo.GetValue(ctx, SettingKeyRiskControlEnabled) if err != nil { diff --git a/backend/internal/service/content_moderation_keyword_matcher.go b/backend/internal/service/content_moderation_keyword_matcher.go new file mode 100644 index 0000000000..4628d3b70a --- /dev/null +++ b/backend/internal/service/content_moderation_keyword_matcher.go @@ -0,0 +1,222 @@ +package service + +import ( + "strings" +) + +type contentModerationKeywordMatcher struct { + nodes []contentModerationKeywordNode + edges []contentModerationKeywordEdge + rootTransitions [256]int32 + keywords []string +} + +type contentModerationKeywordNode struct { + failure int32 + bestKeyword int32 + edgeStart uint32 + edgeCount uint16 +} + +type contentModerationKeywordEdge struct { + target int32 + label byte +} + +type contentModerationKeywordBuildEdge struct { + target int32 + nextSibling int32 + label byte +} + +func newContentModerationKeywordMatcher(keywords []string) *contentModerationKeywordMatcher { + if len(keywords) == 0 { + return nil + } + + buildNodes := []contentModerationKeywordNode{newContentModerationKeywordNode()} + buildEdges := make([]contentModerationKeywordBuildEdge, 0) + originalKeywords := append([]string(nil), keywords...) + + for keywordIndex, keyword := range keywords { + if keyword == "" { + continue + } + state := int32(0) + for _, label := range []byte(strings.ToLower(keyword)) { + next := contentModerationKeywordBuildTransition(buildNodes, buildEdges, state, label) + if next < 0 { + next = int32(len(buildNodes)) + buildNodes = append(buildNodes, newContentModerationKeywordNode()) + buildEdges = append(buildEdges, contentModerationKeywordBuildEdge{ + target: next, + nextSibling: contentModerationKeywordBuildFirstEdge(buildNodes[state]), + label: label, + }) + buildNodes[state].edgeStart = uint32(len(buildEdges)) + } + state = next + } + if current := buildNodes[state].bestKeyword; current < 0 || int32(keywordIndex) < current { + buildNodes[state].bestKeyword = int32(keywordIndex) + } + } + + if len(buildNodes) == 1 { + return nil + } + + queue := make([]int32, 0, len(buildNodes)-1) + var rootTransitions [256]int32 + for edgeIndex := contentModerationKeywordBuildFirstEdge(buildNodes[0]); edgeIndex >= 0; edgeIndex = buildEdges[edgeIndex].nextSibling { + edge := buildEdges[edgeIndex] + rootTransitions[edge.label] = edge.target + queue = append(queue, edge.target) + } + + for queueIndex := 0; queueIndex < len(queue); queueIndex++ { + state := queue[queueIndex] + for edgeIndex := contentModerationKeywordBuildFirstEdge(buildNodes[state]); edgeIndex >= 0; edgeIndex = buildEdges[edgeIndex].nextSibling { + edge := buildEdges[edgeIndex] + failure := buildNodes[state].failure + fallback := contentModerationKeywordBuildTransition(buildNodes, buildEdges, failure, edge.label) + for fallback < 0 && failure != 0 { + failure = buildNodes[failure].failure + fallback = contentModerationKeywordBuildTransition(buildNodes, buildEdges, failure, edge.label) + } + if fallback >= 0 { + buildNodes[edge.target].failure = fallback + } + buildNodes[edge.target].bestKeyword = minKeywordIndex( + buildNodes[edge.target].bestKeyword, + buildNodes[buildNodes[edge.target].failure].bestKeyword, + ) + queue = append(queue, edge.target) + } + } + + edges := make([]contentModerationKeywordEdge, 0, len(buildEdges)) + var outgoing [256]contentModerationKeywordEdge + for nodeIndex := range buildNodes { + count := 0 + for edgeIndex := contentModerationKeywordBuildFirstEdge(buildNodes[nodeIndex]); edgeIndex >= 0; edgeIndex = buildEdges[edgeIndex].nextSibling { + edge := buildEdges[edgeIndex] + outgoing[count] = contentModerationKeywordEdge{target: edge.target, label: edge.label} + count++ + } + for index := 1; index < count; index++ { + current := outgoing[index] + insertAt := index + for insertAt > 0 && current.label < outgoing[insertAt-1].label { + outgoing[insertAt] = outgoing[insertAt-1] + insertAt-- + } + outgoing[insertAt] = current + } + buildNodes[nodeIndex].edgeStart = uint32(len(edges)) + buildNodes[nodeIndex].edgeCount = uint16(count) + edges = append(edges, outgoing[:count]...) + } + + return &contentModerationKeywordMatcher{ + nodes: buildNodes, + edges: edges, + rootTransitions: rootTransitions, + keywords: originalKeywords, + } +} + +func newContentModerationKeywordNode() contentModerationKeywordNode { + return contentModerationKeywordNode{bestKeyword: -1} +} + +func contentModerationKeywordBuildFirstEdge(node contentModerationKeywordNode) int32 { + if node.edgeStart == 0 { + return -1 + } + return int32(node.edgeStart - 1) +} + +func contentModerationKeywordBuildTransition( + nodes []contentModerationKeywordNode, + edges []contentModerationKeywordBuildEdge, + state int32, + label byte, +) int32 { + if state < 0 || int(state) >= len(nodes) { + return -1 + } + for edgeIndex := contentModerationKeywordBuildFirstEdge(nodes[state]); edgeIndex >= 0; edgeIndex = edges[edgeIndex].nextSibling { + if edges[edgeIndex].label == label { + return edges[edgeIndex].target + } + } + return -1 +} + +func minKeywordIndex(left, right int32) int32 { + if left < 0 { + return right + } + if right < 0 || left < right { + return left + } + return right +} + +func (m *contentModerationKeywordMatcher) Match(text string) (string, bool) { + if m == nil || text == "" || len(m.nodes) == 0 || len(m.keywords) == 0 { + return "", false + } + lower := strings.ToLower(text) + state := int32(0) + bestKeyword := int32(-1) + for index := 0; index < len(lower); index++ { + label := lower[index] + for { + next := m.next(state, label) + if next != 0 { + state = next + break + } + if state == 0 { + break + } + state = m.nodes[state].failure + } + bestKeyword = minKeywordIndex(bestKeyword, m.nodes[state].bestKeyword) + if bestKeyword == 0 { + return m.keywords[0], true + } + } + if bestKeyword < 0 || int(bestKeyword) >= len(m.keywords) { + return "", false + } + return m.keywords[bestKeyword], true +} + +func (m *contentModerationKeywordMatcher) next(state int32, label byte) int32 { + if state == 0 { + return m.rootTransitions[label] + } + if state < 0 || int(state) >= len(m.nodes) { + return 0 + } + node := m.nodes[state] + left := int(node.edgeStart) + right := left + int(node.edgeCount) + for left < right { + middle := left + (right-left)/2 + edge := m.edges[middle] + if edge.label < label { + left = middle + 1 + continue + } + right = middle + } + end := int(node.edgeStart) + int(node.edgeCount) + if left < end && m.edges[left].label == label { + return m.edges[left].target + } + return 0 +} diff --git a/backend/internal/service/content_moderation_keyword_matcher_test.go b/backend/internal/service/content_moderation_keyword_matcher_test.go new file mode 100644 index 0000000000..4be70dca3e --- /dev/null +++ b/backend/internal/service/content_moderation_keyword_matcher_test.go @@ -0,0 +1,59 @@ +package service + +import ( + "math/rand" + "strings" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestContentModerationKeywordMatcherMatchesLegacyBehavior(t *testing.T) { + tests := []struct { + name string + text string + keywords []string + }{ + {name: "miss", text: "clean prompt", keywords: []string{"blocked", "secret"}}, + {name: "case insensitive", text: "contains SECRET value", keywords: []string{"secret"}}, + {name: "configured order wins", text: "early appears before later", keywords: []string{"later", "early"}}, + {name: "overlap uses configured order", text: "abc", keywords: []string{"bc", "abc"}}, + {name: "unicode", text: "这里包含敏感词和世界", keywords: []string{"世界", "敏感词"}}, + {name: "duplicates", text: "duplicate", keywords: []string{"duplicate", "DUPLICATE"}}, + {name: "empty entries", text: "blocked", keywords: []string{"", "blocked"}}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + wantKeyword, wantHit := matchBlockedKeyword(tt.text, tt.keywords) + gotKeyword, gotHit := newContentModerationKeywordMatcher(tt.keywords).Match(tt.text) + require.Equal(t, wantHit, gotHit) + require.Equal(t, wantKeyword, gotKeyword) + }) + } +} + +func TestContentModerationKeywordMatcherRandomizedParity(t *testing.T) { + rng := rand.New(rand.NewSource(20260714)) + const alphabet = "abcXYZ" + for iteration := 0; iteration < 1000; iteration++ { + keywords := make([]string, 1+rng.Intn(30)) + for index := range keywords { + length := 1 + rng.Intn(8) + var value strings.Builder + for range length { + _ = value.WriteByte(alphabet[rng.Intn(len(alphabet))]) + } + keywords[index] = value.String() + } + var text strings.Builder + for range 20 + rng.Intn(100) { + _ = text.WriteByte(alphabet[rng.Intn(len(alphabet))]) + } + + wantKeyword, wantHit := matchBlockedKeyword(text.String(), keywords) + gotKeyword, gotHit := newContentModerationKeywordMatcher(keywords).Match(text.String()) + require.Equal(t, wantHit, gotHit, "iteration %d", iteration) + require.Equal(t, wantKeyword, gotKeyword, "iteration %d", iteration) + } +} diff --git a/backend/internal/service/content_moderation_runtime_cache_test.go b/backend/internal/service/content_moderation_runtime_cache_test.go new file mode 100644 index 0000000000..2b34d1fa57 --- /dev/null +++ b/backend/internal/service/content_moderation_runtime_cache_test.go @@ -0,0 +1,451 @@ +package service + +import ( + "context" + "encoding/json" + "errors" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +type contentModerationRuntimeSettingRepo struct { + mu sync.Mutex + values map[string]string + getValueCalls int + getMultipleCalls int + getMultipleErr error + getMultipleStart chan<- struct{} + getMultipleWait <-chan struct{} +} + +func (r *contentModerationRuntimeSettingRepo) Get(_ context.Context, key string) (*Setting, error) { + r.mu.Lock() + defer r.mu.Unlock() + value, ok := r.values[key] + if !ok { + return nil, ErrSettingNotFound + } + return &Setting{Key: key, Value: value}, nil +} + +func (r *contentModerationRuntimeSettingRepo) GetValue(_ context.Context, key string) (string, error) { + r.mu.Lock() + defer r.mu.Unlock() + r.getValueCalls++ + value, ok := r.values[key] + if !ok { + return "", ErrSettingNotFound + } + return value, nil +} + +func (r *contentModerationRuntimeSettingRepo) Set(_ context.Context, key, value string) error { + r.mu.Lock() + defer r.mu.Unlock() + if r.values == nil { + r.values = make(map[string]string) + } + r.values[key] = value + return nil +} + +func (r *contentModerationRuntimeSettingRepo) GetMultiple(_ context.Context, keys []string) (map[string]string, error) { + r.mu.Lock() + r.getMultipleCalls++ + if err := r.getMultipleErr; err != nil { + r.mu.Unlock() + return nil, err + } + out := make(map[string]string, len(keys)) + for _, key := range keys { + if value, ok := r.values[key]; ok { + out[key] = value + } + } + start := r.getMultipleStart + wait := r.getMultipleWait + r.getMultipleStart = nil + r.getMultipleWait = nil + r.mu.Unlock() + if start != nil { + start <- struct{}{} + } + if wait != nil { + <-wait + } + return out, nil +} + +func (r *contentModerationRuntimeSettingRepo) SetMultiple(_ context.Context, values map[string]string) error { + r.mu.Lock() + defer r.mu.Unlock() + if r.values == nil { + r.values = make(map[string]string) + } + for key, value := range values { + r.values[key] = value + } + return nil +} + +func (r *contentModerationRuntimeSettingRepo) GetAll(_ context.Context) (map[string]string, error) { + r.mu.Lock() + defer r.mu.Unlock() + out := make(map[string]string, len(r.values)) + for key, value := range r.values { + out[key] = value + } + return out, nil +} + +func (r *contentModerationRuntimeSettingRepo) Delete(_ context.Context, key string) error { + r.mu.Lock() + defer r.mu.Unlock() + delete(r.values, key) + return nil +} + +func (r *contentModerationRuntimeSettingRepo) calls() (getValue, getMultiple int) { + r.mu.Lock() + defer r.mu.Unlock() + return r.getValueCalls, r.getMultipleCalls +} + +func (r *contentModerationRuntimeSettingRepo) failMultiple(err error) { + r.mu.Lock() + defer r.mu.Unlock() + r.getMultipleErr = err +} + +func (r *contentModerationRuntimeSettingRepo) blockNextMultiple(start chan<- struct{}, wait <-chan struct{}) { + r.mu.Lock() + defer r.mu.Unlock() + r.getMultipleStart = start + r.getMultipleWait = wait +} + +func runtimeCacheTestConfig(t *testing.T, keywords ...string) string { + t.Helper() + cfg := defaultContentModerationConfig() + cfg.Enabled = true + cfg.Mode = ContentModerationModePreBlock + cfg.KeywordBlockingMode = ContentModerationKeywordModeKeywordOnly + cfg.BlockedKeywords = keywords + raw, err := json.Marshal(cfg) + require.NoError(t, err) + return string(raw) +} + +func runtimeCacheTestService(repo *contentModerationRuntimeSettingRepo, ttl time.Duration) *ContentModerationService { + return &ContentModerationService{ + settingRepo: repo, + repo: &contentModerationTestRepo{}, + runtimeCacheTTL: ttl, + } +} + +func runtimeCacheTestInput(text string) ContentModerationCheckInput { + return ContentModerationCheckInput{ + Protocol: ContentModerationProtocolOpenAIChat, + Model: "risk-cache-test", + Body: []byte(`{"messages":[{"role":"user","content":"` + text + `"}]}`), + } +} + +func TestContentModerationRuntimeSnapshotCachesSettings(t *testing.T) { + repo := &contentModerationRuntimeSettingRepo{values: map[string]string{ + SettingKeyRiskControlEnabled: "true", + SettingKeyContentModerationConfig: runtimeCacheTestConfig(t, "blocked"), + }} + svc := runtimeCacheTestService(repo, time.Hour) + + for range 20 { + decision, err := svc.Check(context.Background(), runtimeCacheTestInput("clean prompt")) + require.NoError(t, err) + require.True(t, decision.Allowed) + } + + getValue, getMultiple := repo.calls() + require.Zero(t, getValue) + require.Equal(t, 1, getMultiple) +} + +func TestContentModerationRuntimeSnapshotUpdateConfigIsImmediate(t *testing.T) { + repo := &contentModerationRuntimeSettingRepo{values: map[string]string{ + SettingKeyRiskControlEnabled: "true", + SettingKeyContentModerationConfig: runtimeCacheTestConfig(t, "old-keyword"), + }} + svc := runtimeCacheTestService(repo, time.Hour) + + decision, err := svc.Check(context.Background(), runtimeCacheTestInput("new-keyword")) + require.NoError(t, err) + require.True(t, decision.Allowed) + + keywords := []string{"new-keyword"} + _, err = svc.UpdateConfig(context.Background(), UpdateContentModerationConfigInput{ + BlockedKeywords: &keywords, + }) + require.NoError(t, err) + + decision, err = svc.Check(context.Background(), runtimeCacheTestInput("new-keyword")) + require.NoError(t, err) + require.True(t, decision.Blocked) + + _, getMultiple := repo.calls() + require.Equal(t, 1, getMultiple) +} + +func TestContentModerationRuntimeSnapshotUpdateWinsOverInitialLoad(t *testing.T) { + repo := &contentModerationRuntimeSettingRepo{values: map[string]string{ + SettingKeyRiskControlEnabled: "true", + SettingKeyContentModerationConfig: runtimeCacheTestConfig(t, "old-keyword"), + }} + svc := runtimeCacheTestService(repo, time.Hour) + + refreshStarted := make(chan struct{}, 1) + releaseRefresh := make(chan struct{}) + released := false + defer func() { + if !released { + close(releaseRefresh) + } + }() + repo.blockNextMultiple(refreshStarted, releaseRefresh) + + initialCheckDone := make(chan error, 1) + go func() { + decision, err := svc.Check(context.Background(), runtimeCacheTestInput("clean prompt")) + if err == nil && (decision == nil || !decision.Allowed) { + err = errors.New("unexpected initial moderation decision") + } + initialCheckDone <- err + }() + require.Eventually(t, func() bool { + select { + case <-refreshStarted: + return true + default: + return false + } + }, time.Second, time.Millisecond) + + updateDone := make(chan error, 1) + go func() { + keywords := []string{"new-keyword"} + _, updateErr := svc.UpdateConfig(context.Background(), UpdateContentModerationConfigInput{ + BlockedKeywords: &keywords, + }) + updateDone <- updateErr + }() + select { + case updateErr := <-updateDone: + require.NoError(t, updateErr) + t.Fatal("configuration update completed before the initial load released its lock") + case <-time.After(10 * time.Millisecond): + } + + close(releaseRefresh) + released = true + require.NoError(t, <-initialCheckDone) + require.NoError(t, <-updateDone) + + decision, err := svc.Check(context.Background(), runtimeCacheTestInput("new-keyword")) + require.NoError(t, err) + require.True(t, decision.Blocked) + decision, err = svc.Check(context.Background(), runtimeCacheTestInput("old-keyword")) + require.NoError(t, err) + require.True(t, decision.Allowed) +} + +func TestContentModerationRuntimeSnapshotRefreshFailureKeepsStaleConfig(t *testing.T) { + repo := &contentModerationRuntimeSettingRepo{values: map[string]string{ + SettingKeyRiskControlEnabled: "true", + SettingKeyContentModerationConfig: runtimeCacheTestConfig(t, "blocked"), + }} + svc := runtimeCacheTestService(repo, time.Nanosecond) + input := runtimeCacheTestInput("blocked") + + decision, err := svc.Check(context.Background(), input) + require.NoError(t, err) + require.True(t, decision.Blocked) + + repo.failMultiple(errors.New("database unavailable")) + decision, err = svc.Check(context.Background(), input) + require.NoError(t, err) + require.True(t, decision.Blocked) + require.Eventually(t, func() bool { + _, calls := repo.calls() + return calls >= 2 + }, time.Second, time.Millisecond) +} + +func TestContentModerationRuntimeSnapshotRefreshFailureBacksOff(t *testing.T) { + repo := &contentModerationRuntimeSettingRepo{values: map[string]string{ + SettingKeyRiskControlEnabled: "true", + SettingKeyContentModerationConfig: runtimeCacheTestConfig(t, "blocked"), + }} + svc := runtimeCacheTestService(repo, time.Minute) + input := runtimeCacheTestInput("blocked") + + decision, err := svc.Check(context.Background(), input) + require.NoError(t, err) + require.True(t, decision.Blocked) + + current := svc.runtimeSnapshot.Load() + require.NotNil(t, current) + expired := *current + expired.loadedAt = time.Now().Add(-2 * time.Minute) + svc.runtimeSnapshot.Store(&expired) + repo.failMultiple(errors.New("database unavailable")) + + decision, err = svc.Check(context.Background(), input) + require.NoError(t, err) + require.True(t, decision.Blocked) + require.Eventually(t, func() bool { + _, calls := repo.calls() + return calls == 2 + }, time.Second, time.Millisecond) + + for range 100 { + decision, err = svc.Check(context.Background(), input) + require.NoError(t, err) + require.True(t, decision.Blocked) + } + _, calls := repo.calls() + require.Equal(t, 2, calls) +} + +func TestContentModerationRuntimeSnapshotRefreshReusesUnchangedMatcher(t *testing.T) { + repo := &contentModerationRuntimeSettingRepo{values: map[string]string{ + SettingKeyRiskControlEnabled: "true", + SettingKeyContentModerationConfig: runtimeCacheTestConfig(t, "blocked"), + }} + svc := runtimeCacheTestService(repo, time.Minute) + input := runtimeCacheTestInput("blocked") + + decision, err := svc.Check(context.Background(), input) + require.NoError(t, err) + require.True(t, decision.Blocked) + + current := svc.runtimeSnapshot.Load() + require.NotNil(t, current) + expired := *current + expired.loadedAt = time.Now().Add(-2 * time.Minute) + svc.runtimeSnapshot.Store(&expired) + + decision, err = svc.Check(context.Background(), input) + require.NoError(t, err) + require.True(t, decision.Blocked) + require.Eventually(t, func() bool { + refreshed := svc.runtimeSnapshot.Load() + return refreshed != nil && refreshed.loadedAt.After(expired.loadedAt) + }, time.Second, time.Millisecond) + + refreshed := svc.runtimeSnapshot.Load() + require.Same(t, current.config, refreshed.config) + require.Same(t, current.keywordMatcher, refreshed.keywordMatcher) + _, calls := repo.calls() + require.Equal(t, 2, calls) +} + +func TestContentModerationRuntimeSnapshotUpdateWinsOverInFlightRefresh(t *testing.T) { + repo := &contentModerationRuntimeSettingRepo{values: map[string]string{ + SettingKeyRiskControlEnabled: "true", + SettingKeyContentModerationConfig: runtimeCacheTestConfig(t, "old-keyword"), + }} + svc := runtimeCacheTestService(repo, time.Minute) + + decision, err := svc.Check(context.Background(), runtimeCacheTestInput("old-keyword")) + require.NoError(t, err) + require.True(t, decision.Blocked) + + current := svc.runtimeSnapshot.Load() + require.NotNil(t, current) + expired := *current + expired.loadedAt = time.Now().Add(-2 * time.Minute) + svc.runtimeSnapshot.Store(&expired) + + refreshStarted := make(chan struct{}, 1) + releaseRefresh := make(chan struct{}) + repo.blockNextMultiple(refreshStarted, releaseRefresh) + decision, err = svc.Check(context.Background(), runtimeCacheTestInput("clean prompt")) + require.NoError(t, err) + require.True(t, decision.Allowed) + require.Eventually(t, func() bool { + select { + case <-refreshStarted: + return true + default: + return false + } + }, time.Second, time.Millisecond) + + updateDone := make(chan error, 1) + go func() { + keywords := []string{"new-keyword"} + _, updateErr := svc.UpdateConfig(context.Background(), UpdateContentModerationConfigInput{ + BlockedKeywords: &keywords, + }) + updateDone <- updateErr + }() + select { + case updateErr := <-updateDone: + require.NoError(t, updateErr) + t.Fatal("configuration update completed before the in-flight refresh released its lock") + case <-time.After(10 * time.Millisecond): + } + + close(releaseRefresh) + require.NoError(t, <-updateDone) + decision, err = svc.Check(context.Background(), runtimeCacheTestInput("new-keyword")) + require.NoError(t, err) + require.True(t, decision.Blocked) + decision, err = svc.Check(context.Background(), runtimeCacheTestInput("old-keyword")) + require.NoError(t, err) + require.True(t, decision.Allowed) +} + +func TestContentModerationRuntimeSnapshotConcurrentReadAndReplace(t *testing.T) { + repo := &contentModerationRuntimeSettingRepo{values: map[string]string{ + SettingKeyRiskControlEnabled: "true", + SettingKeyContentModerationConfig: runtimeCacheTestConfig(t, "blocked-0"), + }} + svc := runtimeCacheTestService(repo, time.Hour) + _, err := svc.Check(context.Background(), runtimeCacheTestInput("clean prompt")) + require.NoError(t, err) + + var wg sync.WaitGroup + errs := make(chan error, 8) + for range 8 { + wg.Add(1) + go func() { + defer wg.Done() + for range 100 { + decision, checkErr := svc.Check(context.Background(), runtimeCacheTestInput("clean prompt")) + if checkErr != nil { + errs <- checkErr + return + } + if decision == nil || !decision.Allowed { + errs <- errors.New("unexpected moderation decision") + return + } + } + }() + } + for i := 1; i <= 20; i++ { + keywords := []string{"blocked-" + time.Duration(i).String()} + _, err := svc.UpdateConfig(context.Background(), UpdateContentModerationConfigInput{ + BlockedKeywords: &keywords, + }) + require.NoError(t, err) + } + wg.Wait() + close(errs) + for err := range errs { + require.NoError(t, err) + } +} diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index c151375623..49ef63c0c1 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -558,19 +558,79 @@ type ForwardResult struct { ImageSizeBreakdown map[string]int } -// UpstreamFailoverError indicates an upstream error that should trigger account failover. +// GatewayFailureStage identifies which request stage failed. The zero value is +// intentionally treated as inference so existing UpstreamFailoverError callers +// retain their current behavior. +type GatewayFailureStage string + +const ( + GatewayFailureStageInference GatewayFailureStage = "inference" + GatewayFailureStageAccountAuth GatewayFailureStage = "account_auth" +) + +// GatewayFailureScope identifies whether selecting another account can help. +type GatewayFailureScope string + +const ( + GatewayFailureScopeAccount GatewayFailureScope = "account" + GatewayFailureScopeProvider GatewayFailureScope = "provider" + GatewayFailureScopeRequest GatewayFailureScope = "request" +) + +// NextAccountAction is tri-state for backwards compatibility. The zero value +// means legacy retry behavior; only NextAccountStop explicitly short-circuits. +type NextAccountAction uint8 + +const ( + NextAccountLegacyRetry NextAccountAction = iota + NextAccountRetry + NextAccountStop +) + +type GatewayFailureReason string + +// UpstreamFailoverError indicates an upstream or credential error that may +// trigger account failover. Additive metadata keeps existing composite literals +// source-compatible and preserves their legacy retry-next-account behavior. type UpstreamFailoverError struct { StatusCode int ResponseBody []byte // 上游响应体,用于错误透传规则匹配 ResponseHeaders http.Header // 上游响应头,用于透传 cf-ray/cf-mitigated/content-type 等诊断信息 ForceCacheBilling bool // Antigravity 粘性会话切换时设为 true RetryableOnSameAccount bool // 临时性错误(如 Google 间歇性 400、空响应),应在同一账号上重试 N 次再切换 + Stage GatewayFailureStage + Scope GatewayFailureScope + Reason GatewayFailureReason + NextAccountAction NextAccountAction + ClientStatusCode int + ClientMessage string } func (e *UpstreamFailoverError) Error() string { + if e != nil && e.Stage == GatewayFailureStageAccountAuth { + return fmt.Sprintf("credential failure: %s (failover)", e.Reason) + } return fmt.Sprintf("upstream error: %d (failover)", e.StatusCode) } +func (e *UpstreamFailoverError) ShouldRetryNextAccount() bool { + return e != nil && e.NextAccountAction != NextAccountStop +} + +func (e *UpstreamFailoverError) IsCredentialFailure() bool { + return e != nil && e.Stage == GatewayFailureStageAccountAuth +} + +// ShouldReportAccountScheduleFailure prevents provider- and request-scoped +// credential failures from being misattributed to the selected account. Legacy +// and inference failures retain their existing scheduler-health behavior. +func (e *UpstreamFailoverError) ShouldReportAccountScheduleFailure() bool { + if e == nil { + return false + } + return !e.IsCredentialFailure() || e.Scope == GatewayFailureScopeAccount +} + // sseStreamErrorEventError 表示上游 SSE 流体内出现 event:error 帧。 // RawData 是该事件 data: 行的原始 JSON 字符串 // (Anthropic 标准结构 {"type":"error","error":{"type":"...","message":"..."}})。 diff --git a/backend/internal/service/grok_credential_failure.go b/backend/internal/service/grok_credential_failure.go new file mode 100644 index 0000000000..65dd844072 --- /dev/null +++ b/backend/internal/service/grok_credential_failure.go @@ -0,0 +1,655 @@ +package service + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "log/slog" + "net/http" + "strings" + "time" + + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/gin-gonic/gin" +) + +const ( + grokCredentialFailoverDeadlineKey = "grok_credential_failover_deadline" + grokCredentialFailoverBudget = 15 * time.Second + grokCredentialMutationTimeout = 5 * time.Second + grokCredentialMutationConfirmWait = 250 * time.Millisecond + grokCredentialCacheCleanupTimeout = 500 * time.Millisecond + + GrokCredentialUnavailableClientMessage = "No healthy Grok OAuth account is currently available" + + GrokCredentialReasonRevoked GatewayFailureReason = "grok_oauth_credential_revoked" + GrokCredentialReasonMissing GatewayFailureReason = "grok_oauth_credentials_missing" + GrokCredentialReasonEntitlement GatewayFailureReason = "grok_oauth_entitlement_action_required" + GrokCredentialReasonProxyInvalid GatewayFailureReason = "grok_oauth_proxy_invalid" + GrokCredentialReasonRefreshTransient GatewayFailureReason = "grok_oauth_refresh_transient" + GrokCredentialReasonProviderConfig GatewayFailureReason = "grok_oauth_provider_config" + GrokCredentialReasonProviderDown GatewayFailureReason = "grok_oauth_provider_unavailable" + GrokCredentialReasonAccountChanged GatewayFailureReason = "grok_oauth_account_state_changed" + GrokCredentialReasonStateUpdate GatewayFailureReason = "grok_oauth_account_state_update_failed" + GrokCredentialReasonFailoverTimeout GatewayFailureReason = "grok_oauth_failover_timeout" +) + +var errGrokCredentialStateUpdateFailed = errors.New("grok oauth account state update failed") + +type grokCredentialFailureClass struct { + scope GatewayFailureScope + reason GatewayFailureReason + action NextAccountAction + permanent bool + transient bool + message string + snapshot *GrokCredentialMutationSnapshot +} + +// GrokCredentialMutationSnapshot is the credential identity observed when the +// request selected an account. Repository mutations compare all fields before +// quarantining that account so a concurrent refresh cannot be overwritten. +type GrokCredentialMutationSnapshot struct { + CredentialsJSON string + AccessToken string + RefreshToken string + TokenVersion int64 + ProxyID *int64 +} + +type grokCredentialFailureSnapshotError struct { + cause error + snapshot GrokCredentialMutationSnapshot +} + +func (e *grokCredentialFailureSnapshotError) Error() string { return e.cause.Error() } +func (e *grokCredentialFailureSnapshotError) Unwrap() error { return e.cause } + +func withGrokCredentialFailureSnapshot(err error, account *Account) error { + if err == nil || account == nil || !account.IsGrokOAuth() { + return err + } + var existing *grokCredentialFailureSnapshotError + if errors.As(err, &existing) { + return err + } + return &grokCredentialFailureSnapshotError{cause: err, snapshot: grokCredentialMutationSnapshot(account)} +} + +func grokCredentialFailureSnapshot(err error) (GrokCredentialMutationSnapshot, bool) { + var snapshotErr *grokCredentialFailureSnapshotError + if !errors.As(err, &snapshotErr) || snapshotErr == nil { + return GrokCredentialMutationSnapshot{}, false + } + return snapshotErr.snapshot, true +} + +type grokCredentialConditionalStateRepository interface { + SetGrokCredentialErrorIfMatch(context.Context, int64, GrokCredentialMutationSnapshot, string) (bool, error) + SetGrokCredentialTempUnschedulableIfMatch(context.Context, int64, GrokCredentialMutationSnapshot, time.Time, string) (bool, error) +} + +// GetRequestCredential applies the request-path credential and failover contract +// before any upstream transport is opened. +func (s *OpenAIGatewayService) GetRequestCredential(ctx context.Context, c *gin.Context, account *Account) (string, string, error) { + return s.getRequestCredential(ctx, c, account) +} + +func (s *OpenAIGatewayService) getRequestCredential(ctx context.Context, c *gin.Context, account *Account) (string, string, error) { + if ctx == nil { + ctx = context.Background() + } + if account == nil { + return "", "", errors.New("account is nil") + } + if !account.IsGrokOAuth() { + return s.GetAccessToken(ctx, account) + } + if err := ctx.Err(); err != nil { + return "", "", err + } + if s == nil || s.grokTokenProvider == nil { + return "", "", s.newGrokCredentialFailover(c, account, grokCredentialFailureClass{ + scope: GatewayFailureScopeProvider, + reason: GrokCredentialReasonProviderConfig, + action: NextAccountStop, + message: "Grok OAuth credential provider is unavailable", + }) + } + if s.isOpenAIAccountRuntimeBlocked(account) { + return "", "", s.newGrokCredentialFailover(c, account, grokCredentialFailureClass{ + scope: GatewayFailureScopeAccount, + reason: GrokCredentialReasonAccountChanged, + action: NextAccountRetry, + message: "Grok OAuth account is not currently schedulable", + }) + } + + credentialCtx, cancel, budgetExpired := grokCredentialAcquisitionContext(ctx, c) + if cancel != nil { + defer cancel() + } + if budgetExpired { + return "", "", s.newGrokCredentialFailover(c, account, grokCredentialFailureClass{ + scope: GatewayFailureScopeRequest, + reason: GrokCredentialReasonFailoverTimeout, + action: NextAccountStop, + message: "Grok OAuth credential failover budget exhausted", + }) + } + + token, kind, err := s.GetAccessToken(credentialCtx, account) + if err == nil { + if s.isOpenAIAccountRuntimeBlocked(account) { + return "", "", s.newGrokCredentialFailover(c, account, grokCredentialFailureClass{ + scope: GatewayFailureScopeAccount, + reason: GrokCredentialReasonAccountChanged, + action: NextAccountRetry, + message: "Grok OAuth account is not currently schedulable", + }) + } + return token, kind, nil + } + if parentErr := ctx.Err(); parentErr != nil { + return "", "", parentErr + } + if credentialCtx.Err() != nil { + return "", "", s.newGrokCredentialFailover(c, account, grokCredentialFailureClass{ + scope: GatewayFailureScopeRequest, + reason: GrokCredentialReasonFailoverTimeout, + action: NextAccountStop, + message: "Grok OAuth credential failover budget exhausted", + }) + } + + class := classifyGrokCredentialFailure(account, err) + if snapshot, ok := grokCredentialFailureSnapshot(err); ok { + class.snapshot = &snapshot + } + if ctx.Err() != nil { + return "", "", ctx.Err() + } + if class.permanent || class.transient { + freshToken, mutationErr := s.applyGrokCredentialAccountFailure(credentialCtx, account, class) + if freshToken != "" { + return freshToken, "oauth", nil + } + if mutationErr != nil { + if ctx.Err() != nil { + return "", "", ctx.Err() + } + if credentialCtx.Err() != nil { + return "", "", s.newGrokCredentialFailover(c, account, grokCredentialFailureClass{ + scope: GatewayFailureScopeRequest, + reason: GrokCredentialReasonFailoverTimeout, + action: NextAccountStop, + message: "Grok OAuth credential failover budget exhausted", + }) + } + if errors.Is(mutationErr, errOAuthRefreshAccountStateChanged) { + class = grokCredentialFailureClass{ + scope: GatewayFailureScopeAccount, + reason: GrokCredentialReasonAccountChanged, + action: NextAccountRetry, + message: "Grok OAuth account eligibility changed", + } + } else if errors.Is(mutationErr, errOAuthRefreshAccountRereadFailed) { + class = grokCredentialFailureClass{ + scope: GatewayFailureScopeProvider, + reason: GrokCredentialReasonProviderDown, + action: NextAccountStop, + message: "Grok OAuth account state is temporarily unavailable", + } + } else { + class = grokCredentialFailureClass{ + scope: GatewayFailureScopeProvider, + reason: GrokCredentialReasonStateUpdate, + action: NextAccountStop, + message: "Grok OAuth account state could not be updated safely", + } + } + } + } + return "", "", s.newGrokCredentialFailover(c, account, class) +} + +func grokCredentialAcquisitionContext(ctx context.Context, c *gin.Context) (context.Context, context.CancelFunc, bool) { + if c == nil { + return ctx, nil, false + } + deadline := time.Time{} + if raw, ok := c.Get(grokCredentialFailoverDeadlineKey); ok { + deadline, _ = raw.(time.Time) + } + if deadline.IsZero() { + deadline = time.Now().Add(grokCredentialFailoverBudget) + c.Set(grokCredentialFailoverDeadlineKey, deadline) + } + if !time.Now().Before(deadline) { + return ctx, nil, true + } + acquireCtx, cancel := context.WithDeadline(ctx, deadline) + return acquireCtx, cancel, false +} + +func classifyGrokCredentialFailure(account *Account, err error) grokCredentialFailureClass { + stableReason := strings.ToLower(strings.TrimSpace(infraerrors.Reason(err))) + message := "" + if err != nil { + message = strings.ToLower(err.Error()) + } + contains := func(values ...string) bool { + for _, value := range values { + if strings.Contains(stableReason, value) || strings.Contains(message, value) { + return true + } + } + return false + } + + switch { + case errors.Is(err, errGrokOAuthRefreshTokenMissing), errors.Is(err, errGrokOAuthAccessTokenMissing), errors.Is(err, errGrokOAuthAccessTokenExpired): + return grokCredentialFailureClass{scope: GatewayFailureScopeAccount, reason: GrokCredentialReasonMissing, action: NextAccountRetry, permanent: true, message: "Grok OAuth credentials are missing or expired"} + case contains("invalid_grant", "invalid_refresh_token", "token_expired", "refresh_token_reused", "refresh_token_invalidated", "app_session_terminated"): + return grokCredentialFailureClass{scope: GatewayFailureScopeAccount, reason: GrokCredentialReasonRevoked, action: NextAccountRetry, permanent: true, message: "Grok OAuth credentials require account action"} + case contains("grok_oauth_entitlement_denied", "entitlement_denied", "access_denied", "subscription required", "no active grok subscription"): + return grokCredentialFailureClass{scope: GatewayFailureScopeAccount, reason: GrokCredentialReasonEntitlement, action: NextAccountRetry, permanent: true, message: "Grok OAuth entitlement requires account action"} + case errors.Is(err, errGrokOAuthConfiguredProxyMiss), contains("grok_oauth_proxy_not_found"): + return grokCredentialFailureClass{scope: GatewayFailureScopeAccount, reason: GrokCredentialReasonProxyInvalid, action: NextAccountRetry, permanent: true, message: "Grok OAuth account proxy configuration is invalid"} + case errors.Is(err, errOAuthRefreshAccountRereadFailed): + return grokCredentialFailureClass{scope: GatewayFailureScopeProvider, reason: GrokCredentialReasonProviderDown, action: NextAccountStop, message: "Grok OAuth account state is temporarily unavailable"} + case errors.Is(err, errOAuthRefreshCredentialPersist): + return grokCredentialFailureClass{scope: GatewayFailureScopeProvider, reason: GrokCredentialReasonProviderDown, action: NextAccountStop, message: "Grok OAuth shared credential state is temporarily unavailable"} + case errors.Is(err, errOAuthRefreshAccountStateChanged): + return grokCredentialFailureClass{scope: GatewayFailureScopeAccount, reason: GrokCredentialReasonAccountChanged, action: NextAccountRetry, message: "Grok OAuth account eligibility changed"} + case errors.Is(err, errGrokOAuthRefreshNotConfigured), contains("invalid_client", "unauthorized_client", "invalid_scope", "unknown scope", "grok oauth service is not configured", "grok_oauth_proxy_not_available"): + return grokCredentialFailureClass{scope: GatewayFailureScopeProvider, reason: GrokCredentialReasonProviderConfig, action: NextAccountStop, message: "Grok OAuth provider configuration is unavailable"} + case contains("grok_oauth_proxy_lookup_failed"), + contains("grok_oauth_token_refresh_failed") && contains("status 403") && (account == nil || account.ProxyID == nil): + return grokCredentialFailureClass{scope: GatewayFailureScopeProvider, reason: GrokCredentialReasonProviderDown, action: NextAccountStop, message: "Grok OAuth provider is temporarily unavailable"} + case contains("grok_oauth_client_init_failed") && (account == nil || account.ProxyID == nil): + return grokCredentialFailureClass{scope: GatewayFailureScopeProvider, reason: GrokCredentialReasonProviderConfig, action: NextAccountStop, message: "Grok OAuth provider configuration is unavailable"} + case contains("grok_oauth_request_failed") && (account == nil || account.ProxyID == nil): + return grokCredentialFailureClass{scope: GatewayFailureScopeProvider, reason: GrokCredentialReasonProviderDown, action: NextAccountStop, message: "Grok OAuth provider is temporarily unavailable"} + case contains("status 429", "status 500", "status 502", "status 503", "status 504") && (account == nil || account.ProxyID == nil): + return grokCredentialFailureClass{scope: GatewayFailureScopeProvider, reason: GrokCredentialReasonProviderDown, action: NextAccountStop, message: "Grok OAuth provider is temporarily unavailable"} + default: + return grokCredentialFailureClass{scope: GatewayFailureScopeAccount, reason: GrokCredentialReasonRefreshTransient, action: NextAccountRetry, transient: true, message: "Grok OAuth credential refresh is temporarily unavailable"} + } +} + +func (s *OpenAIGatewayService) applyGrokCredentialAccountFailure(ctx context.Context, account *Account, class grokCredentialFailureClass) (string, error) { + if s == nil || account == nil || ctx == nil || ctx.Err() != nil { + if ctx != nil { + return "", ctx.Err() + } + return "", context.Canceled + } + mutationMu := s.grokCredentialMutationLock(account.ID) + if err := mutationMu.Lock(ctx); err != nil { + return "", err + } + defer mutationMu.Unlock() + stateRepo, hasConditionalStateRepo := s.accountRepo.(grokCredentialConditionalStateRepository) + snapshot := grokCredentialMutationSnapshot(account) + if class.snapshot != nil { + snapshot = *class.snapshot + } + if token, err := s.validateCurrentGrokCredentialFailure(ctx, account.ID, snapshot, class); err != nil || token != "" { + return token, err + } + + if class.permanent { + if token, ok := s.grokCredentialConcurrentlyRefreshedToken(ctx, account.ID, snapshot); ok { + return token, nil + } + if ctx.Err() != nil { + return "", ctx.Err() + } + rollbackRuntime := s.blockGrokCredentialRuntime(account, time.Time{}, string(class.reason)) + keepRuntimeBlock := false + runtimeRollbackDone := false + defer func() { + if !keepRuntimeBlock && !runtimeRollbackDone { + rollbackRuntime() + } + }() + if s.accountRepo == nil { + keepRuntimeBlock = true + return "", fmt.Errorf("%w: account repository is not configured", errGrokCredentialStateUpdateFailed) + } + if !hasConditionalStateRepo { + keepRuntimeBlock = true + return "", fmt.Errorf("%w: conditional account repository is not configured", errGrokCredentialStateUpdateFailed) + } + stateCtx, cancel := context.WithTimeout(ctx, grokCredentialMutationTimeout) + if err := ctx.Err(); err != nil { + cancel() + return "", err + } + updated, err := stateRepo.SetGrokCredentialErrorIfMatch(stateCtx, account.ID, snapshot, string(class.reason)) + requestErr := ctx.Err() + cancel() + if err != nil { + slog.Warn("grok_credential_failure.set_error_failed", "account_id", account.ID, "reason", class.reason, "error", err) + if requestErr != nil || errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + if s.grokCredentialMutationCommitted(account.ID, class, time.Time{}) { + updated = true + } else if requestErr != nil { + return "", requestErr + } else { + keepRuntimeBlock = true + return "", fmt.Errorf("%w: permanent state commit could not be confirmed: %v", errGrokCredentialStateUpdateFailed, err) + } + } else { + keepRuntimeBlock = true + return "", fmt.Errorf("%w: persist permanent state: %v", errGrokCredentialStateUpdateFailed, err) + } + } + if !updated { + rollbackRuntime() + runtimeRollbackDone = true + return s.resolveGrokCredentialCASMiss(ctx, account.ID, snapshot) + } + // SetError is the linearization point: the durable quarantine is now + // committed and this node's runtime block must not be rolled back. + keepRuntimeBlock = true + if s.grokTokenProvider == nil { + if ctx.Err() != nil { + return "", ctx.Err() + } + return "", fmt.Errorf("%w: token provider is not configured", errGrokCredentialStateUpdateFailed) + } + invalidateCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), grokCredentialCacheCleanupTimeout) + err = s.grokTokenProvider.InvalidateToken(invalidateCtx, account) + cancel() + if err != nil { + slog.Warn("grok_credential_failure.invalidate_token_failed", "account_id", account.ID, "reason", class.reason, "error", err) + if ctx.Err() != nil { + return "", ctx.Err() + } + return "", fmt.Errorf("%w: invalidate cached credential: %v", errGrokCredentialStateUpdateFailed, err) + } + if ctx.Err() != nil { + return "", ctx.Err() + } + return "", nil + } + + if class.transient { + until := time.Now().Add(tokenRefreshTempUnschedDuration) + if ctx.Err() != nil { + return "", ctx.Err() + } + rollbackRuntime := s.blockGrokCredentialRuntime(account, until, string(class.reason)) + keepRuntimeBlock := false + runtimeRollbackDone := false + defer func() { + if !keepRuntimeBlock && !runtimeRollbackDone { + rollbackRuntime() + } + }() + stateCtx, cancel := context.WithTimeout(ctx, grokCredentialMutationTimeout) + if s.accountRepo == nil { + cancel() + keepRuntimeBlock = true + return "", fmt.Errorf("%w: account repository is not configured", errGrokCredentialStateUpdateFailed) + } + if !hasConditionalStateRepo { + cancel() + keepRuntimeBlock = true + return "", fmt.Errorf("%w: conditional account repository is not configured", errGrokCredentialStateUpdateFailed) + } + if err := ctx.Err(); err != nil { + cancel() + return "", err + } + updated, err := stateRepo.SetGrokCredentialTempUnschedulableIfMatch(stateCtx, account.ID, snapshot, until, string(class.reason)) + requestErr := ctx.Err() + cancel() + if err != nil { + slog.Warn("grok_credential_failure.set_temp_unschedulable_failed", "account_id", account.ID, "reason", class.reason, "error", err) + if requestErr != nil || errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + if s.grokCredentialMutationCommitted(account.ID, class, until) { + updated = true + } else if requestErr != nil { + return "", requestErr + } else { + keepRuntimeBlock = true + return "", fmt.Errorf("%w: transient state commit could not be confirmed: %v", errGrokCredentialStateUpdateFailed, err) + } + } else { + keepRuntimeBlock = true + return "", fmt.Errorf("%w: persist transient state: %v", errGrokCredentialStateUpdateFailed, err) + } + } + if !updated { + rollbackRuntime() + runtimeRollbackDone = true + return s.resolveGrokCredentialCASMiss(ctx, account.ID, snapshot) + } + // The temporary quarantine is durable after SetTempUnschedulable succeeds. + keepRuntimeBlock = true + if ctx.Err() != nil { + return "", ctx.Err() + } + return "", nil + } + + return "", nil +} + +func (s *OpenAIGatewayService) validateCurrentGrokCredentialFailure( + ctx context.Context, + accountID int64, + snapshot GrokCredentialMutationSnapshot, + class grokCredentialFailureClass, +) (string, error) { + if s == nil || s.accountRepo == nil || accountID <= 0 || ctx == nil || ctx.Err() != nil { + if ctx != nil { + return "", ctx.Err() + } + return "", context.Canceled + } + checkCtx, cancel := context.WithTimeout(ctx, time.Second) + defer cancel() + latest, err := s.accountRepo.GetByID(checkCtx, accountID) + if err != nil { + if errors.Is(err, ErrAccountNotFound) { + return "", errOAuthRefreshAccountStateChanged + } + return "", fmt.Errorf("%w: %v", errOAuthRefreshAccountRereadFailed, err) + } + if latest == nil || !latest.IsGrokOAuth() || !latest.IsSchedulable() || s.isOpenAIAccountRuntimeBlocked(latest) { + return "", errOAuthRefreshAccountStateChanged + } + latestSnapshot := grokCredentialMutationSnapshot(latest) + if latestSnapshot.CredentialsJSON != snapshot.CredentialsJSON || + !grokCredentialProxyIDsEqual(latestSnapshot.ProxyID, snapshot.ProxyID) { + if token, ok := s.grokCredentialConcurrentlyRefreshedToken(ctx, accountID, snapshot); ok { + return token, nil + } + return "", errOAuthRefreshAccountStateChanged + } + + // A configured proxy is external to the account-row CAS identity. Recheck + // the hydrated proxy object so restoring a deleted row under the same ID + // wins over a stale proxy-invalid failure. + if class.reason == GrokCredentialReasonProxyInvalid { + if latest.ProxyID == nil || latest.Proxy != nil { + return "", errOAuthRefreshAccountStateChanged + } + } else if latest.ProxyID != nil && latest.Proxy == nil { + return "", errOAuthRefreshAccountStateChanged + } + + if class.reason == GrokCredentialReasonMissing { + expiresAt := latest.GetCredentialAsTime("expires_at") + credentialsStillMissing := strings.TrimSpace(latest.GetGrokAccessToken()) == "" || + strings.TrimSpace(latest.GetGrokRefreshToken()) == "" || expiresAt == nil || !time.Now().Before(*expiresAt) + if !credentialsStillMissing { + return "", errOAuthRefreshAccountStateChanged + } + } + return "", nil +} + +func (s *OpenAIGatewayService) grokCredentialMutationLock(accountID int64) *oauthRefreshLocalLock { + actual, _ := s.grokCredentialMutationLocks.LoadOrStore(accountID, newOAuthRefreshLocalLock()) + mu, ok := actual.(*oauthRefreshLocalLock) + if !ok { + mu = newOAuthRefreshLocalLock() + s.grokCredentialMutationLocks.Store(accountID, mu) + } + return mu +} + +func (s *OpenAIGatewayService) grokCredentialMutationCommitted(accountID int64, class grokCredentialFailureClass, until time.Time) bool { + if s == nil || s.accountRepo == nil || accountID <= 0 { + return false + } + confirmCtx, cancel := context.WithTimeout(context.Background(), grokCredentialMutationConfirmWait) + defer cancel() + latest, err := s.accountRepo.GetByID(confirmCtx, accountID) + if err != nil || latest == nil { + return false + } + if class.permanent { + return latest.Status == StatusError && !latest.Schedulable && latest.ErrorMessage == string(class.reason) + } + if class.transient { + return latest.TempUnschedulableUntil != nil && !latest.TempUnschedulableUntil.Before(until) && + latest.TempUnschedulableReason == string(class.reason) + } + return false +} + +func grokCredentialMutationSnapshot(account *Account) GrokCredentialMutationSnapshot { + if account == nil { + return GrokCredentialMutationSnapshot{} + } + credentialsJSON := "null" + if encoded, err := json.Marshal(account.Credentials); err == nil { + credentialsJSON = string(encoded) + } + snapshot := GrokCredentialMutationSnapshot{ + CredentialsJSON: credentialsJSON, + AccessToken: strings.TrimSpace(account.GetGrokAccessToken()), + RefreshToken: strings.TrimSpace(account.GetGrokRefreshToken()), + TokenVersion: account.GetCredentialAsInt64("_token_version"), + } + if account.ProxyID != nil { + proxyID := *account.ProxyID + snapshot.ProxyID = &proxyID + } + return snapshot +} + +func (s *OpenAIGatewayService) resolveGrokCredentialCASMiss(ctx context.Context, accountID int64, snapshot GrokCredentialMutationSnapshot) (string, error) { + if ctx.Err() != nil { + return "", ctx.Err() + } + if token, ok := s.grokCredentialConcurrentlyRefreshedToken(ctx, accountID, snapshot); ok { + return token, nil + } + return "", errOAuthRefreshAccountStateChanged +} + +func (s *OpenAIGatewayService) blockGrokCredentialRuntime(account *Account, until time.Time, reason string) func() { + if s == nil || account == nil { + return func() {} + } + mu := s.openAIAccountRuntimeBlockLock(account.ID) + mu.Lock() + before, hadBefore := s.openaiAccountRuntimeBlockUntil.Load(account.ID) + installedGeneration, changed := s.blockAccountSchedulingLocked(account, until, reason) + installed, installedOK := s.openaiAccountRuntimeBlockUntil.Load(account.ID) + installedUntil, isTime := installed.(time.Time) + mu.Unlock() + if !changed || !installedOK || !isTime { + return func() {} + } + if hadBefore { + if beforeUntil, ok := before.(time.Time); ok && beforeUntil.Equal(installedUntil) { + return func() {} + } + } + return func() { + mu.Lock() + defer mu.Unlock() + generation, ok := s.openaiAccountRuntimeBlockGeneration.Load(account.ID) + if !ok || generation != installedGeneration { + return + } + current, ok := s.openaiAccountRuntimeBlockUntil.Load(account.ID) + currentUntil, isTime := current.(time.Time) + if !ok || !isTime || !currentUntil.Equal(installedUntil) { + return + } + if hadBefore { + s.openaiAccountRuntimeBlockUntil.Store(account.ID, before) + s.openaiAccountRuntimeBlockGeneration.Store(account.ID, s.openaiAccountRuntimeBlockSequence.Add(1)) + return + } + s.openaiAccountRuntimeBlockUntil.Delete(account.ID) + s.openaiAccountRuntimeBlockGeneration.Store(account.ID, s.openaiAccountRuntimeBlockSequence.Add(1)) + } +} + +func (s *OpenAIGatewayService) grokCredentialConcurrentlyRefreshedToken(ctx context.Context, accountID int64, baseline GrokCredentialMutationSnapshot) (string, bool) { + if s == nil || s.accountRepo == nil || accountID <= 0 || ctx == nil || ctx.Err() != nil { + return "", false + } + checkCtx, cancel := context.WithTimeout(ctx, time.Second) + defer cancel() + latest, err := s.accountRepo.GetByID(checkCtx, accountID) + if err != nil || latest == nil { + return "", false + } + latestSnapshot := grokCredentialMutationSnapshot(latest) + if !grokCredentialProxyIDsEqual(latestSnapshot.ProxyID, baseline.ProxyID) || + latestSnapshot.CredentialsJSON == baseline.CredentialsJSON || !latest.IsSchedulable() || + (latest.ProxyID != nil && latest.Proxy == nil) || s.isOpenAIAccountRuntimeBlocked(latest) { + return "", false + } + latestToken := strings.TrimSpace(latest.GetGrokAccessToken()) + if latestToken == "" || strings.TrimSpace(latest.GetGrokRefreshToken()) == "" { + return "", false + } + expiresAt := latest.GetCredentialAsTime("expires_at") + if expiresAt == nil || !time.Now().Before(*expiresAt) { + return "", false + } + return latestToken, true +} + +func grokCredentialProxyIDsEqual(left, right *int64) bool { + if left == nil || right == nil { + return left == nil && right == nil + } + return *left == *right +} + +func (s *OpenAIGatewayService) newGrokCredentialFailover(c *gin.Context, account *Account, class grokCredentialFailureClass) error { + if strings.TrimSpace(class.message) == "" { + class.message = "Grok OAuth credentials are unavailable" + } + appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ + Platform: PlatformGrok, + AccountID: account.ID, + Stage: string(GatewayFailureStageAccountAuth), + Scope: string(class.scope), + Reason: string(class.reason), + Kind: "credential_failover", + Message: class.message, + }) + return &UpstreamFailoverError{ + Stage: GatewayFailureStageAccountAuth, + Scope: class.scope, + Reason: class.reason, + NextAccountAction: class.action, + ClientStatusCode: http.StatusServiceUnavailable, + ClientMessage: GrokCredentialUnavailableClientMessage, + } +} diff --git a/backend/internal/service/grok_credential_failure_test.go b/backend/internal/service/grok_credential_failure_test.go new file mode 100644 index 0000000000..c32559ea12 --- /dev/null +++ b/backend/internal/service/grok_credential_failure_test.go @@ -0,0 +1,1622 @@ +//go:build unit + +package service + +import ( + "context" + "errors" + "net/http" + "net/http/httptest" + "sync" + "testing" + "time" + + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +type grokCredentialPersistingRepo struct { + *tokenRefreshAccountRepo +} + +func (r *grokCredentialPersistingRepo) SetError(ctx context.Context, id int64, message string) error { + if err := ctx.Err(); err != nil { + return err + } + if err := r.tokenRefreshAccountRepo.SetError(ctx, id, message); err != nil { + return err + } + if account := r.accountsByID[id]; account != nil { + account.Status = StatusError + account.Schedulable = false + account.ErrorMessage = message + } + return nil +} + +type grokCredentialProxyRepoStub struct { + ProxyRepository + proxy *Proxy + err error +} + +func (r *grokCredentialProxyRepoStub) GetByID(context.Context, int64) (*Proxy, error) { + return r.proxy, r.err +} + +type grokCredentialBlockingRepo struct { + *tokenRefreshAccountRepo + setErrorStarted chan struct{} + setTempStarted chan struct{} + onceError sync.Once + onceTemp sync.Once +} + +type grokCredentialCommitThenCancelRepo struct { + *tokenRefreshAccountRepo + returnErr error +} + +type grokCredentialUncommittedDeadlineRepo struct { + *tokenRefreshAccountRepo +} + +func (r *grokCredentialUncommittedDeadlineRepo) SetGrokCredentialErrorIfMatch( + context.Context, + int64, + GrokCredentialMutationSnapshot, + string, +) (bool, error) { + return false, context.DeadlineExceeded +} + +func (r *grokCredentialUncommittedDeadlineRepo) SetGrokCredentialTempUnschedulableIfMatch( + context.Context, + int64, + GrokCredentialMutationSnapshot, + time.Time, + string, +) (bool, error) { + return false, context.DeadlineExceeded +} + +func (r *grokCredentialCommitThenCancelRepo) SetGrokCredentialErrorIfMatch( + ctx context.Context, + id int64, + _ GrokCredentialMutationSnapshot, + reason string, +) (bool, error) { + account := r.accountsByID[id] + account.Status = StatusError + account.Schedulable = false + account.ErrorMessage = reason + if r.returnErr != nil { + return false, r.returnErr + } + <-ctx.Done() + return false, ctx.Err() +} + +func (r *grokCredentialCommitThenCancelRepo) SetGrokCredentialTempUnschedulableIfMatch( + ctx context.Context, + id int64, + _ GrokCredentialMutationSnapshot, + until time.Time, + reason string, +) (bool, error) { + account := r.accountsByID[id] + account.TempUnschedulableUntil = &until + account.TempUnschedulableReason = reason + if r.returnErr != nil { + return false, r.returnErr + } + <-ctx.Done() + return false, ctx.Err() +} + +func (r *grokCredentialBlockingRepo) SetError(ctx context.Context, _ int64, _ string) error { + r.onceError.Do(func() { close(r.setErrorStarted) }) + <-ctx.Done() + return ctx.Err() +} + +func (r *grokCredentialBlockingRepo) SetTempUnschedulable(ctx context.Context, _ int64, _ time.Time, _ string) error { + r.onceTemp.Do(func() { close(r.setTempStarted) }) + <-ctx.Done() + return ctx.Err() +} + +func (r *grokCredentialBlockingRepo) SetGrokCredentialErrorIfMatch( + ctx context.Context, + _ int64, + _ GrokCredentialMutationSnapshot, + _ string, +) (bool, error) { + r.onceError.Do(func() { close(r.setErrorStarted) }) + <-ctx.Done() + return false, ctx.Err() +} + +func (r *grokCredentialBlockingRepo) SetGrokCredentialTempUnschedulableIfMatch( + ctx context.Context, + _ int64, + _ GrokCredentialMutationSnapshot, + _ time.Time, + _ string, +) (bool, error) { + r.onceTemp.Do(func() { close(r.setTempStarted) }) + <-ctx.Done() + return false, ctx.Err() +} + +type grokCredentialBlockingCache struct { + GrokTokenCache + deleteStarted chan struct{} + releaseDelete chan struct{} + once sync.Once + mu sync.Mutex + deleted bool +} + +type grokCredentialSequencedRepo struct { + *tokenRefreshAccountRepo + mu sync.Mutex + latest *Account + getCall int +} + +type grokCredentialRereadFailureRepo struct { + *tokenRefreshAccountRepo + account *Account + err error +} + +type grokCredentialCountingRefresher struct { + refreshCalls int +} + +func (r *grokCredentialCountingRefresher) CacheKey(account *Account) string { + return GrokTokenCacheKey(account) +} + +func (r *grokCredentialCountingRefresher) CanRefresh(*Account) bool { return true } + +func (r *grokCredentialCountingRefresher) NeedsRefresh(*Account, time.Duration) bool { return true } + +func (r *grokCredentialCountingRefresher) Refresh(context.Context, *Account) (map[string]any, error) { + r.refreshCalls++ + return map[string]any{"access_token": "must-not-be-used"}, nil +} + +func (r *grokCredentialRereadFailureRepo) GetByID(context.Context, int64) (*Account, error) { + return r.account, r.err +} + +func (r *grokCredentialSequencedRepo) GetByID(ctx context.Context, id int64) (*Account, error) { + r.mu.Lock() + defer r.mu.Unlock() + r.getCall++ + if r.getCall > 1 && r.latest != nil { + return r.latest, nil + } + return r.tokenRefreshAccountRepo.GetByID(ctx, id) +} + +func (c *grokCredentialBlockingCache) DeleteAccessToken(ctx context.Context, _ string) error { + c.once.Do(func() { close(c.deleteStarted) }) + select { + case <-c.releaseDelete: + c.mu.Lock() + c.deleted = true + c.mu.Unlock() + return nil + case <-ctx.Done(): + return ctx.Err() + } +} + +func (c *grokCredentialBlockingCache) wasDeleted() bool { + c.mu.Lock() + defer c.mu.Unlock() + return c.deleted +} + +func TestUpstreamFailoverErrorNextAccountActionPreservesLegacyRetry(t *testing.T) { + t.Parallel() + + require.True(t, (&UpstreamFailoverError{}).ShouldRetryNextAccount()) + require.True(t, (&UpstreamFailoverError{NextAccountAction: NextAccountRetry}).ShouldRetryNextAccount()) + require.False(t, (&UpstreamFailoverError{NextAccountAction: NextAccountStop}).ShouldRetryNextAccount()) +} + +func TestGetRequestCredentialMapsPermanentGrokOAuthFailureAndRedactsSecrets(t *testing.T) { + gin.SetMode(gin.TestMode) + account := expiredGrokOAuthAccountForCredentialTest(701) + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{account.ID: account} + cache := &grokTokenCacheForProviderTest{lockResult: true} + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), &tokenRefresherStub{ + err: infraerrors.New(http.StatusBadGateway, "GROK_OAUTH_TOKEN_REFRESH_FAILED", "invalid_grant access_token=leaked-access refresh_token=leaked-refresh"), + }) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + + token, kind, err := svc.getRequestCredential(context.Background(), c, account) + require.Error(t, err) + require.Empty(t, token) + require.Empty(t, kind) + + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, GatewayFailureStageAccountAuth, failoverErr.Stage) + require.Equal(t, GatewayFailureScopeAccount, failoverErr.Scope) + require.Equal(t, GrokCredentialReasonRevoked, failoverErr.Reason) + require.True(t, failoverErr.ShouldRetryNextAccount()) + require.Equal(t, 0, failoverErr.StatusCode) + require.Equal(t, http.StatusServiceUnavailable, failoverErr.ClientStatusCode) + require.NotContains(t, err.Error(), "leaked-access") + require.NotContains(t, err.Error(), "leaked-refresh") + + require.Equal(t, 1, repo.setErrorCalls) + require.Zero(t, repo.setTempUnschedCalls) + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) + require.Equal(t, []string{GrokTokenCacheKey(account)}, cache.deletedKeys) + require.NotContains(t, repo.lastErrorMessage, "leaked-access") + require.NotContains(t, repo.lastErrorMessage, "leaked-refresh") + + rawEvents, ok := c.Get(OpsUpstreamErrorsKey) + require.True(t, ok) + events, ok := rawEvents.([]*OpsUpstreamErrorEvent) + require.True(t, ok) + require.Len(t, events, 1) + require.Equal(t, string(GatewayFailureStageAccountAuth), events[0].Stage) + require.Equal(t, string(GatewayFailureScopeAccount), events[0].Scope) + require.Equal(t, string(GrokCredentialReasonRevoked), events[0].Reason) + require.Zero(t, events[0].UpstreamStatusCode) + require.NotContains(t, events[0].Message, "leaked-access") + require.NotContains(t, events[0].Message, "leaked-refresh") +} + +func TestGetRequestCredentialPermanentMappingsPersistAndInvalidate(t *testing.T) { + gin.SetMode(gin.TestMode) + tests := []struct { + name string + prepare func(*Account) + refreshErr error + wantReason GatewayFailureReason + cachedToken string + }{ + { + name: "missing refresh credential", + prepare: func(account *Account) { + delete(account.Credentials, "refresh_token") + }, + wantReason: GrokCredentialReasonMissing, + }, + { + name: "missing access credential", + prepare: func(account *Account) { + delete(account.Credentials, "access_token") + account.Credentials["expires_at"] = time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339) + }, + wantReason: GrokCredentialReasonMissing, + cachedToken: "stale-cached-access", + }, + { + name: "explicit entitlement action required", + prepare: func(*Account) {}, + refreshErr: infraerrors.New(http.StatusForbidden, "GROK_OAUTH_ENTITLEMENT_DENIED", "access_denied"), + wantReason: GrokCredentialReasonEntitlement, + }, + } + + for index, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(int64(720 + index)) + account.Status = StatusActive + account.Schedulable = true + tt.prepare(account) + baseRepo := &tokenRefreshAccountRepo{} + baseRepo.accountsByID = map[int64]*Account{account.ID: account} + repo := &grokCredentialPersistingRepo{tokenRefreshAccountRepo: baseRepo} + cache := &grokTokenCacheForProviderTest{lockResult: true, token: tt.cachedToken} + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), &tokenRefresherStub{err: tt.refreshErr}) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + + _, _, err := svc.getRequestCredential(context.Background(), c, account) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, tt.wantReason, failoverErr.Reason) + require.Equal(t, GatewayFailureScopeAccount, failoverErr.Scope) + require.Equal(t, 1, baseRepo.setErrorCalls) + require.Equal(t, StatusError, account.Status) + require.False(t, account.Schedulable) + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) + require.Equal(t, []string{GrokTokenCacheKey(account)}, cache.deletedKeys) + }) + } +} + +func TestGetRequestCredentialMissingAccessNeverRefreshesAndPermanentlyFailsOver(t *testing.T) { + gin.SetMode(gin.TestMode) + tests := []struct { + name string + expiresAt *time.Time + }{ + {name: "expiry missing"}, + {name: "expired", expiresAt: func() *time.Time { value := time.Now().Add(-time.Minute); return &value }()}, + {name: "near expiry", expiresAt: func() *time.Time { value := time.Now().Add(30 * time.Minute); return &value }()}, + } + + for index, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(int64(760 + index)) + account.Schedulable = true + delete(account.Credentials, "access_token") + if tt.expiresAt == nil { + delete(account.Credentials, "expires_at") + } else { + account.Credentials["expires_at"] = tt.expiresAt.UTC().Format(time.RFC3339) + } + baseRepo := &tokenRefreshAccountRepo{} + baseRepo.accountsByID = map[int64]*Account{account.ID: account} + repo := &grokCredentialPersistingRepo{tokenRefreshAccountRepo: baseRepo} + cache := &grokTokenCacheForProviderTest{lockResult: true, token: "stale-cache-must-not-win"} + refresher := &grokCredentialCountingRefresher{} + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), refresher) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + + token, kind, err := svc.getRequestCredential(context.Background(), c, account) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Empty(t, token) + require.Empty(t, kind) + require.Equal(t, GrokCredentialReasonMissing, failoverErr.Reason) + require.Equal(t, NextAccountRetry, failoverErr.NextAccountAction) + require.Zero(t, refresher.refreshCalls, "structurally missing access credentials must not reach the token endpoint") + require.Equal(t, 1, baseRepo.setErrorCalls) + require.Equal(t, StatusError, account.Status) + require.False(t, account.Schedulable) + require.Equal(t, []string{GrokTokenCacheKey(account)}, cache.deletedKeys) + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) + }) + } +} + +func TestGetRequestCredentialWarmCachedAccessWithMissingRefreshPermanentlyFailsOver(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(764) + account.Credentials["access_token"] = "valid-access" + account.Credentials["expires_at"] = time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339) + delete(account.Credentials, "refresh_token") + baseRepo := &tokenRefreshAccountRepo{} + baseRepo.accountsByID = map[int64]*Account{account.ID: account} + repo := &grokCredentialPersistingRepo{tokenRefreshAccountRepo: baseRepo} + cache := &grokTokenCacheForProviderTest{lockResult: true, token: "valid-access"} + refresher := &grokCredentialCountingRefresher{} + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), refresher) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + + _, _, err := svc.getRequestCredential(context.Background(), c, account) + + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, GrokCredentialReasonMissing, failoverErr.Reason) + require.Equal(t, NextAccountRetry, failoverErr.NextAccountAction) + require.Zero(t, refresher.refreshCalls) + require.Equal(t, 1, baseRepo.setErrorCalls) + require.Equal(t, []string{GrokTokenCacheKey(account)}, cache.deletedKeys) + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) +} + +func TestGetRequestCredentialMapsTransientAndProviderFailuresSeparately(t *testing.T) { + gin.SetMode(gin.TestMode) + + t.Run("account transient temporarily unschedules", func(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(702) + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{account.ID: account} + cache := &grokTokenCacheForProviderTest{lockResult: true} + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), &tokenRefresherStub{err: errors.New("temporary refresh transport failure")}) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + + _, _, err := svc.getRequestCredential(context.Background(), c, account) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, GatewayFailureScopeAccount, failoverErr.Scope) + require.Equal(t, GrokCredentialReasonRefreshTransient, failoverErr.Reason) + require.True(t, failoverErr.ShouldRetryNextAccount()) + require.Zero(t, repo.setErrorCalls) + require.Equal(t, 1, repo.setTempUnschedCalls) + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) + }) + + t.Run("shared provider configuration stops without mutation", func(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(703) + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{account.ID: account} + provider := NewGrokTokenProvider(repo, nil) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + + _, _, err := svc.getRequestCredential(context.Background(), c, account) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, GatewayFailureScopeProvider, failoverErr.Scope) + require.Equal(t, NextAccountStop, failoverErr.NextAccountAction) + require.Zero(t, repo.setErrorCalls) + require.Zero(t, repo.setTempUnschedCalls) + require.False(t, svc.isOpenAIAccountRuntimeBlocked(account)) + }) + + t.Run("account reread failures preserve shared versus missing-row scope", func(t *testing.T) { + for _, tt := range []struct { + name string + account *Account + err error + }{ + {name: "repository error", err: errors.New("database temporarily unavailable")}, + {name: "missing row"}, + } { + t.Run(tt.name, func(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(712) + baseRepo := &tokenRefreshAccountRepo{} + baseRepo.accountsByID = map[int64]*Account{account.ID: account} + repo := &grokCredentialRereadFailureRepo{tokenRefreshAccountRepo: baseRepo, account: tt.account, err: tt.err} + cache := &grokTokenCacheForProviderTest{lockResult: true} + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), NewGrokTokenRefresher(nil)) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + + _, _, err := svc.getRequestCredential(context.Background(), c, account) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + if tt.err != nil { + require.Equal(t, GatewayFailureScopeProvider, failoverErr.Scope) + require.Equal(t, GrokCredentialReasonProviderDown, failoverErr.Reason) + require.Equal(t, NextAccountStop, failoverErr.NextAccountAction) + } else { + require.Equal(t, GatewayFailureScopeAccount, failoverErr.Scope) + require.Equal(t, GrokCredentialReasonAccountChanged, failoverErr.Reason) + require.Equal(t, NextAccountRetry, failoverErr.NextAccountAction) + } + require.Zero(t, baseRepo.setErrorCalls) + require.Zero(t, baseRepo.setTempUnschedCalls) + require.Empty(t, cache.deletedKeys) + require.False(t, svc.isOpenAIAccountRuntimeBlocked(account)) + }) + } + }) + + t.Run("fresh account eligibility changes retry without mutating stale state", func(t *testing.T) { + for _, tt := range []struct { + name string + mutate func(*Account) + }{ + { + name: "account disabled", + mutate: func(account *Account) { + account.Status = StatusDisabled + }, + }, + { + name: "account converted", + mutate: func(account *Account) { + account.Type = AccountTypeUpstream + }, + }, + { + name: "account manually unschedulable", + mutate: func(account *Account) { + account.Schedulable = false + }, + }, + { + name: "account temporarily unschedulable", + mutate: func(account *Account) { + until := time.Now().Add(time.Minute) + account.TempUnschedulableUntil = &until + }, + }, + } { + t.Run(tt.name, func(t *testing.T) { + staleAccount := expiredGrokOAuthAccountForCredentialTest(713) + freshAccount := *staleAccount + freshAccount.Credentials = shallowCopyMap(staleAccount.Credentials) + tt.mutate(&freshAccount) + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{staleAccount.ID: &freshAccount} + cache := &grokTokenCacheForProviderTest{lockResult: true} + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), NewGrokTokenRefresher(nil)) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + + _, _, err := svc.getRequestCredential(context.Background(), c, staleAccount) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, GatewayFailureScopeAccount, failoverErr.Scope) + require.Equal(t, GrokCredentialReasonAccountChanged, failoverErr.Reason) + require.Equal(t, NextAccountRetry, failoverErr.NextAccountAction) + require.Zero(t, repo.setErrorCalls) + require.Zero(t, repo.setTempUnschedCalls) + require.Empty(t, cache.deletedKeys) + require.False(t, svc.isOpenAIAccountRuntimeBlocked(staleAccount)) + }) + } + }) + + t.Run("fresh missing refresh credential permanently blocks the account", func(t *testing.T) { + staleAccount := expiredGrokOAuthAccountForCredentialTest(714) + staleAccount.Schedulable = true + freshAccount := *staleAccount + freshAccount.Credentials = shallowCopyMap(staleAccount.Credentials) + delete(freshAccount.Credentials, "refresh_token") + baseRepo := &tokenRefreshAccountRepo{} + baseRepo.accountsByID = map[int64]*Account{staleAccount.ID: &freshAccount} + repo := &grokCredentialPersistingRepo{tokenRefreshAccountRepo: baseRepo} + cache := &grokTokenCacheForProviderTest{lockResult: true} + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), NewGrokTokenRefresher(nil)) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + + _, _, err := svc.getRequestCredential(context.Background(), c, staleAccount) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, GatewayFailureScopeAccount, failoverErr.Scope) + require.Equal(t, GrokCredentialReasonMissing, failoverErr.Reason) + require.Equal(t, NextAccountRetry, failoverErr.NextAccountAction) + require.Equal(t, 1, baseRepo.setErrorCalls) + require.Zero(t, baseRepo.setTempUnschedCalls) + require.Equal(t, StatusError, freshAccount.Status) + require.False(t, freshAccount.Schedulable) + require.Equal(t, []string{GrokTokenCacheKey(staleAccount)}, cache.deletedKeys) + require.True(t, svc.isOpenAIAccountRuntimeBlocked(staleAccount)) + }) + + t.Run("refresh added after locked structural failure wins conditional mutation", func(t *testing.T) { + staleAccount := expiredGrokOAuthAccountForCredentialTest(717) + freshAccount := *staleAccount + freshAccount.Credentials = shallowCopyMap(staleAccount.Credentials) + delete(freshAccount.Credentials, "refresh_token") + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{staleAccount.ID: &freshAccount} + repo.beforeConditionalState = func() { + repaired := freshAccount + repaired.Credentials = shallowCopyMap(freshAccount.Credentials) + repaired.Credentials["refresh_token"] = "repaired-refresh-token" + repaired.Credentials["expires_at"] = time.Now().Add(time.Hour).UTC().Format(time.RFC3339) + repo.accountsByID[staleAccount.ID] = &repaired + } + cache := &grokTokenCacheForProviderTest{lockResult: true} + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), NewGrokTokenRefresher(nil)) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + + token, kind, err := svc.getRequestCredential(context.Background(), c, staleAccount) + + require.NoError(t, err) + require.Equal(t, "expired-access-token", token) + require.Equal(t, "oauth", kind) + require.Zero(t, repo.setErrorCalls) + require.Zero(t, repo.setTempUnschedCalls) + require.Empty(t, cache.deletedKeys) + require.False(t, svc.isOpenAIAccountRuntimeBlocked(staleAccount)) + }) + + t.Run("expiry-only repair wins full credential fingerprint CAS", func(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(718) + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{account.ID: account} + repo.beforeConditionalState = func() { + repaired := *account + repaired.Credentials = shallowCopyMap(account.Credentials) + repaired.Credentials["expires_at"] = time.Now().Add(time.Hour).UTC().Format(time.RFC3339) + repo.accountsByID[account.ID] = &repaired + } + cache := &grokTokenCacheForProviderTest{lockResult: true} + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), &tokenRefresherStub{ + err: infraerrors.New(http.StatusBadGateway, "GROK_OAUTH_TOKEN_REFRESH_FAILED", "invalid_grant"), + }) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + + token, kind, err := svc.getRequestCredential(context.Background(), c, account) + + require.NoError(t, err) + require.Equal(t, "expired-access-token", token) + require.Equal(t, "oauth", kind) + require.Zero(t, repo.setErrorCalls) + require.Empty(t, cache.deletedKeys) + require.False(t, svc.isOpenAIAccountRuntimeBlocked(account)) + }) + + t.Run("generic token endpoint 403 stops as shared provider failure", func(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(708) + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{account.ID: account} + cache := &grokTokenCacheForProviderTest{lockResult: true} + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), &tokenRefresherStub{ + err: infraerrors.New(http.StatusBadGateway, "GROK_OAUTH_TOKEN_REFRESH_FAILED", "token refresh failed: status 403, body: forbidden"), + }) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + + _, _, err := svc.getRequestCredential(context.Background(), c, account) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, GatewayFailureScopeProvider, failoverErr.Scope) + require.Equal(t, GrokCredentialReasonProviderDown, failoverErr.Reason) + require.Equal(t, NextAccountStop, failoverErr.NextAccountAction) + require.Zero(t, repo.setErrorCalls) + require.Zero(t, repo.setTempUnschedCalls) + require.Empty(t, cache.deletedKeys) + require.False(t, svc.isOpenAIAccountRuntimeBlocked(account)) + }) + + t.Run("account proxy generic 403 remains bounded account transient", func(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(711) + proxyID := int64(43) + account.ProxyID = &proxyID + account.Proxy = &Proxy{} + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{account.ID: account} + cache := &grokTokenCacheForProviderTest{lockResult: true} + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), &tokenRefresherStub{ + err: infraerrors.New(http.StatusBadGateway, "GROK_OAUTH_TOKEN_REFRESH_FAILED", "token refresh failed: status 403, body: forbidden"), + }) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + + _, _, err := svc.getRequestCredential(context.Background(), c, account) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, GatewayFailureScopeAccount, failoverErr.Scope) + require.Equal(t, GrokCredentialReasonRefreshTransient, failoverErr.Reason) + require.Equal(t, NextAccountRetry, failoverErr.NextAccountAction) + require.Zero(t, repo.setErrorCalls) + require.Equal(t, 1, repo.setTempUnschedCalls) + require.Empty(t, cache.deletedKeys) + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) + }) + + t.Run("proxy repository read failure stops without account mutation", func(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(709) + proxyID := int64(41) + account.ProxyID = &proxyID + account.Proxy = &Proxy{} + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{account.ID: account} + cache := &grokTokenCacheForProviderTest{lockResult: true} + oauthSvc := NewGrokOAuthService(&grokCredentialProxyRepoStub{err: errors.New("database temporarily unavailable")}, &grokOAuthClientStub{}) + defer oauthSvc.Stop() + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), NewGrokTokenRefresher(oauthSvc)) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + + _, _, err := svc.getRequestCredential(context.Background(), c, account) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, GatewayFailureScopeProvider, failoverErr.Scope) + require.Equal(t, GrokCredentialReasonProviderDown, failoverErr.Reason) + require.Equal(t, NextAccountStop, failoverErr.NextAccountAction) + require.Zero(t, repo.setErrorCalls) + require.Zero(t, repo.setTempUnschedCalls) + require.Empty(t, cache.deletedKeys) + require.False(t, svc.isOpenAIAccountRuntimeBlocked(account)) + }) + + t.Run("structurally missing configured proxy permanently blocks only that account", func(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(710) + account.Status = StatusActive + account.Schedulable = true + proxyID := int64(42) + account.ProxyID = &proxyID + baseRepo := &tokenRefreshAccountRepo{} + baseRepo.accountsByID = map[int64]*Account{account.ID: account} + repo := &grokCredentialPersistingRepo{tokenRefreshAccountRepo: baseRepo} + cache := &grokTokenCacheForProviderTest{lockResult: true} + oauthSvc := NewGrokOAuthService(&grokCredentialProxyRepoStub{err: ErrProxyNotFound}, &grokOAuthClientStub{}) + defer oauthSvc.Stop() + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), NewGrokTokenRefresher(oauthSvc)) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + + _, _, err := svc.getRequestCredential(context.Background(), c, account) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, GatewayFailureScopeAccount, failoverErr.Scope) + require.Equal(t, GrokCredentialReasonProxyInvalid, failoverErr.Reason) + require.Equal(t, NextAccountRetry, failoverErr.NextAccountAction) + require.Equal(t, 1, baseRepo.setErrorCalls) + require.Equal(t, StatusError, account.Status) + require.False(t, account.Schedulable) + require.Equal(t, []string{GrokTokenCacheKey(account)}, cache.deletedKeys) + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) + }) +} + +func TestGetRequestCredentialRuntimeBlockWinsBeforeWarmTokenCache(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(716) + account.Credentials["access_token"] = "valid-access" + account.Credentials["expires_at"] = time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339) + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{account.ID: account} + cache := &grokTokenCacheForProviderTest{token: "valid-access"} + provider := NewGrokTokenProvider(repo, cache) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + svc.BlockAccountScheduling(account, time.Now().Add(time.Minute), "independent") + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + + _, _, err := svc.getRequestCredential(context.Background(), c, account) + + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, GrokCredentialReasonAccountChanged, failoverErr.Reason) + require.Equal(t, NextAccountRetry, failoverErr.NextAccountAction) + require.Zero(t, cache.getCalls) + require.Zero(t, repo.setErrorCalls) + require.Zero(t, repo.setTempUnschedCalls) +} + +func TestGetRequestCredentialWarmCachedAccessWithMissingConfiguredProxyPermanentlyFailsOver(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(715) + account.Credentials["access_token"] = "valid-access" + account.Credentials["expires_at"] = time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339) + proxyID := int64(44) + account.ProxyID = &proxyID + account.Proxy = nil + baseRepo := &tokenRefreshAccountRepo{} + baseRepo.accountsByID = map[int64]*Account{account.ID: account} + repo := &grokCredentialPersistingRepo{tokenRefreshAccountRepo: baseRepo} + cache := &grokTokenCacheForProviderTest{lockResult: true, token: "valid-access"} + refresher := &grokCredentialCountingRefresher{} + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), refresher) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + + _, _, err := svc.getRequestCredential(context.Background(), c, account) + + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, GrokCredentialReasonProxyInvalid, failoverErr.Reason) + require.Equal(t, NextAccountRetry, failoverErr.NextAccountAction) + require.Zero(t, refresher.refreshCalls) + require.Equal(t, 1, baseRepo.setErrorCalls) + require.Equal(t, []string{GrokTokenCacheKey(account)}, cache.deletedKeys) + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) +} + +func TestGetRequestCredentialCancellationAndBudgetDoNotMutateAccount(t *testing.T) { + gin.SetMode(gin.TestMode) + account := expiredGrokOAuthAccountForCredentialTest(704) + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{account.ID: account} + provider := NewGrokTokenProvider(repo, nil) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + + t.Run("parent cancellation is returned directly", func(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + + _, _, err := svc.getRequestCredential(ctx, c, account) + require.ErrorIs(t, err, context.Canceled) + var failoverErr *UpstreamFailoverError + require.False(t, errors.As(err, &failoverErr)) + }) + + t.Run("request credential budget stops safely", func(t *testing.T) { + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + c.Set(grokCredentialFailoverDeadlineKey, time.Now().Add(-time.Second)) + + _, _, err := svc.getRequestCredential(context.Background(), c, account) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, GatewayFailureScopeRequest, failoverErr.Scope) + require.Equal(t, GrokCredentialReasonFailoverTimeout, failoverErr.Reason) + require.False(t, failoverErr.ShouldRetryNextAccount()) + }) + + require.Zero(t, repo.setErrorCalls) + require.Zero(t, repo.setTempUnschedCalls) + require.False(t, svc.isOpenAIAccountRuntimeBlocked(account)) +} + +func TestGetRequestCredentialStateMutationFailureStopsAndKeepsRuntimeBlock(t *testing.T) { + gin.SetMode(gin.TestMode) + tests := []struct { + name string + refreshErr error + configure func(*tokenRefreshAccountRepo, *grokTokenCacheForProviderTest) + wantSetError int + wantSetTemp int + wantCacheDelete int + }{ + { + name: "permanent state persistence", + refreshErr: infraerrors.New(http.StatusBadGateway, "GROK_OAUTH_TOKEN_REFRESH_FAILED", "invalid_grant"), + configure: func(repo *tokenRefreshAccountRepo, _ *grokTokenCacheForProviderTest) { + repo.setErrorErr = errors.New("database write failed") + }, + wantSetError: 1, + }, + { + name: "transient state persistence", + refreshErr: errors.New("temporary refresh transport failure"), + configure: func(repo *tokenRefreshAccountRepo, _ *grokTokenCacheForProviderTest) { + repo.setTempUnschedErr = errors.New("database write failed") + }, + wantSetTemp: 1, + }, + { + name: "permanent token cache invalidation", + refreshErr: infraerrors.New(http.StatusBadGateway, "GROK_OAUTH_TOKEN_REFRESH_FAILED", "invalid_grant"), + configure: func(_ *tokenRefreshAccountRepo, cache *grokTokenCacheForProviderTest) { + cache.deleteErr = errors.New("cache delete failed") + }, + wantSetError: 1, + wantCacheDelete: 1, + }, + } + + for index, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(int64(740 + index)) + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{account.ID: account} + cache := &grokTokenCacheForProviderTest{lockResult: true} + tt.configure(repo, cache) + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), &tokenRefresherStub{err: tt.refreshErr}) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + + _, _, err := svc.getRequestCredential(context.Background(), c, account) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, GatewayFailureScopeProvider, failoverErr.Scope) + require.Equal(t, GrokCredentialReasonStateUpdate, failoverErr.Reason) + require.Equal(t, NextAccountStop, failoverErr.NextAccountAction) + require.Equal(t, tt.wantSetError, repo.setErrorCalls) + require.Equal(t, tt.wantSetTemp, repo.setTempUnschedCalls) + require.Len(t, cache.deletedKeys, tt.wantCacheDelete) + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account), "failed mutation must retain the immediate local block") + }) + } +} + +func TestGrokCredentialMutationBoundariesHonorParentCancellation(t *testing.T) { + t.Run("blocked SetError cancellation prevents cache and runtime mutation", func(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(730) + baseRepo := &tokenRefreshAccountRepo{} + baseRepo.accountsByID = map[int64]*Account{account.ID: account} + repo := &grokCredentialBlockingRepo{ + tokenRefreshAccountRepo: baseRepo, + setErrorStarted: make(chan struct{}), + setTempStarted: make(chan struct{}), + } + cache := &grokTokenCacheForProviderTest{} + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: NewGrokTokenProvider(repo, cache)} + ctx, cancel := context.WithCancel(context.Background()) + result := make(chan error, 1) + go func() { + _, err := svc.applyGrokCredentialAccountFailure(ctx, account, grokCredentialFailureClass{ + reason: GrokCredentialReasonRevoked, permanent: true, + }) + result <- err + }() + <-repo.setErrorStarted + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account), "runtime block must precede persistent SetError") + cancel() + + require.ErrorIs(t, <-result, context.Canceled) + require.Zero(t, baseRepo.setErrorCalls) + require.Empty(t, cache.deletedKeys) + require.False(t, svc.isOpenAIAccountRuntimeBlocked(account)) + }) + + t.Run("blocked temporary unschedule cancellation prevents runtime mutation", func(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(731) + baseRepo := &tokenRefreshAccountRepo{} + baseRepo.accountsByID = map[int64]*Account{account.ID: account} + repo := &grokCredentialBlockingRepo{ + tokenRefreshAccountRepo: baseRepo, + setErrorStarted: make(chan struct{}), + setTempStarted: make(chan struct{}), + } + svc := &OpenAIGatewayService{accountRepo: repo} + ctx, cancel := context.WithCancel(context.Background()) + result := make(chan error, 1) + go func() { + _, err := svc.applyGrokCredentialAccountFailure(ctx, account, grokCredentialFailureClass{ + reason: GrokCredentialReasonRefreshTransient, transient: true, + }) + result <- err + }() + <-repo.setTempStarted + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account), "runtime block must precede temporary unscheduling") + cancel() + + require.ErrorIs(t, <-result, context.Canceled) + require.Zero(t, baseRepo.setTempUnschedCalls) + require.False(t, svc.isOpenAIAccountRuntimeBlocked(account)) + }) + + t.Run("post-commit cancellation finishes cache cleanup and retains quarantine", func(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(732) + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{account.ID: account} + cache := &grokCredentialBlockingCache{deleteStarted: make(chan struct{}), releaseDelete: make(chan struct{})} + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: NewGrokTokenProvider(repo, cache)} + ctx, cancel := context.WithCancel(context.Background()) + result := make(chan error, 1) + go func() { + _, err := svc.applyGrokCredentialAccountFailure(ctx, account, grokCredentialFailureClass{ + reason: GrokCredentialReasonRevoked, permanent: true, + }) + result <- err + }() + <-cache.deleteStarted + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account), "runtime block must precede cache invalidation") + cancel() + close(cache.releaseDelete) + + require.ErrorIs(t, <-result, context.Canceled) + require.Equal(t, 1, repo.setErrorCalls) + require.True(t, cache.wasDeleted()) + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) + }) +} + +func TestGrokCredentialMutationLockWaitHonorsCredentialBudget(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(735) + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{account.ID: account} + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: NewGrokTokenProvider(repo, &grokTokenCacheForProviderTest{})} + mutationLock := svc.grokCredentialMutationLock(account.ID) + require.NoError(t, mutationLock.Lock(context.Background())) + defer mutationLock.Unlock() + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Millisecond) + defer cancel() + + startedAt := time.Now() + token, err := svc.applyGrokCredentialAccountFailure(ctx, account, grokCredentialFailureClass{ + reason: GrokCredentialReasonRevoked, permanent: true, + }) + + require.ErrorIs(t, err, context.DeadlineExceeded) + require.Empty(t, token) + require.Less(t, time.Since(startedAt), 500*time.Millisecond) + require.Zero(t, repo.setErrorCalls) + require.Zero(t, repo.setTempUnschedCalls) + require.False(t, svc.isOpenAIAccountRuntimeBlocked(account)) +} + +func TestGetRequestCredentialBudgetBoundsBlockedConditionalMutation(t *testing.T) { + gin.SetMode(gin.TestMode) + account := expiredGrokOAuthAccountForCredentialTest(736) + baseRepo := &tokenRefreshAccountRepo{} + baseRepo.accountsByID = map[int64]*Account{account.ID: account} + repo := &grokCredentialBlockingRepo{ + tokenRefreshAccountRepo: baseRepo, + setErrorStarted: make(chan struct{}), + setTempStarted: make(chan struct{}), + } + cache := &grokTokenCacheForProviderTest{lockResult: true} + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), &tokenRefresherStub{ + err: infraerrors.New(http.StatusBadGateway, "GROK_OAUTH_TOKEN_REFRESH_FAILED", "invalid_grant"), + }) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + c.Set(grokCredentialFailoverDeadlineKey, time.Now().Add(40*time.Millisecond)) + + startedAt := time.Now() + token, kind, err := svc.getRequestCredential(context.Background(), c, account) + + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, GatewayFailureScopeRequest, failoverErr.Scope) + require.Equal(t, GrokCredentialReasonFailoverTimeout, failoverErr.Reason) + require.Equal(t, NextAccountStop, failoverErr.NextAccountAction) + require.Empty(t, token) + require.Empty(t, kind) + require.Less(t, time.Since(startedAt), 500*time.Millisecond) + require.Zero(t, baseRepo.setErrorCalls) + require.Zero(t, baseRepo.setTempUnschedCalls) + require.Empty(t, cache.deletedKeys) + require.False(t, svc.isOpenAIAccountRuntimeBlocked(account)) +} + +func TestGetRequestCredentialLockHeldTimeoutDoesNotQuarantineAccount(t *testing.T) { + gin.SetMode(gin.TestMode) + tests := []struct { + name string + buildRepo func(*Account) AccountRepository + wantScope GatewayFailureScope + wantReason GatewayFailureReason + wantAction NextAccountAction + }{ + { + name: "authoritative row unchanged", + buildRepo: func(account *Account) AccountRepository { + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{account.ID: account} + return repo + }, + wantScope: GatewayFailureScopeAccount, + wantReason: GrokCredentialReasonAccountChanged, + wantAction: NextAccountRetry, + }, + { + name: "selected account was deleted", + buildRepo: func(account *Account) AccountRepository { + base := &tokenRefreshAccountRepo{} + base.accountsByID = map[int64]*Account{account.ID: account} + return &grokCredentialRereadFailureRepo{tokenRefreshAccountRepo: base} + }, + wantScope: GatewayFailureScopeAccount, + wantReason: GrokCredentialReasonAccountChanged, + wantAction: NextAccountRetry, + }, + { + name: "shared account store unavailable", + buildRepo: func(account *Account) AccountRepository { + base := &tokenRefreshAccountRepo{} + base.accountsByID = map[int64]*Account{account.ID: account} + return &grokCredentialRereadFailureRepo{tokenRefreshAccountRepo: base, err: errors.New("database unavailable")} + }, + wantScope: GatewayFailureScopeProvider, + wantReason: GrokCredentialReasonProviderDown, + wantAction: NextAccountStop, + }, + } + + for index, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(int64(7400 + index)) + repo := tt.buildRepo(account) + cache := &grokTokenCacheForProviderTest{lockResult: false} + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), &tokenRefresherStub{}) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + + startedAt := time.Now() + _, _, err := svc.getRequestCredential(context.Background(), c, account) + + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, tt.wantScope, failoverErr.Scope) + require.Equal(t, tt.wantReason, failoverErr.Reason) + require.Equal(t, tt.wantAction, failoverErr.NextAccountAction) + require.Less(t, time.Since(startedAt), 3*time.Second) + switch countingRepo := repo.(type) { + case *tokenRefreshAccountRepo: + require.Zero(t, countingRepo.setErrorCalls) + require.Zero(t, countingRepo.setTempUnschedCalls) + case *grokCredentialRereadFailureRepo: + require.Zero(t, countingRepo.tokenRefreshAccountRepo.setErrorCalls) + require.Zero(t, countingRepo.tokenRefreshAccountRepo.setTempUnschedCalls) + } + require.Empty(t, cache.deletedKeys) + require.False(t, svc.isOpenAIAccountRuntimeBlocked(account)) + }) + } +} + +func TestGrokCredentialMutationCancellationAmbiguityConfirmsDurableCommit(t *testing.T) { + tests := []struct { + name string + class grokCredentialFailureClass + committed func(*Account) bool + }{ + { + name: "permanent quarantine", + class: grokCredentialFailureClass{reason: GrokCredentialReasonRevoked, permanent: true}, + committed: func(account *Account) bool { + return account.Status == StatusError && !account.Schedulable && account.ErrorMessage == string(GrokCredentialReasonRevoked) + }, + }, + { + name: "temporary quarantine", + class: grokCredentialFailureClass{reason: GrokCredentialReasonRefreshTransient, transient: true}, + committed: func(account *Account) bool { + return account.TempUnschedulableUntil != nil && account.TempUnschedulableReason == string(GrokCredentialReasonRefreshTransient) + }, + }, + } + + for index, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(int64(737 + index)) + baseRepo := &tokenRefreshAccountRepo{} + baseRepo.accountsByID = map[int64]*Account{account.ID: account} + repo := &grokCredentialCommitThenCancelRepo{tokenRefreshAccountRepo: baseRepo} + cache := &grokTokenCacheForProviderTest{} + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: NewGrokTokenProvider(repo, cache)} + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Millisecond) + defer cancel() + + token, err := svc.applyGrokCredentialAccountFailure(ctx, account, tt.class) + + require.ErrorIs(t, err, context.DeadlineExceeded) + require.Empty(t, token) + require.True(t, tt.committed(account), "the detached confirmation must recognize the durable mutation") + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account), "a confirmed durable quarantine must retain its runtime block") + if tt.class.permanent { + require.Equal(t, []string{GrokTokenCacheKey(account)}, cache.deletedKeys) + } else { + require.Empty(t, cache.deletedKeys) + } + }) + } +} + +func TestGrokCredentialInnerStateDeadlineAmbiguityConfirmsDurableCommit(t *testing.T) { + tests := []struct { + name string + class grokCredentialFailureClass + }{ + {name: "permanent quarantine", class: grokCredentialFailureClass{reason: GrokCredentialReasonRevoked, permanent: true}}, + {name: "temporary quarantine", class: grokCredentialFailureClass{reason: GrokCredentialReasonRefreshTransient, transient: true}}, + } + + for index, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(int64(7500 + index)) + baseRepo := &tokenRefreshAccountRepo{} + baseRepo.accountsByID = map[int64]*Account{account.ID: account} + repo := &grokCredentialCommitThenCancelRepo{ + tokenRefreshAccountRepo: baseRepo, + returnErr: context.DeadlineExceeded, + } + cache := &grokTokenCacheForProviderTest{} + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: NewGrokTokenProvider(repo, cache)} + + token, err := svc.applyGrokCredentialAccountFailure(context.Background(), account, tt.class) + + require.NoError(t, err, "the detached readback must resolve the inner timeout's commit ambiguity") + require.Empty(t, token) + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) + if tt.class.permanent { + require.Equal(t, StatusError, account.Status) + require.False(t, account.Schedulable) + require.Equal(t, []string{GrokTokenCacheKey(account)}, cache.deletedKeys) + } else { + require.NotNil(t, account.TempUnschedulableUntil) + require.Equal(t, string(GrokCredentialReasonRefreshTransient), account.TempUnschedulableReason) + require.Empty(t, cache.deletedKeys) + } + }) + } +} + +func TestGrokCredentialUnconfirmedInnerStateDeadlineStopsAndRetainsSafetyBlock(t *testing.T) { + tests := []struct { + name string + class grokCredentialFailureClass + }{ + {name: "permanent quarantine", class: grokCredentialFailureClass{reason: GrokCredentialReasonRevoked, permanent: true}}, + {name: "temporary quarantine", class: grokCredentialFailureClass{reason: GrokCredentialReasonRefreshTransient, transient: true}}, + } + + for index, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(int64(7600 + index)) + baseRepo := &tokenRefreshAccountRepo{} + baseRepo.accountsByID = map[int64]*Account{account.ID: account} + repo := &grokCredentialUncommittedDeadlineRepo{tokenRefreshAccountRepo: baseRepo} + cache := &grokTokenCacheForProviderTest{} + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: NewGrokTokenProvider(repo, cache)} + + token, err := svc.applyGrokCredentialAccountFailure(context.Background(), account, tt.class) + + require.ErrorIs(t, err, errGrokCredentialStateUpdateFailed) + require.Empty(t, token) + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account), "an unknown commit outcome must retain the local safety block") + require.Equal(t, StatusActive, account.Status) + require.True(t, account.Schedulable) + require.Nil(t, account.TempUnschedulableUntil) + require.Empty(t, cache.deletedKeys) + }) + } +} + +func TestGrokCredentialRuntimeRollbackOwnership(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(734) + t.Run("later extending block survives", func(t *testing.T) { + svc := &OpenAIGatewayService{} + until := time.Now().Add(time.Minute) + rollbackFirst := svc.blockGrokCredentialRuntime(account, until, "first") + + secondInstalled := make(chan struct{}) + go func() { + svc.BlockAccountScheduling(account, until.Add(time.Minute), "independent") + close(secondInstalled) + }() + <-secondInstalled + rollbackFirst() + + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account), + "rollback owned by the first invocation must not remove a later extending block") + }) + + t.Run("independent shorter block steals rollback ownership", func(t *testing.T) { + svc := &OpenAIGatewayService{} + until := time.Now().Add(2 * time.Minute) + rollbackFirst := svc.blockGrokCredentialRuntime(account, until, "first") + svc.BlockAccountScheduling(account, until.Add(-time.Minute), "shorter-no-op") + + rollbackFirst() + + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) + }) + + t.Run("serialized tentative rollbacks leave no block", func(t *testing.T) { + svc := &OpenAIGatewayService{} + for i := 0; i < 2; i++ { + mu := svc.grokCredentialMutationLock(account.ID) + require.NoError(t, mu.Lock(context.Background())) + rollback := svc.blockGrokCredentialRuntime(account, time.Now().Add(time.Minute), "tentative") + rollback() + mu.Unlock() + } + require.False(t, svc.isOpenAIAccountRuntimeBlocked(account)) + }) +} + +func TestGetRequestCredentialAPIKeyBypassesOAuthFailureMapping(t *testing.T) { + gin.SetMode(gin.TestMode) + account := &Account{ + ID: 705, + Platform: PlatformGrok, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "api_key": "third-party-key", + "base_url": "https://grok.example.test/v1", + }, + } + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + svc := &OpenAIGatewayService{} + + token, kind, err := svc.getRequestCredential(context.Background(), c, account) + require.NoError(t, err) + require.Equal(t, "third-party-key", token) + require.Equal(t, "apikey", kind) + _, hasEvents := c.Get(OpsUpstreamErrorsKey) + require.False(t, hasEvents) +} + +func TestPermanentCredentialFailureDoesNotDisableConcurrentlyRefreshedAccount(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(707) + latest := *account + latest.Credentials = shallowCopyMap(account.Credentials) + latest.Credentials["access_token"] = "fresh-access-token" + latest.Credentials["refresh_token"] = "rotated-refresh-token" + latest.Credentials["expires_at"] = time.Now().Add(time.Hour).UTC().Format(time.RFC3339) + latest.Credentials["_token_version"] = time.Now().UnixMilli() + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{account.ID: &latest} + cache := &grokTokenCacheForProviderTest{} + svc := &OpenAIGatewayService{ + accountRepo: repo, + grokTokenProvider: NewGrokTokenProvider(repo, cache), + } + + _, mutationErr := svc.applyGrokCredentialAccountFailure(context.Background(), account, grokCredentialFailureClass{ + scope: GatewayFailureScopeAccount, + reason: GrokCredentialReasonRevoked, + action: NextAccountRetry, + permanent: true, + }) + + require.NoError(t, mutationErr) + require.Zero(t, repo.setErrorCalls) + require.Empty(t, cache.deletedKeys) + require.False(t, svc.isOpenAIAccountRuntimeBlocked(account)) +} + +func TestCredentialFailureConditionalMutationLosesToConcurrentRefresh(t *testing.T) { + for index, tt := range []struct { + name string + refreshErr error + }{ + {name: "permanent", refreshErr: infraerrors.New(http.StatusBadGateway, "GROK_OAUTH_TOKEN_REFRESH_FAILED", "invalid_grant")}, + {name: "transient", refreshErr: errors.New("temporary refresh transport failure")}, + } { + t.Run(tt.name, func(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(int64(770 + index)) + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{account.ID: account} + repo.beforeConditionalState = func() { + fresh := *account + fresh.Credentials = shallowCopyMap(account.Credentials) + fresh.Credentials["access_token"] = "refresh-won-token" + fresh.Credentials["refresh_token"] = "refresh-won-refresh" + fresh.Credentials["expires_at"] = time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339) + fresh.Credentials["_token_version"] = time.Now().UnixMilli() + repo.accountsByID[account.ID] = &fresh + } + cache := &grokTokenCacheForProviderTest{lockResult: true} + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), &tokenRefresherStub{err: tt.refreshErr}) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + + token, kind, err := svc.getRequestCredential(context.Background(), c, account) + + require.NoError(t, err) + require.Equal(t, "refresh-won-token", token) + require.Equal(t, "oauth", kind) + require.Zero(t, repo.setErrorCalls) + require.Zero(t, repo.setTempUnschedCalls) + require.Empty(t, cache.deletedKeys) + require.False(t, svc.isOpenAIAccountRuntimeBlocked(account)) + }) + } +} + +func TestCredentialFailureConditionalMutationLosesToConcurrentProxyRepair(t *testing.T) { + for index, tt := range []struct { + name string + refreshErr error + }{ + {name: "permanent", refreshErr: infraerrors.New(http.StatusBadGateway, "GROK_OAUTH_TOKEN_REFRESH_FAILED", "invalid_grant")}, + {name: "transient", refreshErr: errors.New("temporary refresh transport failure")}, + } { + t.Run(tt.name, func(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(int64(780 + index)) + oldProxyID := int64(10) + account.ProxyID = &oldProxyID + account.Proxy = &Proxy{} + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{account.ID: account} + repo.beforeConditionalState = func() { + fresh := *account + fresh.Credentials = shallowCopyMap(account.Credentials) + repairedProxyID := int64(11) + fresh.ProxyID = &repairedProxyID + repo.accountsByID[account.ID] = &fresh + } + cache := &grokTokenCacheForProviderTest{lockResult: true} + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), &tokenRefresherStub{err: tt.refreshErr}) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + + _, _, err := svc.getRequestCredential(context.Background(), c, account) + + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, GrokCredentialReasonAccountChanged, failoverErr.Reason) + require.Equal(t, NextAccountRetry, failoverErr.NextAccountAction) + require.Zero(t, repo.setErrorCalls) + require.Zero(t, repo.setTempUnschedCalls) + require.Empty(t, cache.deletedKeys) + require.False(t, svc.isOpenAIAccountRuntimeBlocked(account)) + }) + } +} + +func TestCredentialFailureConditionalMutationLosesToSameIDProxyRestoration(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(790) + proxyID := int64(10) + account.ProxyID = &proxyID + account.Proxy = nil + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{account.ID: account} + repo.beforeConditionalState = func() { + account.Proxy = &Proxy{ID: proxyID} + } + cache := &grokTokenCacheForProviderTest{lockResult: true} + provider := NewGrokTokenProvider(repo, cache) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + + _, _, err := svc.getRequestCredential(context.Background(), c, account) + + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, GrokCredentialReasonAccountChanged, failoverErr.Reason) + require.Equal(t, NextAccountRetry, failoverErr.NextAccountAction) + require.Zero(t, repo.setErrorCalls) + require.Zero(t, repo.setTempUnschedCalls) + require.Empty(t, cache.deletedKeys) + require.False(t, svc.isOpenAIAccountRuntimeBlocked(account)) +} + +func TestCredentialFailureConditionalMutationLosesToConcurrentUnschedulableState(t *testing.T) { + future := time.Now().Add(time.Hour) + states := []struct { + name string + mutate func(*Account) + }{ + {name: "admin schedulable false", mutate: func(account *Account) { account.Schedulable = false }}, + {name: "temporary cooldown", mutate: func(account *Account) { account.TempUnschedulableUntil = &future }}, + {name: "rate limit cooldown", mutate: func(account *Account) { account.RateLimitResetAt = &future }}, + {name: "overload cooldown", mutate: func(account *Account) { account.OverloadUntil = &future }}, + } + classes := []struct { + name string + class grokCredentialFailureClass + }{ + {name: "permanent", class: grokCredentialFailureClass{reason: GrokCredentialReasonRevoked, permanent: true}}, + {name: "transient", class: grokCredentialFailureClass{reason: GrokCredentialReasonRefreshTransient, transient: true}}, + } + + for classIndex, classCase := range classes { + for stateIndex, stateCase := range states { + t.Run(classCase.name+"/"+stateCase.name, func(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(int64(791 + classIndex*10 + stateIndex)) + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{account.ID: account} + repo.beforeConditionalState = func() { stateCase.mutate(account) } + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: NewGrokTokenProvider(repo, &grokTokenCacheForProviderTest{})} + + token, err := svc.applyGrokCredentialAccountFailure(context.Background(), account, classCase.class) + + require.ErrorIs(t, err, errOAuthRefreshAccountStateChanged) + require.Empty(t, token) + require.Zero(t, repo.setErrorCalls) + require.Zero(t, repo.setTempUnschedCalls) + require.False(t, svc.isOpenAIAccountRuntimeBlocked(account)) + }) + } + } +} + +func TestCredentialFailureCASMissDoesNotRecoverIneligibleLatestCredential(t *testing.T) { + future := time.Now().Add(time.Hour) + states := []struct { + name string + mutate func(*OpenAIGatewayService, *Account) + wantRuntimeBlocked bool + }{ + {name: "disabled", mutate: func(_ *OpenAIGatewayService, account *Account) { account.Status = StatusDisabled }}, + {name: "not schedulable", mutate: func(_ *OpenAIGatewayService, account *Account) { account.Schedulable = false }}, + {name: "temporarily unschedulable", mutate: func(_ *OpenAIGatewayService, account *Account) { account.TempUnschedulableUntil = &future }}, + {name: "rate limited", mutate: func(_ *OpenAIGatewayService, account *Account) { account.RateLimitResetAt = &future }}, + {name: "overloaded", mutate: func(_ *OpenAIGatewayService, account *Account) { account.OverloadUntil = &future }}, + { + name: "independently runtime blocked", + mutate: func(svc *OpenAIGatewayService, account *Account) { + svc.BlockAccountScheduling(account, time.Now().Add(24*time.Hour), "independent") + }, + wantRuntimeBlocked: true, + }, + } + classes := []struct { + name string + class grokCredentialFailureClass + }{ + {name: "permanent", class: grokCredentialFailureClass{reason: GrokCredentialReasonRevoked, permanent: true}}, + {name: "transient", class: grokCredentialFailureClass{reason: GrokCredentialReasonRefreshTransient, transient: true}}, + } + + for classIndex, classCase := range classes { + for stateIndex, stateCase := range states { + t.Run(classCase.name+"/"+stateCase.name, func(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(int64(800 + classIndex*20 + stateIndex)) + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{account.ID: account} + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: NewGrokTokenProvider(repo, &grokTokenCacheForProviderTest{})} + repo.beforeConditionalState = func() { + latest := *account + latest.Credentials = shallowCopyMap(account.Credentials) + latest.Credentials["access_token"] = "fresh-but-ineligible-token" + latest.Credentials["refresh_token"] = "fresh-but-ineligible-refresh" + latest.Credentials["expires_at"] = time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339) + latest.Credentials["_token_version"] = time.Now().UnixMilli() + stateCase.mutate(svc, &latest) + repo.accountsByID[account.ID] = &latest + } + + token, err := svc.applyGrokCredentialAccountFailure(context.Background(), account, classCase.class) + + require.ErrorIs(t, err, errOAuthRefreshAccountStateChanged) + require.Empty(t, token) + require.Zero(t, repo.setErrorCalls) + require.Zero(t, repo.setTempUnschedCalls) + require.Equal(t, stateCase.wantRuntimeBlocked, svc.isOpenAIAccountRuntimeBlocked(account)) + }) + } + } +} + +func TestGetRequestCredentialSharedCredentialPersistenceFailureStopsWithoutAccountMutation(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(782) + repo := &tokenRefreshAccountRepo{updateErr: errors.New("database unavailable")} + repo.accountsByID = map[int64]*Account{account.ID: account} + cache := &grokTokenCacheForProviderTest{lockResult: true} + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), &tokenRefresherStub{credentials: map[string]any{ + "access_token": "new-access-token", + "refresh_token": "new-refresh-token", + "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), + }}) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + + _, _, err := svc.getRequestCredential(context.Background(), c, account) + + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, GatewayFailureScopeProvider, failoverErr.Scope) + require.Equal(t, GrokCredentialReasonProviderDown, failoverErr.Reason) + require.Equal(t, NextAccountStop, failoverErr.NextAccountAction) + require.Equal(t, 1, repo.updateCredentialsCalls) + require.Zero(t, repo.setErrorCalls) + require.Zero(t, repo.setTempUnschedCalls) + require.Empty(t, cache.deletedKeys) + require.False(t, svc.isOpenAIAccountRuntimeBlocked(account)) +} + +func TestGetRequestCredentialRecoversConcurrentRefreshWithoutFailover(t *testing.T) { + gin.SetMode(gin.TestMode) + account := expiredGrokOAuthAccountForCredentialTest(733) + latest := *account + latest.Credentials = shallowCopyMap(account.Credentials) + latest.Credentials["access_token"] = "fresh-concurrent-access" + latest.Credentials["refresh_token"] = "fresh-concurrent-refresh" + latest.Credentials["expires_at"] = time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339) + latest.Credentials["_token_version"] = time.Now().UnixMilli() + baseRepo := &tokenRefreshAccountRepo{} + baseRepo.accountsByID = map[int64]*Account{account.ID: account} + repo := &grokCredentialSequencedRepo{tokenRefreshAccountRepo: baseRepo, latest: &latest} + cache := &grokTokenCacheForProviderTest{lockResult: true} + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), &tokenRefresherStub{ + err: infraerrors.New(http.StatusForbidden, "GROK_OAUTH_ENTITLEMENT_DENIED", "access_denied"), + }) + svc := &OpenAIGatewayService{accountRepo: repo, grokTokenProvider: provider} + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + + token, kind, err := svc.getRequestCredential(context.Background(), c, account) + + require.NoError(t, err) + require.Equal(t, "fresh-concurrent-access", token) + require.Equal(t, "oauth", kind) + require.Zero(t, baseRepo.setErrorCalls) + require.Zero(t, baseRepo.setTempUnschedCalls) + require.Empty(t, cache.deletedKeys) + require.False(t, svc.isOpenAIAccountRuntimeBlocked(account)) + _, hasEvents := c.Get(OpsUpstreamErrorsKey) + require.False(t, hasEvents) +} + +func expiredGrokOAuthAccountForCredentialTest(id int64) *Account { + return &Account{ + ID: id, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, + Credentials: map[string]any{ + "access_token": "expired-access-token", + "refresh_token": "refresh-token", + "expires_at": time.Now().Add(-time.Minute).UTC().Format(time.RFC3339), + "base_url": xai.DefaultCLIBaseURL, + }, + } +} diff --git a/backend/internal/service/grok_media.go b/backend/internal/service/grok_media.go index 100d720659..f7c6bd8153 100644 --- a/backend/internal/service/grok_media.go +++ b/backend/internal/service/grok_media.go @@ -13,7 +13,6 @@ import ( "strings" "time" - "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/Wei-Shaw/sub2api/internal/util/responseheaders" "github.com/gin-gonic/gin" "github.com/tidwall/gjson" @@ -268,25 +267,6 @@ func (s *OpenAIGatewayService) BindGrokMediaVideoRequestAccount(ctx context.Cont return s.BindStickySession(ctx, groupID, GrokMediaVideoRequestSessionHash(requestID), accountID) } -func (e GrokMediaEndpoint) upstreamURL(baseURL, requestID string) (string, error) { - switch e { - case GrokMediaEndpointImagesGenerations: - return xai.BuildImagesGenerationsURL(baseURL) - case GrokMediaEndpointImagesEdits: - return xai.BuildImagesEditsURL(baseURL) - case GrokMediaEndpointVideosGenerations: - return xai.BuildVideosGenerationsURL(baseURL) - case GrokMediaEndpointVideosEdits: - return xai.BuildVideosEditsURL(baseURL) - case GrokMediaEndpointVideosExtensions: - return xai.BuildVideosExtensionsURL(baseURL) - case GrokMediaEndpointVideoStatus: - return xai.BuildVideoURL(baseURL, requestID) - default: - return "", fmt.Errorf("unsupported grok media endpoint: %s", e) - } -} - func (s *OpenAIGatewayService) ForwardGrokMedia( ctx context.Context, c *gin.Context, @@ -304,11 +284,11 @@ func (s *OpenAIGatewayService) ForwardGrokMedia( return nil, fmt.Errorf("account platform %s is not supported for grok media", account.Platform) } - token, _, err := s.GetAccessToken(ctx, account) + token, _, err := s.getRequestCredential(ctx, c, account) if err != nil { return nil, err } - targetURL, err := endpoint.upstreamURL(account.GetGrokMediaBaseURL(), requestID) + targetURL, err := buildGrokMediaURL(account, s.cfg, endpoint, requestID) if err != nil { return nil, err } @@ -339,7 +319,9 @@ func (s *OpenAIGatewayService) ForwardGrokMedia( } upstreamReq.Header.Set("Authorization", "Bearer "+token) upstreamReq.Header.Set("Accept", "application/json") - applyGrokCLIHeaders(upstreamReq.Header) + if account.IsGrokOAuth() { + applyGrokCLIHeaders(upstreamReq.Header) + } if endpoint.RequiresRequestBody() { contentType = strings.TrimSpace(contentType) if contentType == "" { @@ -366,7 +348,7 @@ func (s *OpenAIGatewayService) ForwardGrokMedia( return s.handleGrokMediaErrorResponse(ctx, resp, c, account, requestIDHeader, requestModel) } - s.updateGrokUsageSnapshot(ctx, account, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) + s.updateGrokUsageFromResponse(ctx, account, resp.Header, resp.StatusCode) respBody, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, openAITooLargeError) if err != nil { return nil, err @@ -635,6 +617,7 @@ func (s *OpenAIGatewayService) handleGrokMediaErrorResponse( return nil, &UpstreamFailoverError{ StatusCode: resp.StatusCode, ResponseBody: body, + ResponseHeaders: resp.Header.Clone(), RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode), } } diff --git a/backend/internal/service/grok_oauth_service.go b/backend/internal/service/grok_oauth_service.go index 136e5e3b21..92f3937967 100644 --- a/backend/internal/service/grok_oauth_service.go +++ b/backend/internal/service/grok_oauth_service.go @@ -3,6 +3,7 @@ package service import ( "context" "crypto/subtle" + "errors" "net/http" "strings" "time" @@ -314,10 +315,13 @@ func (s *GrokOAuthService) proxyURL(ctx context.Context, proxyID *int64) (string } proxy, err := s.proxyRepo.GetByID(ctx, *proxyID) if err != nil { - return "", infraerrors.Newf(http.StatusBadRequest, "GROK_OAUTH_PROXY_NOT_FOUND", "proxy not found: %v", err) + if errors.Is(err, ErrProxyNotFound) { + return "", infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_PROXY_NOT_FOUND", "configured proxy was not found") + } + return "", infraerrors.New(http.StatusServiceUnavailable, "GROK_OAUTH_PROXY_LOOKUP_FAILED", "proxy lookup is temporarily unavailable") } if proxy == nil { - return "", nil + return "", infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_PROXY_NOT_FOUND", "configured proxy was not found") } return proxy.URL(), nil } diff --git a/backend/internal/service/grok_quota_service.go b/backend/internal/service/grok_quota_service.go index 2219a75c15..a95b778707 100644 --- a/backend/internal/service/grok_quota_service.go +++ b/backend/internal/service/grok_quota_service.go @@ -135,7 +135,7 @@ func (s *GrokQuotaService) probeUsage(ctx context.Context, accountID int64) (*Gr if err != nil { return nil, infraerrors.Newf(http.StatusBadRequest, "GROK_QUOTA_PROBE_BODY_ERROR", "failed to build probe body: %v", err) } - targetURL, err := xai.BuildResponsesURL(account.GetGrokBaseURL()) + targetURL, err := buildGrokResponsesURL(account, nil) if err != nil { return nil, infraerrors.Newf(http.StatusBadRequest, "GROK_QUOTA_BASE_URL_INVALID", "invalid Grok base_url: %v", err) } @@ -149,7 +149,9 @@ func (s *GrokQuotaService) probeUsage(ctx context.Context, accountID int64) (*Gr req.Header.Set("Authorization", "Bearer "+token) req.Header.Set("Content-Type", "application/json") req.Header.Set("Accept", "application/json") - applyGrokCLIHeaders(req.Header) + if account.IsGrokOAuth() { + applyGrokCLIHeaders(req.Header) + } resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, maxInt(account.Concurrency, 1)) if err != nil { @@ -158,7 +160,7 @@ func (s *GrokQuotaService) probeUsage(ctx context.Context, accountID int64) (*Gr defer func() { _ = resp.Body.Close() }() snapshot := xai.ObserveQuotaHeaders(resp.Header, resp.StatusCode, "active_probe") - resetAt, limited := grokRateLimitResetAt(snapshot, time.Now()) + resetAt, limited := grokRateLimitResetAtForAccount(account, snapshot, time.Now()) if limited { normalizeGrokExhaustedWindowResets(snapshot, resetAt, time.Now()) } @@ -167,6 +169,8 @@ func (s *GrokQuotaService) probeUsage(ctx context.Context, accountID int64) (*Gr }) if limited { persistGrokRateLimit(ctx, s.accountRepo, account, resetAt) + } else if isSuccessfulGrokRateLimitRecovery(account, snapshot) { + clearGrokRateLimitAfterRecovery(ctx, s.accountRepo, account) } result := &GrokQuotaProbeResult{ @@ -387,6 +391,7 @@ func (s *GrokQuotaService) prepareProbe(ctx context.Context, accountID int64) (* if err != nil { return nil, "", "", err } + proxyURL := s.resolveProxyURL(ctx, account) token, err := s.tokenProvider.GetAccessToken(ctx, account) if err != nil { @@ -396,7 +401,7 @@ func (s *GrokQuotaService) prepareProbe(ctx context.Context, accountID int64) (* return nil, "", "", infraerrors.New(http.StatusBadGateway, "GROK_QUOTA_TOKEN_UNAVAILABLE", "access token is empty") } - return account, token, s.resolveProxyURL(ctx, account), nil + return account, token, proxyURL, nil } func (s *GrokQuotaService) resolveProxyURL(ctx context.Context, account *Account) string { @@ -408,6 +413,7 @@ func (s *GrokQuotaService) resolveProxyURL(ctx context.Context, account *Account return account.Proxy.URL() case s != nil && s.proxyRepo != nil: if proxy, err := s.proxyRepo.GetByID(ctx, *account.ProxyID); err == nil && proxy != nil { + account.Proxy = proxy return proxy.URL() } } diff --git a/backend/internal/service/grok_quota_service_test.go b/backend/internal/service/grok_quota_service_test.go index ba9b6cbcb8..52def9f101 100644 --- a/backend/internal/service/grok_quota_service_test.go +++ b/backend/internal/service/grok_quota_service_test.go @@ -32,6 +32,10 @@ type grokQuotaAccountRepo struct { lastTempUnschedID int64 lastTempUnschedUntil time.Time lastTempUnschedReason string + recoveryClearCalls int + recoveryObservedAt time.Time + recoveryObservedReset time.Time + recoveryClearResult bool } func (r *grokQuotaAccountRepo) UpdateExtra(_ context.Context, id int64, updates map[string]any) error { @@ -54,6 +58,13 @@ func (r *grokQuotaAccountRepo) SetRateLimitedIfLater(ctx context.Context, id int return r.SetRateLimited(ctx, id, resetAt) } +func (r *grokQuotaAccountRepo) ClearRateLimitIfObserved(_ context.Context, _ int64, observedLimitedAt, observedResetAt time.Time) (bool, error) { + r.recoveryClearCalls++ + r.recoveryObservedAt = observedLimitedAt + r.recoveryObservedReset = observedResetAt + return r.recoveryClearResult, nil +} + func (r *grokQuotaAccountRepo) SetTempUnschedulable(_ context.Context, id int64, until time.Time, reason string) error { r.tempUnschedCalls++ r.lastTempUnschedID = id @@ -180,19 +191,26 @@ func (r *grokQuotaProxyRepo) GetByID(_ context.Context, id int64) (*Proxy, error return r.proxies[id], nil } +func healthyGrokQuotaOAuthAccount(id int64) *Account { + return &Account{ + ID: id, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": "access-token", + "refresh_token": "refresh-token", + "expires_at": time.Now().Add(2 * grokTokenRefreshSkew).UTC().Format(time.RFC3339), + }, + } +} + func TestGrokQuotaServiceProbeUsageStoresHeaders(t *testing.T) { t.Parallel() - account := &Account{ - ID: 42, - Platform: PlatformGrok, - Type: AccountTypeOAuth, - Concurrency: 1, - Credentials: map[string]any{ - "access_token": "access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), - }, - } + account := healthyGrokQuotaOAuthAccount(42) repo := &grokQuotaAccountRepo{ mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ accountsByID: map[int64]*Account{42: account}, @@ -236,19 +254,10 @@ func TestGrokQuotaServiceProbeUsageStoresHeaders(t *testing.T) { func TestGrokQuotaServiceProbeUsageIgnoresAccountGrokMapping(t *testing.T) { t.Parallel() - account := &Account{ - ID: 47, - Platform: PlatformGrok, - Type: AccountTypeOAuth, - Concurrency: 1, - Credentials: map[string]any{ - "access_token": "access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), - "model_mapping": map[string]any{ - "grok": "grok-composer", - "grok-composer": "grok-composer-2.5-fast", - }, - }, + account := healthyGrokQuotaOAuthAccount(47) + account.Credentials["model_mapping"] = map[string]any{ + "grok": "grok-composer", + "grok-composer": "grok-composer-2.5-fast", } repo := &grokQuotaAccountRepo{ mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ @@ -272,16 +281,7 @@ func TestGrokQuotaServiceProbeUsageIgnoresAccountGrokMapping(t *testing.T) { func TestGrokQuotaServiceProbeUsageReportsProbeModelOnUpstreamError(t *testing.T) { t.Parallel() - account := &Account{ - ID: 48, - Platform: PlatformGrok, - Type: AccountTypeOAuth, - Concurrency: 1, - Credentials: map[string]any{ - "access_token": "access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), - }, - } + account := healthyGrokQuotaOAuthAccount(48) repo := &grokQuotaAccountRepo{ mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ accountsByID: map[int64]*Account{48: account}, @@ -302,16 +302,7 @@ func TestGrokQuotaServiceProbeUsageReportsProbeModelOnUpstreamError(t *testing.T func TestGrokQuotaServiceProbeUsageRedactsUpstreamErrorBodyFromErrorAndLogs(t *testing.T) { const upstreamSecret = "upstream-secret-refresh-token" - account := &Account{ - ID: 49, - Platform: PlatformGrok, - Type: AccountTypeOAuth, - Concurrency: 1, - Credentials: map[string]any{ - "access_token": "access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), - }, - } + account := healthyGrokQuotaOAuthAccount(49) repo := &grokQuotaAccountRepo{ mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ accountsByID: map[int64]*Account{49: account}, @@ -352,17 +343,8 @@ func TestGrokQuotaServiceProbeUsageLoadsProxyWhenAccountEdgeMissing(t *testing.T t.Parallel() proxyID := int64(7) - account := &Account{ - ID: 46, - Platform: PlatformGrok, - Type: AccountTypeOAuth, - Concurrency: 1, - ProxyID: &proxyID, - Credentials: map[string]any{ - "access_token": "access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), - }, - } + account := healthyGrokQuotaOAuthAccount(46) + account.ProxyID = &proxyID repo := &grokQuotaAccountRepo{ mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ accountsByID: map[int64]*Account{46: account}, @@ -394,20 +376,16 @@ func TestGrokQuotaServiceProbeUsageLoadsProxyWhenAccountEdgeMissing(t *testing.T func TestGrokQuotaServiceProbeUsageStoresNoHeadersState(t *testing.T) { t.Parallel() - account := &Account{ - ID: 45, - Platform: PlatformGrok, - Type: AccountTypeOAuth, - Concurrency: 1, - Credentials: map[string]any{ - "access_token": "access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), - }, - } + account := healthyGrokQuotaOAuthAccount(45) + observedResetAt := time.Now().Add(-time.Second).UTC().Truncate(time.Second) + observedLimitedAt := observedResetAt.Add(-grokRateLimitRepeatCooldown) + account.RateLimitedAt = &observedLimitedAt + account.RateLimitResetAt = &observedResetAt repo := &grokQuotaAccountRepo{ mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ accountsByID: map[int64]*Account{45: account}, }, + recoveryClearResult: true, } upstream := &httpUpstreamRecorder{resp: &http.Response{ StatusCode: http.StatusOK, @@ -430,20 +408,15 @@ func TestGrokQuotaServiceProbeUsageStoresNoHeadersState(t *testing.T) { require.True(t, ok) require.False(t, stored.HeadersObserved) require.Equal(t, http.StatusOK, stored.StatusCode) + require.Equal(t, 1, repo.recoveryClearCalls) + require.Equal(t, observedLimitedAt, repo.recoveryObservedAt) + require.Equal(t, observedResetAt, repo.recoveryObservedReset) } func TestGrokQuotaServiceProbeUsageReturnsRateLimitedSnapshot(t *testing.T) { t.Parallel() - account := &Account{ - ID: 43, - Platform: PlatformGrok, - Type: AccountTypeOAuth, - Credentials: map[string]any{ - "access_token": "access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), - }, - } + account := healthyGrokQuotaOAuthAccount(43) repo := &grokQuotaAccountRepo{ mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ accountsByID: map[int64]*Account{43: account}, @@ -471,13 +444,7 @@ func TestGrokQuotaServiceProbeUsageReturnsRateLimitedSnapshot(t *testing.T) { func TestGrokQuotaServiceQueryQuotaFreeFallsBackToGrok45(t *testing.T) { t.Parallel() - account := &Account{ - ID: 51, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, - Credentials: map[string]any{ - "access_token": "access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), - }, - } + account := healthyGrokQuotaOAuthAccount(51) repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ accountsByID: map[int64]*Account{account.ID: account}, }} @@ -518,13 +485,7 @@ func TestGrokQuotaServiceQueryQuotaFreeFallsBackToGrok45(t *testing.T) { func TestGrokQuotaServiceQueryQuotaPaidBillingSkipsActiveProbe(t *testing.T) { t.Parallel() - account := &Account{ - ID: 52, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, - Credentials: map[string]any{ - "access_token": "access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), - }, - } + account := healthyGrokQuotaOAuthAccount(52) repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ accountsByID: map[int64]*Account{account.ID: account}, }} @@ -552,13 +513,7 @@ func TestGrokQuotaServiceQueryQuotaPaidBillingSkipsActiveProbe(t *testing.T) { func TestGrokQuotaServiceQueryQuotaCustomPaidMonthlyLimitSkipsActiveProbe(t *testing.T) { t.Parallel() - account := &Account{ - ID: 57, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, - Credentials: map[string]any{ - "access_token": "access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), - }, - } + account := healthyGrokQuotaOAuthAccount(57) repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ accountsByID: map[int64]*Account{account.ID: account}, }} @@ -691,13 +646,7 @@ func TestGrokLocalUsageForBillingOnlyReturnsAvailableWindows(t *testing.T) { func TestAccountUsageServiceGrokRefreshUsesBillingOnly(t *testing.T) { t.Parallel() - account := &Account{ - ID: 54, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, - Credentials: map[string]any{ - "access_token": "access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), - }, - } + account := healthyGrokQuotaOAuthAccount(54) repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ accountsByID: map[int64]*Account{account.ID: account}, }} @@ -732,13 +681,7 @@ func TestAccountUsageServiceGrokRefreshUsesBillingOnly(t *testing.T) { func TestGrokQuotaServiceProbeFlightsDeduplicateBillingAndSeparateActive(t *testing.T) { t.Parallel() - account := &Account{ - ID: 55, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, - Credentials: map[string]any{ - "access_token": "access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), - }, - } + account := healthyGrokQuotaOAuthAccount(55) repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ accountsByID: map[int64]*Account{account.ID: account}, }} @@ -794,13 +737,7 @@ func TestGrokQuotaServiceProbeFlightsDeduplicateBillingAndSeparateActive(t *test func TestGrokQuotaServiceBilling429DoesNotPauseModelScheduling(t *testing.T) { t.Parallel() - account := &Account{ - ID: 56, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, - Credentials: map[string]any{ - "access_token": "access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), - }, - } + account := healthyGrokQuotaOAuthAccount(56) repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ accountsByID: map[int64]*Account{account.ID: account}, }} @@ -820,13 +757,7 @@ func TestGrokQuotaServiceBilling429DoesNotPauseModelScheduling(t *testing.T) { func TestGrokQuotaServiceQueryQuotaFree429PersistsLimitAndKeepsBilling(t *testing.T) { t.Parallel() - account := &Account{ - ID: 53, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1, - Credentials: map[string]any{ - "access_token": "access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), - }, - } + account := healthyGrokQuotaOAuthAccount(53) repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ accountsByID: map[int64]*Account{account.ID: account}, }} diff --git a/backend/internal/service/grok_token_provider.go b/backend/internal/service/grok_token_provider.go index f19ee42b88..ae57cdc59e 100644 --- a/backend/internal/service/grok_token_provider.go +++ b/backend/internal/service/grok_token_provider.go @@ -3,19 +3,25 @@ package service import ( "context" "errors" - "log/slog" + "fmt" "strconv" "strings" "time" - - "github.com/Wei-Shaw/sub2api/internal/util/logredact" ) const ( - grokTokenCacheSkew = 5 * time.Minute - grokRequestRefreshTimeout = 8 * time.Second - grokTokenProviderLogComponent = "grok_token_provider" - grokTempUnschedulableErrorCode = "token_refresh_failed" + grokTokenCacheSkew = 5 * time.Minute + grokRequestRefreshTimeout = 8 * time.Second + grokRefreshLockWaitTimeout = 2 * time.Second + grokRefreshLockPollInterval = 25 * time.Millisecond +) + +var ( + errGrokOAuthRefreshNotConfigured = errors.New("grok oauth refresh is not configured") + errGrokOAuthRefreshTokenMissing = errors.New("grok oauth refresh token is missing") + errGrokOAuthAccessTokenMissing = errors.New("grok oauth access token is missing") + errGrokOAuthAccessTokenExpired = errors.New("grok oauth access token is expired") + errGrokOAuthConfiguredProxyMiss = errors.New("grok oauth configured proxy is missing") ) type GrokTokenCache = GeminiTokenCache @@ -36,7 +42,7 @@ func NewGrokTokenProvider( return &GrokTokenProvider{ accountRepo: accountRepo, tokenCache: tokenCache, - refreshPolicy: AntigravityProviderRefreshPolicy(), + refreshPolicy: GrokProviderRefreshPolicy(), } } @@ -60,32 +66,57 @@ func (p *GrokTokenProvider) GetAccessToken(ctx context.Context, account *Account if account.Platform != PlatformGrok || account.Type != AccountTypeOAuth { return "", errors.New("not a grok oauth account") } - - cacheKey := GrokTokenCacheKey(account) - if p.tokenCache != nil { - if token, err := p.tokenCache.GetAccessToken(ctx, cacheKey); err == nil && strings.TrimSpace(token) != "" { - return token, nil - } + selectedProxyID := cloneGrokProxyID(account.ProxyID) + if eligibilityErr := grokOAuthRequestAccountEligibilityError(account); eligibilityErr != nil { + return "", withGrokCredentialFailureSnapshot(eligibilityErr, account) } expiresAt := account.GetCredentialAsTime("expires_at") - needsRefresh := expiresAt == nil || time.Until(*expiresAt) <= grokTokenRefreshSkew - if needsRefresh && strings.TrimSpace(account.GetGrokRefreshToken()) == "" { - if expiresAt == nil || !time.Now().Before(*expiresAt) { - return "", errors.New("grok access_token expired and refresh_token is missing") - } - needsRefresh = false + accountAccessToken := strings.TrimSpace(account.GetGrokAccessToken()) + if accountAccessToken == "" { + return "", withGrokCredentialFailureSnapshot(errGrokOAuthAccessTokenMissing, account) } - if needsRefresh && p.refreshAPI != nil && p.executor != nil { + if strings.TrimSpace(account.GetGrokRefreshToken()) == "" { + return "", withGrokCredentialFailureSnapshot(errGrokOAuthRefreshTokenMissing, account) + } + cacheKey := GrokTokenCacheKey(account) + if p.tokenCache != nil { + if token, err := p.tokenCache.GetAccessToken(ctx, cacheKey); err == nil { + cachedToken := strings.TrimSpace(token) + if cachedToken != "" && accountAccessToken != "" && cachedToken == accountAccessToken && + expiresAt != nil && time.Until(*expiresAt) > grokTokenRefreshSkew { + return cachedToken, nil + } + } + } + + needsRefresh := expiresAt == nil || time.Until(*expiresAt) <= grokTokenRefreshSkew + if needsRefresh { + if p.refreshAPI == nil || p.executor == nil { + return "", errGrokOAuthRefreshNotConfigured + } refreshCtx, cancel := context.WithTimeout(ctx, grokRequestRefreshTimeout) defer cancel() - result, err := p.refreshAPI.RefreshIfNeeded(refreshCtx, account, p.executor, grokTokenRefreshSkew) + result, err := p.refreshAPI.RefreshIfNeeded(withOAuthRefreshRequestPath(refreshCtx), account, p.executor, grokTokenRefreshSkew) if err != nil { - p.markTempUnschedulable(account, err) if p.refreshPolicy.OnRefreshError == ProviderRefreshErrorReturn { return "", err } - } else if !result.LockHeld && result.Account != nil { + } else if result != nil && result.LockHeld { + if p.refreshPolicy.OnLockHeld == ProviderLockHeldWaitForCache { + token, waitErr := p.waitForRefreshedToken(refreshCtx, account, cacheKey) + return token, withGrokCredentialFailureSnapshot(waitErr, account) + } + if expiresAt == nil || !time.Now().Before(*expiresAt) { + return "", withGrokCredentialFailureSnapshot(errGrokOAuthAccessTokenExpired, account) + } + } else if result != nil && result.Account != nil { + if eligibilityErr := grokOAuthRequestAccountEligibilityError(result.Account); eligibilityErr != nil { + return "", withGrokCredentialFailureSnapshot(eligibilityErr, result.Account) + } + if !grokCredentialProxyIDsEqual(result.Account.ProxyID, selectedProxyID) { + return "", withGrokCredentialFailureSnapshot(errOAuthRefreshAccountStateChanged, result.Account) + } account = result.Account expiresAt = account.GetCredentialAsTime("expires_at") } @@ -93,15 +124,28 @@ func (p *GrokTokenProvider) GetAccessToken(ctx context.Context, account *Account accessToken := account.GetGrokAccessToken() if strings.TrimSpace(accessToken) == "" { - return "", errors.New("access_token not found in credentials") + return "", withGrokCredentialFailureSnapshot(errGrokOAuthAccessTokenMissing, account) + } + if expiresAt != nil && !time.Now().Before(*expiresAt) { + return "", withGrokCredentialFailureSnapshot(errGrokOAuthAccessTokenExpired, account) } if p.tokenCache != nil { latestAccount, isStale := CheckTokenVersion(ctx, account, p.accountRepo) if isStale && latestAccount != nil { + if eligibilityErr := grokOAuthRequestAccountEligibilityError(latestAccount); eligibilityErr != nil { + return "", withGrokCredentialFailureSnapshot(eligibilityErr, latestAccount) + } + if !grokCredentialProxyIDsEqual(latestAccount.ProxyID, selectedProxyID) { + return "", withGrokCredentialFailureSnapshot(errOAuthRefreshAccountStateChanged, latestAccount) + } accessToken = latestAccount.GetGrokAccessToken() if strings.TrimSpace(accessToken) == "" { - return "", errors.New("access_token not found after version check") + return "", withGrokCredentialFailureSnapshot(errGrokOAuthAccessTokenMissing, latestAccount) + } + latestExpiry := latestAccount.GetCredentialAsTime("expires_at") + if latestExpiry == nil || !time.Now().Before(*latestExpiry) { + return "", withGrokCredentialFailureSnapshot(errGrokOAuthAccessTokenExpired, latestAccount) } } else { ttl := 30 * time.Minute @@ -123,40 +167,105 @@ func (p *GrokTokenProvider) GetAccessToken(ctx context.Context, account *Account return accessToken, nil } -func (p *GrokTokenProvider) markTempUnschedulable(account *Account, refreshErr error) { - if p == nil || p.accountRepo == nil || account == nil { - return - } - now := time.Now() - until := now.Add(tokenRefreshTempUnschedDuration) - redactedErr := "unknown error" - if refreshErr != nil { - redactedErr = logredact.RedactText(refreshErr.Error()) - } - if isNonRetryableRefreshError(refreshErr) { - if err := p.accountRepo.SetError(context.Background(), account.ID, "grok token refresh failed (non-retryable): "+redactedErr); err != nil { - slog.Warn(grokTokenProviderLogComponent+".set_error_status_failed", "account_id", account.ID, "error", err) +func (p *GrokTokenProvider) waitForRefreshedToken(ctx context.Context, account *Account, cacheKey string) (string, error) { + waitCtx, cancel := context.WithTimeout(ctx, grokRefreshLockWaitTimeout) + defer cancel() + + initialToken := strings.TrimSpace(account.GetGrokAccessToken()) + initialVersion := account.GetCredentialAsInt64("_token_version") + selectedProxyID := cloneGrokProxyID(account.ProxyID) + sawAuthoritativeState := false + var lastAccountReadErr error + ticker := time.NewTicker(grokRefreshLockPollInterval) + defer ticker.Stop() + + for { + cachedToken := "" + if p.tokenCache != nil { + if token, err := p.tokenCache.GetAccessToken(waitCtx, cacheKey); err == nil { + cachedToken = strings.TrimSpace(token) + } } - return - } - reason := "grok token refresh failed on request path: " + redactedErr - bgCtx := context.Background() - if err := p.accountRepo.SetTempUnschedulable(bgCtx, account.ID, until, reason); err != nil { - slog.Warn(grokTokenProviderLogComponent+".set_temp_unschedulable_failed", "account_id", account.ID, "error", err) - return - } - if p.tempUnschedCache != nil { - state := &TempUnschedState{ - UntilUnix: until.Unix(), - TriggeredAtUnix: now.Unix(), - ErrorMessage: grokTempUnschedulableErrorCode + ": " + reason, + + if p.accountRepo != nil { + latest, err := p.accountRepo.GetByID(waitCtx, account.ID) + if err != nil { + lastAccountReadErr = err + } else if latest == nil { + return "", errOAuthRefreshAccountStateChanged + } else { + sawAuthoritativeState = true + if eligibilityErr := grokOAuthRequestAccountEligibilityError(latest); eligibilityErr != nil { + return "", withGrokCredentialFailureSnapshot(eligibilityErr, latest) + } + if !grokCredentialProxyIDsEqual(latest.ProxyID, selectedProxyID) { + return "", withGrokCredentialFailureSnapshot(errOAuthRefreshAccountStateChanged, latest) + } + token := strings.TrimSpace(latest.GetGrokAccessToken()) + version := latest.GetCredentialAsInt64("_token_version") + expiresAt := latest.GetCredentialAsTime("expires_at") + changed := token != initialToken || (version > 0 && version > initialVersion) + valid := expiresAt != nil && time.Now().Before(*expiresAt) + if token != "" && changed && valid { + // The versioned DB credential is authoritative. A stale cache must + // not hold the request on the old expired token; repair it best-effort. + if cachedToken != "" && cachedToken != token { + ttl := time.Until(*expiresAt) + if ttl > grokTokenCacheSkew { + ttl -= grokTokenCacheSkew + } + _ = p.tokenCache.SetAccessToken(waitCtx, cacheKey, token, ttl) + } + return token, nil + } + } } - if err := p.tempUnschedCache.SetTempUnsched(bgCtx, account.ID, state); err != nil { - slog.Warn(grokTokenProviderLogComponent+".temp_unsched_cache_set_failed", "account_id", account.ID, "error", err) + + select { + case <-waitCtx.Done(): + if ctx.Err() != nil { + return "", ctx.Err() + } + if !sawAuthoritativeState { + if lastAccountReadErr == nil { + lastAccountReadErr = waitCtx.Err() + } + return "", fmt.Errorf("%w: %v", errOAuthRefreshAccountRereadFailed, lastAccountReadErr) + } + // Another worker still owns the refresh and the authoritative row is + // unchanged. Do not quarantine the old credential: its refresh may + // commit immediately after this bounded wait. + return "", errOAuthRefreshAccountStateChanged + case <-ticker.C: } } } +func grokOAuthRequestAccountEligibilityError(account *Account) error { + if account == nil || !account.IsGrokOAuth() || !account.IsSchedulable() { + return errOAuthRefreshAccountStateChanged + } + if account.ProxyID != nil && account.Proxy == nil { + return errGrokOAuthConfiguredProxyMiss + } + return nil +} + +func cloneGrokProxyID(proxyID *int64) *int64 { + if proxyID == nil { + return nil + } + value := *proxyID + return &value +} + +func (p *GrokTokenProvider) InvalidateToken(ctx context.Context, account *Account) error { + if p == nil || p.tokenCache == nil || account == nil { + return nil + } + return p.tokenCache.DeleteAccessToken(ctx, GrokTokenCacheKey(account)) +} + func GrokTokenCacheKey(account *Account) string { if account == nil { return "grok:account:0" diff --git a/backend/internal/service/grok_token_provider_test.go b/backend/internal/service/grok_token_provider_test.go index 1f7e2d8e0e..7c68e55455 100644 --- a/backend/internal/service/grok_token_provider_test.go +++ b/backend/internal/service/grok_token_provider_test.go @@ -5,6 +5,7 @@ package service import ( "context" "errors" + "sync" "testing" "time" @@ -19,9 +20,33 @@ type grokTokenCacheForProviderTest struct { setTTL time.Duration lockResult bool releaseCalls int + deletedKeys []string + deleteErr error + getCalls int + mu sync.Mutex +} + +type grokCredentialRaceRepo struct { + *tokenRefreshAccountRepo + mu sync.RWMutex +} + +func (r *grokCredentialRaceRepo) GetByID(ctx context.Context, id int64) (*Account, error) { + r.mu.RLock() + defer r.mu.RUnlock() + return r.tokenRefreshAccountRepo.GetByID(ctx, id) +} + +func (r *grokCredentialRaceRepo) setAccount(account *Account) { + r.mu.Lock() + defer r.mu.Unlock() + r.accountsByID[account.ID] = account } func (c *grokTokenCacheForProviderTest) GetAccessToken(context.Context, string) (string, error) { + c.mu.Lock() + defer c.mu.Unlock() + c.getCalls++ if c.token == "" { return "", errors.New("not cached") } @@ -29,14 +54,19 @@ func (c *grokTokenCacheForProviderTest) GetAccessToken(context.Context, string) } func (c *grokTokenCacheForProviderTest) SetAccessToken(_ context.Context, key string, token string, ttl time.Duration) error { + c.mu.Lock() + defer c.mu.Unlock() c.setKey = key c.setToken = token c.setTTL = ttl return nil } -func (c *grokTokenCacheForProviderTest) DeleteAccessToken(context.Context, string) error { - return nil +func (c *grokTokenCacheForProviderTest) DeleteAccessToken(_ context.Context, key string) error { + c.mu.Lock() + defer c.mu.Unlock() + c.deletedKeys = append(c.deletedKeys, key) + return c.deleteErr } func (c *grokTokenCacheForProviderTest) AcquireRefreshLock(context.Context, string, time.Duration) (bool, error) { @@ -53,9 +83,11 @@ func TestGrokTokenProviderRefreshesExpiredTokenOnRequestPath(t *testing.T) { expiredAt := time.Now().Add(-time.Minute).UTC().Format(time.RFC3339) account := &Account{ - ID: 54, - Platform: PlatformGrok, - Type: AccountTypeOAuth, + ID: 54, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, Credentials: map[string]any{ "access_token": "expired-access-token", "refresh_token": "refresh-token", @@ -95,9 +127,11 @@ func TestGrokTokenProviderRefreshesExpiredTokenOnRequestPath(t *testing.T) { func TestGrokTokenProviderRefreshFailureUnschedulesWithRedactedReason(t *testing.T) { expiredAt := time.Now().Add(-time.Minute).UTC().Format(time.RFC3339) account := &Account{ - ID: 55, - Platform: PlatformGrok, - Type: AccountTypeOAuth, + ID: 55, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, Credentials: map[string]any{ "access_token": "expired-access-token", "refresh_token": "refresh-token", @@ -108,24 +142,178 @@ func TestGrokTokenProviderRefreshFailureUnschedulesWithRedactedReason(t *testing repo := &tokenRefreshAccountRepo{} repo.accountsByID = map[int64]*Account{55: account} cache := &grokTokenCacheForProviderTest{lockResult: true} - tempCache := &tempUnschedCacheStub{} provider := NewGrokTokenProvider(repo, cache) provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), &tokenRefresherStub{ err: errors.New("temporary refresh failure access_token=leaked-access refresh_token=leaked-refresh"), }) - provider.SetTempUnschedCache(tempCache) token, err := provider.GetAccessToken(context.Background(), account) require.Error(t, err) require.Empty(t, token) - require.Equal(t, 1, repo.setTempUnschedCalls) + require.Equal(t, 0, repo.setTempUnschedCalls) require.Equal(t, 0, repo.setErrorCalls) - require.Contains(t, repo.lastTempUnschedReason, "access_token=***") - require.Contains(t, repo.lastTempUnschedReason, "refresh_token=***") - require.NotContains(t, repo.lastTempUnschedReason, "leaked-access") - require.NotContains(t, repo.lastTempUnschedReason, "leaked-refresh") - require.Equal(t, 1, tempCache.setCalls) - require.NotNil(t, tempCache.lastState) - require.NotContains(t, tempCache.lastState.ErrorMessage, "leaked-access") - require.NotContains(t, tempCache.lastState.ErrorMessage, "leaked-refresh") +} + +func TestGrokTokenProviderLockHeldWaitsForRefreshedCacheAndNeverUsesExpiredToken(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(56) + baseRepo := &tokenRefreshAccountRepo{} + baseRepo.accountsByID = map[int64]*Account{account.ID: account} + repo := &grokCredentialRaceRepo{tokenRefreshAccountRepo: baseRepo} + cache := &grokTokenCacheForProviderTest{lockResult: false, token: "expired-access-token"} + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), &tokenRefresherStub{}) + + go func() { + time.Sleep(40 * time.Millisecond) + refreshed := *account + refreshed.Credentials = shallowCopyMap(account.Credentials) + refreshed.Credentials["access_token"] = "refreshed-after-lock" + refreshed.Credentials["expires_at"] = time.Now().Add(time.Hour).UTC().Format(time.RFC3339) + refreshed.Credentials["_token_version"] = time.Now().UnixMilli() + repo.setAccount(&refreshed) + cache.mu.Lock() + cache.token = "refreshed-after-lock" + cache.mu.Unlock() + }() + + startedAt := time.Now() + token, err := provider.GetAccessToken(context.Background(), account) + require.NoError(t, err) + require.Equal(t, "refreshed-after-lock", token) + require.NotEqual(t, "expired-access-token", token) + require.GreaterOrEqual(t, time.Since(startedAt), 25*time.Millisecond, + "expired account metadata must prevent returning the old cached token") +} + +func TestGrokTokenProviderLockHeldTimeoutDoesNotReturnExpiredToken(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(57) + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{account.ID: account} + cache := &grokTokenCacheForProviderTest{lockResult: false} + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), &tokenRefresherStub{}) + ctx, cancel := context.WithTimeout(context.Background(), 80*time.Millisecond) + defer cancel() + + token, err := provider.GetAccessToken(ctx, account) + require.Error(t, err) + require.Empty(t, token) +} + +func TestGrokTokenProviderLockHeldRejectsChangedTokenWithoutExpiry(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(58) + baseRepo := &tokenRefreshAccountRepo{} + baseRepo.accountsByID = map[int64]*Account{account.ID: account} + repo := &grokCredentialRaceRepo{tokenRefreshAccountRepo: baseRepo} + cache := &grokTokenCacheForProviderTest{lockResult: false, token: "expired-access-token"} + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), &tokenRefresherStub{}) + + go func() { + time.Sleep(30 * time.Millisecond) + refreshed := *account + refreshed.Credentials = shallowCopyMap(account.Credentials) + refreshed.Credentials["access_token"] = "changed-without-expiry" + delete(refreshed.Credentials, "expires_at") + refreshed.Credentials["_token_version"] = time.Now().UnixMilli() + repo.setAccount(&refreshed) + cache.mu.Lock() + cache.token = "changed-without-expiry" + cache.mu.Unlock() + }() + + ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) + defer cancel() + token, err := provider.GetAccessToken(ctx, account) + + require.Error(t, err) + require.Empty(t, token, "an unbounded credential must not win the lock-held race") +} + +func TestGrokTokenProviderLockHeldUsesVersionedDBTokenAndRepairsStaleCache(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(60) + baseRepo := &tokenRefreshAccountRepo{} + baseRepo.accountsByID = map[int64]*Account{account.ID: account} + repo := &grokCredentialRaceRepo{tokenRefreshAccountRepo: baseRepo} + cache := &grokTokenCacheForProviderTest{lockResult: false, token: "expired-access-token"} + provider := NewGrokTokenProvider(repo, cache) + provider.SetRefreshAPI(NewOAuthRefreshAPI(repo, cache), &tokenRefresherStub{}) + + go func() { + time.Sleep(30 * time.Millisecond) + refreshed := *account + refreshed.Credentials = shallowCopyMap(account.Credentials) + refreshed.Credentials["access_token"] = "db-authoritative-token" + refreshed.Credentials["expires_at"] = time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339) + refreshed.Credentials["_token_version"] = time.Now().UnixMilli() + repo.setAccount(&refreshed) + }() + + token, err := provider.GetAccessToken(context.Background(), account) + + require.NoError(t, err) + require.Equal(t, "db-authoritative-token", token) + require.Equal(t, "db-authoritative-token", cache.setToken) + require.Greater(t, cache.setTTL, time.Duration(0)) +} + +func TestGrokTokenProviderRejectsStaleDBTokenWithoutExpiry(t *testing.T) { + expiresAt := time.Now().Add(2 * grokTokenRefreshSkew).UTC().Format(time.RFC3339) + account := &Account{ + ID: 59, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, + Credentials: map[string]any{ + "access_token": "old-access-token", + "refresh_token": "refresh-token", + "expires_at": expiresAt, + }, + } + latest := *account + latest.Credentials = shallowCopyMap(account.Credentials) + latest.Credentials["access_token"] = "new-access-token-without-expiry" + latest.Credentials["_token_version"] = time.Now().UnixMilli() + delete(latest.Credentials, "expires_at") + repo := &tokenRefreshAccountRepo{} + repo.accountsByID = map[int64]*Account{account.ID: &latest} + cache := &grokTokenCacheForProviderTest{} + provider := NewGrokTokenProvider(repo, cache) + + token, err := provider.GetAccessToken(context.Background(), account) + + require.ErrorIs(t, err, errGrokOAuthAccessTokenExpired) + require.Empty(t, token) +} + +func TestGrokTokenProviderRejectsIneligibleSelectedAccountBeforeWarmCache(t *testing.T) { + future := time.Now().Add(time.Hour) + tests := []struct { + name string + mutate func(*Account) + }{ + {name: "disabled", mutate: func(account *Account) { account.Status = StatusDisabled }}, + {name: "not schedulable", mutate: func(account *Account) { account.Schedulable = false }}, + {name: "temporarily unschedulable", mutate: func(account *Account) { account.TempUnschedulableUntil = &future }}, + {name: "rate limited", mutate: func(account *Account) { account.RateLimitResetAt = &future }}, + {name: "overloaded", mutate: func(account *Account) { account.OverloadUntil = &future }}, + } + + for index, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(int64(90 + index)) + account.Credentials["access_token"] = "warm-cache-token" + account.Credentials["expires_at"] = time.Now().Add(2 * grokTokenRefreshSkew).UTC().Format(time.RFC3339) + tt.mutate(account) + cache := &grokTokenCacheForProviderTest{token: "warm-cache-token"} + provider := NewGrokTokenProvider(&tokenRefreshAccountRepo{}, cache) + + token, err := provider.GetAccessToken(context.Background(), account) + + require.ErrorIs(t, err, errOAuthRefreshAccountStateChanged) + require.Empty(t, token) + require.Zero(t, cache.getCalls, "an ineligible selected account must be rejected before cache lookup") + }) + } } diff --git a/backend/internal/service/grok_token_refresher.go b/backend/internal/service/grok_token_refresher.go index 92d88cc058..0b87fcac78 100644 --- a/backend/internal/service/grok_token_refresher.go +++ b/backend/internal/service/grok_token_refresher.go @@ -22,7 +22,8 @@ func (r *GrokTokenRefresher) CacheKey(account *Account) string { } func (r *GrokTokenRefresher) CanRefresh(account *Account) bool { - return account != nil && account.Platform == PlatformGrok && account.Type == AccountTypeOAuth + return account != nil && account.Platform == PlatformGrok && account.Type == AccountTypeOAuth && + strings.TrimSpace(account.GetGrokRefreshToken()) != "" } func (r *GrokTokenRefresher) NeedsRefresh(account *Account, refreshWindow time.Duration) bool { diff --git a/backend/internal/service/grok_upstream_url.go b/backend/internal/service/grok_upstream_url.go new file mode 100644 index 0000000000..ca2ba6b64d --- /dev/null +++ b/backend/internal/service/grok_upstream_url.go @@ -0,0 +1,90 @@ +package service + +import ( + "errors" + "fmt" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/Wei-Shaw/sub2api/internal/util/urlvalidator" +) + +func grokBaseURLValidator(account *Account, cfg *config.Config) (xai.BaseURLValidator, error) { + if account == nil || !account.IsGrok() { + return nil, fmt.Errorf("grok account is required") + } + switch account.Type { + case AccountTypeOAuth: + // Subscription credentials are never governed by the operator's API-key + // URL policy. They stay pinned to the supported CLI gateway. + return redactedGrokBaseURLValidator(xai.ValidateTrustedBaseURL), nil + case AccountTypeAPIKey: + if cfg == nil { + return redactedGrokBaseURLValidator(xai.ValidateBaseURL), nil + } + if !cfg.Security.URLAllowlist.Enabled { + return redactedGrokBaseURLValidator(func(raw string) (string, error) { + return urlvalidator.ValidateURLFormat(raw, cfg.Security.URLAllowlist.AllowInsecureHTTP) + }), nil + } + return redactedGrokBaseURLValidator(func(raw string) (string, error) { + return urlvalidator.ValidateHTTPSURL(raw, urlvalidator.ValidationOptions{ + AllowedHosts: cfg.Security.URLAllowlist.UpstreamHosts, + RequireAllowlist: true, + AllowPrivate: cfg.Security.URLAllowlist.AllowPrivateHosts, + }) + }), nil + default: + return nil, fmt.Errorf("unsupported grok account type: %s", account.Type) + } +} + +func redactedGrokBaseURLValidator(validator xai.BaseURLValidator) xai.BaseURLValidator { + return func(raw string) (string, error) { + validated, err := validator(raw) + if err != nil { + return "", errors.New("base URL rejected by URL security policy") + } + return validated, nil + } +} + +func buildGrokResponsesURL(account *Account, cfg *config.Config) (string, error) { + validator, err := grokBaseURLValidator(account, cfg) + if err != nil { + return "", err + } + return xai.BuildResponsesURLWithValidator(account.GetGrokBaseURL(), validator) +} + +func buildGrokChatCompletionsURL(account *Account, cfg *config.Config) (string, error) { + validator, err := grokBaseURLValidator(account, cfg) + if err != nil { + return "", err + } + return xai.BuildChatCompletionsURLWithValidator(account.GetGrokBaseURL(), validator) +} + +func buildGrokMediaURL(account *Account, cfg *config.Config, endpoint GrokMediaEndpoint, requestID string) (string, error) { + validator, err := grokBaseURLValidator(account, cfg) + if err != nil { + return "", err + } + baseURL := account.GetGrokMediaBaseURL() + switch endpoint { + case GrokMediaEndpointImagesGenerations: + return xai.BuildImagesGenerationsURLWithValidator(baseURL, validator) + case GrokMediaEndpointImagesEdits: + return xai.BuildImagesEditsURLWithValidator(baseURL, validator) + case GrokMediaEndpointVideosGenerations: + return xai.BuildVideosGenerationsURLWithValidator(baseURL, validator) + case GrokMediaEndpointVideosEdits: + return xai.BuildVideosEditsURLWithValidator(baseURL, validator) + case GrokMediaEndpointVideosExtensions: + return xai.BuildVideosExtensionsURLWithValidator(baseURL, validator) + case GrokMediaEndpointVideoStatus: + return xai.BuildVideoURLWithValidator(baseURL, requestID, validator) + default: + return "", fmt.Errorf("unsupported grok media endpoint: %s", endpoint) + } +} diff --git a/backend/internal/service/grok_upstream_url_test.go b/backend/internal/service/grok_upstream_url_test.go new file mode 100644 index 0000000000..7919e47202 --- /dev/null +++ b/backend/internal/service/grok_upstream_url_test.go @@ -0,0 +1,122 @@ +//go:build unit + +package service + +import ( + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/stretchr/testify/require" +) + +func TestGrokAPIKeyURLPolicyFollowsGlobalSecurityConfig(t *testing.T) { + account := &Account{ + Platform: PlatformGrok, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "base_url": "http://grok.example.test/v1", + }, + } + + t.Run("insecure HTTP enabled with allowlist disabled", func(t *testing.T) { + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + cfg.Security.URLAllowlist.AllowInsecureHTTP = true + + responsesURL, err := buildGrokResponsesURL(account, cfg) + require.NoError(t, err) + require.Equal(t, "http://grok.example.test/v1/responses", responsesURL) + + chatURL, err := buildGrokChatCompletionsURL(account, cfg) + require.NoError(t, err) + require.Equal(t, "http://grok.example.test/v1/chat/completions", chatURL) + + mediaURL, err := buildGrokMediaURL(account, cfg, GrokMediaEndpointImagesGenerations, "") + require.NoError(t, err) + require.Equal(t, "http://grok.example.test/v1/images/generations", mediaURL) + }) + + t.Run("insecure HTTP disabled", func(t *testing.T) { + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + cfg.Security.URLAllowlist.AllowInsecureHTTP = false + + _, err := buildGrokResponsesURL(account, cfg) + require.EqualError(t, err, "invalid base url: base URL rejected by URL security policy") + }) + + t.Run("enabled allowlist remains HTTPS only", func(t *testing.T) { + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = true + cfg.Security.URLAllowlist.AllowInsecureHTTP = true + cfg.Security.URLAllowlist.UpstreamHosts = []string{"grok.example.test"} + + _, err := buildGrokResponsesURL(account, cfg) + require.EqualError(t, err, "invalid base url: base URL rejected by URL security policy") + }) +} + +func TestGrokAPIKeyURLPolicyAppliesAllowlistAndPrivateHostControls(t *testing.T) { + account := &Account{ + Platform: PlatformGrok, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "base_url": "https://grok.example.test/v1", + }, + } + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = true + cfg.Security.URLAllowlist.UpstreamHosts = []string{"grok.example.test"} + + target, err := buildGrokResponsesURL(account, cfg) + require.NoError(t, err) + require.Equal(t, "https://grok.example.test/v1/responses", target) + + cfg.Security.URLAllowlist.UpstreamHosts = []string{"other.example.test"} + _, err = buildGrokResponsesURL(account, cfg) + require.EqualError(t, err, "invalid base url: base URL rejected by URL security policy") + + account.Credentials["base_url"] = "https://127.0.0.1/v1" + cfg.Security.URLAllowlist.UpstreamHosts = []string{"127.0.0.1"} + _, err = buildGrokResponsesURL(account, cfg) + require.EqualError(t, err, "invalid base url: base URL rejected by URL security policy") + + cfg.Security.URLAllowlist.AllowPrivateHosts = true + target, err = buildGrokResponsesURL(account, cfg) + require.NoError(t, err) + require.Equal(t, "https://127.0.0.1/v1/responses", target) +} + +func TestGrokAPIKeyURLPolicyRedactsMalformedConfiguredURL(t *testing.T) { + account := &Account{ + Platform: PlatformGrok, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "base_url": "https://%zz:secret@grok.example.test/v1", + }, + } + cfg := &config.Config{} + cfg.Security.URLAllowlist.AllowInsecureHTTP = true + + _, err := buildGrokResponsesURL(account, cfg) + require.EqualError(t, err, "invalid base url: base URL rejected by URL security policy") + require.NotContains(t, err.Error(), "secret") +} + +func TestGrokOAuthURLPolicyIgnoresAPIKeyOverrides(t *testing.T) { + account := &Account{ + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "base_url": "http://attacker.example.test/v1", + }, + } + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + cfg.Security.URLAllowlist.AllowInsecureHTTP = true + + target, err := buildGrokResponsesURL(account, cfg) + require.NoError(t, err) + require.Equal(t, xai.DefaultCLIBaseURL+"/responses", target) +} diff --git a/backend/internal/service/oauth_refresh_api.go b/backend/internal/service/oauth_refresh_api.go index 571e9ecdb7..b80cf84b64 100644 --- a/backend/internal/service/oauth_refresh_api.go +++ b/backend/internal/service/oauth_refresh_api.go @@ -2,6 +2,7 @@ package service import ( "context" + "errors" "fmt" "log/slog" "strconv" @@ -20,6 +21,45 @@ type OAuthRefreshExecutor interface { } const defaultRefreshLockTTL = 60 * time.Second +const oauthRefreshLockCleanupTimeout = 2 * time.Second + +var ( + errOAuthRefreshAccountRereadFailed = errors.New("oauth refresh account reread failed") + errOAuthRefreshAccountStateChanged = errors.New("oauth refresh account state changed") + errOAuthRefreshCredentialPersist = errors.New("oauth refresh credential persistence failed") +) + +type oauthRefreshRequestPathKey struct{} + +func withOAuthRefreshRequestPath(ctx context.Context) context.Context { + return context.WithValue(ctx, oauthRefreshRequestPathKey{}, true) +} + +func isOAuthRefreshRequestPath(ctx context.Context) bool { + requestPath, _ := ctx.Value(oauthRefreshRequestPathKey{}).(bool) + return requestPath +} + +type oauthRefreshLocalLock struct { + semaphore chan struct{} +} + +func newOAuthRefreshLocalLock() *oauthRefreshLocalLock { + return &oauthRefreshLocalLock{semaphore: make(chan struct{}, 1)} +} + +func (l *oauthRefreshLocalLock) Lock(ctx context.Context) error { + select { + case l.semaphore <- struct{}{}: + return nil + case <-ctx.Done(): + return ctx.Err() + } +} + +func (l *oauthRefreshLocalLock) Unlock() { + <-l.semaphore +} // OAuthRefreshResult 统一刷新结果 type OAuthRefreshResult struct { @@ -35,7 +75,7 @@ type OAuthRefreshAPI struct { accountRepo AccountRepository tokenCache GeminiTokenCache // 可选,nil = 无分布式锁 lockTTL time.Duration - localLocks sync.Map // key: cacheKey string -> value: *sync.Mutex + localLocks sync.Map // key: cacheKey string -> value: *oauthRefreshLocalLock } // NewOAuthRefreshAPI 创建统一刷新 API @@ -53,11 +93,11 @@ func NewOAuthRefreshAPI(accountRepo AccountRepository, tokenCache GeminiTokenCac } // getLocalLock 返回指定 cacheKey 的进程内互斥锁 -func (api *OAuthRefreshAPI) getLocalLock(cacheKey string) *sync.Mutex { - actual, _ := api.localLocks.LoadOrStore(cacheKey, &sync.Mutex{}) - mu, ok := actual.(*sync.Mutex) +func (api *OAuthRefreshAPI) getLocalLock(cacheKey string) *oauthRefreshLocalLock { + actual, _ := api.localLocks.LoadOrStore(cacheKey, newOAuthRefreshLocalLock()) + mu, ok := actual.(*oauthRefreshLocalLock) if !ok { - mu = &sync.Mutex{} + mu = newOAuthRefreshLocalLock() api.localLocks.Store(cacheKey, mu) } return mu @@ -78,15 +118,25 @@ func (api *OAuthRefreshAPI) RefreshIfNeeded( executor OAuthRefreshExecutor, refreshWindow time.Duration, ) (*OAuthRefreshResult, error) { + if api == nil || api.accountRepo == nil { + return nil, errors.New("oauth refresh account repository is not configured") + } + if account == nil { + return nil, errors.New("oauth refresh account is nil") + } + if executor == nil { + return nil, errors.New("oauth refresh executor is nil") + } cacheKey := executor.CacheKey(account) // 0. 获取进程内互斥锁(防止同一进程内的并发刷新竞争) localMu := api.getLocalLock(cacheKey) - localMu.Lock() + if err := localMu.Lock(ctx); err != nil { + return nil, fmt.Errorf("oauth refresh local lock: %w", err) + } defer localMu.Unlock() // 1. 获取分布式锁 - lockAcquired := false if api.tokenCache != nil { acquired, lockErr := api.tokenCache.AcquireRefreshLock(ctx, cacheKey, api.lockTTL) if lockErr != nil { @@ -100,22 +150,38 @@ func (api *OAuthRefreshAPI) RefreshIfNeeded( // 锁被其他 worker 持有 return &OAuthRefreshResult{LockHeld: true}, nil } else { - lockAcquired = true - defer func() { _ = api.tokenCache.ReleaseRefreshLock(ctx, cacheKey) }() + defer func() { + cleanupCtx, cancel := context.WithTimeout(context.Background(), oauthRefreshLockCleanupTimeout) + defer cancel() + _ = api.tokenCache.ReleaseRefreshLock(cleanupCtx, cacheKey) + }() } } // 2. 从 DB 重读最新 account(锁保护下,确保使用最新的 refresh_token) freshAccount, err := api.accountRepo.GetByID(ctx, account.ID) if err != nil { - slog.Warn("oauth_refresh_db_reread_failed", - "account_id", account.ID, - "error", err, - ) - // 降级使用传入的 account - freshAccount = account - } else if freshAccount == nil { - freshAccount = account + return nil, fmt.Errorf("%w: %v", errOAuthRefreshAccountRereadFailed, err) + } + if freshAccount == nil { + return nil, fmt.Errorf("%w: account not found", errOAuthRefreshAccountStateChanged) + } + if freshAccount.ID != account.ID { + return nil, fmt.Errorf("%w: account identity mismatch", errOAuthRefreshAccountRereadFailed) + } + if !freshAccount.IsActive() { + return nil, fmt.Errorf("%w: account is not active", errOAuthRefreshAccountStateChanged) + } + if isOAuthRefreshRequestPath(ctx) && freshAccount.Platform == PlatformGrok { + if eligibilityErr := grokOAuthRequestAccountEligibilityError(freshAccount); eligibilityErr != nil { + return nil, withGrokCredentialFailureSnapshot(eligibilityErr, freshAccount) + } + } + if !executor.CanRefresh(freshAccount) { + if freshAccount.IsGrokOAuth() && strings.TrimSpace(freshAccount.GetGrokRefreshToken()) == "" { + return nil, withGrokCredentialFailureSnapshot(errGrokOAuthRefreshTokenMissing, freshAccount) + } + return nil, fmt.Errorf("%w: account is no longer refreshable", errOAuthRefreshAccountStateChanged) } // 3. 二次检查是否仍需刷新(另一条路径可能已刷新) @@ -127,11 +193,19 @@ func (api *OAuthRefreshAPI) RefreshIfNeeded( // 4. 执行平台特定刷新逻辑 newCredentials, refreshErr := executor.Refresh(ctx, freshAccount) + if err := ctx.Err(); err != nil { + return nil, err + } if refreshErr != nil { // 竞争恢复:invalid_grant 可能是另一个 worker 已消费了旧 refresh_token // 重新读取 DB,如果 refresh_token 已更新则说明是竞争,返回成功 if isInvalidGrantError(refreshErr) { if recoveredAccount, recovered := api.tryRecoverFromRefreshRace(ctx, freshAccount); recovered { + if isOAuthRefreshRequestPath(ctx) && recoveredAccount.Platform == PlatformGrok { + if eligibilityErr := grokOAuthRequestAccountEligibilityError(recoveredAccount); eligibilityErr != nil { + return nil, withGrokCredentialFailureSnapshot(eligibilityErr, recoveredAccount) + } + } slog.Info("oauth_refresh_race_recovered", "account_id", freshAccount.ID, "platform", freshAccount.Platform, @@ -141,7 +215,7 @@ func (api *OAuthRefreshAPI) RefreshIfNeeded( }, nil } } - return nil, refreshErr + return nil, withGrokCredentialFailureSnapshot(refreshErr, freshAccount) } // 5. 设置版本号 + 更新 DB @@ -152,16 +226,30 @@ func (api *OAuthRefreshAPI) RefreshIfNeeded( "account_id", freshAccount.ID, "error", updateErr, ) - return nil, fmt.Errorf("oauth refresh succeeded but DB update failed: %w", updateErr) + return nil, withGrokCredentialFailureSnapshot( + fmt.Errorf("%w: %v", errOAuthRefreshCredentialPersist, updateErr), freshAccount, + ) } } - - _ = lockAcquired // suppress unused warning when tokenCache is nil + resultAccount := freshAccount + if isOAuthRefreshRequestPath(ctx) && freshAccount.Platform == PlatformGrok { + latestAccount, rereadErr := api.accountRepo.GetByID(ctx, freshAccount.ID) + if rereadErr != nil { + return nil, fmt.Errorf("%w: %v", errOAuthRefreshAccountRereadFailed, rereadErr) + } + if latestAccount == nil { + return nil, fmt.Errorf("%w: account not found after refresh", errOAuthRefreshAccountStateChanged) + } + if eligibilityErr := grokOAuthRequestAccountEligibilityError(latestAccount); eligibilityErr != nil { + return nil, withGrokCredentialFailureSnapshot(eligibilityErr, latestAccount) + } + resultAccount = latestAccount + } return &OAuthRefreshResult{ Refreshed: true, NewCredentials: newCredentials, - Account: freshAccount, + Account: resultAccount, }, nil } diff --git a/backend/internal/service/oauth_refresh_api_test.go b/backend/internal/service/oauth_refresh_api_test.go index cacffddc59..adf089f0bd 100644 --- a/backend/internal/service/oauth_refresh_api_test.go +++ b/backend/internal/service/oauth_refresh_api_test.go @@ -51,13 +51,14 @@ func (r *refreshAPIAccountRepo) UpdateCredentials(_ context.Context, id int64, c // refreshAPIExecutorStub implements OAuthRefreshExecutor for tests. type refreshAPIExecutorStub struct { - needsRefresh bool - credentials map[string]any - err error - refreshCalls int + needsRefresh bool + cannotRefresh bool + credentials map[string]any + err error + refreshCalls int } -func (e *refreshAPIExecutorStub) CanRefresh(_ *Account) bool { return true } +func (e *refreshAPIExecutorStub) CanRefresh(_ *Account) bool { return !e.cannotRefresh } func (e *refreshAPIExecutorStub) NeedsRefresh(_ *Account, _ time.Duration) bool { return e.needsRefresh @@ -77,9 +78,10 @@ func (e *refreshAPIExecutorStub) CacheKey(account *Account) string { // refreshAPICacheStub implements GeminiTokenCache for OAuthRefreshAPI tests. type refreshAPICacheStub struct { - lockResult bool - lockErr error - releaseCalls int + lockResult bool + lockErr error + releaseCalls int + releaseCtxErr error } func (c *refreshAPICacheStub) GetAccessToken(context.Context, string) (string, error) { @@ -96,15 +98,16 @@ func (c *refreshAPICacheStub) AcquireRefreshLock(context.Context, string, time.D return c.lockResult, c.lockErr } -func (c *refreshAPICacheStub) ReleaseRefreshLock(context.Context, string) error { +func (c *refreshAPICacheStub) ReleaseRefreshLock(ctx context.Context, _ string) error { c.releaseCalls++ + c.releaseCtxErr = ctx.Err() return nil } // ========== RefreshIfNeeded tests ========== func TestRefreshIfNeeded_Success(t *testing.T) { - account := &Account{ID: 1, Platform: PlatformAnthropic, Type: AccountTypeOAuth} + account := &Account{ID: 1, Platform: PlatformAnthropic, Type: AccountTypeOAuth, Status: StatusActive} repo := &refreshAPIAccountRepo{account: account} cache := &refreshAPICacheStub{lockResult: true} executor := &refreshAPIExecutorStub{ @@ -132,6 +135,7 @@ func TestRefreshIfNeeded_UpdateCredentialsPreservesRateLimitState(t *testing.T) ID: 11, Platform: PlatformGemini, Type: AccountTypeOAuth, + Status: StatusActive, RateLimitResetAt: &resetAt, } repo := &refreshAPIAccountRepo{account: account} @@ -152,7 +156,7 @@ func TestRefreshIfNeeded_UpdateCredentialsPreservesRateLimitState(t *testing.T) } func TestRefreshIfNeeded_LockHeld(t *testing.T) { - account := &Account{ID: 2, Platform: PlatformAnthropic} + account := &Account{ID: 2, Platform: PlatformAnthropic, Status: StatusActive} repo := &refreshAPIAccountRepo{account: account} cache := &refreshAPICacheStub{lockResult: false} // lock not acquired executor := &refreshAPIExecutorStub{needsRefresh: true} @@ -168,7 +172,7 @@ func TestRefreshIfNeeded_LockHeld(t *testing.T) { } func TestRefreshIfNeeded_LockErrorDegrades(t *testing.T) { - account := &Account{ID: 3, Platform: PlatformGemini, Type: AccountTypeOAuth} + account := &Account{ID: 3, Platform: PlatformGemini, Type: AccountTypeOAuth, Status: StatusActive} repo := &refreshAPIAccountRepo{account: account} cache := &refreshAPICacheStub{lockErr: errors.New("redis down")} // lock error executor := &refreshAPIExecutorStub{ @@ -187,7 +191,7 @@ func TestRefreshIfNeeded_LockErrorDegrades(t *testing.T) { } func TestRefreshIfNeeded_NoCacheNoLock(t *testing.T) { - account := &Account{ID: 4, Platform: PlatformGemini, Type: AccountTypeOAuth} + account := &Account{ID: 4, Platform: PlatformGemini, Type: AccountTypeOAuth, Status: StatusActive} repo := &refreshAPIAccountRepo{account: account} executor := &refreshAPIExecutorStub{ needsRefresh: true, @@ -203,7 +207,7 @@ func TestRefreshIfNeeded_NoCacheNoLock(t *testing.T) { } func TestRefreshIfNeeded_AlreadyRefreshed(t *testing.T) { - account := &Account{ID: 5, Platform: PlatformAnthropic} + account := &Account{ID: 5, Platform: PlatformAnthropic, Status: StatusActive} repo := &refreshAPIAccountRepo{account: account} cache := &refreshAPICacheStub{lockResult: true} executor := &refreshAPIExecutorStub{needsRefresh: false} // already refreshed @@ -220,7 +224,7 @@ func TestRefreshIfNeeded_AlreadyRefreshed(t *testing.T) { } func TestRefreshIfNeeded_RefreshError(t *testing.T) { - account := &Account{ID: 6, Platform: PlatformAnthropic} + account := &Account{ID: 6, Platform: PlatformAnthropic, Status: StatusActive} repo := &refreshAPIAccountRepo{account: account} cache := &refreshAPICacheStub{lockResult: true} executor := &refreshAPIExecutorStub{ @@ -239,7 +243,7 @@ func TestRefreshIfNeeded_RefreshError(t *testing.T) { } func TestRefreshIfNeeded_DBUpdateError(t *testing.T) { - account := &Account{ID: 7, Platform: PlatformGemini, Type: AccountTypeOAuth} + account := &Account{ID: 7, Platform: PlatformGemini, Type: AccountTypeOAuth, Status: StatusActive} repo := &refreshAPIAccountRepo{ account: account, updateErr: errors.New("db connection lost"), @@ -255,12 +259,12 @@ func TestRefreshIfNeeded_DBUpdateError(t *testing.T) { require.Error(t, err) require.Nil(t, result) - require.Contains(t, err.Error(), "DB update failed") + require.ErrorIs(t, err, errOAuthRefreshCredentialPersist) require.Equal(t, 1, repo.updateCalls) // attempted } func TestRefreshIfNeeded_DBRereadFails(t *testing.T) { - account := &Account{ID: 8, Platform: PlatformAnthropic, Type: AccountTypeOAuth} + account := &Account{ID: 8, Platform: PlatformAnthropic, Type: AccountTypeOAuth, Status: StatusActive} repo := &refreshAPIAccountRepo{ account: nil, // GetByID returns nil getByIDErr: errors.New("db timeout"), @@ -274,13 +278,95 @@ func TestRefreshIfNeeded_DBRereadFails(t *testing.T) { api := NewOAuthRefreshAPI(repo, cache) result, err := api.RefreshIfNeeded(context.Background(), account, executor, 3*time.Minute) - require.NoError(t, err) - require.True(t, result.Refreshed) - require.Equal(t, 1, executor.refreshCalls) // still refreshes using passed-in account + require.ErrorContains(t, err, "oauth refresh account reread") + require.Nil(t, result) + require.Zero(t, executor.refreshCalls, "must not refresh with the stale caller snapshot") + require.Zero(t, repo.updateCalls) + require.Equal(t, 1, cache.releaseCalls) +} + +func TestRefreshIfNeeded_DBRereadNilFailsClosed(t *testing.T) { + account := &Account{ID: 81, Platform: PlatformGrok, Type: AccountTypeOAuth, Status: StatusActive} + repo := &refreshAPIAccountRepo{} + cache := &refreshAPICacheStub{lockResult: true} + executor := &refreshAPIExecutorStub{needsRefresh: true} + + api := NewOAuthRefreshAPI(repo, cache) + result, err := api.RefreshIfNeeded(context.Background(), account, executor, 3*time.Minute) + + require.ErrorIs(t, err, errOAuthRefreshAccountStateChanged) + require.Nil(t, result) + require.Zero(t, executor.refreshCalls) + require.Zero(t, repo.updateCalls) + require.Equal(t, 1, cache.releaseCalls) +} + +func TestRefreshIfNeeded_DBRereadInactiveFailsClosed(t *testing.T) { + account := &Account{ID: 82, Platform: PlatformGrok, Type: AccountTypeOAuth, Status: StatusActive} + freshAccount := &Account{ID: account.ID, Platform: PlatformGrok, Type: AccountTypeOAuth, Status: StatusDisabled} + repo := &refreshAPIAccountRepo{account: freshAccount} + executor := &refreshAPIExecutorStub{needsRefresh: true} + + api := NewOAuthRefreshAPI(repo, nil) + result, err := api.RefreshIfNeeded(context.Background(), account, executor, 3*time.Minute) + + require.ErrorContains(t, err, "account is not active") + require.Nil(t, result) + require.Zero(t, executor.refreshCalls) + require.Zero(t, repo.updateCalls) +} + +func TestRefreshIfNeeded_DBRereadRevalidatesExecutorContract(t *testing.T) { + tests := []struct { + name string + freshPlatform string + freshType string + }{ + {name: "platform changed", freshPlatform: PlatformAnthropic, freshType: AccountTypeOAuth}, + {name: "type changed", freshPlatform: PlatformGrok, freshType: AccountTypeUpstream}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + account := &Account{ID: 83, Platform: PlatformGrok, Type: AccountTypeOAuth, Status: StatusActive} + freshAccount := &Account{ID: account.ID, Platform: tt.freshPlatform, Type: tt.freshType, Status: StatusActive} + repo := &refreshAPIAccountRepo{account: freshAccount} + executor := NewGrokTokenRefresher(nil) + + api := NewOAuthRefreshAPI(repo, nil) + result, err := api.RefreshIfNeeded(context.Background(), account, executor, 3*time.Minute) + + require.ErrorContains(t, err, "no longer refreshable") + require.Nil(t, result) + require.Zero(t, repo.updateCalls) + }) + } +} + +func TestRefreshIfNeeded_DBRereadMissingGrokRefreshCredentialReturnsPermanentSignal(t *testing.T) { + account := &Account{ + ID: 84, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Status: StatusActive, + Credentials: map[string]any{ + "refresh_token": "caller-snapshot-refresh-token", + }, + } + freshAccount := &Account{ID: account.ID, Platform: PlatformGrok, Type: AccountTypeOAuth, Status: StatusActive} + repo := &refreshAPIAccountRepo{account: freshAccount} + executor := NewGrokTokenRefresher(nil) + + api := NewOAuthRefreshAPI(repo, nil) + result, err := api.RefreshIfNeeded(context.Background(), account, executor, 3*time.Minute) + + require.ErrorIs(t, err, errGrokOAuthRefreshTokenMissing) + require.Nil(t, result) + require.Zero(t, repo.updateCalls) } func TestRefreshIfNeeded_NilCredentials(t *testing.T) { - account := &Account{ID: 9, Platform: PlatformGemini, Type: AccountTypeOAuth} + account := &Account{ID: 9, Platform: PlatformGemini, Type: AccountTypeOAuth, Status: StatusActive} repo := &refreshAPIAccountRepo{account: account} cache := &refreshAPICacheStub{lockResult: true} executor := &refreshAPIExecutorStub{ @@ -413,6 +499,7 @@ func TestRefreshIfNeeded_InvalidGrantRaceRecovered(t *testing.T) { ID: 10, Platform: PlatformAnthropic, Type: AccountTypeOAuth, + Status: StatusActive, Credentials: map[string]any{"refresh_token": "old-rt", "access_token": "old-at"}, } // After race, DB has new refresh token from another worker @@ -420,6 +507,7 @@ func TestRefreshIfNeeded_InvalidGrantRaceRecovered(t *testing.T) { ID: 10, Platform: PlatformAnthropic, Type: AccountTypeOAuth, + Status: StatusActive, Credentials: map[string]any{"refresh_token": "new-rt", "access_token": "new-at"}, } repo := &refreshAPIAccountRepoWithRace{ @@ -449,6 +537,7 @@ func TestRefreshIfNeeded_InvalidGrantGenuine(t *testing.T) { ID: 11, Platform: PlatformAnthropic, Type: AccountTypeOAuth, + Status: StatusActive, Credentials: map[string]any{"refresh_token": "revoked-rt", "access_token": "old-at"}, } repo := &refreshAPIAccountRepoWithRace{ @@ -474,6 +563,7 @@ func TestRefreshIfNeeded_InvalidGrantDBRereadFailsOnRecovery(t *testing.T) { ID: 12, Platform: PlatformAnthropic, Type: AccountTypeOAuth, + Status: StatusActive, Credentials: map[string]any{"refresh_token": "old-rt"}, } repo := &refreshAPIAccountRepoWithRace{ @@ -500,6 +590,7 @@ func TestRefreshIfNeeded_LocalMutexSerializesConcurrent(t *testing.T) { ID: 20, Platform: PlatformAnthropic, Type: AccountTypeOAuth, + Status: StatusActive, Credentials: map[string]any{"refresh_token": "new-rt", "access_token": "new-at"}, } callCount := 0 @@ -507,6 +598,7 @@ func TestRefreshIfNeeded_LocalMutexSerializesConcurrent(t *testing.T) { ID: 20, Platform: PlatformAnthropic, Type: AccountTypeOAuth, + Status: StatusActive, Credentials: map[string]any{"refresh_token": "old-rt"}, }} @@ -556,6 +648,78 @@ func TestRefreshIfNeeded_LocalMutexSerializesConcurrent(t *testing.T) { mu.Unlock() } +func TestRefreshIfNeeded_LocalLockWaitHonorsContextCancellation(t *testing.T) { + account := &Account{ID: 21, Platform: PlatformGrok, Type: AccountTypeOAuth, Status: StatusActive} + repo := &refreshAPIAccountRepo{account: account} + refreshStarted := make(chan struct{}) + releaseRefresh := make(chan struct{}) + var once sync.Once + executor := &dynamicRefreshExecutor{ + canRefresh: true, + cacheKey: "test:context-lock:grok", + needsRefreshFunc: func() bool { return true }, + refreshFunc: func(context.Context, *Account) (map[string]any, error) { + once.Do(func() { close(refreshStarted) }) + <-releaseRefresh + return map[string]any{"access_token": "new-at"}, nil + }, + } + api := NewOAuthRefreshAPI(repo, nil) + firstDone := make(chan error, 1) + go func() { + _, err := api.RefreshIfNeeded(context.Background(), account, executor, 3*time.Minute) + firstDone <- err + }() + <-refreshStarted + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Millisecond) + defer cancel() + startedAt := time.Now() + result, err := api.RefreshIfNeeded(ctx, account, executor, 3*time.Minute) + + require.Nil(t, result) + require.ErrorIs(t, err, context.DeadlineExceeded) + require.Less(t, time.Since(startedAt), 500*time.Millisecond) + close(releaseRefresh) + require.NoError(t, <-firstDone) +} + +func TestRefreshIfNeeded_ReleasesDistributedLockWithCleanupContext(t *testing.T) { + account := &Account{ + ID: 22, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Status: StatusActive, + Credentials: map[string]any{ + "access_token": "old-access", + "refresh_token": "old-refresh", + }, + } + repo := &refreshAPIAccountRepo{account: account} + cache := &refreshAPICacheStub{lockResult: true} + ctx, cancel := context.WithCancel(context.Background()) + executor := &dynamicRefreshExecutor{ + canRefresh: true, + cacheKey: "test:cleanup:grok", + needsRefreshFunc: func() bool { return true }, + refreshFunc: func(context.Context, *Account) (map[string]any, error) { + cancel() + return map[string]any{"access_token": "new-at"}, nil + }, + } + api := NewOAuthRefreshAPI(repo, cache) + + result, err := api.RefreshIfNeeded(ctx, account, executor, 3*time.Minute) + + require.ErrorIs(t, err, context.Canceled) + require.Nil(t, result) + require.Zero(t, repo.updateCalls) + require.Equal(t, "old-access", account.GetGrokAccessToken()) + require.Zero(t, account.GetCredentialAsInt64("_token_version")) + require.Equal(t, 1, cache.releaseCalls) + require.NoError(t, cache.releaseCtxErr) +} + // dynamicRefreshExecutor is a test helper with function-based NeedsRefresh and Refresh. type dynamicRefreshExecutor struct { canRefresh bool diff --git a/backend/internal/service/openai_account_runtime_block_fastpath.go b/backend/internal/service/openai_account_runtime_block_fastpath.go index a3f60fee2f..00d971ba22 100644 --- a/backend/internal/service/openai_account_runtime_block_fastpath.go +++ b/backend/internal/service/openai_account_runtime_block_fastpath.go @@ -3,6 +3,7 @@ package service import ( "context" "net/http" + "sync" "time" ) @@ -15,6 +16,13 @@ const ( openAIOAuth429StormMaxAccountSwitches = 1 ) +// OpenAIOAuth429FailoverState tracks the request-local follow-up budget after +// the first Grok OAuth 429. Once that 429 occurs, exactly one different account +// may be attempted; any failure from that follow-up account ends failover. +type OpenAIOAuth429FailoverState struct { + grokOAuth429FollowupPending bool +} + func openAIAccountStateContext(ctx context.Context) (context.Context, context.CancelFunc) { base := context.Background() if ctx != nil { @@ -96,6 +104,25 @@ func (s *OpenAIGatewayService) BlockAccountScheduling(account *Account, until ti if s == nil || !isOpenAIAccount(account) { return } + mu := s.openAIAccountRuntimeBlockLock(account.ID) + mu.Lock() + defer mu.Unlock() + _, _ = s.blockAccountSchedulingLocked(account, until, reason) +} + +func (s *OpenAIGatewayService) openAIAccountRuntimeBlockLock(accountID int64) *sync.Mutex { + actual, _ := s.openaiAccountRuntimeBlockLocks.LoadOrStore(accountID, &sync.Mutex{}) + mu, ok := actual.(*sync.Mutex) + if !ok { + mu = &sync.Mutex{} + s.openaiAccountRuntimeBlockLocks.Store(accountID, mu) + } + return mu +} + +func (s *OpenAIGatewayService) blockAccountSchedulingLocked(account *Account, until time.Time, _ string) (uint64, bool) { + generation := s.openaiAccountRuntimeBlockSequence.Add(1) + s.openaiAccountRuntimeBlockGeneration.Store(account.ID, generation) now := time.Now() blockUntil := until if blockUntil.IsZero() || !blockUntil.After(now) { @@ -107,7 +134,7 @@ func (s *OpenAIGatewayService) BlockAccountScheduling(account *Account, until ti if !loaded { actual, stored := s.openaiAccountRuntimeBlockUntil.LoadOrStore(account.ID, blockUntil) if !stored { - return + return generation, true } current = actual } @@ -115,15 +142,15 @@ func (s *OpenAIGatewayService) BlockAccountScheduling(account *Account, until ti currentUntil, ok := current.(time.Time) if !ok || currentUntil.IsZero() { if s.openaiAccountRuntimeBlockUntil.CompareAndSwap(account.ID, current, blockUntil) { - return + return generation, true } continue } - if currentUntil.After(blockUntil) { - return + if !blockUntil.After(currentUntil) { + return generation, false } if s.openaiAccountRuntimeBlockUntil.CompareAndSwap(account.ID, current, blockUntil) { - return + return generation, true } } } @@ -132,13 +159,20 @@ func (s *OpenAIGatewayService) ClearAccountSchedulingBlock(accountID int64) { if s == nil || accountID <= 0 { return } + mu := s.openAIAccountRuntimeBlockLock(accountID) + mu.Lock() + defer mu.Unlock() s.openaiAccountRuntimeBlockUntil.Delete(accountID) + s.openaiAccountRuntimeBlockGeneration.Store(accountID, s.openaiAccountRuntimeBlockSequence.Add(1)) } func (s *OpenAIGatewayService) isOpenAIAccountRuntimeBlocked(account *Account) bool { if s == nil || !isOpenAIAccount(account) { return false } + mu := s.openAIAccountRuntimeBlockLock(account.ID) + mu.Lock() + defer mu.Unlock() value, ok := s.openaiAccountRuntimeBlockUntil.Load(account.ID) if !ok { return false @@ -146,12 +180,14 @@ func (s *OpenAIGatewayService) isOpenAIAccountRuntimeBlocked(account *Account) b cooldownUntil, ok := value.(time.Time) if !ok || cooldownUntil.IsZero() { s.openaiAccountRuntimeBlockUntil.Delete(account.ID) + s.openaiAccountRuntimeBlockGeneration.Store(account.ID, s.openaiAccountRuntimeBlockSequence.Add(1)) return false } if time.Now().Before(cooldownUntil) { return true } s.openaiAccountRuntimeBlockUntil.Delete(account.ID) + s.openaiAccountRuntimeBlockGeneration.Store(account.ID, s.openaiAccountRuntimeBlockSequence.Add(1)) return false } @@ -181,14 +217,28 @@ func (s *OpenAIGatewayService) isOpenAIOAuth429Storm() bool { return s.openaiOAuth429WindowCount.Load() >= openAIOAuth429StormThreshold } -func (s *OpenAIGatewayService) ShouldStopOpenAIOAuth429Failover(account *Account, statusCode int, failedSwitches int) bool { - if statusCode != http.StatusTooManyRequests || failedSwitches < openAIOAuth429StormMaxAccountSwitches { +func (s *OpenAIGatewayService) ShouldStopOpenAIOAuth429Failover(account *Account, statusCode int, failedSwitches int, state *OpenAIOAuth429FailoverState) bool { + if failedSwitches < openAIOAuth429StormMaxAccountSwitches { return false } - if isGrokOAuthAccount(account) { + if state != nil && state.grokOAuth429FollowupPending { + // The follow-up budget was armed by a Grok OAuth 429. Consume it on + // any failing follow-up account, even if a mixed pool selected an API-key + // account next. return true } - if !isOpenAIOAuthAccount(account) { + if isGrokOAuthAccount(account) { + if state == nil { + // Preserve the old threshold for callers that have not adopted the + // request-local state contract yet. + return statusCode == http.StatusTooManyRequests && failedSwitches >= 2 + } + if statusCode == http.StatusTooManyRequests { + state.grokOAuth429FollowupPending = true + } + return false + } + if statusCode != http.StatusTooManyRequests || !isOpenAIOAuthAccount(account) { return false } return s.isOpenAIOAuth429Storm() diff --git a/backend/internal/service/openai_account_runtime_block_fastpath_test.go b/backend/internal/service/openai_account_runtime_block_fastpath_test.go index a51c13c531..7d7cfbda50 100644 --- a/backend/internal/service/openai_account_runtime_block_fastpath_test.go +++ b/backend/internal/service/openai_account_runtime_block_fastpath_test.go @@ -135,26 +135,45 @@ func TestShouldStopOpenAIOAuth429Failover_OnlyDuringStorm(t *testing.T) { svc := &OpenAIGatewayService{} account := &Account{ID: 42, Platform: PlatformOpenAI, Type: AccountTypeOAuth} apiKeyAccount := &Account{ID: 43, Platform: PlatformOpenAI, Type: AccountTypeAPIKey} + var state OpenAIOAuth429FailoverState - require.False(t, svc.ShouldStopOpenAIOAuth429Failover(account, http.StatusTooManyRequests, 1)) + require.False(t, svc.ShouldStopOpenAIOAuth429Failover(account, http.StatusTooManyRequests, 1, &state)) for i := 0; i < openAIOAuth429StormThreshold; i++ { svc.recordOpenAIOAuth429() } - require.True(t, svc.ShouldStopOpenAIOAuth429Failover(account, http.StatusTooManyRequests, 1)) - require.False(t, svc.ShouldStopOpenAIOAuth429Failover(apiKeyAccount, http.StatusTooManyRequests, 1)) - require.False(t, svc.ShouldStopOpenAIOAuth429Failover(account, http.StatusInternalServerError, 1)) - require.False(t, svc.ShouldStopOpenAIOAuth429Failover(account, http.StatusTooManyRequests, 0)) + require.True(t, svc.ShouldStopOpenAIOAuth429Failover(account, http.StatusTooManyRequests, 1, &state)) + require.False(t, svc.ShouldStopOpenAIOAuth429Failover(apiKeyAccount, http.StatusTooManyRequests, 1, &state)) + require.False(t, svc.ShouldStopOpenAIOAuth429Failover(account, http.StatusInternalServerError, 1, &state)) + require.False(t, svc.ShouldStopOpenAIOAuth429Failover(account, http.StatusTooManyRequests, 0, &state)) } -func TestShouldStopOpenAIOAuth429Failover_StopsGrokAfterFirst429Switch(t *testing.T) { +func TestShouldStopOpenAIOAuth429Failover_TracksOneGrokFollowupAttempt(t *testing.T) { svc := &OpenAIGatewayService{} account := &Account{ID: 44, Platform: PlatformGrok, Type: AccountTypeOAuth} apiKeyAccount := &Account{ID: 45, Platform: PlatformGrok, Type: AccountTypeAPIKey} - require.True(t, svc.ShouldStopOpenAIOAuth429Failover(account, http.StatusTooManyRequests, 1)) - require.False(t, svc.ShouldStopOpenAIOAuth429Failover(account, http.StatusTooManyRequests, 0)) - require.False(t, svc.ShouldStopOpenAIOAuth429Failover(apiKeyAccount, http.StatusTooManyRequests, 1)) - require.False(t, svc.ShouldStopOpenAIOAuth429Failover(account, http.StatusInternalServerError, 1)) + t.Run("429 then 500 stops after one followup", func(t *testing.T) { + var state OpenAIOAuth429FailoverState + require.False(t, svc.ShouldStopOpenAIOAuth429Failover(account, http.StatusTooManyRequests, 1, &state)) + require.True(t, svc.ShouldStopOpenAIOAuth429Failover(account, http.StatusInternalServerError, 2, &state)) + }) + + t.Run("500 then 429 still allows one followup", func(t *testing.T) { + var state OpenAIOAuth429FailoverState + require.False(t, svc.ShouldStopOpenAIOAuth429Failover(account, http.StatusInternalServerError, 1, &state)) + require.False(t, svc.ShouldStopOpenAIOAuth429Failover(account, http.StatusTooManyRequests, 2, &state)) + require.True(t, svc.ShouldStopOpenAIOAuth429Failover(account, http.StatusBadGateway, 3, &state)) + }) + + t.Run("OAuth 429 then API-key failure consumes the same followup", func(t *testing.T) { + var state OpenAIOAuth429FailoverState + require.False(t, svc.ShouldStopOpenAIOAuth429Failover(account, http.StatusTooManyRequests, 1, &state)) + require.True(t, svc.ShouldStopOpenAIOAuth429Failover(apiKeyAccount, http.StatusInternalServerError, 2, &state)) + }) + + var state OpenAIOAuth429FailoverState + require.False(t, svc.ShouldStopOpenAIOAuth429Failover(account, http.StatusTooManyRequests, 0, &state)) + require.False(t, svc.ShouldStopOpenAIOAuth429Failover(apiKeyAccount, http.StatusTooManyRequests, 2, &state)) } diff --git a/backend/internal/service/openai_agent_identity.go b/backend/internal/service/openai_agent_identity.go new file mode 100644 index 0000000000..aa96227f3d --- /dev/null +++ b/backend/internal/service/openai_agent_identity.go @@ -0,0 +1,522 @@ +package service + +import ( + "context" + "crypto" + "crypto/ed25519" + "crypto/sha512" + "crypto/x509" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "strings" + "sync" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/httpclient" + "golang.org/x/crypto/curve25519" + "golang.org/x/crypto/nacl/box" +) + +const ( + OpenAIAuthModeAgentIdentity = "agentIdentity" + agentIdentityAuthAPIBaseURL = "https://auth.openai.com/api/accounts" + agentIdentityTaskRegistrationTimeout = 30 * time.Second +) + +var openAIAgentIdentityAuthAPIBaseURL = agentIdentityAuthAPIBaseURL + +var agentIdentityTaskLocks sync.Map // map[int64]*sync.Mutex + +type agentIdentityWSConnectionInvalidator interface { + InvalidateAgentIdentityWSConnections(accountID int64) +} + +type agentIdentityKey struct { + runtimeID string + privateKey ed25519.PrivateKey + taskID string +} + +type agentIdentityTaskRegistrationResponse struct { + TaskID string `json:"task_id"` + TaskIDCamel string `json:"taskId"` + EncryptedTaskID string `json:"encrypted_task_id"` + EncryptedTaskIDCamel string `json:"encryptedTaskId"` +} + +type agentIdentityTaskRecoveredError struct{} + +func (e *agentIdentityTaskRecoveredError) Error() string { + return "agent identity task recovered" +} + +func (a *Account) IsOpenAIAgentIdentity() bool { + if a == nil || !a.IsOpenAIOAuth() { + return false + } + return strings.EqualFold(strings.TrimSpace(a.GetCredential(openAIAuthModeCredentialKey)), OpenAIAuthModeAgentIdentity) +} + +func agentIdentityPrivateKey(account *Account) (ed25519.PrivateKey, error) { + if account == nil { + return nil, errors.New("agent identity account is nil") + } + raw := strings.TrimSpace(account.GetCredential("agent_private_key")) + if raw == "" { + return nil, errors.New("agent identity private key is missing") + } + der, err := base64.StdEncoding.DecodeString(raw) + if err != nil { + return nil, errors.New("agent identity private key is not valid base64") + } + key, err := x509.ParsePKCS8PrivateKey(der) + if err != nil { + return nil, errors.New("agent identity private key is not valid PKCS#8") + } + privateKey, ok := key.(ed25519.PrivateKey) + if !ok || len(privateKey) != ed25519.PrivateKeySize { + return nil, errors.New("agent identity private key is not Ed25519") + } + return privateKey, nil +} + +// ValidateOpenAIAgentIdentityPrivateKey validates the stored PKCS#8 Ed25519 +// form without returning or logging the key material. +func ValidateOpenAIAgentIdentityPrivateKey(encoded string) error { + account := &Account{Credentials: map[string]any{"agent_private_key": encoded}} + _, err := agentIdentityPrivateKey(account) + return err +} + +func agentIdentityKeyFromAccount(account *Account) (agentIdentityKey, error) { + privateKey, err := agentIdentityPrivateKey(account) + if err != nil { + return agentIdentityKey{}, err + } + runtimeID := strings.TrimSpace(account.GetCredential("agent_runtime_id")) + if runtimeID == "" { + return agentIdentityKey{}, errors.New("agent identity runtime id is missing") + } + return agentIdentityKey{ + runtimeID: runtimeID, + privateKey: privateKey, + taskID: strings.TrimSpace(account.GetCredential("task_id")), + }, nil +} + +func buildAgentAssertion(key agentIdentityKey, now time.Time) (string, error) { + if key.runtimeID == "" || key.taskID == "" { + return "", errors.New("agent identity runtime or task id is missing") + } + timestamp := now.UTC().Format(time.RFC3339) + payload := []byte(key.runtimeID + ":" + key.taskID + ":" + timestamp) + signature, err := key.privateKey.Sign(nil, payload, crypto.Hash(0)) + if err != nil { + return "", errors.New("failed to sign agent assertion") + } + envelope := map[string]string{ + "agent_runtime_id": key.runtimeID, + "task_id": key.taskID, + "timestamp": timestamp, + "signature": base64.StdEncoding.EncodeToString(signature), + } + encoded, err := json.Marshal(envelope) + if err != nil { + return "", errors.New("failed to serialize agent assertion") + } + return "AgentAssertion " + base64.RawURLEncoding.EncodeToString(encoded), nil +} + +func signAgentTaskRegistration(key agentIdentityKey, timestamp time.Time) (string, string, error) { + if key.runtimeID == "" { + return "", "", errors.New("agent identity runtime id is missing") + } + formatted := timestamp.UTC().Format(time.RFC3339) + signature, err := key.privateKey.Sign(nil, []byte(key.runtimeID+":"+formatted), crypto.Hash(0)) + if err != nil { + return "", "", errors.New("failed to sign agent task registration") + } + return formatted, base64.StdEncoding.EncodeToString(signature), nil +} + +func decryptAgentTaskID(key agentIdentityKey, encoded string) (string, error) { + ciphertext, err := base64.StdEncoding.DecodeString(strings.TrimSpace(encoded)) + if err != nil { + return "", errors.New("encrypted agent task id is not valid base64") + } + seed := key.privateKey.Seed() + digest := sha512.Sum512(seed) + var curvePrivate [32]byte + copy(curvePrivate[:], digest[:32]) + curvePrivate[0] &= 248 + curvePrivate[31] &= 127 + curvePrivate[31] |= 64 + curvePublicBytes, err := curve25519.X25519(curvePrivate[:], curve25519.Basepoint) + if err != nil { + return "", errors.New("failed to derive agent identity decryption key") + } + var curvePublic [32]byte + copy(curvePublic[:], curvePublicBytes) + plaintext, ok := box.OpenAnonymous(nil, ciphertext, &curvePublic, &curvePrivate) + if !ok { + return "", errors.New("failed to decrypt encrypted agent task id") + } + taskID := strings.TrimSpace(string(plaintext)) + if taskID == "" { + return "", errors.New("decrypted agent task id is empty") + } + return taskID, nil +} + +func registerAgentIdentityTask(ctx context.Context, account *Account) (string, error) { + key, err := agentIdentityKeyFromAccount(account) + if err != nil { + return "", err + } + timestamp, signature, err := signAgentTaskRegistration(key, time.Now()) + if err != nil { + return "", err + } + proxyURL := "" + if account.ProxyID != nil && account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + client, err := httpclient.GetClient(httpclient.Options{ + ProxyURL: proxyURL, + Timeout: agentIdentityTaskRegistrationTimeout, + ResponseHeaderTimeout: 15 * time.Second, + }) + if err != nil { + return "", errors.New("invalid proxy configuration for agent task registration") + } + body, err := json.Marshal(map[string]string{ + "timestamp": timestamp, + "signature": signature, + }) + if err != nil { + return "", errors.New("failed to serialize agent task registration") + } + url := strings.TrimRight(strings.TrimSpace(openAIAgentIdentityAuthAPIBaseURL), "/") + "/v1/agent/" + key.runtimeID + "/task/register" + req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, strings.NewReader(string(body))) + if err != nil { + return "", errors.New("failed to build agent task registration request") + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Accept", "application/json") + resp, err := client.Do(req) + if err != nil { + return "", errors.New("agent task registration request failed") + } + defer func() { _ = resp.Body.Close() }() + if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices { + return "", fmt.Errorf("agent task registration returned status %d", resp.StatusCode) + } + var result agentIdentityTaskRegistrationResponse + if err := json.NewDecoder(io.LimitReader(resp.Body, 64*1024)).Decode(&result); err != nil { + return "", errors.New("agent task registration response is invalid") + } + if taskID := strings.TrimSpace(result.TaskID); taskID != "" { + return taskID, nil + } + if taskID := strings.TrimSpace(result.TaskIDCamel); taskID != "" { + return taskID, nil + } + encrypted := strings.TrimSpace(result.EncryptedTaskID) + if encrypted == "" { + encrypted = strings.TrimSpace(result.EncryptedTaskIDCamel) + } + if encrypted == "" { + return "", errors.New("agent task registration response omitted task id") + } + return decryptAgentTaskID(key, encrypted) +} + +func ensureAgentIdentityTaskForAccount(ctx context.Context, repo AccountRepository, wsInvalidator agentIdentityWSConnectionInvalidator, taskMu *sync.Mutex, account *Account, expectedTaskID string) error { + if account == nil || !account.IsOpenAIAgentIdentity() { + return nil + } + credAccount := account + if account.IsShadow() { + resolved, err := resolveCredentialAccount(ctx, repo, account) + if err != nil { + return err + } + credAccount = resolved + } + if credAccount == nil || !credAccount.IsOpenAIAgentIdentity() { + return errors.New("agent identity credentials are unavailable") + } + currentTaskID := strings.TrimSpace(credAccount.GetCredential("task_id")) + if currentTaskID != "" && (expectedTaskID == "" || currentTaskID != expectedTaskID) { + return nil + } + if taskMu == nil { + return errors.New("agent identity task lock is unavailable") + } + sharedTaskMu := taskMu + if credAccount.ID > 0 { + candidate := &sync.Mutex{} + actual, _ := agentIdentityTaskLocks.LoadOrStore(credAccount.ID, candidate) + loadedTaskMu, ok := actual.(*sync.Mutex) + if !ok { + return errors.New("agent identity task lock has invalid type") + } + sharedTaskMu = loadedTaskMu + } + sharedTaskMu.Lock() + defer sharedTaskMu.Unlock() + // Re-read inside the shared lock. Different request paths often receive + // independent repository snapshots; checking only the caller's snapshot + // would allow sequential duplicate registrations after the first writer + // has already persisted a new task. + if repo != nil && credAccount.ID > 0 { + if refreshed, refreshErr := repo.GetByID(ctx, credAccount.ID); refreshErr == nil && refreshed != nil { + if refreshed.IsShadow() { + if resolved, resolveErr := resolveCredentialAccount(ctx, repo, refreshed); resolveErr == nil && resolved != nil { + refreshed = resolved + } + } + if refreshed.IsOpenAIAgentIdentity() { + credAccount = refreshed + if !account.IsShadow() { + account.Credentials = shallowCopyMap(credAccount.Credentials) + } + } + } + } + currentTaskID = strings.TrimSpace(credAccount.GetCredential("task_id")) + if currentTaskID != "" && (expectedTaskID == "" || currentTaskID != expectedTaskID) { + return nil + } + newTaskID, err := registerAgentIdentityTask(ctx, credAccount) + if err != nil { + return err + } + credentials := make(map[string]any, len(credAccount.Credentials)+1) + for key, value := range credAccount.Credentials { + credentials[key] = value + } + credentials["task_id"] = newTaskID + if err := persistAccountCredentials(ctx, repo, credAccount, credentials); err != nil { + return err + } + if !account.IsShadow() && account != credAccount { + account.Credentials = shallowCopyMap(credAccount.Credentials) + } + if wsInvalidator != nil { + wsInvalidator.InvalidateAgentIdentityWSConnections(credAccount.ID) + } + return nil +} + +func (s *OpenAIGatewayService) ensureAgentIdentityTask(ctx context.Context, account *Account, expectedTaskID string) error { + if s == nil { + return errors.New("openai gateway service is nil") + } + return ensureAgentIdentityTaskForAccount(ctx, s.accountRepo, s, &s.agentIdentityTaskMu, account, expectedTaskID) +} + +func isAgentIdentityTaskInvalidHTTPResponse(statusCode int, body []byte) bool { + if statusCode != http.StatusUnauthorized { + return false + } + lower := strings.ToLower(string(body)) + compact := strings.NewReplacer(" ", "", "\t", "", "\r", "", "\n", "").Replace(lower) + for _, marker := range []string{ + `"code":"invalid_task_id"`, + `"code":"task_not_found"`, + `"code":"task_expired"`, + `"error":"invalid_task_id"`, + } { + if strings.Contains(compact, marker) { + return true + } + } + for _, marker := range []string{ + "invalid task_id", + "invalid task id", + "task_id is invalid", + "task id is invalid", + "task not found", + "task expired", + "unknown task_id", + "unknown task id", + } { + if strings.Contains(lower, marker) { + return true + } + } + return false +} + +type agentIdentityTaskRecoveryContextKey struct{} + +func markAgentIdentityTaskRecoveryTried(ctx context.Context) context.Context { + return context.WithValue(ctx, agentIdentityTaskRecoveryContextKey{}, true) +} + +func agentIdentityTaskRecoveryWasTried(ctx context.Context) bool { + tried, _ := ctx.Value(agentIdentityTaskRecoveryContextKey{}).(bool) + return tried +} + +func isAgentIdentityTaskInvalidWSDialError(err *openAIWSDialError) bool { + return err != nil && isAgentIdentityTaskInvalidHTTPResponse(err.StatusCode, err.ResponseBody) +} + +func (s *OpenAIGatewayService) buildOpenAIAuthenticationHeaders(ctx context.Context, account *Account, token string) (http.Header, error) { + if account == nil { + return nil, errors.New("account is nil") + } + credAccount := account + if account.IsShadow() { + resolved, err := resolveCredentialAccount(ctx, s.accountRepo, account) + if err != nil { + return nil, err + } + credAccount = resolved + } + headers := make(http.Header) + if credAccount != nil && credAccount.IsOpenAIAgentIdentity() { + agentHeaders, err := buildAgentIdentityAuthenticationHeaders(ctx, s.accountRepo, s, &s.agentIdentityTaskMu, credAccount) + if err != nil { + return nil, err + } + return agentHeaders, nil + } + headers.Set("Authorization", "Bearer "+token) + return headers, nil +} + +func buildAgentIdentityAuthenticationHeaders(ctx context.Context, repo AccountRepository, wsInvalidator agentIdentityWSConnectionInvalidator, taskMu *sync.Mutex, account *Account) (http.Header, error) { + if account == nil || !account.IsOpenAIAgentIdentity() { + return nil, errors.New("agent identity account is required") + } + if err := ensureAgentIdentityTaskForAccount(ctx, repo, wsInvalidator, taskMu, account, ""); err != nil { + return nil, err + } + key, err := agentIdentityKeyFromAccount(account) + if err != nil { + return nil, err + } + assertion, err := buildAgentAssertion(key, time.Now()) + if err != nil { + return nil, err + } + headers := make(http.Header) + headers.Set("Authorization", assertion) + return headers, nil +} + +func (s *OpenAIGatewayService) refreshOpenAIAgentIdentityHeaders(ctx context.Context, account *Account, headers http.Header) (http.Header, error) { + if account == nil { + return cloneHeader(headers), nil + } + credAccount := account + if account.IsShadow() { + resolved, err := resolveCredentialAccount(ctx, s.accountRepo, account) + if err != nil { + return nil, err + } + credAccount = resolved + } + if !credAccount.IsOpenAIAgentIdentity() { + return cloneHeader(headers), nil + } + refreshed := cloneHeader(headers) + if refreshed == nil { + refreshed = make(http.Header) + } + authHeaders, err := buildAgentIdentityAuthenticationHeaders(ctx, s.accountRepo, s, &s.agentIdentityTaskMu, credAccount) + if err != nil { + return nil, err + } + refreshed.Set("Authorization", authHeaders.Get("Authorization")) + return refreshed, nil +} + +func (s *OpenAIGatewayService) recoverAgentIdentityTask(ctx context.Context, account *Account, expectedTaskID string) error { + if account != nil && account.IsShadow() { + if resolved, err := resolveCredentialAccount(ctx, s.accountRepo, account); err == nil && resolved != nil && strings.TrimSpace(expectedTaskID) == "" { + expectedTaskID = strings.TrimSpace(resolved.GetCredential("task_id")) + } + } + return s.ensureAgentIdentityTask(ctx, account, expectedTaskID) +} + +func (s *OpenAIGatewayService) isAgentIdentityAccount(ctx context.Context, account *Account) bool { + if account == nil { + return false + } + credAccount := account + if account.IsShadow() { + resolved, err := resolveCredentialAccount(ctx, s.accountRepo, account) + if err != nil { + return false + } + credAccount = resolved + } + return credAccount != nil && credAccount.IsOpenAIAgentIdentity() +} + +// redactAgentIdentitySensitiveBody removes credential values before an +// upstream error can reach logs, ops events, or returned error text. Agent +// Identity responses should not echo these values, but keeping this boundary +// defensive prevents accidental disclosure if an upstream error does. +func redactAgentIdentitySensitiveBodyForAccount(ctx context.Context, repo AccountRepository, account *Account, body []byte) []byte { + if account == nil || len(body) == 0 { + return body + } + credAccount := account + if account != nil && account.IsShadow() { + if resolved, err := resolveCredentialAccount(ctx, repo, account); err == nil && resolved != nil { + credAccount = resolved + } + } + if credAccount == nil || !credAccount.IsOpenAIAgentIdentity() { + return body + } + redacted := string(body) + for _, key := range []string{ + "agent_private_key", + "agent_runtime_id", + "task_id", + "access_token", + "refresh_token", + "id_token", + "api_key", + "session_key", + "cookie", + } { + if value := strings.TrimSpace(credAccount.GetCredential(key)); value != "" { + redacted = strings.ReplaceAll(redacted, value, "[redacted]") + } + } + const assertionPrefix = "AgentAssertion " + for offset := 0; offset < len(redacted); { + relativeStart := strings.Index(redacted[offset:], assertionPrefix) + if relativeStart < 0 { + break + } + start := offset + relativeStart + valueStart := start + len(assertionPrefix) + end := valueStart + for end < len(redacted) && !strings.ContainsRune(" \t\r\n\"',}", rune(redacted[end])) { + end++ + } + redacted = redacted[:valueStart] + "[redacted]" + redacted[end:] + offset = valueStart + len("[redacted]") + } + return []byte(redacted) +} + +func (s *OpenAIGatewayService) redactAgentIdentitySensitiveBody(ctx context.Context, account *Account, body []byte) []byte { + if !s.isAgentIdentityAccount(ctx, account) { + return body + } + return redactAgentIdentitySensitiveBodyForAccount(ctx, s.accountRepo, account, body) +} diff --git a/backend/internal/service/openai_agent_identity_compat_test.go b/backend/internal/service/openai_agent_identity_compat_test.go new file mode 100644 index 0000000000..5dad35887b --- /dev/null +++ b/backend/internal/service/openai_agent_identity_compat_test.go @@ -0,0 +1,500 @@ +package service + +import ( + "bytes" + "context" + "encoding/base64" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestAccountTestServiceOpenAICompactAgentIdentityUsesFreshAssertion(t *testing.T) { + gin.SetMode(gin.TestMode) + key, privateKey := newTestAgentIdentityKey(t) + account := Account{ + ID: 21, + Name: "agent-identity", + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Credentials: map[string]any{ + "auth_mode": OpenAIAuthModeAgentIdentity, + "agent_runtime_id": key.runtimeID, + "agent_private_key": privateKey, + "task_id": key.taskID, + "chatgpt_account_id": "account-agent-test", + "chatgpt_account_is_fedramp": true, + }, + } + repo := &snapshotUpdateAccountRepo{stubOpenAIAccountRepo: stubOpenAIAccountRepo{accounts: []Account{account}}} + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"id":"compact-agent","status":"completed"}`)), + }} + svc := &AccountTestService{accountRepo: repo, httpUpstream: upstream} + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/21/test", bytes.NewReader(nil)) + + require.NoError(t, svc.TestAccountConnection(c, account.ID, "gpt-5.4", "", AccountTestModeCompact)) + require.Equal(t, "AgentAssertion", strings.SplitN(upstream.lastReq.Header.Get("Authorization"), " ", 2)[0]) + require.Equal(t, "account-agent-test", upstream.lastReq.Header.Get("chatgpt-account-id")) + require.Equal(t, "true", upstream.lastReq.Header.Get("x-openai-fedramp")) + require.NotContains(t, upstream.lastReq.Header.Get("Authorization"), privateKey) +} + +func TestAccountTestServiceOpenAICompactAgentIdentityRecoversInvalidTaskOnce(t *testing.T) { + gin.SetMode(gin.TestMode) + key, privateKey := newTestAgentIdentityKey(t) + account := &Account{ + ID: 22, + Name: "agent-identity-recovery", + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Credentials: map[string]any{ + "auth_mode": OpenAIAuthModeAgentIdentity, + "agent_runtime_id": key.runtimeID, + "agent_private_key": privateKey, + "task_id": "task-compact-old", + "chatgpt_account_id": "account-agent-compact-recovery", + }, + } + repo := &accountTestAgentIdentityRepo{account: account} + registerCalls := 0 + registerServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + registerCalls++ + _, _ = io.WriteString(w, `{"task_id":"task-compact-new"}`) + })) + defer registerServer.Close() + oldBase := openAIAgentIdentityAuthAPIBaseURL + openAIAgentIdentityAuthAPIBaseURL = registerServer.URL + t.Cleanup(func() { openAIAgentIdentityAuthAPIBaseURL = oldBase }) + + upstream := &httpUpstreamRecorder{responses: []*http.Response{ + {StatusCode: http.StatusUnauthorized, Header: http.Header{"Content-Type": []string{"application/json"}}, Body: io.NopCloser(strings.NewReader(`{"error":{"code":"invalid_task_id"}}`))}, + {StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"application/json"}}, Body: io.NopCloser(strings.NewReader(`{"id":"compact-agent","status":"completed"}`))}, + }} + invalidator := &agentIdentityWSInvalidationRecorder{} + svc := &AccountTestService{accountRepo: repo, httpUpstream: upstream, agentIdentityWS: invalidator} + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/22/test", bytes.NewReader(nil)) + + require.NoError(t, svc.TestAccountConnection(c, account.ID, "gpt-5.4", "", AccountTestModeCompact)) + require.Equal(t, 1, registerCalls) + require.Len(t, upstream.requests, 2) + require.Equal(t, "task-compact-new", account.GetCredential("task_id")) + require.Equal(t, 0, repo.setErrorCalls) + require.Equal(t, []int64{account.ID}, invalidator.accountIDs) +} + +func TestOpenAIAgentIdentityPassthroughKeepsSessionAndPromptCacheHeaders(t *testing.T) { + gin.SetMode(gin.TestMode) + key, privateKey := newTestAgentIdentityKey(t) + account := &Account{ + ID: 24, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "auth_mode": OpenAIAuthModeAgentIdentity, + "agent_runtime_id": key.runtimeID, + "agent_private_key": privateKey, + "task_id": key.taskID, + "chatgpt_account_id": "account-agent-passthrough", + }, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + body := []byte(`{"model":"gpt-5.4","instructions":"Reply OK","input":[],"stream":true,"prompt_cache_key":"cache-agent"}`) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body)) + c.Request.Header.Set("session_id", "client-session") + c.Request.Header.Set("conversation_id", "client-conversation") + c.Request.Header.Set("Authorization", "Bearer inbound-must-not-forward") + + svc := &OpenAIGatewayService{} + req, err := svc.buildUpstreamRequestOpenAIPassthrough(context.Background(), c, account, body, "") + require.NoError(t, err) + require.Equal(t, "AgentAssertion", strings.SplitN(req.Header.Get("Authorization"), " ", 2)[0]) + require.Equal(t, "account-agent-passthrough", req.Header.Get("chatgpt-account-id")) + require.NotEqual(t, "client-session", req.Header.Get("session_id")) + require.NotEqual(t, "client-conversation", req.Header.Get("conversation_id")) + require.Equal(t, isolateOpenAISessionID(0, "client-session"), req.Header.Get("session_id")) + require.Equal(t, isolateOpenAISessionID(0, "client-conversation"), req.Header.Get("conversation_id")) + requestBody, err := io.ReadAll(req.Body) + require.NoError(t, err) + require.Contains(t, string(requestBody), `"prompt_cache_key":"cache-agent"`) + + // Authentication mode must not affect session isolation or prompt-cache + // behavior. Compare the same request with the existing OAuth path instead + // of pinning this test to an implementation-specific hash. + oauthAccount := &Account{ + ID: 26, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "chatgpt_account_id": "account-oauth-passthrough", + }, + } + oauthRecorder := httptest.NewRecorder() + oauthContext, _ := gin.CreateTestContext(oauthRecorder) + oauthContext.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body)) + oauthContext.Request.Header.Set("session_id", "client-session") + oauthContext.Request.Header.Set("conversation_id", "client-conversation") + oauthReq, err := svc.buildUpstreamRequestOpenAIPassthrough(context.Background(), oauthContext, oauthAccount, body, "oauth-token") + require.NoError(t, err) + require.Equal(t, oauthReq.Header.Get("session_id"), req.Header.Get("session_id")) + require.Equal(t, oauthReq.Header.Get("conversation_id"), req.Header.Get("conversation_id")) +} + +func TestOpenAIAgentIdentityErrorRedactionDoesNotLeakCredentialValues(t *testing.T) { + key, privateKey := newTestAgentIdentityKey(t) + account := &Account{ + ID: 25, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "auth_mode": OpenAIAuthModeAgentIdentity, + "agent_runtime_id": key.runtimeID, + "agent_private_key": privateKey, + "task_id": key.taskID, + "access_token": key.runtimeID + "-oauth-value", + }, + } + svc := &OpenAIGatewayService{} + oauthValue := account.GetCredential("access_token") + redacted := svc.redactAgentIdentitySensitiveBody(context.Background(), account, []byte(`{"message":"runtime-test task-test `+oauthValue+` AgentAssertion abc123"}`)) + require.NotContains(t, string(redacted), key.runtimeID) + require.NotContains(t, string(redacted), key.taskID) + require.NotContains(t, string(redacted), oauthValue) + require.NotContains(t, string(redacted), "AgentAssertion abc123") + require.Contains(t, string(redacted), "[redacted]") +} + +func TestOpenAIAuthenticationHeadersPreserveOAuthPATAndAPIKeyBearerModes(t *testing.T) { + svc := &OpenAIGatewayService{} + tests := []struct { + name string + account *Account + token string + }{ + {name: "oauth", account: &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth}, token: "oauth-runtime-token"}, + {name: "personal access token", account: &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth, Credentials: map[string]any{"auth_mode": OpenAIAuthModePersonalAccessToken}}, token: "pat-runtime-token"}, + {name: "api key", account: &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey}, token: "api-key-runtime-token"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + headers, err := svc.buildOpenAIAuthenticationHeaders(context.Background(), tt.account, tt.token) + require.NoError(t, err) + require.Equal(t, "Bearer "+tt.token, headers.Get("Authorization")) + }) + } +} + +func TestOpenAIWSAgentIdentityRecoveryRequiresTaskInvalidBody(t *testing.T) { + require.False(t, isAgentIdentityTaskInvalidWSDialError(&openAIWSDialError{ + StatusCode: http.StatusUnauthorized, + ResponseBody: []byte(`{"error":{"code":"invalid_signature"}}`), + })) + require.True(t, isAgentIdentityTaskInvalidWSDialError(&openAIWSDialError{ + StatusCode: http.StatusUnauthorized, + ResponseBody: []byte(`{"error":{"code":"invalid_task_id"}}`), + })) +} + +func TestValidateOpenAIWSBearerTokenAllowsAgentIdentityWithoutStoredToken(t *testing.T) { + t.Run("Given Agent Identity When a WS path receives no bearer token Then dial-time assertion auth is allowed", func(t *testing.T) { + account := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "auth_mode": OpenAIAuthModeAgentIdentity, + }, + } + + require.NoError(t, validateOpenAIWSBearerToken(account, "")) + }) + + t.Run("Given bearer credentials When a WS path receives no token Then the request is rejected", func(t *testing.T) { + accounts := []*Account{ + {Platform: PlatformOpenAI, Type: AccountTypeOAuth}, + {Platform: PlatformOpenAI, Type: AccountTypeOAuth, Credentials: map[string]any{"auth_mode": OpenAIAuthModePersonalAccessToken}}, + {Platform: PlatformOpenAI, Type: AccountTypeAPIKey}, + } + + for _, account := range accounts { + require.EqualError(t, validateOpenAIWSBearerToken(account, ""), "token is empty") + } + }) +} + +func TestOpenAIWSConnPoolHeadersFactoryRunsAtDialAndStalePrewarmIsDiscarded(t *testing.T) { + cfg := &config.Config{} + cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1 + pool := newOpenAIWSConnPool(cfg) + defer pool.Close() + pool.setClientDialerForTest(&openAIWSFakeDialer{}) + + accountID := int64(22) + ap := pool.getOrCreateAccountPool(accountID) + factoryCalls := 0 + latestHeader := "" + req := openAIWSAcquireRequest{ + Account: &Account{ID: accountID, Platform: PlatformOpenAI, Type: AccountTypeOAuth}, + WSURL: "wss://example.com/v1/responses", + HeadersFactory: func(_ context.Context, headers http.Header) (http.Header, error) { + factoryCalls++ + latestHeader = "AgentAssertion dial-" + string(rune('0'+factoryCalls)) + if headers == nil { + headers = make(http.Header) + } + headers.Set("Authorization", latestHeader) + return headers, nil + }, + } + ap.mu.Lock() + ap.lastAcquire = &req + generation := ap.generation + ap.mu.Unlock() + + pool.prewarmConns(accountID, req, 1, generation) + require.Equal(t, 1, factoryCalls, "prewarm must generate authorization inside the actual dial") + require.Equal(t, "AgentAssertion dial-1", latestHeader) + + pool.ClearAccount(accountID) + ap.mu.Lock() + require.Empty(t, ap.conns, "credential recovery must remove pooled connections") + require.Nil(t, ap.lastAcquire, "credential recovery must discard delayed acquire state") + require.Equal(t, generation+1, ap.generation) + ap.mu.Unlock() + + // A prewarm captured before ClearAccount must not be admitted after recovery. + pool.prewarmConns(accountID, req, 1, generation) + ap.mu.Lock() + require.Empty(t, ap.conns) + ap.mu.Unlock() +} + +func TestOpenAIAgentIdentityTaskInvalidRetriesExactlyOnce(t *testing.T) { + gin.SetMode(gin.TestMode) + key, privateKey := newTestAgentIdentityKey(t) + account := &Account{ + ID: 23, + Name: "agent-identity", + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Credentials: map[string]any{ + "auth_mode": OpenAIAuthModeAgentIdentity, + "agent_runtime_id": key.runtimeID, + "agent_private_key": privateKey, + "task_id": "task-old", + "chatgpt_account_id": "account-agent-retry", + }, + } + repo := &agentIdentityForwardRepo{account: account} + registerCalls := 0 + registerServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + registerCalls++ + _, _ = io.WriteString(w, `{"task_id":"task-new"}`) + })) + defer registerServer.Close() + oldBase := openAIAgentIdentityAuthAPIBaseURL + openAIAgentIdentityAuthAPIBaseURL = registerServer.URL + t.Cleanup(func() { openAIAgentIdentityAuthAPIBaseURL = oldBase }) + + successBody := `{"id":"resp-agent-retry","object":"response","model":"gpt-5.4","output":[],"usage":{"input_tokens":1,"output_tokens":1,"total_tokens":2}}` + upstream := &httpUpstreamRecorder{responses: []*http.Response{ + {StatusCode: http.StatusUnauthorized, Header: http.Header{"Content-Type": []string{"application/json"}}, Body: io.NopCloser(strings.NewReader(`{"error":{"code":"invalid_task_id"}}`))}, + {StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"application/json"}}, Body: io.NopCloser(strings.NewReader(successBody))}, + }} + require.True(t, isAgentIdentityTaskInvalidHTTPResponse(http.StatusUnauthorized, []byte(`{"error":{"code":"invalid_task_id"}}`))) + svc := &OpenAIGatewayService{cfg: &config.Config{}, accountRepo: repo, httpUpstream: upstream} + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", strings.NewReader(`{"model":"gpt-5.4","instructions":"Reply OK","input":[],"stream":false}`)) + + _, err := svc.Forward(context.Background(), c, account, []byte(`{"model":"gpt-5.4","instructions":"Reply OK","input":[],"stream":false}`)) + require.NoError(t, err) + require.Equal(t, 1, registerCalls) + require.Len(t, upstream.requests, 2) + require.NotEqual(t, upstream.requests[0].Header.Get("Authorization"), upstream.requests[1].Header.Get("Authorization")) + require.Equal(t, "task-new", decodeAgentAssertionTask(t, upstream.requests[1].Header.Get("Authorization"))) + + // Two consecutive invalid responses still produce only one retry for this + // request; the recovery path must not loop indefinitely. + upstream.responses = []*http.Response{ + {StatusCode: http.StatusUnauthorized, Header: http.Header{"Content-Type": []string{"application/json"}}, Body: io.NopCloser(strings.NewReader(`{"error":{"code":"invalid_task_id"}}`))}, + {StatusCode: http.StatusUnauthorized, Header: http.Header{"Content-Type": []string{"application/json"}}, Body: io.NopCloser(strings.NewReader(`{"error":{"code":"invalid_task_id"}}`))}, + } + rec2 := httptest.NewRecorder() + c2, _ := gin.CreateTestContext(rec2) + c2.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", strings.NewReader(`{"model":"gpt-5.4","instructions":"Reply OK","input":[],"stream":false}`)) + _, err = svc.Forward(context.Background(), c2, account, []byte(`{"model":"gpt-5.4","instructions":"Reply OK","input":[],"stream":false}`)) + require.Error(t, err) + require.Equal(t, 2, registerCalls) + require.Len(t, upstream.requests, 4) + + // Passthrough uses the same one-shot task recovery contract. + account.Extra = map[string]any{"openai_passthrough": true} + account.Credentials["task_id"] = "task-old-passthrough" + upstream.responses = []*http.Response{ + {StatusCode: http.StatusUnauthorized, Header: http.Header{"Content-Type": []string{"application/json"}}, Body: io.NopCloser(strings.NewReader(`{"error":{"code":"invalid_task_id"}}`))}, + {StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader("data: {\"type\":\"response.completed\",\"response\":{\"usage\":{\"input_tokens\":1,\"output_tokens\":1,\"total_tokens\":2}}}\n\ndata: [DONE]\n\n"))}, + } + rec3 := httptest.NewRecorder() + c3, _ := gin.CreateTestContext(rec3) + c3.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", strings.NewReader(`{"model":"gpt-5.4","instructions":"Reply OK","input":[],"stream":false}`)) + _, err = svc.Forward(context.Background(), c3, account, []byte(`{"model":"gpt-5.4","instructions":"Reply OK","input":[],"stream":false}`)) + require.NoError(t, err) + require.Equal(t, 3, registerCalls) + require.Len(t, upstream.requests, 6) +} + +func TestOpenAIAgentIdentityCompatRoutesRecoverInvalidTaskOnce(t *testing.T) { + gin.SetMode(gin.TestMode) + tests := []struct { + name string + path string + body []byte + call func(*OpenAIGatewayService, context.Context, *gin.Context, *Account, []byte) (*OpenAIForwardResult, error) + }{ + { + name: "chat completions", + path: "/v1/chat/completions", + body: []byte(`{"model":"gpt-5.4","stream":false,"messages":[{"role":"user","content":"hi"}]}`), + call: func(s *OpenAIGatewayService, ctx context.Context, c *gin.Context, account *Account, body []byte) (*OpenAIForwardResult, error) { + return s.ForwardAsChatCompletions(ctx, c, account, body, "", "gpt-5.4") + }, + }, + { + name: "anthropic messages", + path: "/v1/messages", + body: []byte(`{"model":"gpt-5.4","stream":false,"max_tokens":32,"messages":[{"role":"user","content":"hi"}]}`), + call: func(s *OpenAIGatewayService, ctx context.Context, c *gin.Context, account *Account, body []byte) (*OpenAIForwardResult, error) { + return s.ForwardAsAnthropic(ctx, c, account, body, "", "gpt-5.4") + }, + }, + } + + for index, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + key, privateKey := newTestAgentIdentityKey(t) + account := &Account{ + ID: int64(40 + index), + Name: "agent-identity-compat", + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Credentials: map[string]any{ + "auth_mode": OpenAIAuthModeAgentIdentity, + "agent_runtime_id": key.runtimeID, + "agent_private_key": privateKey, + "task_id": "task-compat-old", + "chatgpt_account_id": "account-compat-recovery", + }, + } + repo := &agentIdentityForwardRepo{account: account} + registerCalls := 0 + registerServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + registerCalls++ + _, _ = io.WriteString(w, `{"task_id":"task-compat-new"}`) + })) + defer registerServer.Close() + oldBase := openAIAgentIdentityAuthAPIBaseURL + openAIAgentIdentityAuthAPIBaseURL = registerServer.URL + t.Cleanup(func() { openAIAgentIdentityAuthAPIBaseURL = oldBase }) + + upstream := &httpUpstreamRecorder{responses: []*http.Response{ + {StatusCode: http.StatusUnauthorized, Header: http.Header{"Content-Type": []string{"application/json"}}, Body: io.NopCloser(strings.NewReader(`{"error":{"code":"invalid_task_id"}}`))}, + {StatusCode: http.StatusUnauthorized, Header: http.Header{"Content-Type": []string{"application/json"}}, Body: io.NopCloser(strings.NewReader(`{"error":{"code":"invalid_task_id"}}`))}, + }} + svc := &OpenAIGatewayService{cfg: &config.Config{}, accountRepo: repo, httpUpstream: upstream} + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, tt.path, bytes.NewReader(tt.body)) + + _, err := tt.call(svc, context.Background(), c, account, tt.body) + require.Error(t, err) + require.Equal(t, 1, registerCalls) + require.Len(t, upstream.requests, 2) + require.Equal(t, "task-compat-new", account.GetCredential("task_id")) + }) + } +} + +func decodeAgentAssertionTask(t *testing.T, header string) string { + t.Helper() + encoded := strings.TrimPrefix(header, "AgentAssertion ") + decoded, err := base64.RawURLEncoding.DecodeString(encoded) + require.NoError(t, err) + var envelope struct { + TaskID string `json:"task_id"` + } + require.NoError(t, json.Unmarshal(decoded, &envelope)) + return envelope.TaskID +} + +type agentIdentityForwardRepo struct { + AccountRepository + account *Account +} + +type agentIdentityWSInvalidationRecorder struct { + accountIDs []int64 +} + +func (r *agentIdentityWSInvalidationRecorder) InvalidateAgentIdentityWSConnections(accountID int64) { + r.accountIDs = append(r.accountIDs, accountID) +} + +type accountTestAgentIdentityRepo struct { + AccountRepository + account *Account + setErrorCalls int +} + +func (r *accountTestAgentIdentityRepo) GetByID(_ context.Context, _ int64) (*Account, error) { + return r.account, nil +} + +func (r *accountTestAgentIdentityRepo) UpdateCredentials(_ context.Context, _ int64, credentials map[string]any) error { + r.account.Credentials = credentials + return nil +} + +func (r *accountTestAgentIdentityRepo) UpdateExtra(_ context.Context, _ int64, _ map[string]any) error { + return nil +} + +func (r *accountTestAgentIdentityRepo) SetError(_ context.Context, _ int64, _ string) error { + r.setErrorCalls++ + return nil +} + +func (r *agentIdentityForwardRepo) GetByID(_ context.Context, _ int64) (*Account, error) { + return r.account, nil +} + +func (r *agentIdentityForwardRepo) UpdateCredentials(_ context.Context, _ int64, credentials map[string]any) error { + r.account.Credentials = credentials + return nil +} diff --git a/backend/internal/service/openai_agent_identity_test.go b/backend/internal/service/openai_agent_identity_test.go new file mode 100644 index 0000000000..a906f044f0 --- /dev/null +++ b/backend/internal/service/openai_agent_identity_test.go @@ -0,0 +1,229 @@ +package service + +import ( + "context" + "crypto/ed25519" + "crypto/rand" + "crypto/sha512" + "crypto/x509" + "encoding/base64" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "strings" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/require" + "golang.org/x/crypto/curve25519" + "golang.org/x/crypto/nacl/box" +) + +func newTestAgentIdentityKey(t *testing.T) (agentIdentityKey, string) { + t.Helper() + _, privateKey, err := ed25519.GenerateKey(rand.Reader) + require.NoError(t, err) + der, err := x509.MarshalPKCS8PrivateKey(privateKey) + require.NoError(t, err) + return agentIdentityKey{ + runtimeID: "runtime-test", + privateKey: privateKey, + taskID: "task-test", + }, base64.StdEncoding.EncodeToString(der) +} + +func TestBuildAgentAssertionMatchesCodexEnvelopeAndSignature(t *testing.T) { + key, _ := newTestAgentIdentityKey(t) + now := time.Date(2026, 7, 14, 8, 9, 10, 0, time.FixedZone("UTC+8", 8*60*60)) + assertion, err := buildAgentAssertion(key, now) + require.NoError(t, err) + require.True(t, strings.HasPrefix(assertion, "AgentAssertion ")) + + encoded := strings.TrimPrefix(assertion, "AgentAssertion ") + decoded, err := base64.RawURLEncoding.DecodeString(encoded) + require.NoError(t, err) + var envelope struct { + AgentRuntimeID string `json:"agent_runtime_id"` + TaskID string `json:"task_id"` + Timestamp string `json:"timestamp"` + Signature string `json:"signature"` + } + require.NoError(t, json.Unmarshal(decoded, &envelope)) + require.Equal(t, "runtime-test", envelope.AgentRuntimeID) + require.Equal(t, "task-test", envelope.TaskID) + require.Equal(t, "2026-07-14T00:09:10Z", envelope.Timestamp) + signature, err := base64.StdEncoding.DecodeString(envelope.Signature) + require.NoError(t, err) + publicKey, ok := key.privateKey.Public().(ed25519.PublicKey) + require.True(t, ok) + require.True(t, ed25519.Verify(publicKey, []byte("runtime-test:task-test:2026-07-14T00:09:10Z"), signature)) +} + +func TestDecryptAgentTaskIDSupportsCodexSealedBoxResponse(t *testing.T) { + key, _ := newTestAgentIdentityKey(t) + digest := sha512.Sum512(key.privateKey.Seed()) + var curvePrivate [32]byte + copy(curvePrivate[:], digest[:32]) + curvePrivate[0] &= 248 + curvePrivate[31] &= 127 + curvePrivate[31] |= 64 + curvePublicBytes, err := curve25519.X25519(curvePrivate[:], curve25519.Basepoint) + require.NoError(t, err) + var curvePublic [32]byte + copy(curvePublic[:], curvePublicBytes) + ciphertext, err := box.SealAnonymous(nil, []byte("task-sealed"), &curvePublic, rand.Reader) + require.NoError(t, err) + got, err := decryptAgentTaskID(key, base64.StdEncoding.EncodeToString(ciphertext)) + require.NoError(t, err) + require.Equal(t, "task-sealed", got) +} + +func TestRegisterAgentIdentityTaskAcceptsPlaintextAndEncryptedResponses(t *testing.T) { + key, privateKey := newTestAgentIdentityKey(t) + requestCount := 0 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + require.Equal(t, http.MethodPost, r.Method) + require.Equal(t, "/v1/agent/runtime-test/task/register", r.URL.Path) + var request map[string]string + require.NoError(t, json.NewDecoder(r.Body).Decode(&request)) + require.NotEmpty(t, request["timestamp"]) + require.NotEmpty(t, request["signature"]) + requestCount++ + if requestCount == 2 { + digest := sha512.Sum512(key.privateKey.Seed()) + var curvePrivate [32]byte + copy(curvePrivate[:], digest[:32]) + curvePrivate[0] &= 248 + curvePrivate[31] &= 127 + curvePrivate[31] |= 64 + curvePublicBytes, curveErr := curve25519.X25519(curvePrivate[:], curve25519.Basepoint) + require.NoError(t, curveErr) + var curvePublic [32]byte + copy(curvePublic[:], curvePublicBytes) + ciphertext, sealErr := box.SealAnonymous(nil, []byte("task-encrypted"), &curvePublic, rand.Reader) + require.NoError(t, sealErr) + _, _ = fmt.Fprintf(w, `{"encrypted_task_id":%q}`, base64.StdEncoding.EncodeToString(ciphertext)) + return + } + _, _ = w.Write([]byte(`{"task_id":"task-plain"}`)) + })) + defer server.Close() + oldBase := openAIAgentIdentityAuthAPIBaseURL + openAIAgentIdentityAuthAPIBaseURL = server.URL + t.Cleanup(func() { openAIAgentIdentityAuthAPIBaseURL = oldBase }) + + account := &Account{ID: 1, Type: AccountTypeOAuth, Platform: PlatformOpenAI, Credentials: map[string]any{ + "auth_mode": OpenAIAuthModeAgentIdentity, + "agent_runtime_id": key.runtimeID, + "agent_private_key": privateKey, + }} + taskID, err := registerAgentIdentityTask(context.Background(), account) + require.NoError(t, err) + require.Equal(t, "task-plain", taskID) + taskID, err = registerAgentIdentityTask(context.Background(), account) + require.NoError(t, err) + require.Equal(t, "task-encrypted", taskID) +} + +func TestEnsureAgentIdentityTaskPersistsAndRedactsCredentials(t *testing.T) { + key, privateKey := newTestAgentIdentityKey(t) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, _ = w.Write([]byte(`{"task_id":"task-persisted"}`)) + })) + defer server.Close() + oldBase := openAIAgentIdentityAuthAPIBaseURL + openAIAgentIdentityAuthAPIBaseURL = server.URL + t.Cleanup(func() { openAIAgentIdentityAuthAPIBaseURL = oldBase }) + + repo := &agentIdentityCredentialsRepo{} + account := &Account{ID: 7, Type: AccountTypeOAuth, Platform: PlatformOpenAI, Credentials: map[string]any{ + "auth_mode": OpenAIAuthModeAgentIdentity, + "agent_runtime_id": key.runtimeID, + "agent_private_key": privateKey, + "chatgpt_account_id": "account-test", + }} + service := &OpenAIGatewayService{accountRepo: repo} + require.NoError(t, service.ensureAgentIdentityTask(context.Background(), account, "")) + require.Equal(t, "task-persisted", account.GetCredential("task_id")) + require.Equal(t, "task-persisted", repo.credentials["task_id"]) + require.True(t, IsSensitiveCredentialKey("agent_private_key")) + redacted := make(map[string]any) + for key, value := range account.Credentials { + if !IsSensitiveCredentialKey(key) { + redacted[key] = value + } + } + require.NotContains(t, string(mustAgentIdentityJSON(t, redacted)), privateKey) +} + +func TestEnsureAgentIdentityTaskSharesLockAcrossServicesForSameAccount(t *testing.T) { + key, privateKey := newTestAgentIdentityKey(t) + account := &Account{ID: 9001, Type: AccountTypeOAuth, Platform: PlatformOpenAI, Credentials: map[string]any{ + "auth_mode": OpenAIAuthModeAgentIdentity, + "agent_runtime_id": key.runtimeID, + "agent_private_key": privateKey, + }} + repo := &agentIdentityCredentialsRepo{account: account} + registerCalls := 0 + var registerMu sync.Mutex + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + registerMu.Lock() + registerCalls++ + registerMu.Unlock() + _, _ = w.Write([]byte(`{"task_id":"task-shared"}`)) + })) + defer server.Close() + oldBase := openAIAgentIdentityAuthAPIBaseURL + openAIAgentIdentityAuthAPIBaseURL = server.URL + t.Cleanup(func() { openAIAgentIdentityAuthAPIBaseURL = oldBase }) + + start := make(chan struct{}) + errors := make(chan error, 2) + requests := []*Account{cloneAgentIdentityTestAccount(account), cloneAgentIdentityTestAccount(account)} + for _, request := range requests { + go func() { + <-start + errors <- ensureAgentIdentityTaskForAccount(context.Background(), repo, nil, &sync.Mutex{}, request, "") + }() + } + close(start) + require.NoError(t, <-errors) + require.NoError(t, <-errors) + registerMu.Lock() + defer registerMu.Unlock() + require.Equal(t, 1, registerCalls) + require.Equal(t, "task-shared", repo.account.GetCredential("task_id")) +} + +func cloneAgentIdentityTestAccount(account *Account) *Account { + copy := *account + copy.Credentials = shallowCopyMap(account.Credentials) + return © +} + +type agentIdentityCredentialsRepo struct { + AccountRepository + credentials map[string]any + account *Account + mu sync.Mutex +} + +func (r *agentIdentityCredentialsRepo) GetByID(_ context.Context, _ int64) (*Account, error) { + return r.account, nil +} + +func (r *agentIdentityCredentialsRepo) UpdateCredentials(_ context.Context, _ int64, credentials map[string]any) error { + r.mu.Lock() + defer r.mu.Unlock() + r.credentials = credentials + return nil +} + +func mustAgentIdentityJSON(t *testing.T, value any) []byte { + t.Helper() + encoded, err := json.Marshal(value) + require.NoError(t, err) + return encoded +} diff --git a/backend/internal/service/openai_codex_models_service.go b/backend/internal/service/openai_codex_models_service.go index a7331f64c9..1d9356441c 100644 --- a/backend/internal/service/openai_codex_models_service.go +++ b/backend/internal/service/openai_codex_models_service.go @@ -241,7 +241,7 @@ func (s *OpenAIGatewayService) FetchCodexModelsManifest(ctx context.Context, acc switch { case credAccount.IsOpenAIOAuth(): authToken = strings.TrimSpace(credAccount.GetOpenAIAccessToken()) - if authToken == "" { + if authToken == "" && !credAccount.IsOpenAIAgentIdentity() { return nil, infraerrors.New(http.StatusBadGateway, "OPENAI_CODEX_MODELS_TOKEN_MISSING", "account has no Codex backend access token") } case credAccount.IsOpenAIApiKey(): @@ -277,16 +277,25 @@ func (s *OpenAIGatewayService) FetchCodexModelsManifest(ctx context.Context, acc } headers := make(http.Header) - headers.Set("Authorization", "Bearer "+authToken) + if useAPIKeyUpstream { + headers.Set("Authorization", "Bearer "+authToken) + credAccount.ApplyHeaderOverrides(headers) + } else { + authHeaders, authErr := s.buildOpenAIAuthenticationHeaders(ctx, credAccount, authToken) + if authErr != nil { + return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_CODEX_MODELS_AUTH_FAILED", "build Codex models authentication: %v", authErr) + } + for key, values := range authHeaders { + for _, value := range values { + headers.Add(key, value) + } + } + setOpenAIChatGPTAccountHeaders(headers, credAccount) + } headers.Set("Accept", "application/json") headers.Set("Originator", "codex_cli_rs") headers.Set("Version", clientVersion) headers.Set("User-Agent", codexCLIUserAgent) - if useAPIKeyUpstream { - credAccount.ApplyHeaderOverrides(headers) - } else { - setOpenAIChatGPTAccountHeaders(headers, credAccount) - } proxyURL := "" if account.ProxyID != nil && account.Proxy != nil { diff --git a/backend/internal/service/openai_codex_models_service_test.go b/backend/internal/service/openai_codex_models_service_test.go index 2c87861f60..3f4ec96a2f 100644 --- a/backend/internal/service/openai_codex_models_service_test.go +++ b/backend/internal/service/openai_codex_models_service_test.go @@ -199,6 +199,49 @@ func TestFetchCodexModelsManifestPassthrough(t *testing.T) { } } +func TestFetchCodexModelsManifestAgentIdentityUsesAssertionWithoutOAuthToken(t *testing.T) { + key, privateKey := newTestAgentIdentityKey(t) + account := &Account{ + ID: 3, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "auth_mode": OpenAIAuthModeAgentIdentity, + "agent_runtime_id": key.runtimeID, + "agent_private_key": privateKey, + "task_id": key.taskID, + "chatgpt_account_id": "acc-agent", + }, + } + + var gotAuth, gotAccountID string + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + gotAuth = r.Header.Get("Authorization") + gotAccountID = r.Header.Get("chatgpt-account-id") + _, _ = w.Write([]byte(`{"models":[]}`)) + })) + defer server.Close() + + original := chatgptCodexModelsURL + chatgptCodexModelsURL = server.URL + defer func() { chatgptCodexModelsURL = original }() + + s := &OpenAIGatewayService{} + manifest, err := s.FetchCodexModelsManifest(context.Background(), account, "0.137.0", "") + if err != nil { + t.Fatalf("FetchCodexModelsManifest returned error: %v", err) + } + if string(manifest.Body) != `{"models":[]}` { + t.Fatalf("unexpected manifest body: %q", manifest.Body) + } + if !strings.HasPrefix(gotAuth, "AgentAssertion ") { + t.Fatalf("authorization scheme: got %q", strings.SplitN(gotAuth, " ", 2)[0]) + } + if gotAccountID != "acc-agent" { + t.Fatalf("chatgpt-account-id header: got %q", gotAccountID) + } +} + func TestFetchCodexModelsManifestDefaultClientVersion(t *testing.T) { var gotClientVersion string server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { diff --git a/backend/internal/service/openai_content_session_seed.go b/backend/internal/service/openai_content_session_seed.go index fce85f11bd..5f8f1872e8 100644 --- a/backend/internal/service/openai_content_session_seed.go +++ b/backend/internal/service/openai_content_session_seed.go @@ -1,6 +1,7 @@ package service import ( + "bytes" "encoding/json" "strings" @@ -25,31 +26,96 @@ func deriveOpenAIContentSessionSeed(body []byte) string { return "" } + const ( + modelField = iota + toolsField + functionsField + instructionsField + messagesField + inputField + contentSessionSeedFieldCount + allContentSessionSeedFields = 1< len(root) { + return false + } + // Result.ForEach can continue after the root '}' on malformed input. + // The separator range excludes braces inside the preceding parsed value. + if bytes.IndexByte(root[nextKeyOffset:key.Index], '}') >= 0 { + return false + } + nextKeyOffset = value.Index + len(value.Raw) + + field := -1 + switch key.Str { + case "model": + field = modelField + case "tools": + field = toolsField + case "functions": + field = functionsField + case "instructions": + field = instructionsField + case "messages": + field = messagesField + case "input": + field = inputField + } + if field < 0 { + return true + } + mask := uint8(1 << field) + if seen&mask == 0 { + fields[field] = value + seen |= mask + } + return seen != allContentSessionSeedFields + }) + var b strings.Builder - if model := gjson.GetBytes(body, "model").String(); model != "" { + if model := fields[modelField].String(); model != "" { _, _ = b.WriteString("model=") _, _ = b.WriteString(model) } - if tools := gjson.GetBytes(body, "tools"); tools.Exists() && tools.IsArray() && tools.Raw != "[]" { + if tools := fields[toolsField]; tools.Exists() && tools.IsArray() && tools.Raw != "[]" { _, _ = b.WriteString("|tools=") _, _ = b.WriteString(normalizeCompatSeedJSON(json.RawMessage(tools.Raw))) } - if funcs := gjson.GetBytes(body, "functions"); funcs.Exists() && funcs.IsArray() && funcs.Raw != "[]" { + if funcs := fields[functionsField]; funcs.Exists() && funcs.IsArray() && funcs.Raw != "[]" { _, _ = b.WriteString("|functions=") _, _ = b.WriteString(normalizeCompatSeedJSON(json.RawMessage(funcs.Raw))) } - if instr := gjson.GetBytes(body, "instructions").String(); instr != "" { + if instr := fields[instructionsField].String(); instr != "" { _, _ = b.WriteString("|instructions=") _, _ = b.WriteString(instr) } firstUserCaptured := false - msgs := gjson.GetBytes(body, "messages") + msgs := fields[messagesField] if msgs.Exists() && msgs.IsArray() { msgs.ForEach(func(_, msg gjson.Result) bool { role := msg.Get("role").String() @@ -70,7 +136,7 @@ func deriveOpenAIContentSessionSeed(body []byte) string { } return true }) - } else if inp := gjson.GetBytes(body, "input"); inp.Exists() { + } else if inp := fields[inputField]; inp.Exists() { if inp.Type == gjson.String { _, _ = b.WriteString("|input=") _, _ = b.WriteString(inp.String()) diff --git a/backend/internal/service/openai_content_session_seed_benchmark_test.go b/backend/internal/service/openai_content_session_seed_benchmark_test.go new file mode 100644 index 0000000000..eda9deaa82 --- /dev/null +++ b/backend/internal/service/openai_content_session_seed_benchmark_test.go @@ -0,0 +1,29 @@ +package service + +import ( + "strings" + "testing" +) + +var benchmarkOpenAIContentSessionSeed string + +func BenchmarkDeriveOpenAIContentSessionSeedLargeBody(b *testing.B) { + largeHistory := strings.Repeat("payload", 1<<17) + tests := []struct { + name string + body []byte + }{ + {name: "ChatCompletions", body: []byte(`{"model":"gpt-5.4","tools":[{"type":"function","function":{"name":"lookup"}}],"messages":[{"role":"system","content":"Be concise."},{"role":"user","content":"Hello"},{"role":"assistant","content":"` + largeHistory + `"},{"role":"user","content":"Follow-up"}]}`)}, + {name: "Responses", body: []byte(`{"model":"gpt-5.4","instructions":"Be concise.","tools":[{"type":"function","name":"lookup"}],"input":[{"role":"system","content":"System prompt"},{"role":"user","content":"Hello"},{"role":"assistant","content":"` + largeHistory + `"}]}`)}, + } + + for _, test := range tests { + b.Run(test.name, func(b *testing.B) { + b.ReportAllocs() + b.SetBytes(int64(len(test.body))) + for range b.N { + benchmarkOpenAIContentSessionSeed = deriveOpenAIContentSessionSeed(test.body) + } + }) + } +} diff --git a/backend/internal/service/openai_content_session_seed_test.go b/backend/internal/service/openai_content_session_seed_test.go index 6dadc5cf53..0c6c63335c 100644 --- a/backend/internal/service/openai_content_session_seed_test.go +++ b/backend/internal/service/openai_content_session_seed_test.go @@ -1,9 +1,12 @@ package service import ( + "encoding/json" + "strings" "testing" "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" ) func TestDeriveOpenAIContentSessionSeed_EmptyInputs(t *testing.T) { @@ -197,6 +200,184 @@ func TestDeriveOpenAIContentSessionSeed_JSONCanonicalisation(t *testing.T) { require.Equal(t, s1, s2, "different formatting of identical JSON should produce the same seed") } +func TestDeriveOpenAIContentSessionSeed_SingleScanMatchesLegacyBytes(t *testing.T) { + largeValue := strings.Repeat("payload", 1<<17) + tests := []struct { + name string + body []byte + }{ + { + name: "large chat completions", + body: []byte(`{"metadata":"` + largeValue + `","model":"gpt-5.4","tools":[{"type":"function","function":{"name":"lookup"}}],"functions":[{"name":"legacy_lookup"}],"messages":[{"role":"system","content":"System prompt"},{"role":"developer","content":[{"type":"text","text":"Developer prompt"}]},{"role":"user","content":"Hello"}]}`), + }, + { + name: "large responses", + body: []byte(`{"metadata":"` + largeValue + `","model":"gpt-5.4","instructions":"Be concise.","tools":[{"type":"function","name":"lookup"}],"input":[{"role":"system","content":"System prompt"},{"role":"user","content":[{"type":"input_text","text":"Hello"}]}]}`), + }, + { + name: "fields in reverse order", + body: []byte(`{"input":"fallback input","messages":[{"role":"user","content":"chat wins"}],"instructions":"Be concise.","functions":[{"name":"lookup"}],"tools":[{"type":"function","name":"lookup"}],"model":"gpt-5.4"}`), + }, + { + name: "missing and wrong type fields", + body: []byte(`{"tools":[],"functions":null,"instructions":0,"messages":{},"input":[{"type":"input_text","text":"fallback"}]}`), + }, + { + name: "duplicate fields keep first value", + body: []byte(`{"model":"first","model":"second","tools":[{"name":"first"}],"tools":[{"name":"second"}],"functions":[],"functions":[{"name":"second"}],"instructions":"first","instructions":"second","messages":null,"messages":[{"role":"user","content":"second"}],"input":"first input","input":"second input"}`), + }, + { + name: "escaped field names", + body: []byte(`{"mo\u0064el":"gpt-5.4","mess\u0061ges":[{"role":"user","content":"Hello"}]}`), + }, + { + name: "trailing object fields are outside the root", + body: []byte(`{"foo":1}{"model":"trailing","input":"trailing input"}`), + }, + { + name: "trailing quoted fields are outside the root", + body: []byte(`{"model":"root"}"input":"trailing input"`), + }, + { + name: "leading garbage before the root", + body: []byte(`garbage{"model":"gpt-5.4","input":"Hello"}`), + }, + { + name: "escaped braces remain inside string values", + body: []byte(`{"metadata":"escaped } and [ and \" quote","model":"root","input":"Hello"}{"model":"trailing"}`), + }, + { + name: "nested braces do not end the root", + body: []byte(`{"metadata":{"nested":"} ]"},"model":"root","input":"Hello"}{"model":"trailing"}`), + }, + { + name: "root array does not expose nested or trailing object fields", + body: []byte(`[{"model":"nested"}]{"model":"trailing","input":"trailing input"}`), + }, + { + name: "trailing messages do not override root input", + body: []byte(`{"input":"root input"}{"messages":[{"role":"user","content":"trailing"}]}`), + }, + { + name: "truncated string containing a closing brace", + body: []byte(`{"model":"root","metadata":"still } inside`), + }, + { + name: "lenient truncated body", + body: []byte(`{"model":"gpt-5.4","messages":[{"role":"user","content":"Hello"}]`), + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + require.Equal(t, legacyDeriveOpenAIContentSessionSeed(test.body), deriveOpenAIContentSessionSeed(test.body)) + }) + } +} + +func TestDeriveOpenAIContentSessionSeed_AllTruncationOffsetsMatchLegacyBytes(t *testing.T) { + bodies := []string{ + `{"model":"gpt-5.4","tools":[{"type":"function","function":{"name":"lookup"}}],"functions":[{"name":"legacy"}],"instructions":"escaped \" } text","messages":[{"role":"system","content":"System"},{"role":"user","content":[{"type":"text","text":"Hello"}]}],"input":"fallback"}`, + `{"model":"gpt-5.4","instructions":"Be concise.","tools":[{"type":"function","name":"lookup"}],"input":[{"role":"system","content":"System"},{"role":"user","content":[{"type":"input_text","text":"Hello"}]}]}`, + } + for bodyIndex, body := range bodies { + for end := 1; end < len(body); end++ { + truncated := []byte(body[:end]) + require.Equalf(t, legacyDeriveOpenAIContentSessionSeed(truncated), deriveOpenAIContentSessionSeed(truncated), "body %d truncated at byte %d", bodyIndex, end) + } + } +} + +func legacyDeriveOpenAIContentSessionSeed(body []byte) string { + if len(body) == 0 { + return "" + } + + var b strings.Builder + + if model := gjson.GetBytes(body, "model").String(); model != "" { + _, _ = b.WriteString("model=") + _, _ = b.WriteString(model) + } + + if tools := gjson.GetBytes(body, "tools"); tools.Exists() && tools.IsArray() && tools.Raw != "[]" { + _, _ = b.WriteString("|tools=") + _, _ = b.WriteString(normalizeCompatSeedJSON(json.RawMessage(tools.Raw))) + } + + if funcs := gjson.GetBytes(body, "functions"); funcs.Exists() && funcs.IsArray() && funcs.Raw != "[]" { + _, _ = b.WriteString("|functions=") + _, _ = b.WriteString(normalizeCompatSeedJSON(json.RawMessage(funcs.Raw))) + } + + if instr := gjson.GetBytes(body, "instructions").String(); instr != "" { + _, _ = b.WriteString("|instructions=") + _, _ = b.WriteString(instr) + } + + firstUserCaptured := false + + msgs := gjson.GetBytes(body, "messages") + if msgs.Exists() && msgs.IsArray() { + msgs.ForEach(func(_, msg gjson.Result) bool { + role := msg.Get("role").String() + switch role { + case "system", "developer": + _, _ = b.WriteString("|system=") + if c := msg.Get("content"); c.Exists() { + _, _ = b.WriteString(normalizeCompatSeedJSON(json.RawMessage(c.Raw))) + } + case "user": + if !firstUserCaptured { + _, _ = b.WriteString("|first_user=") + if c := msg.Get("content"); c.Exists() { + _, _ = b.WriteString(normalizeCompatSeedJSON(json.RawMessage(c.Raw))) + } + firstUserCaptured = true + } + } + return true + }) + } else if inp := gjson.GetBytes(body, "input"); inp.Exists() { + if inp.Type == gjson.String { + _, _ = b.WriteString("|input=") + _, _ = b.WriteString(inp.String()) + } else if inp.IsArray() { + inp.ForEach(func(_, item gjson.Result) bool { + role := item.Get("role").String() + switch role { + case "system", "developer": + _, _ = b.WriteString("|system=") + if c := item.Get("content"); c.Exists() { + _, _ = b.WriteString(normalizeCompatSeedJSON(json.RawMessage(c.Raw))) + } + case "user": + if !firstUserCaptured { + _, _ = b.WriteString("|first_user=") + if c := item.Get("content"); c.Exists() { + _, _ = b.WriteString(normalizeCompatSeedJSON(json.RawMessage(c.Raw))) + } + firstUserCaptured = true + } + } + if !firstUserCaptured && item.Get("type").String() == "input_text" { + _, _ = b.WriteString("|first_user=") + if text := item.Get("text").String(); text != "" { + _, _ = b.WriteString(text) + } + firstUserCaptured = true + } + return true + }) + } + } + + if b.Len() == 0 { + return "" + } + return contentSessionSeedPrefix + b.String() +} + func TestDeriveOpenAIContentSessionSeed_ResponsesAPI_InputTextTypedItem(t *testing.T) { body := []byte(`{ "model": "gpt-5.4", diff --git a/backend/internal/service/openai_gateway_cc_pipeline.go b/backend/internal/service/openai_gateway_cc_pipeline.go index 690b9991bb..2391e25ba5 100644 --- a/backend/internal/service/openai_gateway_cc_pipeline.go +++ b/backend/internal/service/openai_gateway_cc_pipeline.go @@ -118,6 +118,7 @@ func (s *OpenAIGatewayService) failoverOpenAIUpstreamHTTPError( return &UpstreamFailoverError{ StatusCode: resp.StatusCode, ResponseBody: respBody, + ResponseHeaders: resp.Header.Clone(), RetryableOnSameAccount: account.IsPoolMode() && (account.IsPoolModeRetryableStatus(resp.StatusCode) || isOpenAITransientProcessingError(resp.StatusCode, upstreamMsg, respBody)), } } @@ -197,7 +198,9 @@ func (s *OpenAIGatewayService) sendCCUpstreamRequest( // 账号级请求头覆写(仅 openai api_key 账号启用时生效) account.ApplyHeaderOverrides(upstreamReq.Header) if account.Platform == PlatformGrok { - applyGrokCLIHeaders(upstreamReq.Header) + if account.IsGrokOAuth() { + applyGrokCLIHeaders(upstreamReq.Header) + } applyGrokCacheHeaders(upstreamReq.Header, grokCacheIdentity) } diff --git a/backend/internal/service/openai_gateway_chat_completions.go b/backend/internal/service/openai_gateway_chat_completions.go index 533b562dc9..fe24b3db36 100644 --- a/backend/internal/service/openai_gateway_chat_completions.go +++ b/backend/internal/service/openai_gateway_chat_completions.go @@ -277,6 +277,13 @@ func (s *OpenAIGatewayService) ForwardAsChatCompletions( // 8. Handle error response with failover if resp.StatusCode >= 400 { respBody, upstreamMsg := s.readOpenAIUpstreamError(resp) + if !agentIdentityTaskRecoveryWasTried(ctx) && s.isAgentIdentityAccount(ctx, account) && isAgentIdentityTaskInvalidHTTPResponse(resp.StatusCode, respBody) { + expectedTaskID := account.GetCredential("task_id") + if err := s.recoverAgentIdentityTask(ctx, account, expectedTaskID); err != nil { + return nil, fmt.Errorf("agent identity task recovery failed: %w", err) + } + return s.ForwardAsChatCompletions(markAgentIdentityTaskRecoveryTried(ctx), c, account, body, promptCacheKey, defaultMappedModel) + } if account.Type == AccountTypeAPIKey && openai_compat.ResolveResponsesSupport(account.Extra) == openai_compat.ResponsesSupportUnknown && !isResponsesEndpointSupportedByStatus(resp.StatusCode) { diff --git a/backend/internal/service/openai_gateway_chat_completions_raw.go b/backend/internal/service/openai_gateway_chat_completions_raw.go index 7693986b0d..800d53be32 100644 --- a/backend/internal/service/openai_gateway_chat_completions_raw.go +++ b/backend/internal/service/openai_gateway_chat_completions_raw.go @@ -9,7 +9,6 @@ import ( "time" "github.com/Wei-Shaw/sub2api/internal/pkg/logger" - "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/Wei-Shaw/sub2api/internal/util/responseheaders" "github.com/gin-gonic/gin" "github.com/tidwall/gjson" @@ -109,7 +108,7 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( // Grok Composer does not accept image_url parts directly, but Grok Build // can describe the images first. Bridge only this exact failure mode. - token, tokenKind, err := s.GetAccessToken(ctx, account) + token, tokenKind, err := s.getRequestCredential(ctx, c, account) if err != nil { return nil, err } @@ -162,7 +161,7 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( } SetActualOpenAIUpstreamEndpoint(c, grokChatRawEndpoint) customUA := account.GetOpenAIUserAgent() - if customUA == "" && account.Platform == PlatformGrok { + if customUA == "" && account.IsGrokOAuth() { customUA = "sub2api-grok/1.0" } resp, err := s.sendCCUpstreamRequest(ctx, c, account, targetURL, upstreamBody, clientStream, token, customUA, grokCacheIdentity) @@ -189,6 +188,7 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( return nil, &UpstreamFailoverError{ StatusCode: resp.StatusCode, ResponseBody: respBody, + ResponseHeaders: resp.Header.Clone(), RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode), } } @@ -201,7 +201,7 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( } if account.Platform == PlatformGrok { - s.updateGrokUsageSnapshot(ctx, account, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) + s.updateGrokUsageFromResponse(ctx, account, resp.Header, resp.StatusCode) } // 8. Forward response @@ -221,7 +221,7 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( func (s *OpenAIGatewayService) rawChatCompletionsURL(account *Account) (string, error) { if account.Platform == PlatformGrok { - targetURL, err := xai.BuildChatCompletionsURL(account.GetGrokBaseURL()) + targetURL, err := buildGrokChatCompletionsURL(account, s.cfg) if err != nil { return "", fmt.Errorf("invalid grok base_url: %w", err) } diff --git a/backend/internal/service/openai_gateway_count_tokens.go b/backend/internal/service/openai_gateway_count_tokens.go index 7518a6073a..46b5a3843a 100644 --- a/backend/internal/service/openai_gateway_count_tokens.go +++ b/backend/internal/service/openai_gateway_count_tokens.go @@ -215,7 +215,15 @@ func (s *OpenAIGatewayService) buildInputTokensUpstreamRequest( return nil, err } req = req.WithContext(WithHTTPUpstreamProfile(req.Context(), HTTPUpstreamProfileOpenAI)) - req.Header.Set("authorization", "Bearer "+token) + authHeaders, err := s.buildOpenAIAuthenticationHeaders(ctx, account, token) + if err != nil { + return nil, err + } + for key, values := range authHeaders { + for _, value := range values { + req.Header.Add(key, value) + } + } req.Header.Set("content-type", "application/json") req.Header.Set("accept", "application/json") diff --git a/backend/internal/service/openai_gateway_forward.go b/backend/internal/service/openai_gateway_forward.go index 6c7d45946a..509eab8d71 100644 --- a/backend/internal/service/openai_gateway_forward.go +++ b/backend/internal/service/openai_gateway_forward.go @@ -499,6 +499,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco var wsResult *OpenAIForwardResult var wsErr error wsLastFailureReason := "" + agentTaskRecoveryTried := false wsPrevResponseRecoveryTried := false wsInvalidEncryptedContentRecoveryTried := false recoverPrevResponseNotFound := func(attempt int) bool { @@ -583,6 +584,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco startTime, attempt, wsLastFailureReason, + &agentTaskRecoveryTried, ) if wsErr == nil { break @@ -590,6 +592,10 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco if c != nil && c.Writer != nil && c.Writer.Written() { break } + var taskRecoveredErr *agentIdentityTaskRecoveredError + if errors.As(wsErr, &taskRecoveredErr) { + continue + } reason, retryable := classifyOpenAIWSReconnectReason(wsErr) if reason != "" { @@ -696,6 +702,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco } httpInvalidEncryptedContentRetryTried := false + agentTaskRecoveryTried := false for { // Build upstream request upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) @@ -731,6 +738,16 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(respBody)) upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg) upstreamCode := extractUpstreamErrorCode(respBody) + if !agentTaskRecoveryTried && s.isAgentIdentityAccount(ctx, account) && isAgentIdentityTaskInvalidHTTPResponse(resp.StatusCode, respBody) { + agentTaskRecoveryTried = true + expectedTaskID := account.GetCredential("task_id") + if err := s.recoverAgentIdentityTask(ctx, account, expectedTaskID); err != nil { + return nil, fmt.Errorf("agent identity task recovery failed: %w", err) + } + continue + } + respBody = s.redactAgentIdentitySensitiveBody(ctx, account, respBody) + resp.Body = io.NopCloser(bytes.NewReader(respBody)) if !httpInvalidEncryptedContentRetryTried && resp.StatusCode == http.StatusBadRequest && upstreamCode == "invalid_encrypted_content" { decoded, decodeErr := ensureReqBody() if decodeErr != nil { @@ -881,8 +898,17 @@ func (s *OpenAIGatewayService) buildUpstreamRequest(ctx context.Context, c *gin. } req = req.WithContext(WithHTTPUpstreamProfile(req.Context(), HTTPUpstreamProfileOpenAI)) - // Set authentication header - req.Header.Set("authorization", "Bearer "+token) + // Build authentication for this request. Agent Identity signs a fresh + // assertion here; OAuth/PAT/API-key keep their existing Bearer behavior. + authHeaders, err := s.buildOpenAIAuthenticationHeaders(ctx, account, token) + if err != nil { + return nil, fmt.Errorf("build openai authentication headers: %w", err) + } + for key, values := range authHeaders { + for _, value := range values { + req.Header.Add(key, value) + } + } // Set headers specific to OAuth accounts (ChatGPT internal API) if account.Type == AccountTypeOAuth { diff --git a/backend/internal/service/openai_gateway_grok.go b/backend/internal/service/openai_gateway_grok.go index 8bd1d456dc..92cf2910c6 100644 --- a/backend/internal/service/openai_gateway_grok.go +++ b/backend/internal/service/openai_gateway_grok.go @@ -11,6 +11,7 @@ import ( "strings" "time" + "github.com/Wei-Shaw/sub2api/internal/config" "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/gin-gonic/gin" @@ -25,6 +26,10 @@ const ( grokCLIVersion = "0.2.93" grokDefaultResponsesModel = "grok-4.5" grokRateLimitFallbackCooldown = 2 * time.Minute + grokRateLimitRepeatCooldown = 10 * time.Minute + grokRateLimitSustainedCooldown = 30 * time.Minute + grokRateLimitMaxAdaptiveCooldown = time.Hour + grokRateLimitBackoffQuietPeriod = time.Hour ) func (s *OpenAIGatewayService) forwardGrokResponses( @@ -54,14 +59,14 @@ func (s *OpenAIGatewayService) forwardGrokResponses( return nil, fmt.Errorf("apply grok prompt cache identity: %w", err) } - token, _, err := s.GetAccessToken(ctx, account) + token, _, err := s.getRequestCredential(ctx, c, account) if err != nil { return nil, err } upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) defer releaseUpstreamCtx() - upstreamReq, err := buildGrokResponsesRequest(upstreamCtx, c, account, patchedBody, token, cacheIdentity) + upstreamReq, err := buildGrokResponsesRequest(upstreamCtx, c, account, patchedBody, token, cacheIdentity, s.cfg) if err != nil { return nil, err } @@ -100,13 +105,14 @@ func (s *OpenAIGatewayService) forwardGrokResponses( return nil, &UpstreamFailoverError{ StatusCode: resp.StatusCode, ResponseBody: respBody, + ResponseHeaders: resp.Header.Clone(), RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode), } } return s.handleErrorResponse(ctx, resp, c, account, patchedBody, upstreamModel) } - s.updateGrokUsageSnapshot(ctx, account, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) + s.updateGrokUsageFromResponse(ctx, account, resp.Header, resp.StatusCode) var usage *OpenAIUsage var firstTokenMs *int @@ -573,7 +579,7 @@ func (s *OpenAIGatewayService) describeGrokComposerImage( upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) // Image-description probes are auxiliary requests, not conversation turns. // Do not bind them to the caller's Grok prompt-cache identity. - upstreamReq, err := buildGrokResponsesRequest(upstreamCtx, c, account, body, token, "") + upstreamReq, err := buildGrokResponsesRequest(upstreamCtx, c, account, body, token, "", s.cfg) releaseUpstreamCtx() if err != nil { return "", OpenAIUsage{}, fmt.Errorf("build grok composer image bridge request: %w", err) @@ -610,13 +616,14 @@ func (s *OpenAIGatewayService) describeGrokComposerImage( return "", OpenAIUsage{}, &UpstreamFailoverError{ StatusCode: resp.StatusCode, ResponseBody: respBody, + ResponseHeaders: resp.Header.Clone(), RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode), } } return "", OpenAIUsage{}, fmt.Errorf("grok composer image bridge upstream error: %s", upstreamMsg) } - s.updateGrokUsageSnapshot(ctx, account, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) + s.updateGrokUsageFromResponse(ctx, account, resp.Header, resp.StatusCode) respBody, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, nil) if err != nil { return "", OpenAIUsage{}, fmt.Errorf("read grok composer image bridge response: %w", err) @@ -738,8 +745,8 @@ func addOpenAIUsage(dst *OpenAIUsage, usage OpenAIUsage) { dst.ImageOutputTokens += usage.ImageOutputTokens } -func buildGrokResponsesRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token, cacheIdentity string) (*http.Request, error) { - targetURL, err := xai.BuildResponsesURL(account.GetGrokBaseURL()) +func buildGrokResponsesRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token, cacheIdentity string, cfg *config.Config) (*http.Request, error) { + targetURL, err := buildGrokResponsesURL(account, cfg) if err != nil { return nil, err } @@ -750,7 +757,9 @@ func buildGrokResponsesRequest(ctx context.Context, c *gin.Context, account *Acc req.Header.Set("Authorization", "Bearer "+token) req.Header.Set("Content-Type", "application/json") req.Header.Set("Accept", "application/json, text/event-stream") - applyGrokCLIHeaders(req.Header) + if account.IsGrokOAuth() { + applyGrokCLIHeaders(req.Header) + } applyGrokCacheHeaders(req.Header, cacheIdentity) if c != nil { if v := c.GetHeader("OpenAI-Beta"); strings.TrimSpace(v) != "" { @@ -776,11 +785,12 @@ func (s *OpenAIGatewayService) updateGrokUsageSnapshot(ctx context.Context, acco } accountID := account.ID now := time.Now() - resetAt, hasActiveLimit := grokRateLimitResetAt(snapshot, now) + resetAt, hasActiveLimit := grokRateLimitResetAtForAccount(account, snapshot, now) if hasActiveLimit { normalizeGrokExhaustedWindowResets(snapshot, resetAt, now) } - critical := snapshot.StatusCode == http.StatusTooManyRequests || hasActiveLimit + recovery := isSuccessfulGrokRateLimitRecovery(account, snapshot) + critical := snapshot.StatusCode == http.StatusTooManyRequests || hasActiveLimit || recovery if s.codexSnapshotThrottle != nil { allowed := s.codexSnapshotThrottle.Allow(accountID, now) if !critical && !allowed { @@ -806,6 +816,23 @@ func (s *OpenAIGatewayService) updateGrokUsageSnapshot(ctx context.Context, acco // on the passive snapshot scheduler check. if hasActiveLimit { s.rateLimitGrok(stateCtx, account, resetAt) + } else if recovery { + clearGrokRateLimitAfterRecovery(stateCtx, s.accountRepo, account) + } +} + +func (s *OpenAIGatewayService) updateGrokUsageFromResponse(ctx context.Context, account *Account, headers http.Header, statusCode int) { + snapshot := parseGrokQuotaSnapshot(headers, statusCode, time.Now()) + if snapshot != nil { + s.updateGrokUsageSnapshot(ctx, account, snapshot) + return + } + // Successful responses are recovery evidence even when the upstream omits + // optional quota headers. Do not replace an informative stored snapshot with + // an empty one; only clear the exact observed cooldown generation. + recoverySnapshot := &xai.QuotaSnapshot{StatusCode: statusCode} + if isSuccessfulGrokRateLimitRecovery(account, recoverySnapshot) { + clearGrokRateLimitAfterRecovery(ctx, s.accountRepo, account) } } @@ -896,6 +923,37 @@ func grokRateLimitResetAt(snapshot *xai.QuotaSnapshot, now time.Time) (time.Time return time.Time{}, false } +func grokRateLimitResetAtForAccount(account *Account, snapshot *xai.QuotaSnapshot, now time.Time) (time.Time, bool) { + resetAt, limited := grokRateLimitResetAt(snapshot, now) + if !limited || !isGrokOAuthAccount(account) || snapshot == nil || snapshot.StatusCode != http.StatusTooManyRequests { + return resetAt, limited + } + if account.RateLimitedAt == nil || account.RateLimitResetAt == nil { + return resetAt, true + } + previousResetAt := *account.RateLimitResetAt + if previousResetAt.After(now) || now.Sub(previousResetAt) > grokRateLimitBackoffQuietPeriod { + return resetAt, true + } + previousCooldown := previousResetAt.Sub(*account.RateLimitedAt) + if previousCooldown <= 0 { + return resetAt, true + } + + adaptiveCooldown := grokRateLimitRepeatCooldown + switch { + case previousCooldown >= grokRateLimitSustainedCooldown: + adaptiveCooldown = grokRateLimitMaxAdaptiveCooldown + case previousCooldown >= grokRateLimitRepeatCooldown: + adaptiveCooldown = grokRateLimitSustainedCooldown + } + adaptiveResetAt := now.Add(adaptiveCooldown) + if adaptiveResetAt.After(resetAt) { + resetAt = adaptiveResetAt + } + return resetAt, true +} + func normalizeGrokRateLimitResetAt(account *Account, resetAt, now time.Time) time.Time { if !resetAt.After(now) { resetAt = now.Add(grokRateLimitFallbackCooldown) @@ -910,6 +968,33 @@ type grokRateLimitExtendingRepository interface { SetRateLimitedIfLater(ctx context.Context, id int64, resetAt time.Time) error } +type grokRateLimitRecoveryRepository interface { + ClearRateLimitIfObserved(ctx context.Context, id int64, observedLimitedAt, observedResetAt time.Time) (bool, error) +} + +func isSuccessfulGrokRateLimitRecovery(account *Account, snapshot *xai.QuotaSnapshot) bool { + return isGrokOAuthAccount(account) && + account.RateLimitedAt != nil && + account.RateLimitResetAt != nil && + snapshot != nil && + snapshot.StatusCode >= http.StatusOK && + snapshot.StatusCode < http.StatusMultipleChoices +} + +func clearGrokRateLimitAfterRecovery(ctx context.Context, repo AccountRepository, account *Account) { + if repo == nil || account == nil || account.RateLimitedAt == nil || account.RateLimitResetAt == nil || ctx.Err() != nil { + return + } + recoveryRepo, ok := repo.(grokRateLimitRecoveryRepository) + if !ok { + return + } + _, err := recoveryRepo.ClearRateLimitIfObserved(ctx, account.ID, *account.RateLimitedAt, *account.RateLimitResetAt) + if err != nil { + slog.Warn("grok_rate_limit_recovery_clear_failed", "account_id", account.ID, "error", err) + } +} + func persistGrokRateLimit(ctx context.Context, repo AccountRepository, account *Account, resetAt time.Time) { if repo == nil || account == nil || account.ID <= 0 { return diff --git a/backend/internal/service/openai_gateway_grok_chat_bridge.go b/backend/internal/service/openai_gateway_grok_chat_bridge.go index fd6b426536..6fd58b2db1 100644 --- a/backend/internal/service/openai_gateway_grok_chat_bridge.go +++ b/backend/internal/service/openai_gateway_grok_chat_bridge.go @@ -10,7 +10,6 @@ import ( "time" "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" - "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/gin-gonic/gin" ) @@ -253,12 +252,12 @@ func (s *OpenAIGatewayService) forwardGrokChatCompletionsViaResponses( } responsesBody = updatedBody - token, _, err := s.GetAccessToken(ctx, account) + token, _, err := s.getRequestCredential(ctx, c, account) if err != nil { return nil, fmt.Errorf("get grok access token: %w", err) } upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) - upstreamReq, err := buildGrokResponsesRequest(upstreamCtx, c, account, responsesBody, token, cacheIdentity) + upstreamReq, err := buildGrokResponsesRequest(upstreamCtx, c, account, responsesBody, token, cacheIdentity, s.cfg) releaseUpstreamCtx() if err != nil { return nil, fmt.Errorf("build grok responses bridge request: %w", err) @@ -294,13 +293,14 @@ func (s *OpenAIGatewayService) forwardGrokChatCompletionsViaResponses( return nil, &UpstreamFailoverError{ StatusCode: resp.StatusCode, ResponseBody: respBody, + ResponseHeaders: resp.Header.Clone(), RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode), } } return s.handleChatCompletionsErrorResponse(resp, c, account, billingModel) } - s.updateGrokUsageSnapshot(ctx, account, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) + s.updateGrokUsageFromResponse(ctx, account, resp.Header, resp.StatusCode) var result *OpenAIForwardResult if clientStream { diff --git a/backend/internal/service/openai_gateway_grok_chat_bridge_test.go b/backend/internal/service/openai_gateway_grok_chat_bridge_test.go index 0843f1ea42..4fc31461b4 100644 --- a/backend/internal/service/openai_gateway_grok_chat_bridge_test.go +++ b/backend/internal/service/openai_gateway_grok_chat_bridge_test.go @@ -284,6 +284,7 @@ func TestForwardGrokChatViaResponses429UsesGrokRateLimitPolicy(t *testing.T) { var failoverErr *UpstreamFailoverError require.True(t, errors.As(err, &failoverErr)) require.Equal(t, http.StatusTooManyRequests, failoverErr.StatusCode) + require.Equal(t, "45", failoverErr.ResponseHeaders.Get("Retry-After")) require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String()) require.Equal(t, grokChatResponsesEndpoint, GetActualOpenAIUpstreamEndpoint(c)) require.Equal(t, 1, repo.rateLimitedCalls) @@ -292,6 +293,45 @@ func TestForwardGrokChatViaResponses429UsesGrokRateLimitPolicy(t *testing.T) { require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) } +func TestForwardGrokRawChat429PreservesRetryAfter(t *testing.T) { + gin.SetMode(gin.TestMode) + + body := []byte(`{"model":"grok","messages":[{"role":"user","content":"hi"}],"stream":false,"stop":"done"}`) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, grokChatRawEndpoint, bytes.NewReader(body)) + c.Set("api_key", &APIKey{ID: 7551}) + + account := grokChatBridgeTestAccount(755) + account.Credentials["expires_at"] = time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339) + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusTooManyRequests, + Header: http.Header{ + "Content-Type": []string{"application/json"}, + "Retry-After": []string{"45"}, + }, + Body: io.NopCloser(strings.NewReader(`{"error":{"message":"rate limited"}}`)), + }} + svc := &OpenAIGatewayService{ + httpUpstream: upstream, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + accountRepo: repo, + } + + result, err := svc.ForwardAsChatCompletions(context.Background(), c, account, body, "", "") + + require.Error(t, err) + require.Nil(t, result) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, http.StatusTooManyRequests, failoverErr.StatusCode) + require.Equal(t, "45", failoverErr.ResponseHeaders.Get("Retry-After")) + require.Equal(t, xai.DefaultCLIBaseURL+"/chat/completions", upstream.lastReq.URL.String()) +} + func TestForwardGrokRawChatErrorRecordsActualEndpoint(t *testing.T) { gin.SetMode(gin.TestMode) @@ -329,11 +369,14 @@ func grokChatBridgeTestAccount(id int64) *Account { Name: "grok-cache-bridge", Platform: PlatformGrok, Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, Concurrency: 1, Credentials: map[string]any{ - "access_token": "access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), - "base_url": xai.DefaultCLIBaseURL, + "access_token": "access-token", + "refresh_token": "refresh-token", + "expires_at": time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339), + "base_url": xai.DefaultCLIBaseURL, }, } } diff --git a/backend/internal/service/openai_gateway_grok_test.go b/backend/internal/service/openai_gateway_grok_test.go index 3b64cbc16b..d93415b2e0 100644 --- a/backend/internal/service/openai_gateway_grok_test.go +++ b/backend/internal/service/openai_gateway_grok_test.go @@ -256,10 +256,10 @@ func TestBuildGrokResponsesRequestUsesAccountBaseURLAndBearerToken(t *testing.T) }, } - req, err := buildGrokResponsesRequest(context.Background(), nil, account, []byte(`{"model":"grok-4.3"}`), "access-token", "isolated-cache-id") + req, err := buildGrokResponsesRequest(context.Background(), nil, account, []byte(`{"model":"grok-4.3"}`), "access-token", "isolated-cache-id", nil) require.NoError(t, err) require.Equal(t, http.MethodPost, req.Method) - require.Equal(t, "https://xai.test/v1/responses", req.URL.String()) + require.Equal(t, xai.DefaultCLIBaseURL+"/responses", req.URL.String()) require.Equal(t, "Bearer access-token", req.Header.Get("Authorization")) require.Equal(t, "application/json", req.Header.Get("Content-Type")) require.Contains(t, req.Header.Get("Accept"), "text/event-stream") @@ -280,10 +280,12 @@ func TestBuildGrokResponsesRequestAllowsPublicAPIKeyBaseURLByDefault(t *testing. }, } - req, err := buildGrokResponsesRequest(context.Background(), nil, account, []byte(`{"model":"grok-4.3"}`), "api-key", "") + req, err := buildGrokResponsesRequest(context.Background(), nil, account, []byte(`{"model":"grok-4.3"}`), "api-key", "", nil) require.NoError(t, err) require.Equal(t, "https://grok.example.test/v1/responses", req.URL.String()) require.Equal(t, "Bearer api-key", req.Header.Get("Authorization")) + require.Empty(t, req.Header.Get("X-Grok-Client-Version")) + require.NotEqual(t, grokUpstreamUserAgent, req.Header.Get("User-Agent")) } func TestBuildGrokResponsesRequestPinsOAuthCustomBaseURLByDefault(t *testing.T) { @@ -297,7 +299,7 @@ func TestBuildGrokResponsesRequestPinsOAuthCustomBaseURLByDefault(t *testing.T) }, } - req, err := buildGrokResponsesRequest(context.Background(), nil, account, []byte(`{"model":"grok-4.3"}`), "access-token", "") + req, err := buildGrokResponsesRequest(context.Background(), nil, account, []byte(`{"model":"grok-4.3"}`), "access-token", "", nil) require.NoError(t, err) require.Equal(t, xai.DefaultCLIBaseURL+"/responses", req.URL.String()) } @@ -425,7 +427,8 @@ func TestForwardGrokMediaImagesGenerationNormalizesImagineAlias(t *testing.T) { require.Equal(t, http.MethodPost, upstream.lastReq.Method) require.Equal(t, "Bearer api-key", upstream.lastReq.Header.Get("Authorization")) require.Equal(t, "application/json", upstream.lastReq.Header.Get("Content-Type")) - require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version")) + require.Empty(t, upstream.lastReq.Header.Get("X-Grok-Client-Version")) + require.NotEqual(t, grokUpstreamUserAgent, upstream.lastReq.Header.Get("User-Agent")) require.JSONEq(t, `{"model":"grok-imagine-image-quality","prompt":"draw a cat"}`, string(upstream.lastBody)) require.Equal(t, http.StatusOK, recorder.Code) require.JSONEq(t, `{"data":[]}`, recorder.Body.String()) @@ -611,7 +614,7 @@ func TestForwardGrokMediaVideoGenerationPreservesImageToVideoModel(t *testing.T) require.Equal(t, VideoBillingDefaultDurationSeconds, result.VideoDurationSeconds) } -func TestForwardGrokMediaOAuthImageToVideoUsesOfficialAPIForLargeBody(t *testing.T) { +func TestForwardGrokMediaOAuthImageToVideoKeepsCLIGatewayForLargeBody(t *testing.T) { gin.SetMode(gin.TestMode) recorder := httptest.NewRecorder() @@ -626,10 +629,14 @@ func TestForwardGrokMediaOAuthImageToVideoUsesOfficialAPIForLargeBody(t *testing Name: "grok-oauth", Platform: PlatformGrok, Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, Concurrency: 1, Credentials: map[string]any{ - "access_token": "oauth-access-token", - "base_url": xai.DefaultCLIBaseURL, + "access_token": "oauth-access-token", + "refresh_token": "oauth-refresh-token", + "expires_at": time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339), + "base_url": xai.DefaultCLIBaseURL, }, } upstream := &httpUpstreamRecorder{resp: &http.Response{ @@ -639,11 +646,11 @@ func TestForwardGrokMediaOAuthImageToVideoUsesOfficialAPIForLargeBody(t *testing }, Body: io.NopCloser(strings.NewReader(`{"request_id":"video-request-oauth"}`)), }} - svc := &OpenAIGatewayService{httpUpstream: upstream} + svc := &OpenAIGatewayService{httpUpstream: upstream, grokTokenProvider: NewGrokTokenProvider(nil, nil)} _, err := svc.ForwardGrokMedia(context.Background(), c, account, GrokMediaEndpointVideosGenerations, "", body, "application/json") require.NoError(t, err) - require.Equal(t, xai.DefaultBaseURL+"/videos/generations", upstream.lastReq.URL.String()) + require.Equal(t, xai.DefaultCLIBaseURL+"/videos/generations", upstream.lastReq.URL.String()) require.Equal(t, "data:image/png;base64,"+imageData, gjson.GetBytes(upstream.lastBody, "image.image_url").String()) } @@ -681,6 +688,8 @@ func TestForwardGrokMediaVideoStatusUsesGETWithoutBody(t *testing.T) { require.Equal(t, "https://xai.test/v1/videos/request-123", upstream.lastReq.URL.String()) require.Equal(t, http.MethodGet, upstream.lastReq.Method) require.Equal(t, "Bearer api-key", upstream.lastReq.Header.Get("Authorization")) + require.Empty(t, upstream.lastReq.Header.Get("X-Grok-Client-Version")) + require.NotEqual(t, grokUpstreamUserAgent, upstream.lastReq.Header.Get("User-Agent")) require.Empty(t, upstream.lastReq.Header.Get("Content-Type")) require.Empty(t, upstream.lastBody) require.Equal(t, http.StatusOK, recorder.Code) @@ -792,6 +801,54 @@ func TestForwardGrokMedia429ReconcilesRateLimitBeforeCustomErrorBypass(t *testin require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) } +func TestGrokMedia429FailoverPreservesRetryAfter(t *testing.T) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/images/generations", nil) + account := &Account{ + ID: 641, Name: "grok-oauth", Platform: PlatformGrok, Type: AccountTypeOAuth, + Status: StatusActive, Schedulable: true, + Credentials: map[string]any{ + "custom_error_codes_enabled": true, + "custom_error_codes": []any{float64(http.StatusTooManyRequests)}, + }, + } + repo := &grokQuotaAccountRepo{} + svc := &OpenAIGatewayService{accountRepo: repo} + resp := &http.Response{ + StatusCode: http.StatusTooManyRequests, + Header: http.Header{"Retry-After": []string{"45"}}, + Body: io.NopCloser(strings.NewReader(`{"error":{"message":"rate limited"}}`)), + } + + result, err := svc.handleGrokMediaErrorResponse(context.Background(), resp, c, account, "request-id", "grok-imagine") + + require.Nil(t, result) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, http.StatusTooManyRequests, failoverErr.StatusCode) + require.Equal(t, "45", failoverErr.ResponseHeaders.Get("Retry-After")) +} + +func healthyGrokOAuthGatewayTestAccount(id int64, token string) *Account { + return &Account{ + ID: id, + Name: "grok", + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Status: StatusActive, + Schedulable: true, + Concurrency: 1, + Credentials: map[string]any{ + "access_token": token, + "refresh_token": "refresh-token", + "expires_at": time.Now().Add(2 * grokTokenRefreshSkew).UTC().Format(time.RFC3339), + "base_url": xai.DefaultCLIBaseURL, + }, + } +} + func TestForwardAsChatCompletionsForGrokStopFallsBackToXAIChatCompletions(t *testing.T) { gin.SetMode(gin.TestMode) @@ -801,18 +858,7 @@ func TestForwardAsChatCompletionsForGrokStopFallsBackToXAIChatCompletions(t *tes c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body)) c.Set("api_key", &APIKey{ID: 5101}) - account := &Account{ - ID: 51, - Name: "grok", - Platform: PlatformGrok, - Type: AccountTypeOAuth, - Concurrency: 1, - Credentials: map[string]any{ - "access_token": "access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), - "base_url": xai.DefaultCLIBaseURL, - }, - } + account := healthyGrokOAuthGatewayTestAccount(51, "access-token") repo := &grokQuotaAccountRepo{ mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ accountsByID: map[int64]*Account{51: account}, @@ -864,18 +910,7 @@ func TestForwardGrokResponsesStreamingDefaultsEmptyModelTo45AndSnapshots(t *test c.Request.Header.Set("OpenAI-Beta", "responses=experimental") c.Set("api_key", &APIKey{ID: 5201}) - account := &Account{ - ID: 52, - Name: "grok", - Platform: PlatformGrok, - Type: AccountTypeOAuth, - Concurrency: 1, - Credentials: map[string]any{ - "access_token": "access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), - "base_url": xai.DefaultCLIBaseURL, - }, - } + account := healthyGrokOAuthGatewayTestAccount(52, "access-token") repo := &grokQuotaAccountRepo{ mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ accountsByID: map[int64]*Account{52: account}, @@ -968,12 +1003,46 @@ func TestForwardGrokResponsesAPIKeyUsesXAIResponses(t *testing.T) { require.NoError(t, err) require.Equal(t, "https://api.x.ai/v1/responses", upstream.lastReq.URL.String()) require.Equal(t, "Bearer xai-test-key", upstream.lastReq.Header.Get("Authorization")) + require.Empty(t, upstream.lastReq.Header.Get("X-Grok-Client-Version")) + require.NotEqual(t, grokUpstreamUserAgent, upstream.lastReq.Header.Get("User-Agent")) require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String()) require.Equal(t, "resp_grok_api_key", result.ResponseID) require.Equal(t, 2, result.Usage.InputTokens) require.Equal(t, 1, result.Usage.OutputTokens) } +func TestForwardAsChatCompletionsForGrokAPIKeyUsesConfiguredRawEndpointWithoutOAuthIdentity(t *testing.T) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + body := []byte(`{"model":"grok","messages":[{"role":"user","content":"hi"}],"stream":false}`) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + account := &Account{ + ID: 706, + Platform: PlatformGrok, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "third-party-key", + "base_url": "https://grok.example.test/v1", + }, + } + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"id":"chatcmpl","object":"chat.completion","model":"grok-4.5","choices":[],"usage":{"prompt_tokens":1,"completion_tokens":1}}`)), + }} + svc := &OpenAIGatewayService{cfg: rawChatCompletionsTestConfig(), httpUpstream: upstream} + + _, err := svc.ForwardAsChatCompletions(context.Background(), c, account, body, "", "") + require.NoError(t, err) + require.Equal(t, "https://grok.example.test/v1/chat/completions", upstream.lastReq.URL.String()) + require.Equal(t, "Bearer third-party-key", upstream.lastReq.Header.Get("Authorization")) + require.Empty(t, upstream.lastReq.Header.Get("X-Grok-Client-Version")) + require.NotEqual(t, grokUpstreamUserAgent, upstream.lastReq.Header.Get("User-Agent")) +} + func TestAccountTestServiceGrokAPIKeyUsesXAIResponses(t *testing.T) { gin.SetMode(gin.TestMode) @@ -1008,6 +1077,41 @@ func TestAccountTestServiceGrokAPIKeyUsesXAIResponses(t *testing.T) { require.Contains(t, recorder.Body.String(), `"type":"test_complete"`) } +func TestAccountTestServiceGrokAPIKeyAllowsConfiguredHTTPWhenGlobalPolicyDoes(t *testing.T) { + gin.SetMode(gin.TestMode) + + account := &Account{ + ID: 55, + Name: "grok-api-key-http", + Platform: PlatformGrok, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "third-party-key", + "base_url": "http://grok.example.test/v1", + }, + } + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader( + "data: {\"type\":\"response.output_text.delta\",\"delta\":\"ok\"}\n\n" + + "data: {\"type\":\"response.completed\"}\n\n", + )), + }} + svc := &AccountTestService{cfg: rawChatCompletionsTestConfig(), httpUpstream: upstream} + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/55/test", nil) + + err := svc.testGrokAccountConnection(c, account, "grok") + require.NoError(t, err) + require.Equal(t, "http://grok.example.test/v1/responses", upstream.lastReq.URL.String()) + require.Equal(t, "Bearer third-party-key", upstream.lastReq.Header.Get("Authorization")) + require.Empty(t, upstream.lastReq.Header.Get("X-Grok-Client-Version")) + require.Contains(t, recorder.Body.String(), `"type":"test_complete"`) +} + func TestForwardAsChatCompletionsForGrokStreamingUsesRawXAIChatCompletions(t *testing.T) { gin.SetMode(gin.TestMode) @@ -1017,18 +1121,7 @@ func TestForwardAsChatCompletionsForGrokStreamingUsesRawXAIChatCompletions(t *te c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body)) c.Request.Header.Set("Content-Type", "application/json") - account := &Account{ - ID: 53, - Name: "grok", - Platform: PlatformGrok, - Type: AccountTypeOAuth, - Concurrency: 1, - Credentials: map[string]any{ - "access_token": "access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), - "base_url": xai.DefaultCLIBaseURL, - }, - } + account := healthyGrokOAuthGatewayTestAccount(53, "access-token") repo := &grokQuotaAccountRepo{ mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ accountsByID: map[int64]*Account{53: account}, @@ -1085,22 +1178,16 @@ func TestForwardGrokResponsesNonStreamingUsesCacheIdentityAndCachedUsage(t *test c.Request.Header.Set("Content-Type", "application/json") c.Set("api_key", &APIKey{ID: 5202}) - account := &Account{ - ID: 56, - Name: "grok", - Platform: PlatformGrok, - Type: AccountTypeOAuth, - Concurrency: 1, - Credentials: map[string]any{ - "access_token": "access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), - "base_url": xai.DefaultCLIBaseURL, - }, - } + account := healthyGrokOAuthGatewayTestAccount(56, "access-token") + observedResetAt := time.Now().Add(-time.Second).UTC().Truncate(time.Second) + observedLimitedAt := observedResetAt.Add(-grokRateLimitRepeatCooldown) + account.RateLimitedAt = &observedLimitedAt + account.RateLimitResetAt = &observedResetAt repo := &grokQuotaAccountRepo{ mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ accountsByID: map[int64]*Account{56: account}, }, + recoveryClearResult: true, } upstream := &httpUpstreamRecorder{resp: &http.Response{ StatusCode: http.StatusOK, @@ -1133,6 +1220,9 @@ func TestForwardGrokResponsesNonStreamingUsesCacheIdentityAndCachedUsage(t *test require.False(t, gjson.GetBytes(upstream.lastBody, "tools").Exists()) require.False(t, gjson.GetBytes(upstream.lastBody, "tool_choice").Exists()) require.Equal(t, "resp_grok_non_stream", gjson.Get(recorder.Body.String(), "id").String()) + require.Equal(t, 1, repo.recoveryClearCalls) + require.Equal(t, observedLimitedAt, repo.recoveryObservedAt) + require.Equal(t, observedResetAt, repo.recoveryObservedReset) } func TestForwardGrokResponsesFailoverKeepsCacheIdentityAcrossAccounts(t *testing.T) { @@ -1145,18 +1235,9 @@ func TestForwardGrokResponsesFailoverKeepsCacheIdentityAcrossAccounts(t *testing c.Set("api_key", &APIKey{ID: 5203}) newAccount := func(id int64, token string) *Account { - return &Account{ - ID: id, - Name: fmt.Sprintf("grok-%d", id), - Platform: PlatformGrok, - Type: AccountTypeOAuth, - Concurrency: 1, - Credentials: map[string]any{ - "access_token": token, - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), - "base_url": xai.DefaultCLIBaseURL, - }, - } + account := healthyGrokOAuthGatewayTestAccount(id, token) + account.Name = fmt.Sprintf("grok-%d", id) + return account } firstAccount := newAccount(58, "access-token-a") secondAccount := newAccount(59, "access-token-b") @@ -1213,18 +1294,7 @@ func TestForwardAsChatCompletionsForGrokStreamingStopFallsBackToRawXAIChatComple c.Request.Header.Set(grokConversationIDHeader, "native-client-conversation") c.Set("api_key", &APIKey{ID: 5301}) - account := &Account{ - ID: 53, - Name: "grok", - Platform: PlatformGrok, - Type: AccountTypeOAuth, - Concurrency: 1, - Credentials: map[string]any{ - "access_token": "access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), - "base_url": xai.DefaultCLIBaseURL, - }, - } + account := healthyGrokOAuthGatewayTestAccount(53, "access-token") repo := &grokQuotaAccountRepo{ mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ accountsByID: map[int64]*Account{53: account}, @@ -1284,18 +1354,7 @@ func TestForwardAsChatCompletionsForGrokComposerBridgesImageInput(t *testing.T) c.Request.Header.Set("Content-Type", "application/json") c.Set("api_key", &APIKey{ID: 5501}) - account := &Account{ - ID: 55, - Name: "grok", - Platform: PlatformGrok, - Type: AccountTypeOAuth, - Concurrency: 1, - Credentials: map[string]any{ - "access_token": "access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), - "base_url": xai.DefaultCLIBaseURL, - }, - } + account := healthyGrokOAuthGatewayTestAccount(55, "access-token") repo := &grokQuotaAccountRepo{ mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ accountsByID: map[int64]*Account{55: account}, @@ -1358,18 +1417,7 @@ func TestForwardAsAnthropicForGrokUsesXAIResponses(t *testing.T) { c.Request.Header.Set("OpenAI-Beta", "grok-experimental") c.Request.Header.Set("originator", "opencode") - account := &Account{ - ID: 54, - Name: "grok", - Platform: PlatformGrok, - Type: AccountTypeOAuth, - Concurrency: 1, - Credentials: map[string]any{ - "access_token": "access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), - "base_url": xai.DefaultCLIBaseURL, - }, - } + account := healthyGrokOAuthGatewayTestAccount(54, "access-token") repo := &grokQuotaAccountRepo{ mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ accountsByID: map[int64]*Account{54: account}, @@ -1419,18 +1467,7 @@ func TestForwardAsAnthropicForGrokStreamingPreservesCacheUsage(t *testing.T) { c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(body)) c.Set("api_key", &APIKey{ID: 5402}) - account := &Account{ - ID: 57, - Name: "grok", - Platform: PlatformGrok, - Type: AccountTypeOAuth, - Concurrency: 1, - Credentials: map[string]any{ - "access_token": "access-token", - "expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339), - "base_url": xai.DefaultCLIBaseURL, - }, - } + account := healthyGrokOAuthGatewayTestAccount(57, "access-token") repo := &grokQuotaAccountRepo{ mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ accountsByID: map[int64]*Account{57: account}, @@ -1571,6 +1608,102 @@ func TestHandleGrokAccountUpstreamError429UsesFallbackReset(t *testing.T) { require.Zero(t, repo.tempUnschedCalls) } +func TestGrokRateLimitResetAtForAccountEscalatesRepeated429s(t *testing.T) { + now := time.Now().UTC().Truncate(time.Second) + retryAfter := 45 + snapshot := &xai.QuotaSnapshot{ + StatusCode: http.StatusTooManyRequests, + RetryAfterSeconds: &retryAfter, + UpdatedAt: now.Format(time.RFC3339), + } + tests := []struct { + name string + previousCooldown time.Duration + wantCooldown time.Duration + }{ + {name: "repeat after short boundary", previousCooldown: 45 * time.Second, wantCooldown: grokRateLimitRepeatCooldown}, + {name: "sustained repeat", previousCooldown: grokRateLimitRepeatCooldown, wantCooldown: grokRateLimitSustainedCooldown}, + {name: "capped repeat", previousCooldown: grokRateLimitSustainedCooldown, wantCooldown: grokRateLimitMaxAdaptiveCooldown}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + previousReset := now.Add(-time.Second) + previousLimited := previousReset.Add(-tt.previousCooldown) + account := &Account{ + ID: 630, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + RateLimitedAt: &previousLimited, + RateLimitResetAt: &previousReset, + } + + resetAt, limited := grokRateLimitResetAtForAccount(account, snapshot, now) + + require.True(t, limited) + require.WithinDuration(t, now.Add(tt.wantCooldown), resetAt, time.Second) + }) + } +} + +func TestGrokRateLimitResetAtForAccountPreservesAuthoritativeAndQuietRecovery(t *testing.T) { + now := time.Now().UTC().Truncate(time.Second) + retryAfter := 45 + previousReset := now.Add(-grokRateLimitBackoffQuietPeriod - time.Second) + previousLimited := previousReset.Add(-grokRateLimitSustainedCooldown) + account := &Account{ + ID: 631, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + RateLimitedAt: &previousLimited, + RateLimitResetAt: &previousReset, + } + snapshot := &xai.QuotaSnapshot{ + StatusCode: http.StatusTooManyRequests, + RetryAfterSeconds: &retryAfter, + UpdatedAt: now.Format(time.RFC3339), + } + + resetAt, limited := grokRateLimitResetAtForAccount(account, snapshot, now) + require.True(t, limited) + require.WithinDuration(t, now.Add(45*time.Second), resetAt, time.Second) + + authoritativeReset := now.Add(2 * time.Hour) + remaining := int64(0) + snapshot.Requests = &xai.QuotaWindow{Remaining: &remaining, ResetUnix: grokInt64PtrForTest(authoritativeReset.Unix())} + recentReset := now.Add(-time.Second) + recentLimited := recentReset.Add(-grokRateLimitSustainedCooldown) + account.RateLimitResetAt = &recentReset + account.RateLimitedAt = &recentLimited + + resetAt, limited = grokRateLimitResetAtForAccount(account, snapshot, now) + require.True(t, limited) + require.WithinDuration(t, authoritativeReset, resetAt, time.Second) +} + +func TestGrokRateLimitResetAtForAccountLeavesAPIKey429PolicyUnchanged(t *testing.T) { + now := time.Now().UTC().Truncate(time.Second) + retryAfter := 45 + previousReset := now.Add(-time.Second) + previousLimited := previousReset.Add(-grokRateLimitSustainedCooldown) + account := &Account{ + ID: 632, + Platform: PlatformGrok, + Type: AccountTypeAPIKey, + RateLimitedAt: &previousLimited, + RateLimitResetAt: &previousReset, + } + snapshot := &xai.QuotaSnapshot{ + StatusCode: http.StatusTooManyRequests, + RetryAfterSeconds: &retryAfter, + UpdatedAt: now.Format(time.RFC3339), + } + + resetAt, limited := grokRateLimitResetAtForAccount(account, snapshot, now) + require.True(t, limited) + require.WithinDuration(t, now.Add(45*time.Second), resetAt, time.Second) +} + func TestGrokRateLimitResetAtUsesFutureWindowAfterRetryAfterExpires(t *testing.T) { now := time.Now().UTC().Truncate(time.Second) observedAt := now.Add(-2 * time.Minute) @@ -1672,6 +1805,72 @@ func TestUpdateGrokUsageSnapshotAvailableSuccessDoesNotSetRateLimited(t *testing require.Zero(t, repo.rateLimitedCalls) } +func TestUpdateGrokUsageFromResponseHeaderlessSuccessClearsObservedCooldown(t *testing.T) { + now := time.Now().UTC().Truncate(time.Second) + limitedAt := now.Add(-grokRateLimitRepeatCooldown) + observedResetAt := now.Add(-time.Second) + account := &Account{ + ID: 660, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + RateLimitedAt: &limitedAt, + RateLimitResetAt: &observedResetAt, + } + repo := &grokQuotaAccountRepo{recoveryClearResult: true} + svc := &OpenAIGatewayService{ + accountRepo: repo, + codexSnapshotThrottle: newAccountWriteThrottle(time.Hour), + } + + svc.updateGrokUsageFromResponse(context.Background(), account, nil, http.StatusOK) + + require.Zero(t, repo.updateCalls, "headerless success must not overwrite an informative quota snapshot") + require.Equal(t, 1, repo.recoveryClearCalls) + require.Equal(t, limitedAt, repo.recoveryObservedAt) + require.Equal(t, observedResetAt, repo.recoveryObservedReset) + require.Same(t, &observedResetAt, account.RateLimitResetAt, "shared account snapshots must not be mutated in place") +} + +func TestUpdateGrokUsageFromResponseRecoveryRespectsCancellationAndAPIKeyBoundary(t *testing.T) { + now := time.Now().UTC().Truncate(time.Second) + observedResetAt := now.Add(-time.Second) + observedLimitedAt := observedResetAt.Add(-grokRateLimitRepeatCooldown) + + t.Run("parent cancellation does not mutate account state", func(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + account := &Account{ + ID: 661, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + RateLimitedAt: &observedLimitedAt, + RateLimitResetAt: &observedResetAt, + } + repo := &grokQuotaAccountRepo{recoveryClearResult: true} + svc := &OpenAIGatewayService{accountRepo: repo} + + svc.updateGrokUsageFromResponse(ctx, account, nil, http.StatusOK) + + require.Zero(t, repo.recoveryClearCalls) + }) + + t.Run("API key success does not alter OAuth cooldown state", func(t *testing.T) { + account := &Account{ + ID: 662, + Platform: PlatformGrok, + Type: AccountTypeAPIKey, + RateLimitedAt: &observedLimitedAt, + RateLimitResetAt: &observedResetAt, + } + repo := &grokQuotaAccountRepo{recoveryClearResult: true} + svc := &OpenAIGatewayService{accountRepo: repo} + + svc.updateGrokUsageFromResponse(context.Background(), account, nil, http.StatusOK) + + require.Zero(t, repo.recoveryClearCalls) + }) +} + func TestUpdateGrokUsageSnapshotExhaustedSuccessWithoutResetUsesFallback(t *testing.T) { repo := &grokQuotaAccountRepo{} svc := &OpenAIGatewayService{accountRepo: repo} diff --git a/backend/internal/service/openai_gateway_messages.go b/backend/internal/service/openai_gateway_messages.go index 8cf62bdc65..cbfd218a22 100644 --- a/backend/internal/service/openai_gateway_messages.go +++ b/backend/internal/service/openai_gateway_messages.go @@ -14,7 +14,6 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/claude" "github.com/Wei-Shaw/sub2api/internal/pkg/logger" "github.com/Wei-Shaw/sub2api/internal/pkg/openai_compat" - "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/Wei-Shaw/sub2api/internal/util/responseheaders" "github.com/gin-gonic/gin" "go.uber.org/zap" @@ -258,7 +257,7 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( } // 5. Get access token - token, _, err := s.GetAccessToken(ctx, account) + token, _, err := s.getRequestCredential(ctx, c, account) if err != nil { return nil, fmt.Errorf("get access token: %w", err) } @@ -273,7 +272,7 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) var upstreamReq *http.Request if account.Platform == PlatformGrok { - upstreamReq, err = buildGrokResponsesRequest(upstreamCtx, c, account, responsesBody, token, grokCacheIdentity) + upstreamReq, err = buildGrokResponsesRequest(upstreamCtx, c, account, responsesBody, token, grokCacheIdentity, s.cfg) } else { upstreamReq, err = s.buildUpstreamRequest(upstreamCtx, c, account, responsesBody, token, isStream, promptCacheKey, false) } @@ -324,6 +323,13 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( // 8. Handle error response with failover if resp.StatusCode >= 400 { respBody, upstreamMsg := s.readOpenAIUpstreamError(resp) + if !agentIdentityTaskRecoveryWasTried(ctx) && s.isAgentIdentityAccount(ctx, account) && isAgentIdentityTaskInvalidHTTPResponse(resp.StatusCode, respBody) { + expectedTaskID := account.GetCredential("task_id") + if err := s.recoverAgentIdentityTask(ctx, account, expectedTaskID); err != nil { + return nil, fmt.Errorf("agent identity task recovery failed: %w", err) + } + return s.ForwardAsAnthropic(markAgentIdentityTaskRecoveryTried(ctx), c, account, body, promptCacheKey, defaultMappedModel) + } if previousResponseID != "" && (isOpenAICompatPreviousResponseNotFound(resp.StatusCode, upstreamMsg, respBody) || isOpenAICompatPreviousResponseUnsupported(resp.StatusCode, upstreamMsg, respBody)) { if isOpenAICompatPreviousResponseUnsupported(resp.StatusCode, upstreamMsg, respBody) { s.disableOpenAICompatSessionContinuation(ctx, c, account, promptCacheKey) @@ -344,7 +350,7 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( return s.handleAnthropicErrorResponse(resp, c, account, billingModel) } if account.Platform == PlatformGrok && account.Type == AccountTypeOAuth && !account.IsShadow() { - s.updateGrokUsageSnapshot(ctx, account, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) + s.updateGrokUsageFromResponse(ctx, account, resp.Header, resp.StatusCode) } if account.Type == AccountTypeOAuth && promptCacheKey != "" { diff --git a/backend/internal/service/openai_gateway_passthrough.go b/backend/internal/service/openai_gateway_passthrough.go index cbdefc8ac7..b8d74a463e 100644 --- a/backend/internal/service/openai_gateway_passthrough.go +++ b/backend/internal/service/openai_gateway_passthrough.go @@ -11,6 +11,7 @@ import ( "encoding/json" "errors" "fmt" + "io" "net/http" "sort" "strings" @@ -161,13 +162,6 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough( return nil, err } - upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) - upstreamReq, err := s.buildUpstreamRequestOpenAIPassthrough(upstreamCtx, c, account, body, token) - releaseUpstreamCtx() - if err != nil { - return nil, err - } - proxyURL := "" if account.ProxyID != nil && account.Proxy != nil { proxyURL = account.Proxy.URL() @@ -177,26 +171,51 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough( c.Set("openai_passthrough", true) } - upstreamStart := time.Now() - resp, err := s.httpUpstream.Do(upstreamReq, proxyURL, account.ID, account.Concurrency) - SetOpsLatencyMs(c, OpsUpstreamLatencyMsKey, time.Since(upstreamStart).Milliseconds()) - if err != nil { - // Transport-level failure (proxy/DNS/TCP/TLS — no HTTP response). Convert to - // a failover so the handler switches to a healthy account, and temporarily - // unschedule the account on durable faults (e.g. rejected proxy credentials). - return nil, s.handleOpenAIUpstreamTransportError(ctx, c, account, err, true) - } - defer func() { _ = resp.Body.Close() }() + agentTaskRecoveryTried := false + var resp *http.Response + for { + upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) + upstreamReq, buildErr := s.buildUpstreamRequestOpenAIPassthrough(upstreamCtx, c, account, body, token) + releaseUpstreamCtx() + if buildErr != nil { + return nil, buildErr + } + + upstreamStart := time.Now() + resp, err = s.httpUpstream.Do(upstreamReq, proxyURL, account.ID, account.Concurrency) + SetOpsLatencyMs(c, OpsUpstreamLatencyMsKey, time.Since(upstreamStart).Milliseconds()) + if err != nil { + // Transport-level failure (proxy/DNS/TCP/TLS — no HTTP response). Convert to + // a failover so the handler switches to a healthy account. + return nil, s.handleOpenAIUpstreamTransportError(ctx, c, account, err, true) + } + if resp.StatusCode < 400 { + break + } + + // Peek only to identify an invalid task. Restore the body so the existing + // passthrough error handling sees the same response after recovery fails. + probeBody := s.readUpstreamErrorBody(resp) + _ = resp.Body.Close() + resp.Body = io.NopCloser(bytes.NewReader(probeBody)) + if !agentTaskRecoveryTried && s.isAgentIdentityAccount(ctx, account) && isAgentIdentityTaskInvalidHTTPResponse(resp.StatusCode, probeBody) { + agentTaskRecoveryTried = true + expectedTaskID := account.GetCredential("task_id") + if recoveryErr := s.recoverAgentIdentityTask(ctx, account, expectedTaskID); recoveryErr != nil { + return nil, fmt.Errorf("agent identity task recovery failed: %w", recoveryErr) + } + continue + } - if resp.StatusCode >= 400 { - responseBody := s.readUpstreamErrorBody(resp) // 透传模式默认保持原样代理;容量错误以及 API-key 上游的瞬时 // 5xx 应先触发多账号 failover,且此时尚未写入下游响应。 - if shouldFailoverOpenAIPassthroughResponse(account, resp.StatusCode, responseBody) { - return nil, s.handleFailoverErrorResponsePassthrough(ctx, resp, c, account, body, responseBody) + // probeBody 已在上方任务探测时读取过一次,直接复用避免重复读取。 + if shouldFailoverOpenAIPassthroughResponse(account, resp.StatusCode, probeBody) { + return nil, s.handleFailoverErrorResponsePassthrough(ctx, resp, c, account, body, probeBody) } - return nil, s.handleErrorResponsePassthrough(ctx, resp, c, account, body, responseBody) + return nil, s.handleErrorResponsePassthrough(ctx, resp, c, account, body, probeBody) } + defer func() { _ = resp.Body.Close() }() serviceTier := extractOpenAIServiceTierFromBody(body) @@ -339,7 +358,15 @@ func (s *OpenAIGatewayService) buildUpstreamRequestOpenAIPassthrough( req.Header.Del("authorization") req.Header.Del("x-api-key") req.Header.Del("x-goog-api-key") - req.Header.Set("authorization", "Bearer "+token) + authHeaders, err := s.buildOpenAIAuthenticationHeaders(ctx, account, token) + if err != nil { + return nil, fmt.Errorf("build openai authentication headers: %w", err) + } + for key, values := range authHeaders { + for _, value := range values { + req.Header.Add(key, value) + } + } // OAuth 透传到 ChatGPT internal API 时补齐必要头。 if account.Type == AccountTypeOAuth { @@ -449,7 +476,7 @@ func (s *OpenAIGatewayService) handleFailoverErrorResponsePassthrough( requestBody []byte, responseBody []byte, ) error { - body := responseBody + body := s.redactAgentIdentitySensitiveBody(ctx, account, responseBody) upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(body)) upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg) @@ -494,7 +521,7 @@ func (s *OpenAIGatewayService) handleErrorResponsePassthrough( responseBody []byte, ) error { MarkResponseCommitted(c) - body := responseBody + body := s.redactAgentIdentitySensitiveBody(ctx, account, responseBody) // cyber_policy:透传账号本就把原始 body 回给客户端(下方 c.Data),此处仅打标记, // 供 handler 事后写风控/邮件。cyber 是上游网络安全策略拦截,不冷却账号, diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index 11a8255bfa..956f167d1b 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -393,6 +393,7 @@ type OpenAIGatewayService struct { openaiWSStateStoreOnce sync.Once openaiSchedulerOnce sync.Once openaiWSPassthroughDialerOnce sync.Once + agentIdentityTaskMu sync.Mutex openaiWSPool *openAIWSConnPool openaiWSStateStore OpenAIWSStateStore openaiScheduler OpenAIAccountScheduler @@ -401,6 +402,10 @@ type OpenAIGatewayService struct { openaiWSFallbackUntil sync.Map // key: int64(accountID), value: time.Time openaiAccountRuntimeBlockUntil sync.Map // key: int64(accountID), value: time.Time + openaiAccountRuntimeBlockLocks sync.Map // key: int64(accountID), value: *sync.Mutex + openaiAccountRuntimeBlockGeneration sync.Map // key: int64(accountID), value: uint64 + openaiAccountRuntimeBlockSequence atomic.Uint64 + grokCredentialMutationLocks sync.Map // key: int64(accountID), value: *sync.Mutex openaiOAuth429WindowStartUnixNano atomic.Int64 openaiOAuth429WindowCount atomic.Int64 openaiWSRetryMetrics openAIWSRetryMetrics @@ -591,6 +596,12 @@ func (s *OpenAIGatewayService) CloseOpenAIWSPool() { } } +func (s *OpenAIGatewayService) InvalidateAgentIdentityWSConnections(accountID int64) { + if pool := s.getOpenAIWSConnPool(); pool != nil { + pool.ClearAccount(accountID) + } +} + func (s *OpenAIGatewayService) logOpenAIWSModeBootstrap() { if s == nil || s.cfg == nil { return @@ -1090,6 +1101,9 @@ func (s *OpenAIGatewayService) GetAccessToken(ctx context.Context, account *Acco } switch account.Type { case AccountTypeOAuth: + if account.IsOpenAIAgentIdentity() { + return "", OpenAIAuthModeAgentIdentity, nil + } if account.Platform == PlatformGrok { if s.grokTokenProvider != nil { accessToken, err := s.grokTokenProvider.GetAccessToken(ctx, account) diff --git a/backend/internal/service/openai_gateway_upstream_errors.go b/backend/internal/service/openai_gateway_upstream_errors.go index 2e661e7417..d8823fe4ed 100644 --- a/backend/internal/service/openai_gateway_upstream_errors.go +++ b/backend/internal/service/openai_gateway_upstream_errors.go @@ -279,6 +279,7 @@ func (s *OpenAIGatewayService) handleErrorResponse( requestedModel ...string, ) (*OpenAIForwardResult, error) { body := s.readUpstreamErrorBody(resp) + body = s.redactAgentIdentitySensitiveBody(ctx, account, body) // cyber_policy 硬阻断:透传上游原始错误体给客户端(不重包成通用 502),不冷却账号。 // 当前请求恒透传(需求1);标记供 handler 事后写风控/邮件。400 cyber 不可 failover @@ -470,6 +471,7 @@ func (s *OpenAIGatewayService) handleCompatErrorResponse( requestedModel ...string, ) (*OpenAIForwardResult, error) { body := s.readUpstreamErrorBody(resp) + body = s.redactAgentIdentitySensitiveBody(context.Background(), account, body) // cyber_policy:兼容路径(Chat Completions / Anthropic)以各自格式回写错误, // 不原样透传 responses 格式的 cyber body(否则对下游格式不合法)。cyber 是上游网络 diff --git a/backend/internal/service/openai_image_dimensions.go b/backend/internal/service/openai_image_dimensions.go new file mode 100644 index 0000000000..b297aeaf67 --- /dev/null +++ b/backend/internal/service/openai_image_dimensions.go @@ -0,0 +1,92 @@ +package service + +import ( + "bufio" + "encoding/base64" + "encoding/binary" + "fmt" + "image" + _ "image/jpeg" + _ "image/png" + "io" + "strings" +) + +const maxOpenAIImageDimensionProbeBytes int64 = 1 << 20 + +func detectOpenAIImageResultSize(encoded string) string { + payload := strings.TrimSpace(encoded) + if strings.HasPrefix(strings.ToLower(payload), "data:") { + comma := strings.IndexByte(payload, ',') + if comma < 0 || comma+1 >= len(payload) { + return "" + } + payload = strings.TrimSpace(payload[comma+1:]) + } + if payload == "" { + return "" + } + + for _, encoding := range []*base64.Encoding{base64.StdEncoding, base64.RawStdEncoding} { + decoded := base64.NewDecoder(encoding, strings.NewReader(payload)) + buffered := bufio.NewReader(io.LimitReader(decoded, maxOpenAIImageDimensionProbeBytes)) + prefix, _ := buffered.Peek(30) + if width, height, ok := detectOpenAIWebPDimensions(prefix); ok { + return fmt.Sprintf("%dx%d", width, height) + } + cfg, _, err := image.DecodeConfig(buffered) + if err != nil || cfg.Width <= 0 || cfg.Height <= 0 { + continue + } + return fmt.Sprintf("%dx%d", cfg.Width, cfg.Height) + } + return "" +} + +func detectOpenAIWebPDimensions(header []byte) (int, int, bool) { + if len(header) < 16 || string(header[:4]) != "RIFF" || string(header[8:12]) != "WEBP" { + return 0, 0, false + } + + switch string(header[12:16]) { + case "VP8X": + if len(header) < 30 { + return 0, 0, false + } + width := 1 + int(header[24]) + int(header[25])<<8 + int(header[26])<<16 + height := 1 + int(header[27]) + int(header[28])<<8 + int(header[29])<<16 + return width, height, width > 0 && height > 0 + case "VP8 ": + if len(header) < 30 || string(header[23:26]) != "\x9d\x01\x2a" { + return 0, 0, false + } + width := int(binary.LittleEndian.Uint16(header[26:28]) & 0x3fff) + height := int(binary.LittleEndian.Uint16(header[28:30]) & 0x3fff) + return width, height, width > 0 && height > 0 + case "VP8L": + if len(header) < 25 || header[20] != 0x2f { + return 0, 0, false + } + width := 1 + int(header[21]) + int(header[22]&0x3f)<<8 + height := 1 + int(header[22]>>6) + int(header[23])<<2 + int(header[24]&0x0f)<<10 + return width, height, width > 0 && height > 0 + default: + return 0, 0, false + } +} + +func reconcileOpenAIResponsesImageResultSizes(results []openAIResponsesImageResult, firstMeta *openAIResponsesImageResult) { + for i := range results { + // ChatGPT OAuth can normalize requested controls to "auto". The final + // image bytes are authoritative for response metadata and tier billing. + if actualSize := detectOpenAIImageResultSize(results[i].Result); actualSize != "" { + results[i].Size = actualSize + } + } + if firstMeta == nil || len(results) == 0 { + return + } + if size := strings.TrimSpace(results[0].Size); size != "" { + firstMeta.Size = size + } +} diff --git a/backend/internal/service/openai_images.go b/backend/internal/service/openai_images.go index fc7d37a173..081b5e741b 100644 --- a/backend/internal/service/openai_images.go +++ b/backend/internal/service/openai_images.go @@ -632,6 +632,7 @@ func (s *OpenAIGatewayService) forwardOpenAIImagesAPIKey( if resp.StatusCode >= 400 { respBody := s.readUpstreamErrorBody(resp) _ = resp.Body.Close() + respBody = s.redactAgentIdentitySensitiveBody(upstreamCtx, account, respBody) resp.Body = io.NopCloser(bytes.NewReader(respBody)) upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(respBody)) upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg) @@ -752,7 +753,15 @@ func (s *OpenAIGatewayService) buildOpenAIImagesRequest( return nil, err } req = req.WithContext(WithHTTPUpstreamProfile(req.Context(), HTTPUpstreamProfileOpenAI)) - req.Header.Set("Authorization", "Bearer "+token) + authHeaders, err := s.buildOpenAIAuthenticationHeaders(ctx, account, token) + if err != nil { + return nil, fmt.Errorf("build openai authentication headers: %w", err) + } + for key, values := range authHeaders { + for _, value := range values { + req.Header.Add(key, value) + } + } for key, values := range c.Request.Header { if !openaiPassthroughAllowedHeaders[strings.ToLower(key)] { continue diff --git a/backend/internal/service/openai_images_actual_size_test.go b/backend/internal/service/openai_images_actual_size_test.go new file mode 100644 index 0000000000..4455e415a0 --- /dev/null +++ b/backend/internal/service/openai_images_actual_size_test.go @@ -0,0 +1,174 @@ +package service + +import ( + "bytes" + "context" + "encoding/base64" + "encoding/binary" + "fmt" + "image" + "image/color" + "image/jpeg" + "image/png" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +func TestDetectOpenAIImageResultSize(t *testing.T) { + pngEncoded := encodeOpenAIImageTestPNG(t, 1672, 941) + jpegEncoded := encodeOpenAIImageTestJPEG(t, 640, 360) + webpVP8XEncoded := encodeOpenAIImageTestWebPVP8X(1920, 1080) + webpVP8Encoded := encodeOpenAIImageTestWebPVP8(1280, 720) + webpVP8LEncoded := encodeOpenAIImageTestWebPVP8L(640, 480) + + require.Equal(t, "1672x941", detectOpenAIImageResultSize(pngEncoded)) + require.Equal(t, "1672x941", detectOpenAIImageResultSize(strings.TrimRight(pngEncoded, "="))) + require.Equal(t, "1672x941", detectOpenAIImageResultSize("data:image/png;base64,"+pngEncoded)) + require.Equal(t, "640x360", detectOpenAIImageResultSize(jpegEncoded)) + require.Equal(t, "1920x1080", detectOpenAIImageResultSize(webpVP8XEncoded)) + require.Equal(t, "1280x720", detectOpenAIImageResultSize(webpVP8Encoded)) + require.Equal(t, "640x480", detectOpenAIImageResultSize(webpVP8LEncoded)) + require.Empty(t, detectOpenAIImageResultSize("data:image/png;base64")) + require.Empty(t, detectOpenAIImageResultSize("not-image-data")) +} + +func TestOpenAIGatewayServiceForwardImages_OAuthUsesDecodedOutputDimensions(t *testing.T) { + run := runOpenAIOAuthImageActualSizeTest(t, false) + + require.Equal(t, "3840x2160", gjson.GetBytes(run.upstream.lastBody, "tools.0.size").String()) + require.Equal(t, "low", gjson.GetBytes(run.upstream.lastBody, "tools.0.quality").String()) + require.Equal(t, "1672x941", gjson.Get(run.recorder.Body.String(), "size").String()) + require.Equal(t, "auto", gjson.Get(run.recorder.Body.String(), "quality").String()) + require.Equal(t, []string{"1672x941"}, run.result.ImageOutputSizes) + + ApplyOpenAIImageBillingResolution(run.result) + require.Equal(t, ImageBillingSize2K, run.result.ImageSize) + require.Equal(t, "1672x941", run.result.ImageOutputSize) + require.Equal(t, ImageSizeSourceOutput, run.result.ImageSizeSource) +} + +func TestOpenAIGatewayServiceForwardImages_OAuthStreamingUsesDecodedOutputDimensions(t *testing.T) { + run := runOpenAIOAuthImageActualSizeTest(t, true) + + events := parseOpenAIImageTestSSEEvents(run.recorder.Body.String()) + completed, ok := findOpenAIImageTestSSEEvent(events, "image_generation.completed") + require.True(t, ok) + require.Equal(t, "1672x941", gjson.Get(completed.Data, "size").String()) + require.Equal(t, "auto", gjson.Get(completed.Data, "quality").String()) + require.Equal(t, []string{"1672x941"}, run.result.ImageOutputSizes) +} + +type openAIOAuthImageActualSizeTestRun struct { + result *OpenAIForwardResult + recorder *httptest.ResponseRecorder + upstream *httpUpstreamRecorder +} + +func runOpenAIOAuthImageActualSizeTest(t *testing.T, stream bool) openAIOAuthImageActualSizeTestRun { + t.Helper() + gin.SetMode(gin.TestMode) + body := []byte(fmt.Sprintf(`{"model":"gpt-image-2","prompt":"draw a test chart","size":"3840x2160","quality":"low","output_format":"png","stream":%t}`, stream)) + req := httptest.NewRequest(http.MethodPost, "/v1/images/generations", bytes.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = req + c.Set("api_key", &APIKey{ID: 42}) + + encoded := encodeOpenAIImageTestPNG(t, 1672, 941) + upstreamBody := fmt.Sprintf( + "data: {\"type\":\"response.created\",\"response\":{\"created_at\":1710000000,\"tools\":[{\"type\":\"image_generation\",\"model\":\"gpt-image-2\",\"size\":\"auto\",\"quality\":\"auto\",\"output_format\":\"png\"}]}}\n\n"+ + "data: {\"type\":\"response.completed\",\"response\":{\"created_at\":1710000000,\"tools\":[{\"type\":\"image_generation\",\"model\":\"gpt-image-2\",\"size\":\"auto\",\"quality\":\"auto\",\"output_format\":\"png\"}],\"output\":[{\"id\":\"ig_actual_size\",\"type\":\"image_generation_call\",\"result\":%q}]}}\n\n"+ + "data: [DONE]\n\n", + encoded, + ) + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "Content-Type": []string{"text/event-stream"}, + "X-Request-Id": []string{"req_img_actual_size"}, + }, + Body: io.NopCloser(strings.NewReader(upstreamBody)), + }} + svc := &OpenAIGatewayService{httpUpstream: upstream} + parsed, err := svc.ParseOpenAIImagesRequest(c, body) + require.NoError(t, err) + + account := &Account{ + ID: 1, + Name: "openai-oauth", + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "access_token": "token-123", + "chatgpt_account_id": "acct-123", + }, + } + result, err := svc.ForwardImages(context.Background(), c, account, body, parsed, "") + require.NoError(t, err) + require.NotNil(t, result) + return openAIOAuthImageActualSizeTestRun{result: result, recorder: rec, upstream: upstream} +} + +func encodeOpenAIImageTestPNG(t *testing.T, width, height int) string { + t.Helper() + img := image.NewNRGBA(image.Rect(0, 0, width, height)) + img.SetNRGBA(0, 0, color.NRGBA{R: 0xff, A: 0xff}) + var buf bytes.Buffer + require.NoError(t, png.Encode(&buf, img)) + return base64.StdEncoding.EncodeToString(buf.Bytes()) +} + +func encodeOpenAIImageTestJPEG(t *testing.T, width, height int) string { + t.Helper() + img := image.NewNRGBA(image.Rect(0, 0, width, height)) + img.SetNRGBA(0, 0, color.NRGBA{G: 0xff, A: 0xff}) + var buf bytes.Buffer + require.NoError(t, jpeg.Encode(&buf, img, nil)) + return base64.StdEncoding.EncodeToString(buf.Bytes()) +} + +func encodeOpenAIImageTestWebPVP8X(width, height int) string { + header := make([]byte, 30) + copy(header[0:4], "RIFF") + copy(header[8:12], "WEBP") + copy(header[12:16], "VP8X") + width-- + height-- + header[24], header[25], header[26] = byte(width), byte(width>>8), byte(width>>16) + header[27], header[28], header[29] = byte(height), byte(height>>8), byte(height>>16) + return base64.StdEncoding.EncodeToString(header) +} + +func encodeOpenAIImageTestWebPVP8(width, height int) string { + header := make([]byte, 30) + copy(header[0:4], "RIFF") + copy(header[8:12], "WEBP") + copy(header[12:16], "VP8 ") + copy(header[23:26], "\x9d\x01\x2a") + binary.LittleEndian.PutUint16(header[26:28], uint16(width)) + binary.LittleEndian.PutUint16(header[28:30], uint16(height)) + return base64.StdEncoding.EncodeToString(header) +} + +func encodeOpenAIImageTestWebPVP8L(width, height int) string { + header := make([]byte, 25) + copy(header[0:4], "RIFF") + copy(header[8:12], "WEBP") + copy(header[12:16], "VP8L") + header[20] = 0x2f + width-- + height-- + header[21] = byte(width) + header[22] = byte(width>>8)&0x3f | byte(height&0x03)<<6 + header[23] = byte(height >> 2) + header[24] = byte(height>>10) & 0x0f + return base64.StdEncoding.EncodeToString(header) +} diff --git a/backend/internal/service/openai_images_responses.go b/backend/internal/service/openai_images_responses.go index 363ba9f591..151de2c4b1 100644 --- a/backend/internal/service/openai_images_responses.go +++ b/backend/internal/service/openai_images_responses.go @@ -565,12 +565,14 @@ func collectOpenAIImagesFromResponsesBody(body []byte) ([]openAIResponsesImageRe return nil, 0, nil, openAIResponsesImageResult{}, false, collectErr } if len(finalResults) > 0 { + reconcileOpenAIResponsesImageResultSizes(finalResults, &finalMeta) return finalResults, createdAt, usageRaw, finalMeta, true, nil } if len(fallbackResults) > 0 { firstMeta := fallbackResults[0] mergeOpenAIResponsesImageMeta(&firstMeta, responseMeta) + reconcileOpenAIResponsesImageResultSizes(fallbackResults, &firstMeta) return fallbackResults, createdAt, usageRaw, firstMeta, foundFinal, nil } return nil, createdAt, usageRaw, openAIResponsesImageResult{}, foundFinal, nil @@ -1262,6 +1264,7 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthStreamingResponse( mergeOpenAIResponsesImageMeta(&img, streamMeta) appendOpenAIResponsesImageResultDedup(&finalResults, finalSeen, "", img) } + reconcileOpenAIResponsesImageResultSizes(finalResults, nil) if len(finalResults) == 0 { outputErr := fmt.Errorf("upstream did not return image output") // 软失败:response.completed 事件里没有图片。记录上游诊断摘要到 ops, @@ -1324,8 +1327,12 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthStreamingResponse( } if len(pendingResults) > 0 { eventName := streamPrefix + ".completed" - for _, img := range pendingResults { - mergeOpenAIResponsesImageMeta(&img, streamMeta) + finalResults := append([]openAIResponsesImageResult(nil), pendingResults...) + for i := range finalResults { + mergeOpenAIResponsesImageMeta(&finalResults[i], streamMeta) + } + reconcileOpenAIResponsesImageResultSizes(finalResults, nil) + for _, img := range finalResults { key := openAIResponsesImageResultKey("", img) if _, exists := emitted[key]; exists { continue @@ -1335,7 +1342,7 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthStreamingResponse( s.tryWriteOpenAIImagesStreamEvent(c, flusher, &clientDisconnected, &lastDownstreamWriteAt, eventName, payload) } imageCount = len(emitted) - imageOutputSizes = openAIResponsesImageResultSizes(pendingResults) + imageOutputSizes = openAIResponsesImageResultSizes(finalResults) return nil } @@ -1564,6 +1571,14 @@ func (s *OpenAIGatewayService) forwardOpenAIImagesOAuth( if resp.StatusCode >= 400 { respBody := s.readUpstreamErrorBody(resp) _ = resp.Body.Close() + respBody = s.redactAgentIdentitySensitiveBody(upstreamCtx, account, respBody) + if !agentIdentityTaskRecoveryWasTried(ctx) && s.isAgentIdentityAccount(ctx, account) && isAgentIdentityTaskInvalidHTTPResponse(resp.StatusCode, respBody) { + expectedTaskID := account.GetCredential("task_id") + if err := s.recoverAgentIdentityTask(ctx, account, expectedTaskID); err != nil { + return nil, fmt.Errorf("agent identity task recovery failed: %w", err) + } + return s.forwardOpenAIImagesOAuth(markAgentIdentityTaskRecoveryTried(ctx), c, account, parsed, channelMappedModel) + } resp.Body = io.NopCloser(bytes.NewReader(respBody)) upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(respBody)) upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg) diff --git a/backend/internal/service/openai_quota_service.go b/backend/internal/service/openai_quota_service.go index 91b511122e..63ce173639 100644 --- a/backend/internal/service/openai_quota_service.go +++ b/backend/internal/service/openai_quota_service.go @@ -8,6 +8,7 @@ import ( "log/slog" "net/http" "strings" + "sync" "time" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" @@ -21,6 +22,8 @@ import ( // errors.Is still matches it by identity since ResetCredit returns this var. var ErrSparkShadowResetNotSupported = infraerrors.New(http.StatusConflict, "SPARK_SHADOW_RESET_NOT_SUPPORTED", "spark shadow account does not support credit reset; reset the parent account") +var ErrAgentIdentityResetNotSupported = infraerrors.New(http.StatusConflict, "AGENT_IDENTITY_RESET_NOT_SUPPORTED", "agent identity does not support rate-limit reset credit consumption") + // Endpoints used by the OpenAI/ChatGPT/Codex quota query and reset feature. const ( chatGPTUsageURL = "https://chatgpt.com/backend-api/wham/usage" @@ -116,6 +119,8 @@ type OpenAIQuotaService struct { proxyRepo ProxyRepository tokenProvider *OpenAITokenProvider privacyClientFactory PrivacyClientFactory + agentIdentityTaskMu sync.Mutex + agentIdentityWS agentIdentityWSConnectionInvalidator } // NewOpenAIQuotaService constructs a quota service. token provider is required — @@ -151,21 +156,36 @@ func (s *OpenAIQuotaService) QueryUsage(ctx context.Context, accountID int64) (* callCtx, cancel := context.WithTimeout(ctx, openaiQuotaUpstreamTimeout) defer cancel() + agentIdentity := s.isAgentIdentityAccount(ctx, accountID) var payload OpenAIQuotaUsage - resp, err := client.R(). - SetContext(callCtx). - SetHeaders(buildCodexCommonHeaders(accessToken, chatGPTAccountID, fedRAMP)). - SetSuccessResult(&payload). - Get(chatGPTUsageURL) - if err != nil { - return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_QUOTA_REQUEST_FAILED", "upstream request failed: %v", err) - } - if !resp.IsSuccessState() { - status := resp.StatusCode - body := truncate(resp.String(), 240) - slog.Warn("openai_quota_query_failed", "account_id", accountID, "status", status, "body", body) - return nil, infraerrors.Newf(mapUpstreamStatus(status), "OPENAI_QUOTA_UPSTREAM_ERROR", "upstream returned %d: %s", status, body) + for recovered := false; ; { + quotaHeaders, headerErr := s.buildCodexQuotaHeaders(callCtx, accountID, accessToken, chatGPTAccountID, fedRAMP) + if headerErr != nil { + return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_QUOTA_AUTH_FAILED", "failed to build upstream authentication: %v", headerErr) + } + resp, err := client.R(). + SetContext(callCtx). + SetHeaders(quotaHeaders). + SetSuccessResult(&payload). + Get(chatGPTUsageURL) + if err != nil { + return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_QUOTA_REQUEST_FAILED", "upstream request failed: %v", err) + } + if !resp.IsSuccessState() { + if agentIdentity && !recovered && isAgentIdentityTaskInvalidHTTPResponse(resp.StatusCode, []byte(resp.String())) { + recovered = true + if err := s.recoverAgentIdentityTask(ctx, accountID); err != nil { + return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_QUOTA_AUTH_FAILED", "agent identity task recovery failed: %v", err) + } + continue + } + status := resp.StatusCode + body := truncate(s.redactQuotaErrorBody(ctx, accountID, resp.String()), 240) + slog.Warn("openai_quota_query_failed", "account_id", accountID, "status", status, "body", body) + return nil, infraerrors.Newf(mapUpstreamStatus(status), "OPENAI_QUOTA_UPSTREAM_ERROR", "upstream returned %d: %s", status, body) + } + break } payload.FetchedAt = time.Now().Unix() @@ -189,9 +209,14 @@ func (s *OpenAIQuotaService) QueryUsage(ctx context.Context, accountID int64) (* } func (s *OpenAIQuotaService) queryResetCreditDetails(ctx context.Context, client *req.Client, accessToken, chatGPTAccountID string, fedRAMP bool, accountID int64) *openAIRateLimitResetCreditDetails { + quotaHeaders, headerErr := s.buildCodexQuotaHeaders(ctx, accountID, accessToken, chatGPTAccountID, fedRAMP) + if headerErr != nil { + slog.Warn("openai_quota_reset_credit_details_auth_failed", "account_id", accountID, "error", headerErr) + return nil + } resp, err := client.R(). SetContext(ctx). - SetHeaders(buildCodexCommonHeaders(accessToken, chatGPTAccountID, fedRAMP)). + SetHeaders(quotaHeaders). Get(chatGPTRateLimitCreditsURL) if err != nil { slog.Warn("openai_quota_reset_credit_details_failed", "account_id", accountID, "error", err) @@ -233,6 +258,9 @@ func (s *OpenAIQuotaService) ResetCredit(ctx context.Context, accountID int64) ( if acc.IsShadow() { return nil, ErrSparkShadowResetNotSupported } + if acc.IsOpenAIAgentIdentity() { + return nil, ErrAgentIdentityResetNotSupported + } } accessToken, chatGPTAccountID, proxyURL, fedRAMP, err := s.prepareUpstreamCall(ctx, accountID) @@ -252,11 +280,12 @@ func (s *OpenAIQuotaService) ResetCredit(ctx context.Context, accountID int64) ( callCtx, cancel := context.WithTimeout(ctx, openaiQuotaUpstreamTimeout) defer cancel() - - headers := buildCodexCommonHeaders(accessToken, chatGPTAccountID, fedRAMP) - headers["content-type"] = "application/json" - var payload OpenAIQuotaResetResult + headers, headerErr := s.buildCodexQuotaHeaders(callCtx, accountID, accessToken, chatGPTAccountID, fedRAMP) + if headerErr != nil { + return nil, infraerrors.Newf(http.StatusBadGateway, "OPENAI_QUOTA_AUTH_FAILED", "failed to build upstream authentication: %v", headerErr) + } + headers["content-type"] = "application/json" resp, err := client.R(). SetContext(callCtx). SetHeaders(headers). @@ -268,7 +297,7 @@ func (s *OpenAIQuotaService) ResetCredit(ctx context.Context, accountID int64) ( } if !resp.IsSuccessState() { status := resp.StatusCode - body := truncate(resp.String(), 240) + body := truncate(s.redactQuotaErrorBody(callCtx, accountID, resp.String()), 240) slog.Warn("openai_quota_reset_failed", "account_id", accountID, "status", status, "body", body) return nil, infraerrors.Newf(mapUpstreamStatus(status), "OPENAI_QUOTA_RESET_UPSTREAM_ERROR", "upstream returned %d: %s", status, body) } @@ -285,7 +314,7 @@ func (s *OpenAIQuotaService) ResetCredit(ctx context.Context, accountID int64) ( // token via the shared TokenProvider, and resolves the chatgpt-account-id and // proxy URL. Centralized so QueryUsage / ResetCredit share validation. func (s *OpenAIQuotaService) prepareUpstreamCall(ctx context.Context, accountID int64) (accessToken, chatGPTAccountID, proxyURL string, fedRAMP bool, err error) { - if s == nil || s.accountRepo == nil || s.tokenProvider == nil || s.privacyClientFactory == nil { + if s == nil || s.accountRepo == nil || s.privacyClientFactory == nil { return "", "", "", false, infraerrors.New(http.StatusInternalServerError, "OPENAI_QUOTA_NOT_CONFIGURED", "openai quota service is not configured") } @@ -323,12 +352,17 @@ func (s *OpenAIQuotaService) prepareUpstreamCall(ctx context.Context, accountID return "", "", "", false, infraerrors.New(http.StatusBadRequest, "OPENAI_QUOTA_MISSING_ACCOUNT_ID", "chatgpt_account_id is missing; please re-authorize this account") } - accessToken, err = s.tokenProvider.GetAccessToken(ctx, account) - if err != nil { - return "", "", "", false, infraerrors.Newf(http.StatusBadGateway, "OPENAI_QUOTA_TOKEN_UNAVAILABLE", "failed to acquire access token: %v", err) - } - if strings.TrimSpace(accessToken) == "" { - return "", "", "", false, infraerrors.New(http.StatusBadGateway, "OPENAI_QUOTA_TOKEN_UNAVAILABLE", "access token is empty") + if !account.IsOpenAIAgentIdentity() { + if s.tokenProvider == nil { + return "", "", "", false, infraerrors.New(http.StatusInternalServerError, "OPENAI_QUOTA_NOT_CONFIGURED", "openai quota token provider is not configured") + } + accessToken, err = s.tokenProvider.GetAccessToken(ctx, account) + if err != nil { + return "", "", "", false, infraerrors.Newf(http.StatusBadGateway, "OPENAI_QUOTA_TOKEN_UNAVAILABLE", "failed to acquire access token: %v", err) + } + if strings.TrimSpace(accessToken) == "" { + return "", "", "", false, infraerrors.New(http.StatusBadGateway, "OPENAI_QUOTA_TOKEN_UNAVAILABLE", "access token is empty") + } } fedRAMP = account.IsChatGPTAccountFedRAMP() @@ -351,6 +385,91 @@ func (s *OpenAIQuotaService) prepareUpstreamCall(ctx context.Context, accountID return accessToken, chatGPTAccountID, proxyURL, fedRAMP, nil } +func (s *OpenAIQuotaService) recoverAgentIdentityTask(ctx context.Context, accountID int64) error { + if s == nil || s.accountRepo == nil { + return fmt.Errorf("account repository is unavailable") + } + account, err := s.accountRepo.GetByID(ctx, accountID) + if err != nil || account == nil { + return fmt.Errorf("account is unavailable") + } + if account.IsShadow() { + account, err = resolveCredentialAccount(ctx, s.accountRepo, account) + if err != nil || account == nil { + return fmt.Errorf("credential account is unavailable") + } + } + if !account.IsOpenAIAgentIdentity() { + return nil + } + return ensureAgentIdentityTaskForAccount(ctx, s.accountRepo, s.agentIdentityWS, &s.agentIdentityTaskMu, account, account.GetCredential("task_id")) +} + +func (s *OpenAIQuotaService) isAgentIdentityAccount(ctx context.Context, accountID int64) bool { + if s == nil || s.accountRepo == nil { + return false + } + account, err := s.accountRepo.GetByID(ctx, accountID) + if err != nil || account == nil { + return false + } + if account.IsShadow() { + account, err = resolveCredentialAccount(ctx, s.accountRepo, account) + if err != nil || account == nil { + return false + } + } + return account.IsOpenAIAgentIdentity() +} + +func (s *OpenAIQuotaService) buildCodexQuotaHeaders(ctx context.Context, accountID int64, accessToken, chatGPTAccountID string, fedRAMP bool) (map[string]string, error) { + headers := buildCodexCommonHeaders(accessToken, chatGPTAccountID, fedRAMP) + if s == nil || s.accountRepo == nil { + return headers, nil + } + account, err := s.accountRepo.GetByID(ctx, accountID) + if err != nil || account == nil { + if strings.TrimSpace(accessToken) == "" { + return nil, fmt.Errorf("agent identity account credentials are unavailable") + } + return headers, nil + } + if account.IsShadow() { + if resolved, resolveErr := resolveCredentialAccount(ctx, s.accountRepo, account); resolveErr == nil && resolved != nil { + account = resolved + } else if strings.TrimSpace(accessToken) == "" { + return nil, fmt.Errorf("agent identity shadow credentials are unavailable") + } + } + if !account.IsOpenAIAgentIdentity() { + return headers, nil + } + if err := ensureAgentIdentityTaskForAccount(ctx, s.accountRepo, s.agentIdentityWS, &s.agentIdentityTaskMu, account, ""); err != nil { + return nil, err + } + key, err := agentIdentityKeyFromAccount(account) + if err != nil { + return nil, err + } + assertion, err := buildAgentAssertion(key, time.Now()) + if err != nil { + return nil, err + } + headers["authorization"] = assertion + return headers, nil +} + +func (s *OpenAIQuotaService) redactQuotaErrorBody(ctx context.Context, accountID int64, body string) string { + if s == nil || s.accountRepo == nil { + return body + } + account, err := s.accountRepo.GetByID(ctx, accountID) + if err != nil || account == nil { + return body + } + return string(redactAgentIdentitySensitiveBodyForAccount(ctx, s.accountRepo, account, []byte(body))) +} + // buildCodexCommonHeaders sets the request headers expected by the chatgpt.com // backend so calls succeed past Cloudflare/WASM checks. func buildCodexCommonHeaders(accessToken, chatGPTAccountID string, fedRAMP bool) map[string]string { diff --git a/backend/internal/service/openai_quota_spark_window_test.go b/backend/internal/service/openai_quota_spark_window_test.go index 213669f5e6..c521b149e9 100644 --- a/backend/internal/service/openai_quota_spark_window_test.go +++ b/backend/internal/service/openai_quota_spark_window_test.go @@ -2,6 +2,10 @@ package service import ( "context" + "crypto/ed25519" + "crypto/rand" + "crypto/x509" + "encoding/base64" "encoding/json" "errors" "fmt" @@ -34,6 +38,15 @@ func (r *stubQuotaAccountRepo) GetByID(_ context.Context, id int64) (*Account, e return acc, nil } +func (r *stubQuotaAccountRepo) UpdateCredentials(_ context.Context, id int64, credentials map[string]any) error { + acc, ok := r.accounts[id] + if !ok { + return fmt.Errorf("account %d not found", id) + } + acc.Credentials = credentials + return nil +} + // stubQuotaTokenCache 实现 OpenAITokenCache,返回预设静态 token。 type stubQuotaTokenCache struct { tokens map[string]string @@ -163,6 +176,23 @@ func TestResetCreditShadowRejected(t *testing.T) { "shadow ResetCredit 应映射为 409 Conflict 而非 500") } +func TestResetCreditAgentIdentityRejectedBeforeUpstream(t *testing.T) { + account := &Account{ + ID: 201, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "auth_mode": OpenAIAuthModeAgentIdentity, + }, + } + repo := &stubQuotaAccountRepo{accounts: map[int64]*Account{account.ID: account}} + svc := &OpenAIQuotaService{accountRepo: repo} + + _, err := svc.ResetCredit(context.Background(), account.ID) + require.ErrorIs(t, err, ErrAgentIdentityResetNotSupported) + require.Equal(t, http.StatusConflict, infraerrors.Code(err)) +} + // ── Part B: prepareUpstreamCall 影子 resolve ────────────────────────────── // TestPrepareUpstreamCallShadowResolve 验证影子账号(200)QueryUsage 时: @@ -214,6 +244,101 @@ func TestPrepareUpstreamCallShadowResolve(t *testing.T) { "prepareUpstreamCall should use parent's chatgpt_account_id after shadow resolve") } +func TestQueryUsageAgentIdentityUsesAssertionWithoutOAuthToken(t *testing.T) { + _, privateKey, err := ed25519.GenerateKey(rand.Reader) + require.NoError(t, err) + der, err := x509.MarshalPKCS8PrivateKey(privateKey) + require.NoError(t, err) + account := &Account{ + ID: 300, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "auth_mode": OpenAIAuthModeAgentIdentity, + "agent_runtime_id": "runtime-quota", + "agent_private_key": base64.StdEncoding.EncodeToString(der), + "task_id": "task-quota", + "chatgpt_account_id": "account-quota", + "chatgpt_account_is_fedramp": true, + }, + } + repo := &stubQuotaAccountRepo{accounts: map[int64]*Account{account.ID: account}} + var authorization string + var accountHeader string + var fedrampHeader string + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + authorization = r.Header.Get("authorization") + accountHeader = r.Header.Get("chatgpt-account-id") + fedrampHeader = r.Header.Get("x-openai-fedramp") + w.Header().Set("content-type", "application/json") + _, _ = w.Write([]byte(`{"plan_type":"pro","rate_limit":{"allowed":true}}`)) + })) + defer srv.Close() + svc := NewOpenAIQuotaService(repo, nil, nil, newQuotaRedirectingFactory(srv)) + usage, err := svc.QueryUsage(context.Background(), account.ID) + require.NoError(t, err) + require.NotNil(t, usage) + require.True(t, strings.HasPrefix(authorization, "AgentAssertion ")) + require.Equal(t, "account-quota", accountHeader) + require.Equal(t, "true", fedrampHeader) +} + +func TestQueryUsageAgentIdentityRecoversInvalidTaskOnce(t *testing.T) { + _, privateKey, err := ed25519.GenerateKey(rand.Reader) + require.NoError(t, err) + der, err := x509.MarshalPKCS8PrivateKey(privateKey) + require.NoError(t, err) + account := &Account{ + ID: 301, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "auth_mode": OpenAIAuthModeAgentIdentity, + "agent_runtime_id": "runtime-quota-recovery", + "agent_private_key": base64.StdEncoding.EncodeToString(der), + "task_id": "task-quota-old", + "chatgpt_account_id": "account-quota-recovery", + }, + } + repo := &stubQuotaAccountRepo{accounts: map[int64]*Account{account.ID: account}} + usageCalls := 0 + registerCalls := 0 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("content-type", "application/json") + if strings.Contains(r.URL.Path, "/task/register") { + registerCalls++ + _, _ = w.Write([]byte(`{"task_id":"task-quota-new"}`)) + return + } + if strings.Contains(r.URL.Path, "rate-limit-reset-credits") { + _, _ = w.Write([]byte(`{}`)) + return + } + usageCalls++ + if usageCalls == 1 { + w.WriteHeader(http.StatusUnauthorized) + _, _ = w.Write([]byte(`{"error":{"code":"invalid_task_id"}}`)) + return + } + _, _ = w.Write([]byte(`{"plan_type":"pro","rate_limit":{"allowed":true}}`)) + })) + defer srv.Close() + oldBase := openAIAgentIdentityAuthAPIBaseURL + openAIAgentIdentityAuthAPIBaseURL = srv.URL + t.Cleanup(func() { openAIAgentIdentityAuthAPIBaseURL = oldBase }) + + invalidator := &agentIdentityWSInvalidationRecorder{} + svc := NewOpenAIQuotaService(repo, nil, nil, newQuotaRedirectingFactory(srv)) + svc.agentIdentityWS = invalidator + usage, err := svc.QueryUsage(context.Background(), account.ID) + require.NoError(t, err) + require.NotNil(t, usage) + require.Equal(t, 2, usageCalls) + require.Equal(t, 1, registerCalls) + require.Equal(t, "task-quota-new", account.GetCredential("task_id")) + require.Equal(t, []int64{account.ID}, invalidator.accountIDs) +} + func TestParseOpenAIRateLimitResetCreditDetails_CompatibleContainers(t *testing.T) { tests := []struct { name string diff --git a/backend/internal/service/openai_ws_client.go b/backend/internal/service/openai_ws_client.go index e336abcbd6..c9b92d22e8 100644 --- a/backend/internal/service/openai_ws_client.go +++ b/backend/internal/service/openai_ws_client.go @@ -4,6 +4,7 @@ import ( "context" "errors" "fmt" + "io" "net/http" "net/url" "strings" @@ -68,6 +69,28 @@ type coderOpenAIWSClientDialer struct { proxyMisses atomic.Int64 } +// openAIWSHandshakeError keeps a bounded, non-logged HTTP error body so the +// Agent Identity recovery path can distinguish an invalid task from other +// 401 handshake failures. +type openAIWSHandshakeError struct { + Body []byte + Err error +} + +func (e *openAIWSHandshakeError) Error() string { + if e == nil || e.Err == nil { + return "openai ws handshake failed" + } + return e.Err.Error() +} + +func (e *openAIWSHandshakeError) Unwrap() error { + if e == nil { + return nil + } + return e.Err +} + type openAIWSProxyClientEntry struct { client *http.Client lastUsedUnixNano int64 @@ -104,7 +127,12 @@ func (d *coderOpenAIWSClientDialer) Dial( status = resp.StatusCode respHeaders = cloneHeader(resp.Header) } - return nil, status, respHeaders, err + var body []byte + if resp != nil && resp.Body != nil { + body, _ = io.ReadAll(io.LimitReader(resp.Body, 8<<10)) + _ = resp.Body.Close() + } + return nil, status, respHeaders, &openAIWSHandshakeError{Body: body, Err: err} } // coder/websocket 默认单消息读取上限为 32KB,Codex WS 事件(如 rate_limits/大 delta) // 可能超过该阈值,需显式提高上限,避免本地 read_fail(message too big)。 diff --git a/backend/internal/service/openai_ws_forwarder_ingress.go b/backend/internal/service/openai_ws_forwarder_ingress.go index b632c451df..635312738b 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress.go +++ b/backend/internal/service/openai_ws_forwarder_ingress.go @@ -46,8 +46,8 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( if account == nil { return errors.New("account is nil") } - if strings.TrimSpace(token) == "" { - return errors.New("token is empty") + if err := validateOpenAIWSBearerToken(account, token); err != nil { + return err } // 预取一次 OpenAI Fast Policy settings,绑定到 ctx,让该 WS session @@ -577,6 +577,9 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( Account: account, WSURL: wsURL, Headers: wsHeaders, + HeadersFactory: func(factoryCtx context.Context, headers http.Header) (http.Header, error) { + return s.refreshOpenAIAgentIdentityHeaders(factoryCtx, account, headers) + }, ProxyURL: func() string { if account.ProxyID != nil && account.Proxy != nil { return account.Proxy.URL() @@ -640,7 +643,9 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( acquireTimeout = 30 * time.Second } - acquireTurnLease := func(turn int, preferred string, forcePreferredConn bool) (*openAIWSConnLease, error) { + agentTaskRecoveryTried := false + var acquireTurnLease func(int, string, bool) (*openAIWSConnLease, error) + acquireTurnLease = func(turn int, preferred string, forcePreferredConn bool) (*openAIWSConnLease, error) { req := cloneOpenAIWSAcquireRequest(baseAcquireReq) req.PreferredConnID = strings.TrimSpace(preferred) req.ForcePreferredConn = forcePreferredConn @@ -649,6 +654,14 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( acquireCtx, acquireCancel := context.WithTimeout(ctx, acquireTimeout) lease, acquireErr := pool.Acquire(acquireCtx, req) acquireCancel() + var dialErr *openAIWSDialError + if acquireErr != nil && s.isAgentIdentityAccount(ctx, account) && errors.As(acquireErr, &dialErr) && isAgentIdentityTaskInvalidWSDialError(dialErr) && !agentTaskRecoveryTried { + agentTaskRecoveryTried = true + if recoveryErr := s.recoverAgentIdentityTask(ctx, account, account.GetCredential("task_id")); recoveryErr != nil { + return nil, fmt.Errorf("agent identity task recovery failed: %w", recoveryErr) + } + return acquireTurnLease(turn, preferred, forcePreferredConn) + } if acquireErr != nil { dialStatus, dialClass, dialCloseStatus, dialCloseReason, dialRespServer, dialRespVia, dialRespCFRay, dialRespReqID := summarizeOpenAIWSDialError(acquireErr) logOpenAIWSModeInfo( diff --git a/backend/internal/service/openai_ws_forwarder_payload.go b/backend/internal/service/openai_ws_forwarder_payload.go index 5830444815..511b28d0ba 100644 --- a/backend/internal/service/openai_ws_forwarder_payload.go +++ b/backend/internal/service/openai_ws_forwarder_payload.go @@ -15,6 +15,16 @@ import ( "github.com/tidwall/sjson" ) +func validateOpenAIWSBearerToken(account *Account, token string) error { + if account == nil { + return errors.New("account is nil") + } + if strings.TrimSpace(token) == "" && !account.IsOpenAIAgentIdentity() { + return errors.New("token is empty") + } + return nil +} + func (s *OpenAIGatewayService) buildOpenAIResponsesWSURL(account *Account) (string, error) { if account == nil { return "", errors.New("account is nil") @@ -67,7 +77,9 @@ func (s *OpenAIGatewayService) buildOpenAIWSHeaders( promptCacheKey string, ) (http.Header, openAIWSSessionHeaderResolution, error) { headers := make(http.Header) - headers.Set("authorization", "Bearer "+token) + if account == nil || !account.IsOpenAIAgentIdentity() { + headers.Set("authorization", "Bearer "+token) + } sessionResolution := resolveOpenAIWSSessionHeaders(c, promptCacheKey) if c != nil && c.Request != nil { diff --git a/backend/internal/service/openai_ws_forwarder_v2.go b/backend/internal/service/openai_ws_forwarder_v2.go index f0f71648dd..90151be93a 100644 --- a/backend/internal/service/openai_ws_forwarder_v2.go +++ b/backend/internal/service/openai_ws_forwarder_v2.go @@ -30,6 +30,7 @@ func (s *OpenAIGatewayService) forwardOpenAIWSV2( startTime time.Time, attempt int, lastFailureReason string, + agentTaskRecoveryTried *bool, ) (*OpenAIForwardResult, error) { if s == nil || account == nil { return nil, wrapOpenAIWSFallback("invalid_state", errors.New("service or account is nil")) @@ -174,9 +175,12 @@ func (s *OpenAIGatewayService) forwardOpenAIWSV2( defer acquireCancel() lease, err := s.getOpenAIWSConnPool().Acquire(acquireCtx, openAIWSAcquireRequest{ - Account: account, - WSURL: wsURL, - Headers: wsHeaders, + Account: account, + WSURL: wsURL, + Headers: wsHeaders, + HeadersFactory: func(factoryCtx context.Context, headers http.Header) (http.Header, error) { + return s.refreshOpenAIAgentIdentityHeaders(factoryCtx, account, headers) + }, PreferredConnID: preferredConnID, ForceNewConn: forceNewConn, ProxyURL: func() string { @@ -187,6 +191,14 @@ func (s *OpenAIGatewayService) forwardOpenAIWSV2( }(), }) if err != nil { + var agentDialErr *openAIWSDialError + if s.isAgentIdentityAccount(ctx, account) && errors.As(err, &agentDialErr) && isAgentIdentityTaskInvalidWSDialError(agentDialErr) && agentTaskRecoveryTried != nil && !*agentTaskRecoveryTried { + *agentTaskRecoveryTried = true + if recoveryErr := s.recoverAgentIdentityTask(ctx, account, account.GetCredential("task_id")); recoveryErr != nil { + return nil, fmt.Errorf("agent identity task recovery failed: %w", recoveryErr) + } + return nil, &agentIdentityTaskRecoveredError{} + } dialStatus, dialClass, dialCloseStatus, dialCloseReason, dialRespServer, dialRespVia, dialRespCFRay, dialRespReqID := summarizeOpenAIWSDialError(err) logOpenAIWSModeInfo( "acquire_fail account_id=%d account_type=%s transport=%s reason=%s dial_status=%d dial_class=%s dial_close_status=%s dial_close_reason=%s dial_resp_server=%s dial_resp_via=%s dial_resp_cf_ray=%s dial_resp_x_request_id=%s cause=%s preferred_conn_id=%s force_new_conn=%v ws_host=%s ws_path=%s proxy_enabled=%v", diff --git a/backend/internal/service/openai_ws_http_bridge.go b/backend/internal/service/openai_ws_http_bridge.go index dbc24c850a..77a5e59329 100644 --- a/backend/internal/service/openai_ws_http_bridge.go +++ b/backend/internal/service/openai_ws_http_bridge.go @@ -12,7 +12,6 @@ import ( "time" "github.com/Wei-Shaw/sub2api/internal/config" - "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/gin-gonic/gin" "github.com/tidwall/gjson" ) @@ -193,7 +192,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( releaseUpstreamCtx() return nil, fmt.Errorf("apply grok prompt cache identity: %w", err) } - upstreamReq, err = buildGrokResponsesRequest(upstreamCtx, c, account, body, token, grokCacheIdentity) + upstreamReq, err = buildGrokResponsesRequest(upstreamCtx, c, account, body, token, grokCacheIdentity, s.cfg) } else { upstreamReq, err = s.buildUpstreamRequestOpenAIPassthrough(upstreamCtx, c, account, body, token) } @@ -236,7 +235,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( return nil, fmt.Errorf("upstream http bridge error: status=%d message=%s", resp.StatusCode, upstreamMsg) } if account.Platform == PlatformGrok { - s.updateGrokUsageSnapshot(ctx, account, xai.ParseQuotaHeaders(resp.Header, resp.StatusCode)) + s.updateGrokUsageFromResponse(ctx, account, resp.Header, resp.StatusCode) } responseID := "" diff --git a/backend/internal/service/openai_ws_pool.go b/backend/internal/service/openai_ws_pool.go index 8affe2a930..8579cc58d6 100644 --- a/backend/internal/service/openai_ws_pool.go +++ b/backend/internal/service/openai_ws_pool.go @@ -39,6 +39,7 @@ var ( type openAIWSDialError struct { StatusCode int ResponseHeaders http.Header + ResponseBody []byte Err error } @@ -60,9 +61,13 @@ func (e *openAIWSDialError) Unwrap() error { } type openAIWSAcquireRequest struct { - Account *Account - WSURL string - Headers http.Header + Account *Account + WSURL string + Headers http.Header + // HeadersFactory is evaluated inside dialConn. It exists so credentials + // whose authorization is per-dial (Agent Identity) are never cached in + // lastAcquire or delayed prewarm state. + HeadersFactory func(context.Context, http.Header) (http.Header, error) ProxyURL string PreferredConnID string // ForceNewConn: 强制本次获取新连接(避免复用导致连接内续链状态互相污染)。 @@ -544,6 +549,7 @@ type openAIWSAccountPool struct { pinnedConns map[string]int changedCh chan struct{} creating int + generation uint64 lastCleanupAt time.Time lastAcquire *openAIWSAcquireRequest prewarmActive bool @@ -1391,6 +1397,7 @@ func (p *openAIWSConnPool) ensureTargetIdleAsync(accountID int64) { } var req openAIWSAcquireRequest + generation := uint64(0) need := 0 ap, ok := p.getAccountPool(accountID) if !ok || ap == nil { @@ -1425,6 +1432,7 @@ func (p *openAIWSConnPool) ensureTargetIdleAsync(accountID int64) { return } req = cloneOpenAIWSAcquireRequest(*ap.lastAcquire) + generation = ap.generation ap.prewarmActive = true if cooldown := p.prewarmCooldown(); cooldown > 0 { ap.prewarmUntil = now.Add(cooldown) @@ -1432,7 +1440,7 @@ func (p *openAIWSConnPool) ensureTargetIdleAsync(accountID int64) { ap.creating += need p.metrics.scaleUpTotal.Add(int64(need)) - go p.prewarmConns(accountID, req, need) + go p.prewarmConns(accountID, req, need, generation) } func (p *openAIWSConnPool) targetConnCountLocked(ap *openAIWSAccountPool, maxConns int) int { @@ -1475,7 +1483,11 @@ func (p *openAIWSConnPool) targetConnCountLocked(ap *openAIWSAccountPool, maxCon return target } -func (p *openAIWSConnPool) prewarmConns(accountID int64, req openAIWSAcquireRequest, total int) { +func (p *openAIWSConnPool) prewarmConns(accountID int64, req openAIWSAcquireRequest, total int, generations ...uint64) { + generation := uint64(0) + if len(generations) > 0 { + generation = generations[0] + } defer func() { if ap, ok := p.getAccountPool(accountID); ok && ap != nil { ap.mu.Lock() @@ -1507,6 +1519,11 @@ func (p *openAIWSConnPool) prewarmConns(accountID int64, req openAIWSAcquireRequ ap.mu.Unlock() continue } + if ap.generation != generation || ap.lastAcquire == nil { + ap.mu.Unlock() + conn.close() + continue + } if len(ap.conns) >= p.effectiveMaxConnsByAccount(req.Account) { ap.signalChangedLocked() ap.mu.Unlock() @@ -1521,6 +1538,35 @@ func (p *openAIWSConnPool) prewarmConns(accountID int64, req openAIWSAcquireRequ } } +// ClearAccount closes all pooled connections and discards delayed prewarm +// state for one account. The generation guard prevents an in-flight prewarm +// started before credential recovery from re-entering the pool afterwards. +func (p *openAIWSConnPool) ClearAccount(accountID int64) { + if p == nil || accountID <= 0 { + return + } + ap, ok := p.getAccountPool(accountID) + if !ok || ap == nil { + return + } + ap.mu.Lock() + ap.generation++ + conns := make([]*openAIWSConn, 0, len(ap.conns)) + for id, conn := range ap.conns { + delete(ap.conns, id) + delete(ap.pinnedConns, id) + if conn != nil { + conns = append(conns, conn) + } + } + ap.lastAcquire = nil + ap.prewarmUntil = time.Time{} + ap.prewarmFails = 0 + ap.prewarmFailAt = time.Time{} + ap.mu.Unlock() + closeOpenAIWSConns(conns) +} + func (p *openAIWSConnPool) evictConn(accountID int64, connID string) { if p == nil || accountID <= 0 || stringsTrim(connID) == "" { return @@ -1599,11 +1645,25 @@ func (p *openAIWSConnPool) dialConn(ctx context.Context, req openAIWSAcquireRequ if p == nil || p.clientDialer == nil { return nil, errors.New("openai ws client dialer is nil") } - conn, status, handshakeHeaders, err := p.clientDialer.Dial(ctx, req.WSURL, req.Headers, req.ProxyURL) + headers := cloneHeader(req.Headers) + var err error + if req.HeadersFactory != nil { + headers, err = req.HeadersFactory(ctx, headers) + if err != nil { + return nil, err + } + } + conn, status, handshakeHeaders, err := p.clientDialer.Dial(ctx, req.WSURL, headers, req.ProxyURL) if err != nil { + var handshakeErr *openAIWSHandshakeError + var responseBody []byte + if errors.As(err, &handshakeErr) && handshakeErr != nil { + responseBody = append([]byte(nil), handshakeErr.Body...) + } return nil, &openAIWSDialError{ StatusCode: status, ResponseHeaders: cloneHeader(handshakeHeaders), + ResponseBody: responseBody, Err: err, } } diff --git a/backend/internal/service/openai_ws_v2_passthrough_adapter.go b/backend/internal/service/openai_ws_v2_passthrough_adapter.go index d822d181cd..72fbbd0f35 100644 --- a/backend/internal/service/openai_ws_v2_passthrough_adapter.go +++ b/backend/internal/service/openai_ws_v2_passthrough_adapter.go @@ -243,8 +243,8 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( if account == nil { return errors.New("account is nil") } - if strings.TrimSpace(token) == "" { - return errors.New("token is empty") + if err := validateOpenAIWSBearerToken(account, token); err != nil { + return err } requestModel := strings.TrimSpace(gjson.GetBytes(firstClientMessage, "model").String()) requestPreviousResponseID := strings.TrimSpace(gjson.GetBytes(firstClientMessage, "previous_response_id").String()) @@ -359,10 +359,34 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( return errors.New("openai ws passthrough dialer is nil") } - dialCtx, cancelDial := context.WithTimeout(ctx, s.openAIWSDialTimeout()) - defer cancelDial() - upstreamConn, statusCode, handshakeHeaders, err := dialer.Dial(dialCtx, wsURL, headers, proxyURL) - if err != nil { + agentTaskRecoveryTried := false + var upstreamConn openAIWSClientConn + statusCode := 0 + var handshakeHeaders http.Header + for { + headers, err = s.refreshOpenAIAgentIdentityHeaders(ctx, account, headers) + if err != nil { + return fmt.Errorf("refresh ws authentication headers: %w", err) + } + dialCtx, cancelDial := context.WithTimeout(ctx, s.openAIWSDialTimeout()) + upstreamConn, statusCode, handshakeHeaders, err = dialer.Dial(dialCtx, wsURL, headers, proxyURL) + cancelDial() + if err == nil { + break + } + var handshakeErr *openAIWSHandshakeError + responseBody := []byte(nil) + if errors.As(err, &handshakeErr) && handshakeErr != nil { + responseBody = handshakeErr.Body + } + dialErr := &openAIWSDialError{StatusCode: statusCode, ResponseBody: responseBody, Err: err} + if s.isAgentIdentityAccount(ctx, account) && isAgentIdentityTaskInvalidWSDialError(dialErr) && !agentTaskRecoveryTried { + agentTaskRecoveryTried = true + if recoveryErr := s.recoverAgentIdentityTask(ctx, account, account.GetCredential("task_id")); recoveryErr != nil { + return fmt.Errorf("agent identity task recovery failed: %w", recoveryErr) + } + continue + } logOpenAIWSV2Passthrough( "relay_dial_failed account_id=%d status_code=%d err=%s", account.ID, @@ -678,9 +702,15 @@ func (s *OpenAIGatewayService) mapOpenAIWSPassthroughDialError( wrappedErr := err var dialErr *openAIWSDialError if !errors.As(err, &dialErr) { + var handshakeErr *openAIWSHandshakeError + var responseBody []byte + if errors.As(err, &handshakeErr) && handshakeErr != nil { + responseBody = append([]byte(nil), handshakeErr.Body...) + } wrappedErr = &openAIWSDialError{ StatusCode: statusCode, ResponseHeaders: cloneHeader(handshakeHeaders), + ResponseBody: responseBody, Err: err, } } diff --git a/backend/internal/service/ops_metrics_collector.go b/backend/internal/service/ops_metrics_collector.go index b654190cfb..d4f5c32929 100644 --- a/backend/internal/service/ops_metrics_collector.go +++ b/backend/internal/service/ops_metrics_collector.go @@ -39,6 +39,10 @@ const ( var opsMetricsCollectorAdvisoryLockID = hashAdvisoryLockID(opsMetricsCollectorLeaderLockKey) +type opsSchedulableAccountLoadRepository interface { + ListSchedulableAccountLoads(ctx context.Context) ([]AccountWithConcurrency, error) +} + type OpsMetricsCollector struct { opsRepo OpsRepository settingRepo SettingRepository @@ -375,31 +379,16 @@ func (c *OpsMetricsCollector) collectConcurrencyQueueDepth(parentCtx context.Con ctx, cancel := context.WithTimeout(parentCtx, 2*time.Second) defer cancel() - accounts, err := c.accountRepo.ListSchedulable(ctx) + accountLoads, err := c.listSchedulableAccountLoads(ctx) if err != nil { return nil } - if len(accounts) == 0 { + if len(accountLoads) == 0 { zero := 0 return &zero } - batch := make([]AccountWithConcurrency, 0, len(accounts)) - for _, acc := range accounts { - if acc.ID <= 0 { - continue - } - batch = append(batch, AccountWithConcurrency{ - ID: acc.ID, - MaxConcurrency: acc.EffectiveLoadFactor(), - }) - } - if len(batch) == 0 { - zero := 0 - return &zero - } - - loadMap, err := c.concurrencyService.GetAccountsLoadBatch(ctx, batch) + loadMap, err := c.concurrencyService.GetAccountsLoadBatch(ctx, accountLoads) if err != nil { return nil } @@ -423,6 +412,28 @@ func (c *OpsMetricsCollector) collectConcurrencyQueueDepth(parentCtx context.Con return &v } +func (c *OpsMetricsCollector) listSchedulableAccountLoads(ctx context.Context) ([]AccountWithConcurrency, error) { + if repo, ok := c.accountRepo.(opsSchedulableAccountLoadRepository); ok { + return repo.ListSchedulableAccountLoads(ctx) + } + + accounts, err := c.accountRepo.ListSchedulable(ctx) + if err != nil { + return nil, err + } + loads := make([]AccountWithConcurrency, 0, len(accounts)) + for _, account := range accounts { + if account.ID <= 0 { + continue + } + loads = append(loads, AccountWithConcurrency{ + ID: account.ID, + MaxConcurrency: account.EffectiveLoadFactor(), + }) + } + return loads, nil +} + type opsCollectedPercentiles struct { p50 *int p90 *int diff --git a/backend/internal/service/ops_metrics_collector_projection_test.go b/backend/internal/service/ops_metrics_collector_projection_test.go new file mode 100644 index 0000000000..e0d1da774d --- /dev/null +++ b/backend/internal/service/ops_metrics_collector_projection_test.go @@ -0,0 +1,128 @@ +package service + +import ( + "context" + "testing" + + "github.com/stretchr/testify/require" +) + +type opsMetricsProjectionRepo struct { + AccountRepository + accounts []Account + accountLoads []AccountWithConcurrency + listCalls int + projectionCalls int +} + +func (r *opsMetricsProjectionRepo) ListSchedulable(context.Context) ([]Account, error) { + r.listCalls++ + return r.accounts, nil +} + +func (r *opsMetricsProjectionRepo) ListSchedulableAccountLoads(context.Context) ([]AccountWithConcurrency, error) { + r.projectionCalls++ + return r.accountLoads, nil +} + +type opsMetricsFallbackRepo struct { + AccountRepository + accounts []Account + listCalls int +} + +func (r *opsMetricsFallbackRepo) ListSchedulable(context.Context) ([]Account, error) { + r.listCalls++ + return r.accounts, nil +} + +type opsMetricsLoadCache struct { + ConcurrencyCache + loads map[int64]*AccountLoadInfo + got []AccountWithConcurrency +} + +func (c *opsMetricsLoadCache) GetAccountsLoadBatch(_ context.Context, accounts []AccountWithConcurrency) (map[int64]*AccountLoadInfo, error) { + c.got = accounts + return c.loads, nil +} + +func TestCollectConcurrencyQueueDepthUsesProjectionAndPreservesFallbackResult(t *testing.T) { + loadFactor := 7 + accounts := []Account{ + {ID: 11, Concurrency: 2, LoadFactor: &loadFactor}, + {ID: 12, Concurrency: 3}, + {ID: 13}, + } + accountLoads := []AccountWithConcurrency{ + {ID: 11, MaxConcurrency: 7}, + {ID: 12, MaxConcurrency: 3}, + {ID: 13, MaxConcurrency: 1}, + } + loads := map[int64]*AccountLoadInfo{ + 11: {AccountID: 11, WaitingCount: 2}, + 12: {AccountID: 12, WaitingCount: 3}, + 13: {AccountID: 13, WaitingCount: 0}, + } + + projectionRepo := &opsMetricsProjectionRepo{accounts: accounts, accountLoads: accountLoads} + projectionCache := &opsMetricsLoadCache{loads: loads} + projectionConcurrency := NewConcurrencyService(projectionCache) + projectionConcurrency.SetAccountLoadBatchCacheTTL(0) + projectionCollector := &OpsMetricsCollector{ + accountRepo: projectionRepo, + concurrencyService: projectionConcurrency, + } + + fallbackRepo := &opsMetricsFallbackRepo{accounts: accounts} + fallbackCache := &opsMetricsLoadCache{loads: loads} + fallbackConcurrency := NewConcurrencyService(fallbackCache) + fallbackConcurrency.SetAccountLoadBatchCacheTTL(0) + fallbackCollector := &OpsMetricsCollector{ + accountRepo: fallbackRepo, + concurrencyService: fallbackConcurrency, + } + + projectionDepth := projectionCollector.collectConcurrencyQueueDepth(context.Background()) + fallbackDepth := fallbackCollector.collectConcurrencyQueueDepth(context.Background()) + + require.NotNil(t, projectionDepth) + require.NotNil(t, fallbackDepth) + require.Equal(t, 5, *projectionDepth) + require.Equal(t, *fallbackDepth, *projectionDepth) + require.Equal(t, 1, projectionRepo.projectionCalls) + require.Zero(t, projectionRepo.listCalls) + require.Equal(t, 1, fallbackRepo.listCalls) + require.Equal(t, accountLoads, projectionCache.got) + require.Equal(t, fallbackCache.got, projectionCache.got) +} + +func BenchmarkOpsMetricsCollectorCollectConcurrencyQueueDepth(b *testing.B) { + const accountCount = 1000 + loadFactor := 8 + accounts := make([]Account, accountCount) + accountLoads := make([]AccountWithConcurrency, accountCount) + for i := range accountCount { + id := int64(i + 1) + accounts[i] = Account{ + ID: id, + Concurrency: 4, + LoadFactor: &loadFactor, + } + accountLoads[i] = AccountWithConcurrency{ID: id, MaxConcurrency: loadFactor} + } + + repo := &opsMetricsProjectionRepo{accounts: accounts, accountLoads: accountLoads} + cache := &opsMetricsLoadCache{loads: map[int64]*AccountLoadInfo{}} + concurrency := NewConcurrencyService(cache) + concurrency.SetAccountLoadBatchCacheTTL(0) + collector := &OpsMetricsCollector{accountRepo: repo, concurrencyService: concurrency} + + b.ReportAllocs() + b.ResetTimer() + for range b.N { + if depth := collector.collectConcurrencyQueueDepth(context.Background()); depth == nil || *depth != 0 { + b.Fatalf("unexpected queue depth: %v", depth) + } + } +} diff --git a/backend/internal/service/ops_models.go b/backend/internal/service/ops_models.go index d95bbadbb8..411ca7595b 100644 --- a/backend/internal/service/ops_models.go +++ b/backend/internal/service/ops_models.go @@ -27,7 +27,7 @@ type OpsErrorLog struct { CreatedAt time.Time `json:"created_at"` // Standardized classification - // - phase: request|auth|routing|upstream|network|internal + // - phase: request|auth|account_auth|routing|upstream|network|internal // - owner: client|provider|platform // - source: client_request|upstream_http|gateway Phase string `json:"phase"` @@ -121,7 +121,7 @@ type OpsErrorLogFilter struct { StatusCodes []int StatusCodesOther bool - Phase string // Special: Phase=="upstream" bypasses status>=400 clause; do not set together with ErrorPhasesAny. + Phase string // Recovered provider rows bypass status>=400 only with the explicit opt-in below. Owner string Source string Resolved *bool @@ -150,14 +150,15 @@ 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 explicitly exempts provider-health phases + // (upstream and account_auth) from the status>=400 guard. Ops provider + // health lists need status<400 recovered rows; request-error endpoints do + // not set this flag and retain client-error semantics. IncludeRecoveredUpstream bool // ErrorPhasesAny / ErrorTypesAny add plain ANY() filters WITHOUT touching the - // special-cased single `Phase` field (only Phase=="upstream" with - // IncludeRecoveredUpstream bypasses the status>=400 clause). + // special-cased single `Phase` field. With IncludeRecoveredUpstream, an ANY + // list containing only upstream/account_auth also bypasses status>=400. // 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. diff --git a/backend/internal/service/ops_service.go b/backend/internal/service/ops_service.go index 61f85ef904..3b27252665 100644 --- a/backend/internal/service/ops_service.go +++ b/backend/internal/service/ops_service.go @@ -222,6 +222,33 @@ func (s *OpsService) prepareErrorLogInput(ctx context.Context, entry *OpsInsertE entry.ErrorType = "api_error" } + // Credential acquisition is a gateway/account-auth stage, not an inference + // HTTP attempt. Enforce that ownership at the persistence boundary so an + // earlier inference attempt cannot leak its status or text into top-level + // auth fields even if a caller supplied stale single-value context. + for i := len(entry.UpstreamErrors) - 1; i >= 0; i-- { + last := entry.UpstreamErrors[i] + if last == nil { + continue + } + if last.Stage == string(GatewayFailureStageAccountAuth) { + entry.ErrorPhase = string(GatewayFailureStageAccountAuth) + entry.ErrorOwner = "provider" + entry.ErrorSource = "gateway" + code := 0 + entry.UpstreamStatusCode = &code + entry.UpstreamErrorMessage = nil + if message := strings.TrimSpace(last.Message); message != "" { + entry.UpstreamErrorMessage = &message + } + entry.UpstreamErrorDetail = nil + if detail := strings.TrimSpace(last.Detail); detail != "" { + entry.UpstreamErrorDetail = &detail + } + } + break + } + // Sanitize + truncate error_body to avoid storing sensitive data. if strings.TrimSpace(entry.ErrorBody) != "" { sanitized, _ := sanitizeErrorBodyForStorage(entry.ErrorBody, opsMaxStoredErrorBodyBytes) @@ -229,7 +256,7 @@ func (s *OpsService) prepareErrorLogInput(ctx context.Context, entry *OpsInsertE } // Sanitize upstream error context if provided by gateway services. - if entry.UpstreamStatusCode != nil && *entry.UpstreamStatusCode <= 0 { + if entry.UpstreamStatusCode != nil && *entry.UpstreamStatusCode <= 0 && entry.ErrorPhase != string(GatewayFailureStageAccountAuth) { entry.UpstreamStatusCode = nil } if entry.UpstreamErrorMessage != nil { diff --git a/backend/internal/service/ops_service_batch_test.go b/backend/internal/service/ops_service_batch_test.go index a9419ad7bf..a027f48737 100644 --- a/backend/internal/service/ops_service_batch_test.go +++ b/backend/internal/service/ops_service_batch_test.go @@ -97,6 +97,54 @@ func TestOpsServiceRecordErrorBatch_FallsBackToSingleInsert(t *testing.T) { require.Equal(t, 2, singleCalls) } +func TestOpsServiceRecordErrorPersistsExplicitAccountAuthStatusZero(t *testing.T) { + t.Parallel() + + var captured *OpsInsertErrorLogInput + repo := &opsRepoMock{ + InsertErrorLogFn: func(_ context.Context, input *OpsInsertErrorLogInput) (int64, error) { + captured = input + return 1, nil + }, + } + svc := NewOpsService(repo, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + staleStatus := 403 + staleMessage := "stale inference message" + staleDetail := "stale inference detail" + + err := svc.RecordError(context.Background(), &OpsInsertErrorLogInput{ + ErrorPhase: "upstream", + ErrorType: "upstream_error", + ErrorOwner: "provider", + ErrorSource: "upstream_http", + UpstreamStatusCode: &staleStatus, + UpstreamErrorMessage: &staleMessage, + UpstreamErrorDetail: &staleDetail, + UpstreamErrors: []*OpsUpstreamErrorEvent{ + {Stage: string(GatewayFailureStageInference), UpstreamStatusCode: 403, Message: staleMessage, Detail: staleDetail}, + { + Stage: string(GatewayFailureStageAccountAuth), Scope: string(GatewayFailureScopeAccount), + Reason: string(GrokCredentialReasonRevoked), Message: "Grok OAuth credentials require account action", + }, + }, + }) + + require.NoError(t, err) + require.NotNil(t, captured) + require.Equal(t, "account_auth", captured.ErrorPhase) + require.Equal(t, "provider", captured.ErrorOwner) + require.Equal(t, "gateway", captured.ErrorSource) + require.NotNil(t, captured.UpstreamStatusCode) + require.Zero(t, *captured.UpstreamStatusCode) + require.NotNil(t, captured.UpstreamErrorMessage) + require.Equal(t, "Grok OAuth credentials require account action", *captured.UpstreamErrorMessage) + require.Nil(t, captured.UpstreamErrorDetail) + require.Nil(t, captured.UpstreamErrors) + require.NotNil(t, captured.UpstreamErrorsJSON) + require.Contains(t, *captured.UpstreamErrorsJSON, `"upstream_status_code":403`) + require.Contains(t, *captured.UpstreamErrorsJSON, `"stage":"account_auth"`) +} + func strPtr(v string) *string { return &v } diff --git a/backend/internal/service/ops_upstream_context.go b/backend/internal/service/ops_upstream_context.go index 6389f1ecb5..2d28f1a168 100644 --- a/backend/internal/service/ops_upstream_context.go +++ b/backend/internal/service/ops_upstream_context.go @@ -188,6 +188,11 @@ type OpsUpstreamErrorEvent struct { // Kind: http_error | request_error | retry_exhausted | failover Kind string `json:"kind,omitempty"` + // Stage/Scope/Reason distinguish credential acquisition from inference + // without overloading upstream_status_code with a synthetic HTTP status. + Stage string `json:"stage,omitempty"` + Scope string `json:"scope,omitempty"` + Reason string `json:"reason,omitempty"` Message string `json:"message,omitempty"` Detail string `json:"detail,omitempty"` @@ -204,6 +209,9 @@ func appendOpsUpstreamError(c *gin.Context, ev OpsUpstreamErrorEvent) { ev.UpstreamRequestID = strings.TrimSpace(ev.UpstreamRequestID) ev.UpstreamResponseBody = strings.TrimSpace(ev.UpstreamResponseBody) ev.Kind = strings.TrimSpace(ev.Kind) + ev.Stage = strings.TrimSpace(ev.Stage) + ev.Scope = strings.TrimSpace(ev.Scope) + ev.Reason = strings.TrimSpace(ev.Reason) ev.UpstreamURL = strings.TrimSpace(ev.UpstreamURL) ev.Message = strings.TrimSpace(ev.Message) ev.Detail = strings.TrimSpace(ev.Detail) diff --git a/backend/internal/service/ops_user_error.go b/backend/internal/service/ops_user_error.go index e3055c2392..22d160bf34 100644 --- a/backend/internal/service/ops_user_error.go +++ b/backend/internal/service/ops_user_error.go @@ -44,7 +44,7 @@ func MapUserErrorCategory(phase, errType string) string { return "auth" case "routing": return "service_unavailable" - case "upstream", "network": + case "account_auth", "upstream", "network": return "upstream" case "internal": return "internal" @@ -73,7 +73,7 @@ func CategoryToFilter(category string) (phases []string, errorTypes []string) { case "service_unavailable": return []string{"routing"}, nil case "upstream": - return []string{"upstream", "network"}, nil + return []string{"account_auth", "upstream", "network"}, nil case "internal": return []string{"internal"}, nil case "rate_limit": diff --git a/backend/internal/service/ops_user_error_test.go b/backend/internal/service/ops_user_error_test.go index 9e0bc164b4..c72cf5dbfb 100644 --- a/backend/internal/service/ops_user_error_test.go +++ b/backend/internal/service/ops_user_error_test.go @@ -17,6 +17,7 @@ func TestMapUserErrorCategory(t *testing.T) { {"request", "subscription_error", "quota"}, {"request", "invalid_request_error", "invalid_request"}, {"routing", "api_error", "service_unavailable"}, + {"account_auth", "upstream_error", "upstream"}, {"upstream", "upstream_error", "upstream"}, {"network", "api_error", "upstream"}, {"internal", "api_error", "internal"}, @@ -43,7 +44,7 @@ func TestCategoryToFilter(t *testing.T) { t.Fatalf("service_unavailable => phases=%v types=%v", phases, types) } phases, types = CategoryToFilter("upstream") - if len(phases) != 2 || phases[0] != "upstream" || phases[1] != "network" || len(types) != 0 { + if len(phases) != 3 || phases[0] != "account_auth" || phases[1] != "upstream" || phases[2] != "network" || len(types) != 0 { t.Fatalf("upstream => phases=%v types=%v", phases, types) } phases, types = CategoryToFilter("internal") diff --git a/backend/internal/service/refresh_policy.go b/backend/internal/service/refresh_policy.go index 7f299be019..6b90077e3b 100644 --- a/backend/internal/service/refresh_policy.go +++ b/backend/internal/service/refresh_policy.go @@ -61,6 +61,14 @@ func AntigravityProviderRefreshPolicy() ProviderRefreshPolicy { } } +func GrokProviderRefreshPolicy() ProviderRefreshPolicy { + return ProviderRefreshPolicy{ + OnRefreshError: ProviderRefreshErrorReturn, + OnLockHeld: ProviderLockHeldWaitForCache, + FailureTTL: 0, + } +} + // BackgroundSkipAction 定义后台刷新服务在“未实际刷新”场景的计数方式。 type BackgroundSkipAction int diff --git a/backend/internal/service/token_refresh_service_test.go b/backend/internal/service/token_refresh_service_test.go index bd6a521576..c4c0b35ceb 100644 --- a/backend/internal/service/token_refresh_service_test.go +++ b/backend/internal/service/token_refresh_service_test.go @@ -26,6 +26,9 @@ type tokenRefreshAccountRepo struct { lastExtraUpdates map[string]any lastAccount *Account updateErr error + setErrorErr error + setTempUnschedErr error + beforeConditionalState func() } func (r *tokenRefreshAccountRepo) Update(ctx context.Context, account *Account) error { @@ -56,7 +59,7 @@ func (r *tokenRefreshAccountRepo) UpdateCredentials(ctx context.Context, id int6 func (r *tokenRefreshAccountRepo) SetError(ctx context.Context, id int64, errorMsg string) error { r.setErrorCalls++ r.lastErrorMessage = errorMsg - return nil + return r.setErrorErr } func (r *tokenRefreshAccountRepo) ClearTempUnschedulable(ctx context.Context, id int64) error { @@ -67,7 +70,66 @@ func (r *tokenRefreshAccountRepo) ClearTempUnschedulable(ctx context.Context, id func (r *tokenRefreshAccountRepo) SetTempUnschedulable(ctx context.Context, id int64, until time.Time, reason string) error { r.setTempUnschedCalls++ r.lastTempUnschedReason = reason - return nil + return r.setTempUnschedErr +} + +func (r *tokenRefreshAccountRepo) SetGrokCredentialErrorIfMatch( + _ context.Context, + id int64, + snapshot GrokCredentialMutationSnapshot, + errorMsg string, +) (bool, error) { + if r.beforeConditionalState != nil { + hook := r.beforeConditionalState + r.beforeConditionalState = nil + hook() + } + account := r.accountsByID[id] + if !grokCredentialSnapshotMatchesAccount(account, snapshot) || + (errorMsg == string(GrokCredentialReasonProxyInvalid) && account.Proxy != nil) { + return false, nil + } + r.setErrorCalls++ + r.lastErrorMessage = errorMsg + if r.setErrorErr != nil { + return false, r.setErrorErr + } + account.Status = StatusError + account.Schedulable = false + account.ErrorMessage = errorMsg + return true, nil +} + +func (r *tokenRefreshAccountRepo) SetGrokCredentialTempUnschedulableIfMatch( + _ context.Context, + id int64, + snapshot GrokCredentialMutationSnapshot, + until time.Time, + reason string, +) (bool, error) { + if r.beforeConditionalState != nil { + hook := r.beforeConditionalState + r.beforeConditionalState = nil + hook() + } + account := r.accountsByID[id] + if !grokCredentialSnapshotMatchesAccount(account, snapshot) { + return false, nil + } + r.setTempUnschedCalls++ + r.lastTempUnschedReason = reason + if r.setTempUnschedErr != nil { + return false, r.setTempUnschedErr + } + value := until + account.TempUnschedulableUntil = &value + return true, nil +} + +func grokCredentialSnapshotMatchesAccount(account *Account, snapshot GrokCredentialMutationSnapshot) bool { + return account != nil && account.IsGrokOAuth() && account.IsSchedulable() && + grokCredentialMutationSnapshot(account).CredentialsJSON == snapshot.CredentialsJSON && + grokCredentialProxyIDsEqual(account.ProxyID, snapshot.ProxyID) } func (r *tokenRefreshAccountRepo) UpdateExtra(ctx context.Context, id int64, updates map[string]any) error { @@ -750,6 +812,7 @@ func TestPathA_Success(t *testing.T) { ID: 100, Platform: PlatformGemini, Type: AccountTypeOAuth, + Status: StatusActive, } repo := &tokenRefreshAccountRepo{} repo.accountsByID = map[int64]*Account{account.ID: account} @@ -771,6 +834,7 @@ func TestPathA_LockHeld(t *testing.T) { ID: 101, Platform: PlatformGemini, Type: AccountTypeOAuth, + Status: StatusActive, } repo := &tokenRefreshAccountRepo{} invalidator := &tokenCacheInvalidatorStub{} @@ -791,6 +855,7 @@ func TestPathA_AlreadyRefreshed(t *testing.T) { ID: 102, Platform: PlatformGemini, Type: AccountTypeOAuth, + Status: StatusActive, } repo := &tokenRefreshAccountRepo{} repo.accountsByID = map[int64]*Account{account.ID: account} @@ -830,6 +895,7 @@ func TestPathA_NonRetryableError(t *testing.T) { ID: 103, Platform: PlatformGemini, Type: AccountTypeOAuth, + Status: StatusActive, } repo := &tokenRefreshAccountRepo{} repo.accountsByID = map[int64]*Account{account.ID: account} @@ -855,6 +921,7 @@ func TestPathA_RetryableErrorExhausted(t *testing.T) { ID: 104, Platform: PlatformGemini, Type: AccountTypeOAuth, + Status: StatusActive, } repo := &tokenRefreshAccountRepo{} repo.accountsByID = map[int64]*Account{account.ID: account} @@ -888,6 +955,7 @@ func TestPathA_DBUpdateFailed(t *testing.T) { ID: 105, Platform: PlatformGemini, Type: AccountTypeOAuth, + Status: StatusActive, } repo := &tokenRefreshAccountRepo{updateErr: errors.New("db connection lost")} repo.accountsByID = map[int64]*Account{account.ID: account} @@ -898,7 +966,7 @@ func TestPathA_DBUpdateFailed(t *testing.T) { err := service.refreshWithRetry(context.Background(), account, refresher, refresher, time.Hour) require.Error(t, err) - require.Contains(t, err.Error(), "DB update failed") + require.ErrorIs(t, err, errOAuthRefreshCredentialPersist) require.Equal(t, 1, repo.updateCalls) // DB 更新被尝试 require.Equal(t, 0, invalidator.calls) // DB 失败时不应触发缓存失效 } diff --git a/backend/internal/service/wire.go b/backend/internal/service/wire.go index d5d9124ec2..a8674382f2 100644 --- a/backend/internal/service/wire.go +++ b/backend/internal/service/wire.go @@ -131,8 +131,67 @@ func ProvideOpenAIQuotaService( proxyRepo ProxyRepository, tokenProvider *OpenAITokenProvider, privacyClientFactory PrivacyClientFactory, + openAIGatewayService *OpenAIGatewayService, ) *OpenAIQuotaService { - return NewOpenAIQuotaService(accountRepo, proxyRepo, tokenProvider, privacyClientFactory) + service := NewOpenAIQuotaService(accountRepo, proxyRepo, tokenProvider, privacyClientFactory) + service.agentIdentityWS = openAIGatewayService + return service +} + +func ProvideAccountUsageService( + accountRepo AccountRepository, + usageLogRepo UsageLogRepository, + usageFetcher ClaudeUsageFetcher, + geminiQuotaService *GeminiQuotaService, + antigravityQuotaFetcher *AntigravityQuotaFetcher, + grokQuotaFetcher *GrokQuotaFetcher, + grokQuotaService *GrokQuotaService, + openAIQuotaService *OpenAIQuotaService, + cache *UsageCache, + identityCache IdentityCache, + tlsFPProfileService *TLSFingerprintProfileService, + openAIGatewayService *OpenAIGatewayService, +) *AccountUsageService { + service := NewAccountUsageService( + accountRepo, + usageLogRepo, + usageFetcher, + geminiQuotaService, + antigravityQuotaFetcher, + grokQuotaFetcher, + grokQuotaService, + openAIQuotaService, + cache, + identityCache, + tlsFPProfileService, + ) + service.agentIdentityWS = openAIGatewayService + return service +} + +func ProvideAccountTestService( + accountRepo AccountRepository, + geminiTokenProvider *GeminiTokenProvider, + claudeTokenProvider *ClaudeTokenProvider, + grokTokenProvider *GrokTokenProvider, + antigravityGatewayService *AntigravityGatewayService, + httpUpstream HTTPUpstream, + cfg *config.Config, + tlsFPProfileService *TLSFingerprintProfileService, + openAIGatewayService *OpenAIGatewayService, +) *AccountTestService { + service := NewAccountTestService( + accountRepo, + geminiTokenProvider, + claudeTokenProvider, + grokTokenProvider, + antigravityGatewayService, + httpUpstream, + cfg, + tlsFPProfileService, + ) + service.agentIdentityWS = openAIGatewayService + return service } func ProvideGrokQuotaService( @@ -186,7 +245,7 @@ func ProvideGrokTokenProvider( p := NewGrokTokenProvider(accountRepo, tokenCache) executor := NewGrokTokenRefresher(grokOAuthService) p.SetRefreshAPI(refreshAPI, executor) - p.SetRefreshPolicy(AntigravityProviderRefreshPolicy()) + p.SetRefreshPolicy(GrokProviderRefreshPolicy()) p.SetTempUnschedCache(tempUnschedCache) return p } @@ -602,8 +661,8 @@ var ProviderSet = wire.NewSet( ProvideClaudeTokenProvider, NewAntigravityGatewayService, ProvideRateLimitService, - NewAccountUsageService, - NewAccountTestService, + ProvideAccountUsageService, + ProvideAccountTestService, ProvideSettingService, NewDataManagementService, ProvideBackupService, diff --git a/frontend/src/components/account/OAuthAuthorizationFlow.vue b/frontend/src/components/account/OAuthAuthorizationFlow.vue index a23ec798d2..395b4ae923 100644 --- a/frontend/src/components/account/OAuthAuthorizationFlow.vue +++ b/frontend/src/components/account/OAuthAuthorizationFlow.vue @@ -276,7 +276,7 @@ - +
([ const errorPhaseOptions = computed(() => [ { value: null, label: t('admin.usage.allTypes') }, { value: 'upstream', label: t('admin.ops.errorLog.typeUpstream') }, + { value: 'account_auth', label: t('admin.ops.errorLog.typeAccountAuth') }, { value: 'request', label: t('admin.ops.errorLog.typeRequest') }, { value: 'auth', label: t('admin.ops.errorLog.typeAuth') }, { value: 'routing', label: t('admin.ops.errorLog.typeRouting') }, diff --git a/frontend/src/components/common/PlatformTypeBadge.vue b/frontend/src/components/common/PlatformTypeBadge.vue index dbcc54fce3..1e664d51bc 100644 --- a/frontend/src/components/common/PlatformTypeBadge.vue +++ b/frontend/src/components/common/PlatformTypeBadge.vue @@ -77,6 +77,7 @@ const { t } = useI18n() interface Props { platform: AccountPlatform type: AccountType + authMode?: string planType?: string privacyMode?: string subscriptionExpiresAt?: string @@ -92,7 +93,15 @@ const platformLabel = computed(() => { return 'Gemini' }) +const normalizedAuthMode = computed(() => + (props.authMode || '').trim().toLowerCase().replace(/[\s_-]+/g, '') +) + const typeLabel = computed(() => { + if (props.platform === 'openai' && props.type === 'oauth') { + if (normalizedAuthMode.value === 'agentidentity') return 'Agent Identity' + if (normalizedAuthMode.value === 'personalaccesstoken') return 'PAT' + } switch (props.type) { case 'oauth': return 'OAuth' diff --git a/frontend/src/components/common/__tests__/PlatformTypeBadge.grok.spec.ts b/frontend/src/components/common/__tests__/PlatformTypeBadge.grok.spec.ts index 753fb1a5f2..0bad3d9551 100644 --- a/frontend/src/components/common/__tests__/PlatformTypeBadge.grok.spec.ts +++ b/frontend/src/components/common/__tests__/PlatformTypeBadge.grok.spec.ts @@ -62,3 +62,24 @@ describe('PlatformTypeBadge Grok plans', () => { expect(wrapper.findAll('path')).toHaveLength(2) }) }) + +describe('PlatformTypeBadge OpenAI authentication modes', () => { + it('distinguishes Agent Identity, PAT, and OAuth accounts', async () => { + const wrapper = mount(PlatformTypeBadge, { + props: { + platform: 'openai', + type: 'oauth', + authMode: 'agentIdentity', + }, + }) + + expect(wrapper.text()).toContain('Agent Identity') + + await wrapper.setProps({ authMode: 'personalAccessToken' }) + expect(wrapper.text()).toContain('PAT') + expect(wrapper.text()).not.toContain('Agent Identity') + + await wrapper.setProps({ authMode: undefined }) + expect(wrapper.text()).toContain('OAuth') + }) +}) diff --git a/frontend/src/components/keys/UseKeyModal.vue b/frontend/src/components/keys/UseKeyModal.vue index e6bce556da..53ca732ebc 100644 --- a/frontend/src/components/keys/UseKeyModal.vue +++ b/frontend/src/components/keys/UseKeyModal.vue @@ -50,6 +50,57 @@
+ +
+
+

+ {{ t('keys.useKeyModal.openai.authModeTitle') }} +

+

+ {{ t('keys.useKeyModal.openai.authModeDescription') }} +

+
+
+ + +
+
+