From e86063155f456491c4951fd9dde06ea879a87a85 Mon Sep 17 00:00:00 2001 From: Heatherm Huang Date: Fri, 17 Jul 2026 19:15:16 +0800 Subject: [PATCH] fix(grok): gate OAuth media on paid eligibility --- README.md | 4 +- backend/cmd/server/wire_gen.go | 2 +- .../internal/handler/admin/account_handler.go | 2 +- .../handler/admin/grok_import_probe.go | 12 ++-- .../handler/admin/grok_import_probe_test.go | 4 +- .../handler/admin/grok_oauth_handler.go | 2 +- backend/internal/handler/grok_media.go | 37 ++++++++++- backend/internal/handler/grok_media_test.go | 66 +++++++++++++++++++ .../handler/openai_gateway_handler.go | 29 ++++---- backend/internal/handler/wire.go | 2 + backend/internal/service/account.go | 22 +++++-- .../account_grok_media_eligibility_test.go | 42 +++++++++++- backend/internal/service/grok_media.go | 19 +++--- .../internal/service/grok_quota_service.go | 17 +++++ .../service/grok_quota_service_test.go | 59 +++++++++++++++++ .../service/openai_gateway_grok_test.go | 45 +++++++++++-- 16 files changed, 317 insertions(+), 47 deletions(-) diff --git a/README.md b/README.md index dd415c307..9695d7d5a 100644 --- a/README.md +++ b/README.md @@ -788,9 +788,9 @@ xAI quota is passive. Sub2API does not invent subscription quota values; it reco `401` responses temporarily remove accounts with invalid credentials from scheduling. `403` responses are treated as access or entitlement failures instead of token-refresh loops. `429` responses use `Retry-After` or a short cooldown to temporarily remove the account from scheduling. -New Grok image and video generation requests use a media-specific eligibility check. An OAuth account is excluded from new media generation when its recorded weekly or monthly billing probe returns `403`; chat requests and video status lookups are not affected by this media-only quarantine. If no eligible account remains, the media endpoint returns HTTP `503` with error type `grok_media_no_eligible_account` instead of forwarding the request to a known-ineligible account. +New Grok image and video generation requests use a media-specific eligibility check. API-key accounts remain eligible. OAuth accounts require positive paid-entitlement evidence from the xAI billing probe; Free, forbidden, missing, malformed, and inconclusive billing observations are excluded from new media generation. Unobserved OAuth accounts are probed before the first media request is forwarded, and imports run the billing-first quota probe proactively. Chat requests and video status lookups are not affected by this media-only quarantine. If no eligible account remains, the media endpoint returns HTTP `503` with error type `grok_media_no_eligible_account`. -Administrators can override automatic media eligibility through the account create/update API by setting `extra.grok_media_eligible` to `false` (exclude) or `true` (force eligible). On update, set it to `null` to remove the override and return to automatic probe-based behavior; omitting the field preserves the current override. A missing billing observation does not block legacy routing, and a weekly allowance period by itself is not treated as evidence that the account is ineligible. +Administrators can override automatic media eligibility through the account create/update API by setting `extra.grok_media_eligible` to `false` (exclude) or `true` (force eligible). On update, set it to `null` to remove the override and return to automatic probe-based behavior; omitting the field preserves the current override. A weekly allowance period alone is not treated as a paid tier signal. Successful image responses must contain at least one actual image output; empty HTTP `200` responses trigger account failover instead of being counted and returned as successful generations. --- diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index df6659d68..17ae478dd 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -270,7 +270,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { legacyEngine := securityaudit.NewLegacyModerationAdapter(contentModerationService) coordinator := securityaudit.NewCoordinator(legacyEngine, promptService) gatewayHandler := handler.ProvideGatewayHandler(gatewayService, openAIGatewayService, geminiMessagesCompatService, antigravityGatewayService, userService, concurrencyService, billingCacheService, usageService, apiKeyService, usageRecordWorkerPool, errorPassthroughService, contentModerationService, userMessageQueueService, configConfig, settingService, coordinator) - openAIGatewayHandler := handler.ProvideOpenAIGatewayHandler(openAIGatewayService, concurrencyService, billingCacheService, apiKeyService, usageRecordWorkerPool, errorPassthroughService, contentModerationService, opsService, configConfig, coordinator) + openAIGatewayHandler := handler.ProvideOpenAIGatewayHandler(openAIGatewayService, concurrencyService, billingCacheService, apiKeyService, usageRecordWorkerPool, errorPassthroughService, contentModerationService, opsService, grokQuotaService, configConfig, coordinator) handlerSettingHandler := handler.ProvideSettingHandler(settingService, buildInfo, notificationEmailService) totpHandler := handler.NewTotpHandler(totpService) handlerPaymentHandler := handler.NewPaymentHandler(paymentService, paymentConfigService) diff --git a/backend/internal/handler/admin/account_handler.go b/backend/internal/handler/admin/account_handler.go index d9b742003..46cc43ad5 100644 --- a/backend/internal/handler/admin/account_handler.go +++ b/backend/internal/handler/admin/account_handler.go @@ -61,7 +61,7 @@ type AccountHandler struct { sessionLimitCache service.SessionLimitCache rpmCache service.RPMCache tokenCacheInvalidator service.TokenCacheInvalidator - grokImportProber grokUsageProber + grokImportProber grokImportProber upstreamBillingProbe *service.UpstreamBillingProbeService } diff --git a/backend/internal/handler/admin/grok_import_probe.go b/backend/internal/handler/admin/grok_import_probe.go index b42957ce6..9ef4ce8cd 100644 --- a/backend/internal/handler/admin/grok_import_probe.go +++ b/backend/internal/handler/admin/grok_import_probe.go @@ -15,12 +15,12 @@ const ( grokImportProbeTimeout = 25 * time.Second ) -type grokUsageProber interface { - ProbeUsage(ctx context.Context, accountID int64) (*service.GrokQuotaProbeResult, error) +type grokImportProber interface { + QueryQuota(ctx context.Context, accountID int64) (*service.GrokQuotaProbeResult, error) } type grokImportProbeTask struct { - prober grokUsageProber + prober grokImportProber accountID int64 } @@ -51,7 +51,7 @@ func newGrokImportProbeScheduler(concurrency int, timeout time.Duration) *grokIm } } -func (s *grokImportProbeScheduler) schedule(prober grokUsageProber, account *service.Account) { +func (s *grokImportProbeScheduler) schedule(prober grokImportProber, account *service.Account) { if s == nil || prober == nil || account == nil || account.ID <= 0 { return } @@ -97,7 +97,7 @@ func (s *grokImportProbeScheduler) nextTask() (grokImportProbeTask, bool) { return task, true } -func (s *grokImportProbeScheduler) run(prober grokUsageProber, accountID int64) { +func (s *grokImportProbeScheduler) run(prober grokImportProber, accountID int64) { defer func() { if recovered := recover(); recovered != nil { slog.Error( @@ -112,7 +112,7 @@ func (s *grokImportProbeScheduler) run(prober grokUsageProber, accountID int64) // while this timeout only bounds the actual upstream probe execution. ctx, cancel := context.WithTimeout(context.Background(), s.timeout) defer cancel() - result, err := prober.ProbeUsage(ctx, accountID) + result, err := prober.QueryQuota(ctx, accountID) if err != nil { slog.Warn( "grok_import_active_probe_failed", diff --git a/backend/internal/handler/admin/grok_import_probe_test.go b/backend/internal/handler/admin/grok_import_probe_test.go index 3b8fc0ca6..255f04d0d 100644 --- a/backend/internal/handler/admin/grok_import_probe_test.go +++ b/backend/internal/handler/admin/grok_import_probe_test.go @@ -36,7 +36,7 @@ func newGrokImportProbeStub(buffer int) *grokImportProbeStub { } } -func (s *grokImportProbeStub) ProbeUsage(ctx context.Context, accountID int64) (*service.GrokQuotaProbeResult, error) { +func (s *grokImportProbeStub) QueryQuota(ctx context.Context, accountID int64) (*service.GrokQuotaProbeResult, error) { _, deadlineSeen := ctx.Deadline() s.mu.Lock() s.calls[accountID]++ @@ -69,7 +69,7 @@ func (s *grokImportProbeStub) ProbeUsage(ctx context.Context, accountID int64) ( return nil, failure } return &service.GrokQuotaProbeResult{ - Source: "active_probe", + Source: "hybrid_probe", Model: "grok-4.5", StatusCode: 200, ResetSupported: false, diff --git a/backend/internal/handler/admin/grok_oauth_handler.go b/backend/internal/handler/admin/grok_oauth_handler.go index ac05fd8a6..1f679b956 100644 --- a/backend/internal/handler/admin/grok_oauth_handler.go +++ b/backend/internal/handler/admin/grok_oauth_handler.go @@ -23,7 +23,7 @@ type GrokOAuthHandler struct { grokOAuthService *service.GrokOAuthService adminService service.AdminService quotaService *service.GrokQuotaService - importProber grokUsageProber + importProber grokImportProber reconciler service.GrokOAuthReconciler } diff --git a/backend/internal/handler/grok_media.go b/backend/internal/handler/grok_media.go index 6ae712945..9b96a9a12 100644 --- a/backend/internal/handler/grok_media.go +++ b/backend/internal/handler/grok_media.go @@ -167,6 +167,7 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service. sameAccountRetryCount := make(map[int64]int) var lastFailoverErr *service.UpstreamFailoverError var oauth429FailoverState service.OpenAIOAuth429FailoverState + mediaEligibilityRejected := false switchCount := 0 maxAccountSwitches := h.maxAccountSwitches if maxAccountSwitches <= 0 { @@ -202,7 +203,8 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service. zap.Error(err), zap.Int("excluded_account_count", len(failedAccountIDs)), ) - if endpoint.IsGenerationRequest() && len(failedAccountIDs) == 0 && errors.Is(err, service.ErrNoAvailableAccounts) { + if endpoint.IsGenerationRequest() && errors.Is(err, service.ErrNoAvailableAccounts) && + (len(failedAccountIDs) == 0 || (mediaEligibilityRejected && lastFailoverErr == nil)) { markOpsRoutingCapacityLimitedIfNoAvailable(c, err) h.errorResponse(c, http.StatusServiceUnavailable, "grok_media_no_eligible_account", "No eligible Grok media accounts") return @@ -246,6 +248,25 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service. ) account := selection.Account + if endpoint.IsGenerationRequest() { + eligible, eligibilityReason, eligibilityErr := h.ensureGrokMediaAccountEligibility(requestCtx, account) + if !eligible { + mediaEligibilityRejected = true + failedAccountIDs[account.ID] = struct{}{} + reqLog.Warn("grok_media.account_eligibility_rejected", + zap.Int64("account_id", account.ID), + zap.String("reason", eligibilityReason), + zap.Bool("probe_failed", eligibilityErr != nil), + ) + if switchCount >= maxAccountSwitches { + markOpsRoutingCapacityLimited(c) + h.errorResponse(c, http.StatusServiceUnavailable, "grok_media_no_eligible_account", "No eligible Grok media accounts") + return + } + switchCount++ + continue + } + } sessionHash = ensureOpenAIPoolModeSessionHash(sessionHash, account) setOpsSelectedAccount(c, account.ID, account.Platform) @@ -365,6 +386,20 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service. } } +func (h *OpenAIGatewayHandler) ensureGrokMediaAccountEligibility(ctx context.Context, account *service.Account) (bool, string, error) { + if account == nil { + return false, "missing_account", errors.New("grok media account is required") + } + eligible, reason := account.GrokMediaGenerationEligibility() + if eligible || reason != "billing_unobserved" { + return eligible, reason, nil + } + if h == nil || h.grokMediaEligibilityProber == nil { + return false, "billing_probe_unavailable", errors.New("grok media eligibility probe is not configured") + } + return h.grokMediaEligibilityProber.ProbeMediaEligibility(ctx, account.ID) +} + func grokMediaRequiredCapability(endpoint service.GrokMediaEndpoint) service.OpenAIEndpointCapability { if endpoint.IsGenerationRequest() { return service.OpenAIEndpointCapabilityGrokMediaGeneration diff --git a/backend/internal/handler/grok_media_test.go b/backend/internal/handler/grok_media_test.go index 3c58742c7..53348958b 100644 --- a/backend/internal/handler/grok_media_test.go +++ b/backend/internal/handler/grok_media_test.go @@ -1,12 +1,26 @@ package handler import ( + "context" + "errors" "testing" "github.com/Wei-Shaw/sub2api/internal/service" "github.com/stretchr/testify/require" ) +type grokMediaEligibilityProberStub struct { + eligible bool + reason string + err error + calls int +} + +func (s *grokMediaEligibilityProberStub) ProbeMediaEligibility(context.Context, int64) (bool, string, error) { + s.calls++ + return s.eligible, s.reason, s.err +} + func TestShouldRecordGrokMediaUsage(t *testing.T) { tests := []struct { name string @@ -73,3 +87,55 @@ func TestGrokMediaRequiredCapability(t *testing.T) { }) } } + +func TestEnsureGrokMediaAccountEligibility(t *testing.T) { + t.Run("non oauth account does not probe", func(t *testing.T) { + prober := &grokMediaEligibilityProberStub{} + h := &OpenAIGatewayHandler{grokMediaEligibilityProber: prober} + account := &service.Account{Platform: service.PlatformGrok, Type: service.AccountTypeAPIKey} + + eligible, reason, err := h.ensureGrokMediaAccountEligibility(context.Background(), account) + + require.NoError(t, err) + require.True(t, eligible) + require.Equal(t, "non_oauth", reason) + require.Zero(t, prober.calls) + }) + + t.Run("unobserved oauth is probed before forwarding", func(t *testing.T) { + prober := &grokMediaEligibilityProberStub{eligible: true, reason: "eligible"} + h := &OpenAIGatewayHandler{grokMediaEligibilityProber: prober} + account := &service.Account{ID: 7, Platform: service.PlatformGrok, Type: service.AccountTypeOAuth} + + eligible, reason, err := h.ensureGrokMediaAccountEligibility(context.Background(), account) + + require.NoError(t, err) + require.True(t, eligible) + require.Equal(t, "eligible", reason) + require.Equal(t, 1, prober.calls) + }) + + t.Run("missing prober fails closed", func(t *testing.T) { + h := &OpenAIGatewayHandler{} + account := &service.Account{ID: 8, Platform: service.PlatformGrok, Type: service.AccountTypeOAuth} + + eligible, reason, err := h.ensureGrokMediaAccountEligibility(context.Background(), account) + + require.Error(t, err) + require.False(t, eligible) + require.Equal(t, "billing_probe_unavailable", reason) + }) + + t.Run("probe failure fails closed", func(t *testing.T) { + probeErr := errors.New("probe failed") + prober := &grokMediaEligibilityProberStub{reason: "billing_unobserved", err: probeErr} + h := &OpenAIGatewayHandler{grokMediaEligibilityProber: prober} + account := &service.Account{ID: 9, Platform: service.PlatformGrok, Type: service.AccountTypeOAuth} + + eligible, reason, err := h.ensureGrokMediaAccountEligibility(context.Background(), account) + + require.ErrorIs(t, err, probeErr) + require.False(t, eligible) + require.Equal(t, "billing_unobserved", reason) + }) +} diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index ec4e06dd8..3dcc170cc 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -28,18 +28,23 @@ import ( // OpenAIGatewayHandler handles OpenAI API gateway requests type OpenAIGatewayHandler struct { - gatewayService *service.OpenAIGatewayService - billingCacheService *service.BillingCacheService - apiKeyService *service.APIKeyService - usageRecordWorkerPool *service.UsageRecordWorkerPool - errorPassthroughService *service.ErrorPassthroughService - contentModerationService *service.ContentModerationService - securityAuditCoordinator *securityaudit.Coordinator - opsService *service.OpsService - concurrencyHelper *ConcurrencyHelper - imageLimiter *imageConcurrencyLimiter - maxAccountSwitches int - cfg *config.Config + gatewayService *service.OpenAIGatewayService + billingCacheService *service.BillingCacheService + apiKeyService *service.APIKeyService + usageRecordWorkerPool *service.UsageRecordWorkerPool + errorPassthroughService *service.ErrorPassthroughService + contentModerationService *service.ContentModerationService + securityAuditCoordinator *securityaudit.Coordinator + grokMediaEligibilityProber grokMediaEligibilityProber + opsService *service.OpsService + concurrencyHelper *ConcurrencyHelper + imageLimiter *imageConcurrencyLimiter + maxAccountSwitches int + cfg *config.Config +} + +type grokMediaEligibilityProber interface { + ProbeMediaEligibility(ctx context.Context, accountID int64) (bool, string, error) } const maxOpenAIFirstOutputTimeoutSwitches = 1 diff --git a/backend/internal/handler/wire.go b/backend/internal/handler/wire.go index 57bf64d02..888df38b1 100644 --- a/backend/internal/handler/wire.go +++ b/backend/internal/handler/wire.go @@ -120,12 +120,14 @@ func ProvideOpenAIGatewayHandler( errorPassthroughService *service.ErrorPassthroughService, contentModerationService *service.ContentModerationService, opsService *service.OpsService, + grokQuotaService *service.GrokQuotaService, cfg *config.Config, coordinator *securityaudit.Coordinator, ) *OpenAIGatewayHandler { h := NewOpenAIGatewayHandler(gatewayService, concurrencyService, billingCacheService, apiKeyService, usageRecordWorkerPool, errorPassthroughService, contentModerationService, opsService, cfg) h.securityAuditCoordinator = coordinator + h.grokMediaEligibilityProber = grokQuotaService return h } diff --git a/backend/internal/service/account.go b/backend/internal/service/account.go index 724925631..5b3707d2a 100644 --- a/backend/internal/service/account.go +++ b/backend/internal/service/account.go @@ -1423,8 +1423,12 @@ func (a *Account) SupportsOpenAIEndpointCapability(capability OpenAIEndpointCapa case OpenAIEndpointCapabilityChatCompletions: return true case OpenAIEndpointCapabilityGrokMediaGeneration: - eligible, _ := a.GrokMediaGenerationEligibility() - return eligible + eligible, reason := a.GrokMediaGenerationEligibility() + // Unobserved OAuth accounts remain scheduler candidates only so the + // request path can run the billing probe before forwarding. The + // forwarding gate itself fails closed if that probe is unavailable or + // cannot produce positive paid-entitlement evidence. + return eligible || reason == "billing_unobserved" default: return false } @@ -1469,9 +1473,9 @@ func (a *Account) SupportsOpenAIEndpointCapability(capability OpenAIEndpointCapa } // GrokMediaGenerationEligibility reports whether a Grok account may receive -// new image/video generation requests. Missing observations preserve legacy -// routing; operators can fail closed for known-bad accounts with the explicit -// override. A successful override takes precedence over stale probe data. +// new image/video generation requests. OAuth media fails closed unless billing +// observations provide positive paid-entitlement evidence. An explicit +// operator override takes precedence over probe data. func (a *Account) GrokMediaGenerationEligibility() (bool, string) { if a == nil || !a.IsGrok() { return false, "not_grok" @@ -1488,11 +1492,17 @@ func (a *Account) GrokMediaGenerationEligibility() (bool, string) { billing, err := grokBillingSnapshotFromExtra(a.Extra) if err != nil || billing == nil { - return true, "billing_unobserved" + return false, "billing_unobserved" } if billing.StatusCode == 403 || billing.WeeklyStatusCode == 403 || billing.MonthlyStatusCode == 403 { return false, "billing_forbidden" } + if isKnownGrokFreeAccount(a) { + return false, "billing_free_tier" + } + if !grokBillingHasAuthoritativeQuota(billing) { + return false, "billing_inconclusive" + } return true, "eligible" } diff --git a/backend/internal/service/account_grok_media_eligibility_test.go b/backend/internal/service/account_grok_media_eligibility_test.go index be7b4759a..a79f9c02e 100644 --- a/backend/internal/service/account_grok_media_eligibility_test.go +++ b/backend/internal/service/account_grok_media_eligibility_test.go @@ -13,6 +13,7 @@ import ( ) func TestGrokMediaGenerationEligibility(t *testing.T) { + weeklyUsagePercent := 12.5 forbiddenBilling := &xai.BillingSummary{ StatusCode: http.StatusForbidden, WeeklyStatusCode: http.StatusForbidden, @@ -20,9 +21,24 @@ func TestGrokMediaGenerationEligibility(t *testing.T) { } weeklyAllowance := &xai.BillingSummary{ PeriodType: "weekly", + UsagePercent: &weeklyUsagePercent, StatusCode: http.StatusOK, WeeklyStatusCode: http.StatusOK, } + freeBilling := &xai.BillingSummary{ + PeriodType: "monthly", + StatusCode: http.StatusOK, + WeeklyStatusCode: http.StatusOK, + MonthlyStatusCode: http.StatusOK, + MonthlyUpdatedAt: "2026-07-17T00:00:00Z", + } + inconclusiveBilling := &xai.BillingSummary{ + StatusCode: http.StatusOK, + WeeklyStatusCode: http.StatusOK, + MonthlyStatusCode: http.StatusBadGateway, + Partial: true, + FailedWindows: []string{"monthly"}, + } weeklyForbidden := &xai.BillingSummary{ StatusCode: http.StatusOK, WeeklyStatusCode: http.StatusForbidden, @@ -43,12 +59,14 @@ func TestGrokMediaGenerationEligibility(t *testing.T) { {name: "nil account", account: nil, want: false, wantReason: "not_grok"}, {name: "non grok account", account: &Account{Platform: PlatformOpenAI}, want: false, wantReason: "not_grok"}, {name: "non oauth grok account stays eligible", account: &Account{Platform: PlatformGrok, Type: AccountTypeAPIKey}, want: true, wantReason: "non_oauth"}, - {name: "unobserved oauth preserves legacy routing", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth}, want: true, wantReason: "billing_unobserved"}, - {name: "weekly allowance is not treated as weekly subscription", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: weeklyAllowance}}, want: true, wantReason: "eligible"}, + {name: "unobserved oauth fails closed", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth}, want: false, wantReason: "billing_unobserved"}, + {name: "weekly paid usage is eligible without inferring from period type", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: weeklyAllowance}}, want: true, wantReason: "eligible"}, + {name: "observed free account is rejected", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: freeBilling}}, want: false, wantReason: "billing_free_tier"}, + {name: "inconclusive billing fails closed", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: inconclusiveBilling}}, want: false, wantReason: "billing_inconclusive"}, {name: "billing forbidden is rejected", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: forbiddenBilling}}, want: false, wantReason: "billing_forbidden"}, {name: "weekly billing forbidden is rejected after partial success", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: weeklyForbidden}}, want: false, wantReason: "billing_forbidden"}, {name: "monthly billing forbidden is rejected after partial success", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: monthlyForbidden}}, want: false, wantReason: "billing_forbidden"}, - {name: "malformed billing observation preserves legacy routing", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: make(chan int)}}, want: true, wantReason: "billing_unobserved"}, + {name: "malformed billing observation fails closed", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: make(chan int)}}, want: false, wantReason: "billing_unobserved"}, {name: "malformed override falls back to observations", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{GrokMediaEligibleExtraKey: "false", grokBillingExtraKey: weeklyAllowance}}, want: true, wantReason: "eligible"}, {name: "explicit disable wins", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{GrokMediaEligibleExtraKey: false}}, want: false, wantReason: "override_disabled"}, {name: "explicit enable wins over forbidden probe", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{GrokMediaEligibleExtraKey: true, grokBillingExtraKey: forbiddenBilling}}, want: true, wantReason: "override_enabled"}, @@ -63,6 +81,24 @@ func TestGrokMediaGenerationEligibility(t *testing.T) { } } +func TestGrokMediaCapabilityKeepsOnlyUnobservedOAuthAsProbeCandidate(t *testing.T) { + unobserved := &Account{Platform: PlatformGrok, Type: AccountTypeOAuth} + eligible, reason := unobserved.GrokMediaGenerationEligibility() + require.False(t, eligible) + require.Equal(t, "billing_unobserved", reason) + require.True(t, unobserved.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityGrokMediaGeneration)) + + inconclusive := &Account{ + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Extra: map[string]any{grokBillingExtraKey: &xai.BillingSummary{ + StatusCode: http.StatusOK, + Partial: true, + }}, + } + require.False(t, inconclusive.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityGrokMediaGeneration)) +} + func TestGrokMediaCapabilityFiltersOnlyGeneration(t *testing.T) { account := &Account{ ID: 1, diff --git a/backend/internal/service/grok_media.go b/backend/internal/service/grok_media.go index c75bd271f..abfe0df26 100644 --- a/backend/internal/service/grok_media.go +++ b/backend/internal/service/grok_media.go @@ -364,6 +364,16 @@ func (s *OpenAIGatewayService) ForwardGrokMedia( if err != nil { return nil, err } + if endpoint == GrokMediaEndpointImagesGenerations || endpoint == GrokMediaEndpointImagesEdits { + if countOpenAIResponseImageOutputsFromJSONBytes(respBody) <= 0 { + setOpsUpstreamError(c, http.StatusBadGateway, "xAI upstream returned no image output", truncateString(string(respBody), 512)) + return nil, &UpstreamFailoverError{ + StatusCode: http.StatusBadGateway, + ResponseBody: respBody, + ResponseHeaders: resp.Header.Clone(), + } + } + } writeGrokMediaResponse(c, resp, respBody, s.responseHeaderFilter) usage := grokMediaUsageFromResponse(endpoint, requestInfo, respBody) return &OpenAIForwardResult{ @@ -584,14 +594,7 @@ func grokMediaUsageFromResponse(endpoint GrokMediaEndpoint, requestInfo GrokMedi meta := grokMediaUsageMetadata{Usage: usage} switch endpoint { case GrokMediaEndpointImagesGenerations, GrokMediaEndpointImagesEdits: - imageCount := countOpenAIResponseImageOutputsFromJSONBytes(responseBody) - if imageCount <= 0 { - imageCount = requestInfo.N - } - if imageCount <= 0 { - imageCount = 1 - } - meta.ImageCount = imageCount + meta.ImageCount = countOpenAIResponseImageOutputsFromJSONBytes(responseBody) meta.ImageSize = requestInfo.SizeTier meta.ImageInputSize = requestInfo.Size meta.ImageOutputSizes = collectOpenAIResponseImageOutputSizesFromJSONBytes(responseBody) diff --git a/backend/internal/service/grok_quota_service.go b/backend/internal/service/grok_quota_service.go index 45e7be83a..db11fe6cc 100644 --- a/backend/internal/service/grok_quota_service.go +++ b/backend/internal/service/grok_quota_service.go @@ -221,6 +221,23 @@ func (s *GrokQuotaService) ProbeBilling(ctx context.Context, accountID int64) (* }) } +// ProbeMediaEligibility refreshes billing state and evaluates the persisted +// account snapshot used by media scheduling. Probe failures remain fail-closed; +// deterministic persisted states such as forbidden or Free are returned as +// normal ineligibility decisions rather than transport errors. +func (s *GrokQuotaService) ProbeMediaEligibility(ctx context.Context, accountID int64) (bool, string, error) { + _, probeErr := s.ProbeBilling(ctx, accountID) + account, err := s.loadGrokOAuthAccount(ctx, accountID) + if err != nil { + return false, "billing_probe_failed", err + } + eligible, reason := account.GrokMediaGenerationEligibility() + if reason == "billing_unobserved" && probeErr != nil { + return false, reason, probeErr + } + return eligible, reason, nil +} + func (s *GrokQuotaService) probeBilling(ctx context.Context, accountID int64) (*GrokQuotaProbeResult, error) { account, token, proxyURL, err := s.prepareProbe(ctx, accountID) if err != nil { diff --git a/backend/internal/service/grok_quota_service_test.go b/backend/internal/service/grok_quota_service_test.go index 59b19ef27..728277e62 100644 --- a/backend/internal/service/grok_quota_service_test.go +++ b/backend/internal/service/grok_quota_service_test.go @@ -44,6 +44,18 @@ func (r *grokQuotaAccountRepo) UpdateExtra(_ context.Context, id int64, updates r.updates = make(map[int64]map[string]any) } r.updates[id] = updates + if r.mockAccountRepoForPlatform != nil { + account := r.accountsByID[id] + if account == nil { + return nil + } + if account.Extra == nil { + account.Extra = make(map[string]any) + } + for key, value := range updates { + account.Extra[key] = value + } + } return nil } @@ -824,6 +836,53 @@ func TestGrokQuotaServicePartialBilling403PersistsMediaEligibilitySignal(t *test require.Equal(t, "billing_forbidden", reason) } +func TestGrokQuotaServiceProbeMediaEligibility(t *testing.T) { + t.Run("positive paid evidence enables media", func(t *testing.T) { + usagePercent := 10.0 + monthlyLimit := 15_000.0 + account := healthyGrokQuotaOAuthAccount(60) + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + upstream := &grokHybridUpstream{weeklyUsagePercent: &usagePercent, monthlyLimitCents: &monthlyLimit} + svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream, nil) + + eligible, reason, err := svc.ProbeMediaEligibility(context.Background(), account.ID) + + require.NoError(t, err) + require.True(t, eligible) + require.Equal(t, "eligible", reason) + }) + + t.Run("successful empty billing identifies free account", func(t *testing.T) { + account := healthyGrokQuotaOAuthAccount(61) + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), &grokHybridUpstream{}, nil) + + eligible, reason, err := svc.ProbeMediaEligibility(context.Background(), account.ID) + + require.NoError(t, err) + require.False(t, eligible) + require.Equal(t, "billing_free_tier", reason) + }) + + t.Run("forbidden billing is deterministic ineligibility", func(t *testing.T) { + account := healthyGrokQuotaOAuthAccount(62) + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), &grokHybridUpstream{billingStatus: http.StatusForbidden}, nil) + + eligible, reason, err := svc.ProbeMediaEligibility(context.Background(), account.ID) + + require.NoError(t, err) + require.False(t, eligible) + require.Equal(t, "billing_forbidden", reason) + }) +} + func TestPreferBillingObservationStatus(t *testing.T) { t.Parallel() diff --git a/backend/internal/service/openai_gateway_grok_test.go b/backend/internal/service/openai_gateway_grok_test.go index 67ad71214..c238e69d3 100644 --- a/backend/internal/service/openai_gateway_grok_test.go +++ b/backend/internal/service/openai_gateway_grok_test.go @@ -551,7 +551,7 @@ func TestForwardGrokMediaImagesGenerationNormalizesImagineAlias(t *testing.T) { "Content-Type": []string{"application/json"}, "Xai-Request-Id": []string{"xai-image-req"}, }, - Body: io.NopCloser(strings.NewReader(`{"data":[]}`)), + Body: io.NopCloser(strings.NewReader(`{"data":[{"url":"https://images.test/cat.png"}]}`)), }} svc := &OpenAIGatewayService{httpUpstream: upstream} @@ -565,7 +565,7 @@ func TestForwardGrokMediaImagesGenerationNormalizesImagineAlias(t *testing.T) { 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()) + require.JSONEq(t, `{"data":[{"url":"https://images.test/cat.png"}]}`, recorder.Body.String()) require.Equal(t, "xai-image-req", result.RequestID) require.Equal(t, "grok-imagine-image-quality", result.Model) require.Equal(t, "grok-imagine-image-quality", result.BillingModel) @@ -573,6 +573,43 @@ func TestForwardGrokMediaImagesGenerationNormalizesImagineAlias(t *testing.T) { require.Equal(t, ImageBillingSize2K, result.ImageSize) } +func TestForwardGrokMediaImagesGenerationRejectsEmptySuccessfulResponse(t *testing.T) { + t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true") + gin.SetMode(gin.TestMode) + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + body := []byte(`{"model":"grok-imagine-image","prompt":"draw a cat"}`) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/images/generations", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + + account := &Account{ + ID: 66, + Name: "grok", + Platform: PlatformGrok, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "api-key", + "base_url": "https://xai.test/v1", + }, + } + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"data":[]}`)), + }} + svc := &OpenAIGatewayService{httpUpstream: upstream} + + result, err := svc.ForwardGrokMedia(context.Background(), c, account, GrokMediaEndpointImagesGenerations, "", body, "application/json") + require.Nil(t, result) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, http.StatusBadGateway, failoverErr.StatusCode) + require.JSONEq(t, `{"data":[]}`, string(failoverErr.ResponseBody)) + require.Empty(t, recorder.Body.String()) +} + func TestForwardGrokMediaImagesGenerationStripsUnsupportedSize(t *testing.T) { t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true") gin.SetMode(gin.TestMode) @@ -599,7 +636,7 @@ func TestForwardGrokMediaImagesGenerationStripsUnsupportedSize(t *testing.T) { Header: http.Header{ "Content-Type": []string{"application/json"}, }, - Body: io.NopCloser(strings.NewReader(`{"data":[]}`)), + Body: io.NopCloser(strings.NewReader(`{"data":[{"url":"https://images.test/cat.png"}]}`)), }} svc := &OpenAIGatewayService{httpUpstream: upstream} @@ -648,7 +685,7 @@ func TestForwardGrokMediaImagesEditMultipartConvertsToJSON(t *testing.T) { Header: http.Header{ "Content-Type": []string{"application/json"}, }, - Body: io.NopCloser(strings.NewReader(`{"data":[]}`)), + Body: io.NopCloser(strings.NewReader(`{"data":[{"url":"https://images.test/edited.png"}]}`)), }} svc := &OpenAIGatewayService{httpUpstream: upstream}