diff --git a/backend/internal/service/openai_gateway_chat_completions.go b/backend/internal/service/openai_gateway_chat_completions.go index ddf9a288a8..d2533e486d 100644 --- a/backend/internal/service/openai_gateway_chat_completions.go +++ b/backend/internal/service/openai_gateway_chat_completions.go @@ -460,20 +460,14 @@ func (s *OpenAIGatewayService) handleChatBufferedStreamingResponse( return nil, fmt.Errorf("upstream response failed: %s", message) } + if requiresBillableGrokChatUsage(account, billingModel, upstreamModel, finalResponse.Model) && !hasBillableGrokChatUsage(usage) { + upstreamRequestID := firstNonEmpty(requestID, resp.Header.Get("xai-request-id")) + return nil, newGrokMissingUsageFailoverError(c, account, upstreamRequestID) + } + // When the terminal event has an empty output array, reconstruct from // accumulated delta events so the client receives the full content. acc.SupplementResponseOutput(finalResponse) - if requiresBillableGrokChatUsage(account, originalModel, billingModel, upstreamModel, finalResponse.Model) && !hasBillableOpenAIUsage(usage) { - return nil, s.newOpenAIStreamFailoverError( - c, - account, - false, - firstNonEmpty(requestID, resp.Header.Get("xai-request-id")), - nil, - grokMissingUsageMessage, - resp.Header, - ) - } chatResp := apicompat.ResponsesToChatCompletions(finalResponse, originalModel) diff --git a/backend/internal/service/openai_gateway_chat_completions_raw.go b/backend/internal/service/openai_gateway_chat_completions_raw.go index cdafc30d1b..81ef3ac39a 100644 --- a/backend/internal/service/openai_gateway_chat_completions_raw.go +++ b/backend/internal/service/openai_gateway_chat_completions_raw.go @@ -432,16 +432,9 @@ func (s *OpenAIGatewayService) bufferRawChatCompletions( usage = parsedUsage } responseModel := gjson.GetBytes(respBody, "model").String() - if requiresBillableGrokChatUsage(account, originalModel, billingModel, upstreamModel, responseModel) && !hasBillableOpenAIUsage(usage) { - return nil, s.newOpenAIStreamFailoverError( - c, - account, - false, - firstNonEmpty(requestID, resp.Header.Get("xai-request-id")), - nil, - grokMissingUsageMessage, - resp.Header, - ) + if requiresBillableGrokChatUsage(account, billingModel, upstreamModel, responseModel) && !hasBillableGrokChatUsage(usage) { + upstreamRequestID := firstNonEmpty(requestID, resp.Header.Get("xai-request-id")) + return nil, newGrokMissingUsageFailoverError(c, account, upstreamRequestID) } if s.responseHeaderFilter != nil { diff --git a/backend/internal/service/openai_gateway_chat_completions_raw_test.go b/backend/internal/service/openai_gateway_chat_completions_raw_test.go index 0a9741889c..a80234913a 100644 --- a/backend/internal/service/openai_gateway_chat_completions_raw_test.go +++ b/backend/internal/service/openai_gateway_chat_completions_raw_test.go @@ -158,6 +158,92 @@ func TestForwardAsChatCompletions_OpenAICompatibleGrokRawMissingUsageFailsBefore require.Empty(t, recorder.Body.String()) } +func TestForwardAsChatCompletions_OpenAICompatibleRawUsageGuard(t *testing.T) { + tests := []struct { + name string + model string + upstreamResponse string + modelMapping map[string]any + wantGuarded bool + }{ + { + name: "Grok response without usage", + model: "grok-4.5", + upstreamResponse: `{"id":"resp_missing","object":"chat.completion","model":"grok-4.5","choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}]}`, + wantGuarded: true, + }, + { + name: "namespaced Grok response without usage", + model: "x-ai/grok-4.5", + upstreamResponse: `{"id":"resp_namespaced","object":"chat.completion","model":"x-ai/grok-4.5","choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}]}`, + wantGuarded: true, + }, + { + name: "Grok response with aggregate usage passes", + model: "grok-4.5", + upstreamResponse: `{"id":"resp_usage","object":"chat.completion","model":"grok-4.5","choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}],"usage":{"prompt_tokens":9,"completion_tokens":3,"total_tokens":12}}`, + wantGuarded: false, + }, + { + name: "Grok alias mapped to non-Grok remains unchanged", + model: "grok-alias", + upstreamResponse: `{"id":"resp_mapped","object":"chat.completion","model":"gpt-5.4","choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}]}`, + modelMapping: map[string]any{"grok-alias": "gpt-5.4"}, + wantGuarded: false, + }, + { + name: "Grok response with detail-only usage", + model: "grok-4.5", + upstreamResponse: `{"id":"resp_detail_only","object":"chat.completion","model":"grok-4.5","choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}],"usage":{"input_tokens_details":{"text_tokens":9,"image_tokens":2},"output_tokens_details":{"image_tokens":1}}}`, + wantGuarded: true, + }, + { + name: "non-Grok response without usage remains unchanged", + model: "gpt-5.4", + upstreamResponse: `{"id":"resp_openai","object":"chat.completion","model":"gpt-5.4","choices":[{"index":0,"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}]}`, + wantGuarded: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + gin.SetMode(gin.TestMode) + body := []byte(`{"model":"` + tt.model + `","messages":[{"role":"user","content":"hello"}],"stream":false}`) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body)) + + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}, "X-Request-Id": []string{"rid-openai-compatible"}}, + Body: io.NopCloser(strings.NewReader(tt.upstreamResponse)), + }} + svc := &OpenAIGatewayService{cfg: rawChatCompletionsTestConfig(), httpUpstream: upstream} + account := rawChatCompletionsTestAccount() + account.Name = "openai-compatible" + account.Extra = map[string]any{openai_compat.ExtraKeyResponsesSupported: false} + if tt.modelMapping != nil { + account.Credentials["model_mapping"] = tt.modelMapping + } + + result, err := svc.ForwardAsChatCompletions(context.Background(), c, account, body, "", "") + + if !tt.wantGuarded { + require.NoError(t, err) + require.NotNil(t, result) + require.True(t, c.Writer.Written()) + return + } + require.Nil(t, result) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, http.StatusBadGateway, failoverErr.StatusCode) + require.Equal(t, "grok_missing_usage", gjson.GetBytes(failoverErr.ResponseBody, "error.code").String()) + require.False(t, c.Writer.Written(), "unbilled Grok content must not be returned") + }) + } +} + func TestForwardAsRawChatCompletions_PreservesMappedGPT56MaxEffort(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/service/openai_gateway_grok_chat_bridge_test.go b/backend/internal/service/openai_gateway_grok_chat_bridge_test.go index abb21906aa..6c4137d9d8 100644 --- a/backend/internal/service/openai_gateway_grok_chat_bridge_test.go +++ b/backend/internal/service/openai_gateway_grok_chat_bridge_test.go @@ -308,6 +308,8 @@ func TestForwardGrokChatViaResponsesNonStreamingRejectsCompletedResponseWithoutU require.Nil(t, result) var failoverErr *UpstreamFailoverError require.ErrorAs(t, err, &failoverErr) + require.Equal(t, http.StatusBadGateway, failoverErr.StatusCode) + require.Equal(t, grokMissingUsageErrorCode, gjson.GetBytes(failoverErr.ResponseBody, "error.code").String()) require.False(t, c.Writer.Written(), "an unbillable response must not be committed to the client") require.Empty(t, recorder.Body.String()) } diff --git a/backend/internal/service/openai_gateway_grok_test.go b/backend/internal/service/openai_gateway_grok_test.go index 24f87fed40..8bfeeb9edb 100644 --- a/backend/internal/service/openai_gateway_grok_test.go +++ b/backend/internal/service/openai_gateway_grok_test.go @@ -1968,6 +1968,8 @@ func TestForwardAsChatCompletionsForGrokAPIKeyRejectsNonStreamingResponseWithout require.Nil(t, result) var failoverErr *UpstreamFailoverError require.ErrorAs(t, err, &failoverErr) + require.Equal(t, http.StatusBadGateway, failoverErr.StatusCode) + require.Equal(t, grokMissingUsageErrorCode, gjson.GetBytes(failoverErr.ResponseBody, "error.code").String()) require.False(t, c.Writer.Written(), "an unbillable response must not be committed to the client") require.Empty(t, recorder.Body.String()) } diff --git a/backend/internal/service/openai_gateway_usage_integrity.go b/backend/internal/service/openai_gateway_usage_integrity.go index a70357cc6f..e0344d4f44 100644 --- a/backend/internal/service/openai_gateway_usage_integrity.go +++ b/backend/internal/service/openai_gateway_usage_integrity.go @@ -1,18 +1,26 @@ package service -const grokMissingUsageMessage = "Grok upstream returned a successful response without billable usage" +import ( + "encoding/json" + "net/http" + "strings" -// hasBillableOpenAIUsage distinguishes a usable accounting result from the -// zero value produced when an upstream omits usage entirely. A successful Grok -// completion without any positive usage cannot be safely charged, so callers -// must fail over before committing the response to the client. -func hasBillableOpenAIUsage(usage OpenAIUsage) bool { + "github.com/gin-gonic/gin" +) + +const ( + grokMissingUsageErrorCode = "grok_missing_usage" + grokMissingUsageMessage = "xAI upstream returned a successful chat completion without billable usage" +) + +// hasBillableGrokChatUsage stays aligned with the aggregate token buckets used +// to account for chat completions. Detail fields alone do not prove that the +// successful response can be settled safely. +func hasBillableGrokChatUsage(usage OpenAIUsage) bool { return usage.InputTokens > 0 || - usage.ImageInputTokens > 0 || usage.OutputTokens > 0 || usage.CacheCreationInputTokens > 0 || - usage.CacheReadInputTokens > 0 || - usage.ImageOutputTokens > 0 + usage.CacheReadInputTokens > 0 } // requiresBillableGrokChatUsage identifies Grok traffic by both account @@ -24,9 +32,50 @@ func requiresBillableGrokChatUsage(account *Account, models ...string) bool { return true } for _, model := range models { - if platform, ok := DetectModelPlatform(model); ok && platform == PlatformGrok { + normalized := strings.ToLower(strings.TrimSpace(model)) + if separator := strings.LastIndex(normalized, "/"); separator >= 0 { + normalized = strings.TrimSpace(normalized[separator+1:]) + } + if normalized == "grok" || strings.HasPrefix(normalized, "grok-") { return true } } return false } + +func newGrokMissingUsageFailoverError(c *gin.Context, account *Account, upstreamRequestID string) *UpstreamFailoverError { + accountID := int64(0) + accountName := "" + if account != nil { + accountID = account.ID + accountName = account.Name + } + + setOpsUpstreamError(c, http.StatusBadGateway, grokMissingUsageMessage, "") + appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ + Platform: PlatformGrok, + AccountID: accountID, + AccountName: accountName, + UpstreamStatusCode: http.StatusBadGateway, + UpstreamRequestID: strings.TrimSpace(upstreamRequestID), + Kind: "failover", + Message: grokMissingUsageMessage, + }) + + body, _ := json.Marshal(gin.H{ + "error": gin.H{ + "type": "upstream_error", + "code": grokMissingUsageErrorCode, + "message": grokMissingUsageMessage, + }, + }) + headers := http.Header{} + if requestID := strings.TrimSpace(upstreamRequestID); requestID != "" { + headers.Set("x-request-id", requestID) + } + return &UpstreamFailoverError{ + StatusCode: http.StatusBadGateway, + ResponseBody: body, + ResponseHeaders: headers, + } +}