fix(grok): gate OAuth media on paid eligibility

This commit is contained in:
Heatherm Huang
2026-07-17 19:15:16 +08:00
parent 57914967cb
commit e86063155f
16 changed files with 317 additions and 47 deletions
+2 -2
View File
@@ -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.
---
+1 -1
View File
@@ -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)
@@ -61,7 +61,7 @@ type AccountHandler struct {
sessionLimitCache service.SessionLimitCache
rpmCache service.RPMCache
tokenCacheInvalidator service.TokenCacheInvalidator
grokImportProber grokUsageProber
grokImportProber grokImportProber
upstreamBillingProbe *service.UpstreamBillingProbeService
}
@@ -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",
@@ -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,
@@ -23,7 +23,7 @@ type GrokOAuthHandler struct {
grokOAuthService *service.GrokOAuthService
adminService service.AdminService
quotaService *service.GrokQuotaService
importProber grokUsageProber
importProber grokImportProber
reconciler service.GrokOAuthReconciler
}
+36 -1
View File
@@ -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
@@ -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)
})
}
@@ -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
+2
View File
@@ -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
}
+16 -6
View File
@@ -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"
}
@@ -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,
+11 -8
View File
@@ -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)
@@ -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 {
@@ -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()
@@ -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}