Merge pull request #5573 from IanShaw027/fix/grok-long-context-and-media-fallback

fix(grok): 长上下文阶梯不再被 OpenAI 账号开关否决,媒体 ID 不继承文本价
This commit is contained in:
Wesley Liddick
2026-08-13 10:32:07 +08:00
committed by GitHub
8 changed files with 109 additions and 15 deletions
@@ -862,6 +862,9 @@ func (s *BillingService) grokUnknownTextFamilyFallback(model string) *ModelPrici
func isGrokUnknownTextFamilyModel(model string) bool {
native := strings.ToLower(strings.TrimSpace(xai.StripGrokProviderPrefix(model)))
if isGrokMediaFamilyModel(native) {
return false
}
switch {
case native == "grok", native == "grok-latest":
return true
@@ -877,6 +880,19 @@ func isGrokUnknownTextFamilyModel(model string) bool {
}
}
// isGrokMediaFamilyModel matches ids that are billed per image/video/audio unit
// rather than per token, so version-numbered media ids (grok-2-image-1212,
// grok-5-video) cannot slip into the unknown-text fallback and pick up a token
// card. "vision" is deliberately absent: multimodal chat models are token billed.
func isGrokMediaFamilyModel(native string) bool {
for _, marker := range []string{"imagine", "image", "video", "audio", "speech", "tts", "transcribe", "realtime"} {
if strings.Contains(native, marker) {
return true
}
}
return false
}
// HasIdentifiedTokenPricing 判断模型能否在价格表中被"确定性识别"出 token 价格。
//
// 与 GetModelPricing 的关键区别:本函数拒绝按子串猜系列的兜底。GetModelPricing 会
@@ -1221,6 +1221,14 @@ func TestGetModelPricing_UnknownGrokTextFallsBackToGrok45(t *testing.T) {
require.InDelta(t, baseline.CacheReadPricePerToken, pricing.CacheReadPricePerToken, 1e-12, model)
}
// Per-unit media ids must not inherit the text card just because they carry
// a version number; they are billed by the image/video/audio paths instead.
for _, model := range []string{"grok-2-image-1212", "grok-2-audio", "grok-5-video", "x-ai/grok-6-image"} {
require.False(t, isGrokUnknownTextFamilyModel(model), "model %s", model)
}
// Multimodal chat models stay token billed.
require.True(t, isGrokUnknownTextFamilyModel("grok-2-vision-1212"))
for _, model := range []string{
"grok-imagine-image-3.0",
"grok-imagine-video-2",
@@ -50,7 +50,7 @@ func TestCalculateOpenAIRecordUsageCostWebSearchPerCall(t *testing.T) {
// 即使 token 倍率(含高峰,3.0)更高也不采用。
apiKey := &APIKey{ID: 1, GroupID: &groupID, Group: &Group{ID: groupID, Platform: PlatformOpenAI}}
result := &OpenAIForwardResult{Model: "gpt-5.6-sol", UpstreamModel: "gpt-5.6-sol", WebSearchCalls: 1}
cost, err := svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 3.0, 1.0, 1.0, 2.0, UsageTokens{}, "", false)
cost, err := svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 3.0, 1.0, 1.0, 2.0, UsageTokens{}, "", boolPtr(false))
require.NoError(t, err)
require.Equal(t, string(BillingModePerRequest), cost.BillingMode)
require.InDelta(t, 0.01, cost.TotalCost, 1e-12)
@@ -58,7 +58,7 @@ func TestCalculateOpenAIRecordUsageCostWebSearchPerCall(t *testing.T) {
// 分组配置单价 0.005
apiKey.Group.WebSearchPricePerCall = float64Ptr(0.005)
cost, err = svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 1.0, 1.0, 1.0, 1.0, UsageTokens{}, "", false)
cost, err = svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 1.0, 1.0, 1.0, 1.0, UsageTokens{}, "", boolPtr(false))
require.NoError(t, err)
require.InDelta(t, 0.005, cost.TotalCost, 1e-12)
require.InDelta(t, 0.005, cost.ActualCost, 1e-12)
@@ -66,7 +66,7 @@ func TestCalculateOpenAIRecordUsageCostWebSearchPerCall(t *testing.T) {
// WebSearchCalls = 0 时不得走按次分支(无定价数据会返回 pricing 错误,
// 证明回落到了 token 路径而不是被按次分支吞掉)。
result.WebSearchCalls = 0
_, err = svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 1.0, 1.0, 1.0, 1.0, UsageTokens{InputTokens: 10}, "", false)
_, err = svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 1.0, 1.0, 1.0, 1.0, UsageTokens{InputTokens: 10}, "", boolPtr(false))
require.Error(t, err)
}
@@ -1210,6 +1210,58 @@ func TestOpenAIGatewayServiceRecordUsage_GroupAndAccountLongContextMustBothAllow
})
}
// openai_long_context_billing_enabled is an OpenAI-only account setting, so it
// must not veto the official Grok >=200k ladder: a Grok account has no way to
// ever set that flag, which would make the group toggle unreachable.
func TestOpenAIGatewayServiceRecordUsage_GrokLongContextFollowsGroupToggleOnly(t *testing.T) {
baseInput := 250000 * 2e-6
baseOutput := 1000 * 6e-6
grokAccount := func(id int64) *Account {
return &Account{ID: id, Platform: PlatformGrok, Type: AccountTypeOAuth}
}
t.Run("group on applies the official ladder", func(t *testing.T) {
usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
svc := newOpenAIRecordUsageServiceForTest(usageRepo, &openAIRecordUsageUserRepoStub{}, &openAIRecordUsageSubRepoStub{}, nil)
err := svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{
Result: &OpenAIForwardResult{
RequestID: "resp_grok_longctx_on",
Usage: OpenAIUsage{InputTokens: 250000, OutputTokens: 1000},
Model: "grok-4.5",
Duration: time.Second,
},
APIKey: openAIRecordUsageAPIKeyWithGroup(svc, 1030, true),
User: &User{ID: 2030},
Account: grokAccount(3030),
})
require.NoError(t, err)
require.True(t, usageRepo.lastLog.LongContextBillingApplied)
require.InDelta(t, baseInput*2, usageRepo.lastLog.InputCost, 1e-10)
require.InDelta(t, baseOutput*2, usageRepo.lastLog.OutputCost, 1e-10)
})
t.Run("group off keeps the base card", func(t *testing.T) {
usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
svc := newOpenAIRecordUsageServiceForTest(usageRepo, &openAIRecordUsageUserRepoStub{}, &openAIRecordUsageSubRepoStub{}, nil)
err := svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{
Result: &OpenAIForwardResult{
RequestID: "resp_grok_longctx_off",
Usage: OpenAIUsage{InputTokens: 250000, OutputTokens: 1000},
Model: "grok-4.5",
Duration: time.Second,
},
APIKey: openAIRecordUsageAPIKeyWithGroup(svc, 1031, false),
User: &User{ID: 2031},
Account: grokAccount(3031),
})
require.NoError(t, err)
require.False(t, usageRepo.lastLog.LongContextBillingApplied)
require.InDelta(t, baseInput, usageRepo.lastLog.InputCost, 1e-10)
require.InDelta(t, baseOutput, usageRepo.lastLog.OutputCost, 1e-10)
})
}
func TestOpenAIGatewayServiceRecordUsage_SparkShadowUsesCurrentParentBillingSetting(t *testing.T) {
tests := []struct {
name string
@@ -36,7 +36,7 @@ func TestCalculateOpenAIRecordUsageCost_SearchIsAdditiveToTokens(t *testing.T) {
1.0,
UsageTokens{InputTokens: 1000, OutputTokens: 500},
"",
false,
boolPtr(false),
)
require.NoError(t, err)
require.NotNil(t, cost)
@@ -66,7 +66,7 @@ func TestCalculateOpenAIRecordUsageCost_SearchOnlyWhenNoTokenPricing(t *testing.
1.0,
UsageTokens{},
"",
false,
boolPtr(false),
)
require.NoError(t, err)
require.NotNil(t, cost)
@@ -111,7 +111,7 @@ func TestCalculateOpenAIRecordUsageCost_TokenPricingErrorNotSwallowedBySearch(t
1.0,
UsageTokens{InputTokens: 1000, OutputTokens: 500},
"",
false,
boolPtr(false),
)
require.Error(t, err)
require.Nil(t, cost)
@@ -211,7 +211,7 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec
return err
}
}
longContextBillingEnabled := billingAccount.IsOpenAILongContextBillingEnabled()
longContextBillingGate := openAILongContextBillingGate(billingAccount)
cost, err = s.calculateOpenAIRecordUsageCost(
ctx,
result,
@@ -223,7 +223,7 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec
baseMultiplier,
tokens,
serviceTier,
longContextBillingEnabled,
longContextBillingGate,
)
if err != nil {
if !isUsagePricingUnavailableError(err) {
@@ -256,7 +256,7 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec
responseModels := usageBillingModelCandidates(responseModel)
responseCost, responseErr := s.calculateOpenAIRecordUsageCost(
ctx, result, apiKey, responseModels, multiplier, imageMultiplier,
videoMultiplier, baseMultiplier, tokens, serviceTier, longContextBillingEnabled,
videoMultiplier, baseMultiplier, tokens, serviceTier, longContextBillingGate,
)
// 基线定价源以 baselineBillingModel 为准:它正是 calculateOpenAIRecordUsageCost
// 内部做渠道定价判断时使用的模型,且"首候选有渠道价"必然意味着首候选就是实际
@@ -480,6 +480,19 @@ func (s *OpenAIGatewayService) hasIdentifiedOpenAIResponsePricing(ctx context.Co
return s.billingService.HasIdentifiedTokenPricing(model), false
}
// openAILongContextBillingGate returns the per-account long-context opt-in.
// The flag is an OpenAI-only account setting, so other platforms (Grok) return
// nil — "no per-account gate" — and are governed by the group toggle alone.
// Returning a hardcoded false for them would veto the official model ladders
// (e.g. the Grok >=200k 2x card) that no account setting can ever re-enable.
func openAILongContextBillingGate(account *Account) *bool {
if account == nil || !account.IsOpenAI() {
return nil
}
enabled := account.IsOpenAILongContextBillingEnabled()
return &enabled
}
func (s *OpenAIGatewayService) calculateOpenAIRecordUsageCost(
ctx context.Context,
result *OpenAIForwardResult,
@@ -491,7 +504,7 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageCost(
webSearchMultiplier float64,
tokens UsageTokens,
serviceTier string,
longContextBillingEnabled bool,
longContextBillingGate *bool,
) (*CostBreakdown, error) {
billingModel := firstUsageBillingModel(billingModels)
if result != nil && result.WebSearchCalls > 0 {
@@ -543,7 +556,7 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageCost(
multiplier,
tokens,
serviceTier,
longContextBillingEnabled,
longContextBillingGate,
)
if err == nil {
tokenCost = cost
@@ -632,7 +645,7 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageTokenCost(
multiplier float64,
tokens UsageTokens,
serviceTier string,
longContextBillingEnabled bool,
longContextBillingGate *bool,
) (*CostBreakdown, error) {
if s.resolver != nil && apiKey.Group != nil {
gid := apiKey.Group.ID
@@ -640,7 +653,7 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageTokenCost(
Ctx: ctx, Model: billingModel, GroupID: &gid, Group: apiKey.Group,
Tokens: tokens, RequestCount: 1, RateMultiplier: multiplier,
ServiceTier: serviceTier, Resolver: s.resolver,
LongContextBillingEnabled: &longContextBillingEnabled,
LongContextBillingEnabled: longContextBillingGate,
})
}
return s.billingService.calculateCostWithServiceTierPolicy(
@@ -648,7 +661,7 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageTokenCost(
tokens,
multiplier,
serviceTier,
longContextBillingEnabled,
longContextBillingGate == nil || *longContextBillingGate,
)
}
@@ -372,6 +372,12 @@ func TestBillingServiceHasIdentifiedTokenPricing_RejectsFamilyGuesses(t *testing
}
require.False(t, billing.HasIdentifiedTokenPricing(""))
require.False(t, billing.HasIdentifiedTokenPricing("zz-unpriced-response-model"))
// Versioned media ids may inherit a text card via GetModelPricing; the
// identified-token gate must still reject them so response-model billing
// cannot adopt grok-4.5 rates for image/audio/video ids.
require.False(t, billing.HasIdentifiedTokenPricing("grok-2-image-1212"))
require.False(t, billing.HasIdentifiedTokenPricing("grok-2-audio"))
require.False(t, billing.HasIdentifiedTokenPricing("grok-5-video"))
}
func TestGatewayServiceRecordUsage_ResponseModelRejectsUnidentifiedFamilyName(t *testing.T) {
@@ -2440,7 +2440,6 @@ const handleClickOutside = (event: MouseEvent) => {
onMounted(async () => {
if (typeof window !== 'undefined') {
loadSavedAutoRefresh()
desktopViewportMediaQuery = window.matchMedia(desktopViewportQuery)
isDesktopViewport.value = desktopViewportMediaQuery.matches
desktopViewportListener = (event: MediaQueryListEvent) => {