mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 15:38:22 +08:00
fix(grok): gate OAuth media on paid eligibility
This commit is contained in:
@@ -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.
|
||||
|
||||
---
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user