fix grok chat usage guard parity

This commit is contained in:
Rain
2026-08-07 11:36:34 +08:00
parent ba92d70422
commit 8ea68bd689
6 changed files with 157 additions and 31 deletions
@@ -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)
@@ -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 {
@@ -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)
@@ -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())
}
@@ -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())
}
@@ -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,
}
}