mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 15:57:53 +08:00
Merge pull request #4474 from heathermhuang/codex/fix-grok-media-image-url-4420
v0.1.159 fix(grok): normalize media payloads and quarantine ineligible accounts
This commit is contained in:
@@ -721,6 +721,7 @@ Sub2API supports both Grok subscription accounts through xAI OAuth and standard
|
||||
- Text models: `grok-4.5`, `grok-4.3`, `grok-build-0.1`, `grok-composer-2.5-fast`, `grok-4.20-0309-reasoning`, `grok-4.20-0309-non-reasoning`, and `grok-4.20-multi-agent-0309`
|
||||
- Media targets for Grok groups: `/v1/images/generations`, `/images/generations`, `/v1/images/edits`, `/images/edits`, `/v1/videos/generations`, `/videos/generations`, `/v1/videos/edits`, `/videos/edits`, `/v1/videos/extensions`, `/videos/extensions`, `/v1/videos/{request_id}`, and `/videos/{request_id}`. Generation, editing, and extension requests require the group image-generation permission.
|
||||
- Media models: `grok-imagine`, `grok-imagine-image-quality`, `grok-imagine-image`, `grok-imagine-edit`, `grok-imagine-video`, and `grok-imagine-video-1.5`
|
||||
- JSON image-edit and video-generation requests accept image references in `image`, `images`, `reference_images`, and `mask` objects. Use `url` for xAI-compatible payloads; the legacy `image_url` field remains accepted and is normalized to `url` before forwarding.
|
||||
- Out of scope for this provider: TTS, transcription, browser automation, cookies, and Grok web scraping
|
||||
|
||||
### OAuth Configuration
|
||||
@@ -787,6 +788,10 @@ 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.
|
||||
|
||||
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.
|
||||
|
||||
---
|
||||
|
||||
## Antigravity Support
|
||||
|
||||
@@ -173,6 +173,7 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
|
||||
maxAccountSwitches = 3
|
||||
}
|
||||
routingStart := time.Now()
|
||||
requiredCapability := grokMediaRequiredCapability(endpoint)
|
||||
|
||||
for {
|
||||
if failoverClientGone(c) {
|
||||
@@ -186,7 +187,7 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
|
||||
requestModel,
|
||||
failedAccountIDs,
|
||||
service.OpenAIUpstreamTransportHTTPSSE,
|
||||
"",
|
||||
requiredCapability,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
@@ -201,6 +202,11 @@ 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) {
|
||||
markOpsRoutingCapacityLimitedIfNoAvailable(c, err)
|
||||
h.errorResponse(c, http.StatusServiceUnavailable, "grok_media_no_eligible_account", "No eligible Grok media accounts")
|
||||
return
|
||||
}
|
||||
if len(failedAccountIDs) == 0 {
|
||||
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, requestModel, requestModel, service.PlatformGrok)
|
||||
if !cls.ModelNotFound {
|
||||
@@ -217,6 +223,11 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
|
||||
return
|
||||
}
|
||||
if selection == nil || selection.Account == nil {
|
||||
if endpoint.IsGenerationRequest() {
|
||||
markOpsRoutingCapacityLimited(c)
|
||||
h.errorResponse(c, http.StatusServiceUnavailable, "grok_media_no_eligible_account", "No eligible Grok media accounts")
|
||||
return
|
||||
}
|
||||
cls := classifyNoAccountErrorFromGin(c, h.gatewayService, apiKey, requestModel, requestModel, service.PlatformGrok)
|
||||
if !cls.ModelNotFound {
|
||||
markOpsRoutingCapacityLimited(c)
|
||||
@@ -354,6 +365,13 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
|
||||
}
|
||||
}
|
||||
|
||||
func grokMediaRequiredCapability(endpoint service.GrokMediaEndpoint) service.OpenAIEndpointCapability {
|
||||
if endpoint.IsGenerationRequest() {
|
||||
return service.OpenAIEndpointCapabilityGrokMediaGeneration
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func shouldRecordGrokMediaUsage(endpoint service.GrokMediaEndpoint, requestModel string) bool {
|
||||
return endpoint.IsGenerationRequest() && strings.TrimSpace(requestModel) != ""
|
||||
}
|
||||
|
||||
@@ -52,3 +52,24 @@ func TestShouldRecordGrokMediaUsage(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestGrokMediaRequiredCapability(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
endpoint service.GrokMediaEndpoint
|
||||
want service.OpenAIEndpointCapability
|
||||
}{
|
||||
{name: "image generation", endpoint: service.GrokMediaEndpointImagesGenerations, want: service.OpenAIEndpointCapabilityGrokMediaGeneration},
|
||||
{name: "image edit", endpoint: service.GrokMediaEndpointImagesEdits, want: service.OpenAIEndpointCapabilityGrokMediaGeneration},
|
||||
{name: "video generation", endpoint: service.GrokMediaEndpointVideosGenerations, want: service.OpenAIEndpointCapabilityGrokMediaGeneration},
|
||||
{name: "video edit", endpoint: service.GrokMediaEndpointVideosEdits, want: service.OpenAIEndpointCapabilityGrokMediaGeneration},
|
||||
{name: "video extension", endpoint: service.GrokMediaEndpointVideosExtensions, want: service.OpenAIEndpointCapabilityGrokMediaGeneration},
|
||||
{name: "video status preserves lookup", endpoint: service.GrokMediaEndpointVideoStatus, want: ""},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
require.Equal(t, tt.want, grokMediaRequiredCapability(tt.endpoint))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -76,6 +76,8 @@ type BillingSummary struct {
|
||||
UsedPercent *float64 `json:"used_percent,omitempty"`
|
||||
Plan string `json:"plan,omitempty"` // SuperGrok | SuperGrok Heavy | ""
|
||||
StatusCode int `json:"status_code,omitempty"`
|
||||
WeeklyStatusCode int `json:"weekly_status_code,omitempty"`
|
||||
MonthlyStatusCode int `json:"monthly_status_code,omitempty"`
|
||||
Source string `json:"source,omitempty"`
|
||||
FetchedAt string `json:"fetched_at,omitempty"`
|
||||
UpdatedAt string `json:"updated_at,omitempty"`
|
||||
|
||||
@@ -89,6 +89,11 @@ const (
|
||||
OpenAIEndpointCapabilityChatCompletions OpenAIEndpointCapability = "chat_completions"
|
||||
OpenAIEndpointCapabilityEmbeddings OpenAIEndpointCapability = "embeddings"
|
||||
OpenAIEndpointCapabilityAlphaSearch OpenAIEndpointCapability = "alpha_search"
|
||||
// OpenAIEndpointCapabilityGrokMediaGeneration keeps image/video generation
|
||||
// away from Grok accounts that are explicitly disabled or whose billing
|
||||
// entitlement probe was forbidden. Video status lookups intentionally do not
|
||||
// require this capability so already-submitted requests remain queryable.
|
||||
OpenAIEndpointCapabilityGrokMediaGeneration OpenAIEndpointCapability = "grok_media_generation"
|
||||
// OpenAIEndpointCapabilityResponses 表示上游确实提供 /v1/responses 端点。
|
||||
// 与其他能力不同:支持状态来自 accounts.extra 的自动探测标记
|
||||
// (openai_responses_supported / openai_responses_mode),而非
|
||||
@@ -99,6 +104,11 @@ const (
|
||||
|
||||
const openAIEndpointCapabilitiesCredentialKey = "openai_capabilities"
|
||||
|
||||
// GrokMediaEligibleExtraKey is an optional per-account override stored in
|
||||
// accounts.extra. true forces media routing on, false disables it, and an
|
||||
// absent/null value uses provider observations.
|
||||
const GrokMediaEligibleExtraKey = "grok_media_eligible"
|
||||
|
||||
const (
|
||||
OpenAIAuthModePersonalAccessToken = "personalAccessToken"
|
||||
openAIAuthModeCredentialKey = "auth_mode"
|
||||
@@ -1409,7 +1419,15 @@ func (a *Account) SupportsOpenAIEndpointCapability(capability OpenAIEndpointCapa
|
||||
return false
|
||||
}
|
||||
if a.IsGrok() {
|
||||
return capability == OpenAIEndpointCapabilityChatCompletions
|
||||
switch capability {
|
||||
case OpenAIEndpointCapabilityChatCompletions:
|
||||
return true
|
||||
case OpenAIEndpointCapabilityGrokMediaGeneration:
|
||||
eligible, _ := a.GrokMediaGenerationEligibility()
|
||||
return eligible
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
switch capability {
|
||||
case OpenAIEndpointCapabilityChatCompletions:
|
||||
@@ -1450,6 +1468,46 @@ func (a *Account) SupportsOpenAIEndpointCapability(capability OpenAIEndpointCapa
|
||||
return configured[string(capability)]
|
||||
}
|
||||
|
||||
// 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.
|
||||
func (a *Account) GrokMediaGenerationEligibility() (bool, string) {
|
||||
if a == nil || !a.IsGrok() {
|
||||
return false, "not_grok"
|
||||
}
|
||||
if override, ok := grokMediaEligibilityOverride(a.Extra); ok {
|
||||
if override {
|
||||
return true, "override_enabled"
|
||||
}
|
||||
return false, "override_disabled"
|
||||
}
|
||||
if a.Type != AccountTypeOAuth {
|
||||
return true, "non_oauth"
|
||||
}
|
||||
|
||||
billing, err := grokBillingSnapshotFromExtra(a.Extra)
|
||||
if err != nil || billing == nil {
|
||||
return true, "billing_unobserved"
|
||||
}
|
||||
if billing.StatusCode == 403 || billing.WeeklyStatusCode == 403 || billing.MonthlyStatusCode == 403 {
|
||||
return false, "billing_forbidden"
|
||||
}
|
||||
return true, "eligible"
|
||||
}
|
||||
|
||||
func grokMediaEligibilityOverride(extra map[string]any) (bool, bool) {
|
||||
if extra == nil {
|
||||
return false, false
|
||||
}
|
||||
raw, exists := extra[GrokMediaEligibleExtraKey]
|
||||
if !exists || raw == nil {
|
||||
return false, false
|
||||
}
|
||||
value, ok := raw.(bool)
|
||||
return value, ok
|
||||
}
|
||||
|
||||
func (a *Account) openAIEndpointCapabilitySet() (map[string]bool, bool) {
|
||||
if a == nil || a.Credentials == nil {
|
||||
return nil, false
|
||||
|
||||
@@ -0,0 +1,160 @@
|
||||
//go:build unit
|
||||
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"testing"
|
||||
|
||||
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/xai"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestGrokMediaGenerationEligibility(t *testing.T) {
|
||||
forbiddenBilling := &xai.BillingSummary{
|
||||
StatusCode: http.StatusForbidden,
|
||||
WeeklyStatusCode: http.StatusForbidden,
|
||||
MonthlyStatusCode: http.StatusForbidden,
|
||||
}
|
||||
weeklyAllowance := &xai.BillingSummary{
|
||||
PeriodType: "weekly",
|
||||
StatusCode: http.StatusOK,
|
||||
WeeklyStatusCode: http.StatusOK,
|
||||
}
|
||||
weeklyForbidden := &xai.BillingSummary{
|
||||
StatusCode: http.StatusOK,
|
||||
WeeklyStatusCode: http.StatusForbidden,
|
||||
MonthlyStatusCode: http.StatusOK,
|
||||
}
|
||||
monthlyForbidden := &xai.BillingSummary{
|
||||
StatusCode: http.StatusOK,
|
||||
WeeklyStatusCode: http.StatusOK,
|
||||
MonthlyStatusCode: http.StatusForbidden,
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
account *Account
|
||||
want bool
|
||||
wantReason string
|
||||
}{
|
||||
{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: "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 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"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got, reason := tt.account.GrokMediaGenerationEligibility()
|
||||
require.Equal(t, tt.want, got)
|
||||
require.Equal(t, tt.wantReason, reason)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestGrokMediaCapabilityFiltersOnlyGeneration(t *testing.T) {
|
||||
account := &Account{
|
||||
ID: 1,
|
||||
Platform: PlatformGrok,
|
||||
Type: AccountTypeOAuth,
|
||||
Status: StatusActive,
|
||||
Schedulable: true,
|
||||
Concurrency: 1,
|
||||
Extra: map[string]any{GrokMediaEligibleExtraKey: false},
|
||||
}
|
||||
|
||||
require.True(t, account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityChatCompletions))
|
||||
require.False(t, account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityGrokMediaGeneration))
|
||||
require.False(t, isOpenAICompatibleAccountEligibleForRequest(
|
||||
context.Background(), account, PlatformGrok, "grok-imagine-video", false,
|
||||
OpenAIEndpointCapabilityGrokMediaGeneration,
|
||||
))
|
||||
}
|
||||
|
||||
func TestNormalizeGrokMediaEligibilityExtra(t *testing.T) {
|
||||
t.Run("boolean override is accepted", func(t *testing.T) {
|
||||
extra, err := normalizeGrokMediaEligibilityExtra(PlatformGrok, map[string]any{GrokMediaEligibleExtraKey: false})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, false, extra[GrokMediaEligibleExtraKey])
|
||||
})
|
||||
|
||||
t.Run("null clears override", func(t *testing.T) {
|
||||
extra, err := normalizeGrokMediaEligibilityExtra(PlatformGrok, map[string]any{GrokMediaEligibleExtraKey: nil})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotContains(t, extra, GrokMediaEligibleExtraKey)
|
||||
})
|
||||
|
||||
t.Run("malformed override is rejected", func(t *testing.T) {
|
||||
_, err := normalizeGrokMediaEligibilityExtra(PlatformGrok, map[string]any{GrokMediaEligibleExtraKey: "false"})
|
||||
|
||||
require.Error(t, err)
|
||||
require.Equal(t, http.StatusBadRequest, infraerrors.Code(err))
|
||||
})
|
||||
|
||||
t.Run("other platforms ignore provider owned value", func(t *testing.T) {
|
||||
extra := map[string]any{GrokMediaEligibleExtraKey: "provider-owned"}
|
||||
normalized, err := normalizeGrokMediaEligibilityExtra(PlatformOpenAI, extra)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, extra, normalized)
|
||||
})
|
||||
}
|
||||
|
||||
func TestNormalizeGrokMediaEligibilityUpdateExtra(t *testing.T) {
|
||||
account := &Account{Platform: PlatformGrok, Extra: map[string]any{GrokMediaEligibleExtraKey: false}}
|
||||
|
||||
t.Run("omitted override preserves current value", func(t *testing.T) {
|
||||
input := &UpdateAccountInput{Extra: map[string]any{"quota_used": float64(1)}}
|
||||
normalized, err := normalizeGrokMediaEligibilityUpdateExtra(account, input, map[string]any{"quota_used": float64(1)})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, false, normalized[GrokMediaEligibleExtraKey])
|
||||
})
|
||||
|
||||
t.Run("null removes current override", func(t *testing.T) {
|
||||
input := &UpdateAccountInput{Extra: map[string]any{GrokMediaEligibleExtraKey: nil}}
|
||||
normalized, err := normalizeGrokMediaEligibilityUpdateExtra(account, input, map[string]any{GrokMediaEligibleExtraKey: nil})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotContains(t, normalized, GrokMediaEligibleExtraKey)
|
||||
require.Contains(t, input.Extra, GrokMediaEligibleExtraKey)
|
||||
})
|
||||
|
||||
t.Run("provided boolean replaces current override", func(t *testing.T) {
|
||||
input := &UpdateAccountInput{Extra: map[string]any{GrokMediaEligibleExtraKey: true}}
|
||||
normalized, err := normalizeGrokMediaEligibilityUpdateExtra(account, input, map[string]any{GrokMediaEligibleExtraKey: true})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, true, normalized[GrokMediaEligibleExtraKey])
|
||||
})
|
||||
|
||||
t.Run("malformed override is rejected on update", func(t *testing.T) {
|
||||
input := &UpdateAccountInput{Extra: map[string]any{GrokMediaEligibleExtraKey: "false"}}
|
||||
_, err := normalizeGrokMediaEligibilityUpdateExtra(account, input, nil)
|
||||
|
||||
require.Error(t, err)
|
||||
require.Equal(t, http.StatusBadRequest, infraerrors.Code(err))
|
||||
})
|
||||
|
||||
t.Run("non grok update is unchanged", func(t *testing.T) {
|
||||
input := &UpdateAccountInput{Extra: map[string]any{GrokMediaEligibleExtraKey: "provider-owned"}}
|
||||
normalized := map[string]any{GrokMediaEligibleExtraKey: "provider-owned"}
|
||||
got, err := normalizeGrokMediaEligibilityUpdateExtra(&Account{Platform: PlatformOpenAI}, input, normalized)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, normalized, got)
|
||||
})
|
||||
}
|
||||
@@ -394,6 +394,64 @@ func normalizeOpenAILongContextBillingUpdateExtra(account *Account, input *Updat
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
// ValidateGrokMediaEligibilityExtra validates the optional media-routing
|
||||
// override. null removes the override and returns the account to automatic
|
||||
// provider-observation based routing.
|
||||
func ValidateGrokMediaEligibilityExtra(platform string, extra map[string]any) error {
|
||||
if platform != PlatformGrok || extra == nil {
|
||||
return nil
|
||||
}
|
||||
raw, exists := extra[GrokMediaEligibleExtraKey]
|
||||
if !exists || raw == nil {
|
||||
return nil
|
||||
}
|
||||
if _, ok := raw.(bool); !ok {
|
||||
return infraerrors.BadRequest(
|
||||
"GROK_MEDIA_ELIGIBILITY_INVALID",
|
||||
"grok_media_eligible must be a boolean or null",
|
||||
)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func normalizeGrokMediaEligibilityExtra(platform string, extra map[string]any) (map[string]any, error) {
|
||||
if platform != PlatformGrok {
|
||||
return extra, nil
|
||||
}
|
||||
if err := ValidateGrokMediaEligibilityExtra(platform, extra); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
normalized := maps.Clone(extra)
|
||||
if normalized != nil && normalized[GrokMediaEligibleExtraKey] == nil {
|
||||
delete(normalized, GrokMediaEligibleExtraKey)
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func normalizeGrokMediaEligibilityUpdateExtra(account *Account, input *UpdateAccountInput, normalized map[string]any) (map[string]any, error) {
|
||||
if account == nil || account.Platform != PlatformGrok {
|
||||
return normalized, nil
|
||||
}
|
||||
if err := ValidateGrokMediaEligibilityExtra(account.Platform, input.Extra); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
normalized = maps.Clone(normalized)
|
||||
if normalized == nil {
|
||||
normalized = make(map[string]any)
|
||||
}
|
||||
raw, provided := input.Extra[GrokMediaEligibleExtraKey]
|
||||
if provided {
|
||||
if raw == nil {
|
||||
delete(normalized, GrokMediaEligibleExtraKey)
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
if current, ok := account.Extra[GrokMediaEligibleExtraKey].(bool); ok {
|
||||
normalized[GrokMediaEligibleExtraKey] = current
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func buildAccountForCreate(input *CreateAccountInput, accountExtra map[string]any) (*Account, error) {
|
||||
// Probe state is system-managed. New accounts always start with auto probe disabled.
|
||||
delete(accountExtra, UpstreamBillingProbeEnabledExtraKey)
|
||||
@@ -448,6 +506,10 @@ func (s *adminServiceImpl) CreateAccount(ctx context.Context, input *CreateAccou
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
accountExtra, err = normalizeGrokMediaEligibilityExtra(input.Platform, accountExtra)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 绑定分组
|
||||
groupIDs := input.GroupIDs
|
||||
@@ -535,6 +597,10 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
normalizedExtra, err = normalizeGrokMediaEligibilityUpdateExtra(account, input, normalizedExtra)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
previousProbeIdentity := upstreamBillingProbeIdentity(account)
|
||||
// 安全/身份不变量(影子账号):通用更新路径被 edit/re-auth/refresh/batch 共用,
|
||||
@@ -607,6 +673,7 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U
|
||||
"quota_daily_start",
|
||||
"quota_weekly_used",
|
||||
"quota_weekly_start",
|
||||
grokBillingExtraKey,
|
||||
UpstreamBillingProbeEnabledExtraKey,
|
||||
UpstreamBillingProbeExtraKey,
|
||||
} {
|
||||
|
||||
@@ -2,8 +2,10 @@ package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"testing"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/xai"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
@@ -62,6 +64,34 @@ func TestUpdateAccountPreservesManagedUpstreamBillingProbeStateForUnrelatedEdit(
|
||||
require.Equal(t, "value", updated.Extra["custom"])
|
||||
}
|
||||
|
||||
func TestUpdateAccountPreservesGrokBillingSnapshotForUnrelatedEdit(t *testing.T) {
|
||||
accountID := int64(112)
|
||||
billing := &xai.BillingSummary{
|
||||
StatusCode: http.StatusForbidden,
|
||||
WeeklyStatusCode: http.StatusForbidden,
|
||||
}
|
||||
repo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{
|
||||
accountID: {
|
||||
ID: accountID,
|
||||
Platform: PlatformGrok,
|
||||
Type: AccountTypeOAuth,
|
||||
Status: StatusActive,
|
||||
Extra: map[string]any{grokBillingExtraKey: billing},
|
||||
},
|
||||
}}
|
||||
|
||||
updated, err := (&adminServiceImpl{accountRepo: repo}).UpdateAccount(context.Background(), accountID, &UpdateAccountInput{
|
||||
Extra: map[string]any{"custom": "value"},
|
||||
})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, billing, updated.Extra[grokBillingExtraKey])
|
||||
require.Equal(t, "value", updated.Extra["custom"])
|
||||
eligible, reason := updated.GrokMediaGenerationEligibility()
|
||||
require.False(t, eligible)
|
||||
require.Equal(t, "billing_forbidden", reason)
|
||||
}
|
||||
|
||||
func TestUpdateAccountPreservesProbeSnapshotWhenIdentityValuesAreUnchanged(t *testing.T) {
|
||||
accountID := int64(119)
|
||||
repo := &upstreamBillingProbeAccountRepo{accounts: map[int64]*Account{
|
||||
|
||||
@@ -147,7 +147,7 @@ func parseGrokMediaJSONRequest(body []byte, info *GrokMediaRequestInfo) {
|
||||
switch {
|
||||
case value.IsArray():
|
||||
for _, item := range value.Array() {
|
||||
if imageURL := strings.TrimSpace(item.Get("image_url").String()); imageURL != "" {
|
||||
if imageURL := grokMediaJSONImageURL(item); imageURL != "" {
|
||||
info.InputImageURLs = append(info.InputImageURLs, imageURL)
|
||||
continue
|
||||
}
|
||||
@@ -160,7 +160,7 @@ func parseGrokMediaJSONRequest(body []byte, info *GrokMediaRequestInfo) {
|
||||
}
|
||||
}
|
||||
default:
|
||||
if imageURL := strings.TrimSpace(value.Get("image_url").String()); imageURL != "" {
|
||||
if imageURL := grokMediaJSONImageURL(value); imageURL != "" {
|
||||
info.InputImageURLs = append(info.InputImageURLs, imageURL)
|
||||
return
|
||||
}
|
||||
@@ -175,7 +175,15 @@ func parseGrokMediaJSONRequest(body []byte, info *GrokMediaRequestInfo) {
|
||||
}
|
||||
appendJSONImageURLs(gjson.GetBytes(body, "image"))
|
||||
appendJSONImageURLs(gjson.GetBytes(body, "images"))
|
||||
info.MaskImageURL = strings.TrimSpace(gjson.GetBytes(body, "mask.image_url").String())
|
||||
appendJSONImageURLs(gjson.GetBytes(body, "reference_images"))
|
||||
info.MaskImageURL = grokMediaJSONImageURL(gjson.GetBytes(body, "mask"))
|
||||
}
|
||||
|
||||
func grokMediaJSONImageURL(value gjson.Result) string {
|
||||
if imageURL := strings.TrimSpace(value.Get("url").String()); imageURL != "" {
|
||||
return imageURL
|
||||
}
|
||||
return strings.TrimSpace(value.Get("image_url").String())
|
||||
}
|
||||
|
||||
func parseGrokMediaMultipartRequest(contentType string, body []byte, info *GrokMediaRequestInfo) {
|
||||
@@ -409,7 +417,7 @@ func prepareGrokMediaForwardBody(endpoint GrokMediaEndpoint, body []byte, conten
|
||||
images := make([]map[string]string, 0, len(info.InputImageURLs)+len(info.Uploads))
|
||||
for _, imageURL := range info.InputImageURLs {
|
||||
if imageURL = strings.TrimSpace(imageURL); imageURL != "" {
|
||||
images = append(images, map[string]string{"image_url": imageURL})
|
||||
images = append(images, map[string]string{"url": imageURL})
|
||||
}
|
||||
}
|
||||
for _, upload := range info.Uploads {
|
||||
@@ -417,7 +425,7 @@ func prepareGrokMediaForwardBody(endpoint GrokMediaEndpoint, body []byte, conten
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
images = append(images, map[string]string{"image_url": dataURL})
|
||||
images = append(images, map[string]string{"url": dataURL})
|
||||
}
|
||||
if len(images) > 0 {
|
||||
payload["image"] = images[0]
|
||||
@@ -435,7 +443,7 @@ func prepareGrokMediaForwardBody(endpoint GrokMediaEndpoint, body []byte, conten
|
||||
maskImageURL = dataURL
|
||||
}
|
||||
if maskImageURL != "" {
|
||||
payload["mask"] = map[string]string{"image_url": maskImageURL}
|
||||
payload["mask"] = map[string]string{"url": maskImageURL}
|
||||
}
|
||||
|
||||
out, err := marshalOpenAIUpstreamJSON(payload)
|
||||
@@ -449,6 +457,18 @@ func normalizeGrokMediaForwardBody(endpoint GrokMediaEndpoint, body []byte, cont
|
||||
if !endpoint.RequiresRequestBody() || !gjson.ValidBytes(body) {
|
||||
return body, contentType, nil
|
||||
}
|
||||
var imageFields []string
|
||||
switch endpoint {
|
||||
case GrokMediaEndpointImagesEdits:
|
||||
imageFields = []string{"image", "images", "mask"}
|
||||
case GrokMediaEndpointVideosGenerations:
|
||||
imageFields = []string{"image", "images", "reference_images"}
|
||||
}
|
||||
var err error
|
||||
body, err = canonicalizeGrokMediaImageURLFields(body, imageFields...)
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
info := ParseGrokMediaRequest(contentType, body)
|
||||
upstreamModel := normalizeGrokMediaModelForEndpoint(endpoint, info.Model, info.HasInputImage())
|
||||
if upstreamModel == "" || upstreamModel == info.Model {
|
||||
@@ -461,6 +481,54 @@ func normalizeGrokMediaForwardBody(endpoint GrokMediaEndpoint, body []byte, cont
|
||||
return out, contentType, nil
|
||||
}
|
||||
|
||||
func canonicalizeGrokMediaImageURLFields(body []byte, fields ...string) ([]byte, error) {
|
||||
out := body
|
||||
for _, field := range fields {
|
||||
value := gjson.GetBytes(out, field)
|
||||
if !value.Exists() {
|
||||
continue
|
||||
}
|
||||
if value.IsArray() {
|
||||
for index := range value.Array() {
|
||||
var err error
|
||||
out, err = canonicalizeGrokMediaImageURLObject(out, fmt.Sprintf("%s.%d", field, index))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
continue
|
||||
}
|
||||
var err error
|
||||
out, err = canonicalizeGrokMediaImageURLObject(out, field)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func canonicalizeGrokMediaImageURLObject(body []byte, path string) ([]byte, error) {
|
||||
legacyPath := path + ".image_url"
|
||||
legacy := gjson.GetBytes(body, legacyPath)
|
||||
if !legacy.Exists() {
|
||||
return body, nil
|
||||
}
|
||||
|
||||
out := body
|
||||
if strings.TrimSpace(gjson.GetBytes(out, path+".url").String()) == "" {
|
||||
var err error
|
||||
out, err = sjson.SetBytes(out, path+".url", legacy.Value())
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("normalize grok media image url: %w", err)
|
||||
}
|
||||
}
|
||||
out, err := sjson.DeleteBytes(out, legacyPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("remove legacy grok media image url: %w", err)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func sanitizeGrokMediaForwardBody(endpoint GrokMediaEndpoint, body []byte, contentType string) ([]byte, string, error) {
|
||||
if !endpoint.RequiresRequestBody() || !gjson.ValidBytes(body) {
|
||||
return body, contentType, nil
|
||||
|
||||
@@ -249,12 +249,25 @@ func (s *GrokQuotaService) probeBilling(ctx context.Context, accountID int64) (*
|
||||
|
||||
weeklyOK := weekly.summary != nil
|
||||
monthlyOK := monthly.summary != nil
|
||||
previous, _ := grokBillingSnapshotFromExtra(account.Extra)
|
||||
if !weeklyOK && !monthlyOK {
|
||||
return nil, mergeGrokBillingProbeErrors(weekly.status, monthly.status, weekly.err, monthly.err)
|
||||
probeErr := mergeGrokBillingProbeErrors(weekly.status, monthly.status, weekly.err, monthly.err)
|
||||
billing := xai.MergeBillingProbeResult(previous, nil, nil, false, false)
|
||||
if billing == nil {
|
||||
billing = &xai.BillingSummary{Partial: true, FailedWindows: []string{"weekly", "monthly"}}
|
||||
}
|
||||
billing.WeeklyStatusCode = weekly.status
|
||||
billing.MonthlyStatusCode = monthly.status
|
||||
billing = xai.StampBillingSummary(billing, preferBillingObservationStatus(weekly.status, monthly.status), "billing_probe")
|
||||
if persistErr := s.accountRepo.UpdateExtra(ctx, account.ID, map[string]any{grokBillingExtraKey: billing}); persistErr != nil {
|
||||
slog.Warn("grok_billing_failure_persist_failed", "account_id", account.ID, "error", persistErr)
|
||||
}
|
||||
return nil, probeErr
|
||||
}
|
||||
statusCode := preferSuccessfulBillingStatus(weekly.status, monthly.status, weeklyOK, monthlyOK)
|
||||
previous, _ := grokBillingSnapshotFromExtra(account.Extra)
|
||||
billing := xai.MergeBillingProbeResult(previous, weekly.summary, monthly.summary, weeklyOK, monthlyOK)
|
||||
billing.WeeklyStatusCode = weekly.status
|
||||
billing.MonthlyStatusCode = monthly.status
|
||||
billing = xai.StampBillingSummary(billing, statusCode, "billing_probe")
|
||||
persistErr := s.accountRepo.UpdateExtra(ctx, account.ID, map[string]any{
|
||||
grokBillingExtraKey: billing,
|
||||
@@ -276,6 +289,16 @@ func (s *GrokQuotaService) probeBilling(ctx context.Context, accountID int64) (*
|
||||
}, nil
|
||||
}
|
||||
|
||||
func preferBillingObservationStatus(weeklyStatus, monthlyStatus int) int {
|
||||
if weeklyStatus == http.StatusForbidden || monthlyStatus == http.StatusForbidden {
|
||||
return http.StatusForbidden
|
||||
}
|
||||
if weeklyStatus != 0 {
|
||||
return weeklyStatus
|
||||
}
|
||||
return monthlyStatus
|
||||
}
|
||||
|
||||
func (s *GrokQuotaService) runProbeFlight(
|
||||
ctx context.Context,
|
||||
key string,
|
||||
|
||||
@@ -99,18 +99,20 @@ func (r *grokQuotaUsageLogRepo) GetAccountTodayStats(context.Context, int64) (*u
|
||||
|
||||
type grokHybridUpstream struct {
|
||||
httpUpstreamRecorder
|
||||
mu sync.Mutex
|
||||
requests []*http.Request
|
||||
bodies [][]byte
|
||||
weeklyUsagePercent *float64
|
||||
monthlyLimitCents *float64
|
||||
activeStatus int
|
||||
activeHeaders http.Header
|
||||
billingStarted chan struct{}
|
||||
billingRelease <-chan struct{}
|
||||
billingStartOnce sync.Once
|
||||
billingStatus int
|
||||
billingHeaders http.Header
|
||||
mu sync.Mutex
|
||||
requests []*http.Request
|
||||
bodies [][]byte
|
||||
weeklyUsagePercent *float64
|
||||
monthlyLimitCents *float64
|
||||
activeStatus int
|
||||
activeHeaders http.Header
|
||||
billingStarted chan struct{}
|
||||
billingRelease <-chan struct{}
|
||||
billingStartOnce sync.Once
|
||||
billingStatus int
|
||||
weeklyBillingStatus int
|
||||
monthlyBillingStatus int
|
||||
billingHeaders http.Header
|
||||
}
|
||||
|
||||
func (u *grokHybridUpstream) Do(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) {
|
||||
@@ -147,9 +149,16 @@ func (u *grokHybridUpstream) Do(req *http.Request, _ string, _ int64, _ int) (*h
|
||||
return nil, req.Context().Err()
|
||||
}
|
||||
}
|
||||
if u.billingStatus != 0 && u.billingStatus != http.StatusOK {
|
||||
billingStatus := u.billingStatus
|
||||
if req.URL.RawQuery == "format=credits" && u.weeklyBillingStatus != 0 {
|
||||
billingStatus = u.weeklyBillingStatus
|
||||
}
|
||||
if req.URL.RawQuery != "format=credits" && u.monthlyBillingStatus != 0 {
|
||||
billingStatus = u.monthlyBillingStatus
|
||||
}
|
||||
if billingStatus != 0 && billingStatus != http.StatusOK {
|
||||
return &http.Response{
|
||||
StatusCode: u.billingStatus,
|
||||
StatusCode: billingStatus,
|
||||
Header: u.billingHeaders,
|
||||
Body: io.NopCloser(strings.NewReader(`{"error":{"message":"billing limited"}}`)),
|
||||
}, nil
|
||||
@@ -755,6 +764,88 @@ func TestGrokQuotaServiceBilling429DoesNotPauseModelScheduling(t *testing.T) {
|
||||
require.Zero(t, repo.rateLimitedCalls)
|
||||
}
|
||||
|
||||
func TestGrokQuotaServiceBilling403PersistsMediaEligibilitySignal(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
account := healthyGrokQuotaOAuthAccount(58)
|
||||
repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
|
||||
accountsByID: map[int64]*Account{account.ID: account},
|
||||
}}
|
||||
upstream := &grokHybridUpstream{billingStatus: http.StatusForbidden}
|
||||
svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream, nil)
|
||||
|
||||
result, err := svc.ProbeBilling(context.Background(), account.ID)
|
||||
|
||||
require.Error(t, err)
|
||||
require.Nil(t, result)
|
||||
require.Equal(t, 1, repo.updateCalls)
|
||||
raw := repo.updates[account.ID][grokBillingExtraKey]
|
||||
billing, ok := raw.(*xai.BillingSummary)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, http.StatusForbidden, billing.StatusCode)
|
||||
require.Equal(t, http.StatusForbidden, billing.WeeklyStatusCode)
|
||||
require.Equal(t, http.StatusForbidden, billing.MonthlyStatusCode)
|
||||
require.True(t, billing.Partial)
|
||||
|
||||
account.Extra = map[string]any{grokBillingExtraKey: billing}
|
||||
eligible, reason := account.GrokMediaGenerationEligibility()
|
||||
require.False(t, eligible)
|
||||
require.Equal(t, "billing_forbidden", reason)
|
||||
}
|
||||
|
||||
func TestGrokQuotaServicePartialBilling403PersistsMediaEligibilitySignal(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
account := healthyGrokQuotaOAuthAccount(59)
|
||||
repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
|
||||
accountsByID: map[int64]*Account{account.ID: account},
|
||||
}}
|
||||
upstream := &grokHybridUpstream{
|
||||
weeklyBillingStatus: http.StatusForbidden,
|
||||
monthlyBillingStatus: http.StatusOK,
|
||||
}
|
||||
svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream, nil)
|
||||
|
||||
result, err := svc.ProbeBilling(context.Background(), account.ID)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
require.NotNil(t, result.Billing)
|
||||
require.Equal(t, http.StatusOK, result.StatusCode)
|
||||
require.Equal(t, http.StatusForbidden, result.Billing.WeeklyStatusCode)
|
||||
require.Equal(t, http.StatusOK, result.Billing.MonthlyStatusCode)
|
||||
require.True(t, result.Billing.Partial)
|
||||
require.Contains(t, result.Billing.FailedWindows, "weekly")
|
||||
require.Equal(t, 1, repo.updateCalls)
|
||||
|
||||
account.Extra = map[string]any{grokBillingExtraKey: result.Billing}
|
||||
eligible, reason := account.GrokMediaGenerationEligibility()
|
||||
require.False(t, eligible)
|
||||
require.Equal(t, "billing_forbidden", reason)
|
||||
}
|
||||
|
||||
func TestPreferBillingObservationStatus(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
weeklyStatus int
|
||||
monthlyStatus int
|
||||
want int
|
||||
}{
|
||||
{name: "weekly forbidden wins", weeklyStatus: http.StatusForbidden, monthlyStatus: http.StatusBadGateway, want: http.StatusForbidden},
|
||||
{name: "monthly forbidden wins", weeklyStatus: http.StatusBadGateway, monthlyStatus: http.StatusForbidden, want: http.StatusForbidden},
|
||||
{name: "weekly observation otherwise wins", weeklyStatus: http.StatusTooManyRequests, monthlyStatus: http.StatusBadGateway, want: http.StatusTooManyRequests},
|
||||
{name: "monthly observation is fallback", monthlyStatus: http.StatusBadGateway, want: http.StatusBadGateway},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
require.Equal(t, tt.want, preferBillingObservationStatus(tt.weeklyStatus, tt.monthlyStatus))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestGrokQuotaServiceQueryQuotaFree429PersistsLimitAndKeepsBilling(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -761,6 +761,71 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_DefaultDisabled_AllowsG
|
||||
require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer)
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayService_SelectAccountWithScheduler_GrokMediaCapabilityFiltersIneligibleAccounts(t *testing.T) {
|
||||
resetOpenAIAdvancedSchedulerSettingCacheForTest()
|
||||
|
||||
ctx := context.Background()
|
||||
groupID := int64(10114)
|
||||
ineligible := Account{
|
||||
ID: 36051, Platform: PlatformGrok, Type: AccountTypeOAuth,
|
||||
Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5,
|
||||
Extra: map[string]any{GrokMediaEligibleExtraKey: false},
|
||||
}
|
||||
eligible := Account{
|
||||
ID: 36052, Platform: PlatformGrok, Type: AccountTypeOAuth,
|
||||
Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 0,
|
||||
Extra: map[string]any{GrokMediaEligibleExtraKey: true},
|
||||
}
|
||||
newService := func(accounts []Account) *OpenAIGatewayService {
|
||||
cfg := &config.Config{}
|
||||
cfg.Gateway.Scheduling.LoadBatchEnabled = false
|
||||
return &OpenAIGatewayService{
|
||||
accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts},
|
||||
cache: &schedulerTestGatewayCache{},
|
||||
cfg: cfg,
|
||||
concurrencyService: NewConcurrencyService(schedulerTestConcurrencyCache{}),
|
||||
}
|
||||
}
|
||||
|
||||
t.Run("media generation skips higher priority ineligible account", func(t *testing.T) {
|
||||
selection, _, err := newService([]Account{ineligible, eligible}).SelectAccountWithSchedulerForCapability(
|
||||
ctx, &groupID, "", "", "grok-imagine-video", nil,
|
||||
OpenAIUpstreamTransportHTTPSSE, OpenAIEndpointCapabilityGrokMediaGeneration,
|
||||
false, false, false, PlatformGrok,
|
||||
)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, selection)
|
||||
require.NotNil(t, selection.Account)
|
||||
require.Equal(t, eligible.ID, selection.Account.ID)
|
||||
})
|
||||
|
||||
t.Run("media generation fails closed when all accounts are ineligible", func(t *testing.T) {
|
||||
selection, _, err := newService([]Account{ineligible}).SelectAccountWithSchedulerForCapability(
|
||||
ctx, &groupID, "", "", "grok-imagine-video", nil,
|
||||
OpenAIUpstreamTransportHTTPSSE, OpenAIEndpointCapabilityGrokMediaGeneration,
|
||||
false, false, false, PlatformGrok,
|
||||
)
|
||||
|
||||
require.Error(t, err)
|
||||
require.ErrorIs(t, err, ErrNoAvailableAccounts)
|
||||
require.Nil(t, selection)
|
||||
})
|
||||
|
||||
t.Run("chat remains routable on media-ineligible account", func(t *testing.T) {
|
||||
selection, _, err := newService([]Account{ineligible}).SelectAccountWithSchedulerForCapability(
|
||||
ctx, &groupID, "", "", "grok-4.3", nil,
|
||||
OpenAIUpstreamTransportHTTPSSE, OpenAIEndpointCapabilityChatCompletions,
|
||||
false, false, false, PlatformGrok,
|
||||
)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, selection)
|
||||
require.NotNil(t, selection.Account)
|
||||
require.Equal(t, ineligible.ID, selection.Account.ID)
|
||||
})
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayService_SelectAccountWithScheduler_EnabledUsesAdvancedPreviousResponseRouting(t *testing.T) {
|
||||
resetOpenAIAdvancedSchedulerSettingCacheForTest()
|
||||
|
||||
|
||||
@@ -418,6 +418,88 @@ func TestParseGrokMediaVideoRequestResolution(t *testing.T) {
|
||||
require.Equal(t, "720p", info.Resolution)
|
||||
}
|
||||
|
||||
func TestParseGrokMediaRequestAcceptsOfficialImageURLFields(t *testing.T) {
|
||||
body := []byte(`{
|
||||
"model":"grok-imagine-video-1.5",
|
||||
"image":{"url":"https://example.com/source.png"},
|
||||
"reference_images":[{"url":"https://example.com/reference.png"}]
|
||||
}`)
|
||||
|
||||
info := ParseGrokMediaRequest("application/json", body)
|
||||
|
||||
require.Equal(t, []string{
|
||||
"https://example.com/source.png",
|
||||
"https://example.com/reference.png",
|
||||
}, info.InputImageURLs)
|
||||
require.True(t, info.HasInputImage())
|
||||
}
|
||||
|
||||
func TestNormalizeGrokMediaForwardBodyCanonicalizesImageURLAlias(t *testing.T) {
|
||||
body := []byte(`{
|
||||
"model":"grok-imagine-video-1.5",
|
||||
"prompt":"animate",
|
||||
"image":{"image_url":"https://example.com/source.png"},
|
||||
"duration":8
|
||||
}`)
|
||||
|
||||
out, contentType, err := normalizeGrokMediaForwardBody(GrokMediaEndpointVideosGenerations, body, "application/json")
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "application/json", contentType)
|
||||
require.Equal(t, "grok-imagine-video-1.5", gjson.GetBytes(out, "model").String())
|
||||
require.Equal(t, "https://example.com/source.png", gjson.GetBytes(out, "image.url").String())
|
||||
require.False(t, gjson.GetBytes(out, "image.image_url").Exists())
|
||||
}
|
||||
|
||||
func TestNormalizeGrokMediaForwardBodyPreservesImageToVideoModelForOfficialURL(t *testing.T) {
|
||||
body := []byte(`{
|
||||
"model":"grok-imagine-video-1.5",
|
||||
"prompt":"animate",
|
||||
"image":{"url":"https://example.com/source.png"}
|
||||
}`)
|
||||
|
||||
out, _, err := normalizeGrokMediaForwardBody(GrokMediaEndpointVideosGenerations, body, "application/json")
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "grok-imagine-video-1.5", gjson.GetBytes(out, "model").String())
|
||||
require.Equal(t, "https://example.com/source.png", gjson.GetBytes(out, "image.url").String())
|
||||
}
|
||||
|
||||
func TestCanonicalizeGrokMediaImageURLFieldsPreservesOfficialURL(t *testing.T) {
|
||||
body := []byte(`{
|
||||
"image":{"url":"https://example.com/official.png","image_url":"https://example.com/legacy.png"},
|
||||
"images":[
|
||||
{"image_url":"https://example.com/first.png"},
|
||||
{"url":"https://example.com/second.png"}
|
||||
],
|
||||
"reference_images":[{"image_url":"https://example.com/reference.png"}],
|
||||
"mask":{"image_url":"https://example.com/mask.png"}
|
||||
}`)
|
||||
|
||||
out, err := canonicalizeGrokMediaImageURLFields(body, "image", "images", "reference_images", "mask")
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "https://example.com/official.png", gjson.GetBytes(out, "image.url").String())
|
||||
require.False(t, gjson.GetBytes(out, "image.image_url").Exists())
|
||||
require.Equal(t, "https://example.com/first.png", gjson.GetBytes(out, "images.0.url").String())
|
||||
require.False(t, gjson.GetBytes(out, "images.0.image_url").Exists())
|
||||
require.Equal(t, "https://example.com/second.png", gjson.GetBytes(out, "images.1.url").String())
|
||||
require.Equal(t, "https://example.com/reference.png", gjson.GetBytes(out, "reference_images.0.url").String())
|
||||
require.False(t, gjson.GetBytes(out, "reference_images.0.image_url").Exists())
|
||||
require.Equal(t, "https://example.com/mask.png", gjson.GetBytes(out, "mask.url").String())
|
||||
require.False(t, gjson.GetBytes(out, "mask.image_url").Exists())
|
||||
}
|
||||
|
||||
func TestCanonicalizeGrokMediaImageURLFieldsReplacesEmptyOfficialURL(t *testing.T) {
|
||||
body := []byte(`{"image":{"url":" ","image_url":"https://example.com/legacy.png"}}`)
|
||||
|
||||
out, err := canonicalizeGrokMediaImageURLFields(body, "image")
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "https://example.com/legacy.png", gjson.GetBytes(out, "image.url").String())
|
||||
require.False(t, gjson.GetBytes(out, "image.image_url").Exists())
|
||||
}
|
||||
|
||||
func TestNormalizeGrokMediaModelForEndpoint(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
@@ -577,7 +659,8 @@ func TestForwardGrokMediaImagesEditMultipartConvertsToJSON(t *testing.T) {
|
||||
require.True(t, json.Valid(upstream.lastBody))
|
||||
require.Equal(t, "grok-imagine-edit", gjson.GetBytes(upstream.lastBody, "model").String())
|
||||
require.Equal(t, "edit this private image", gjson.GetBytes(upstream.lastBody, "prompt").String())
|
||||
require.True(t, strings.HasPrefix(gjson.GetBytes(upstream.lastBody, "image.image_url").String(), "data:image/png;base64,"))
|
||||
require.True(t, strings.HasPrefix(gjson.GetBytes(upstream.lastBody, "image.url").String(), "data:image/png;base64,"))
|
||||
require.False(t, gjson.GetBytes(upstream.lastBody, "image.image_url").Exists())
|
||||
}
|
||||
|
||||
func TestForwardGrokMediaVideoGenerationReturnsUsageAndResponseID(t *testing.T) {
|
||||
@@ -659,7 +742,7 @@ func TestForwardGrokMediaVideoGenerationPreservesImageToVideoModel(t *testing.T)
|
||||
result, err := svc.ForwardGrokMedia(context.Background(), c, account, GrokMediaEndpointVideosGenerations, "", body, "application/json")
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "https://xai.test/v1/videos/generations", upstream.lastReq.URL.String())
|
||||
require.JSONEq(t, `{"model":"grok-imagine-video-1.5","prompt":"animate","image":{"image_url":"data:image/png;base64,aW1n"}}`, string(upstream.lastBody))
|
||||
require.JSONEq(t, `{"model":"grok-imagine-video-1.5","prompt":"animate","image":{"url":"data:image/png;base64,aW1n"}}`, string(upstream.lastBody))
|
||||
require.Equal(t, "video-request-456", result.ResponseID)
|
||||
require.Equal(t, "grok-imagine-video-1.5", result.BillingModel)
|
||||
// 未指定 duration 时按上游默认 8 秒计费。
|
||||
@@ -705,7 +788,8 @@ func TestForwardGrokMediaOAuthImageToVideoUsesOfficialAPIForLargeBody(t *testing
|
||||
require.Equal(t, xai.DefaultBaseURL+"/videos/generations", upstream.lastReq.URL.String())
|
||||
require.Empty(t, upstream.lastReq.Header.Get("X-XAI-Token-Auth"))
|
||||
require.Empty(t, upstream.lastReq.Header.Get("x-grok-client-version"))
|
||||
require.Equal(t, "data:image/png;base64,"+imageData, gjson.GetBytes(upstream.lastBody, "image.image_url").String())
|
||||
require.Equal(t, "data:image/png;base64,"+imageData, gjson.GetBytes(upstream.lastBody, "image.url").String())
|
||||
require.False(t, gjson.GetBytes(upstream.lastBody, "image.image_url").Exists())
|
||||
}
|
||||
|
||||
func TestForwardGrokMediaVideoStatusUsesGETWithoutBody(t *testing.T) {
|
||||
|
||||
@@ -6,7 +6,6 @@ package service
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"sort"
|
||||
@@ -176,9 +175,21 @@ func noAvailableOpenAISelectionError(requestedModel string, compactBlocked bool)
|
||||
return ErrNoAvailableCompactAccounts
|
||||
}
|
||||
if requestedModel != "" {
|
||||
return fmt.Errorf("no available OpenAI accounts supporting model: %s", requestedModel)
|
||||
return openAINoAvailableSelectionError{message: fmt.Sprintf("no available OpenAI accounts supporting model: %s", requestedModel)}
|
||||
}
|
||||
return errors.New("no available OpenAI accounts")
|
||||
return openAINoAvailableSelectionError{message: "no available OpenAI accounts"}
|
||||
}
|
||||
|
||||
type openAINoAvailableSelectionError struct {
|
||||
message string
|
||||
}
|
||||
|
||||
func (e openAINoAvailableSelectionError) Error() string {
|
||||
return e.message
|
||||
}
|
||||
|
||||
func (e openAINoAvailableSelectionError) Unwrap() error {
|
||||
return ErrNoAvailableAccounts
|
||||
}
|
||||
|
||||
// openAICompactSupportTier classifies an OpenAI account by compact capability.
|
||||
@@ -235,6 +246,10 @@ func isOpenAICompatibleAccountEligibleForRequest(ctx context.Context, account *A
|
||||
return false
|
||||
}
|
||||
if !account.SupportsOpenAIEndpointCapability(requiredCapability) {
|
||||
if account.IsGrok() && requiredCapability == OpenAIEndpointCapabilityGrokMediaGeneration {
|
||||
_, reason := account.GrokMediaGenerationEligibility()
|
||||
slog.Debug("grok_media_account_ineligible", "account_id", account.ID, "reason", reason)
|
||||
}
|
||||
return false
|
||||
}
|
||||
if requireCompact && (!account.IsOpenAI() || openAICompactSupportTier(account) == 0) {
|
||||
|
||||
Reference in New Issue
Block a user