mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 13:28:39 +08:00
fix grok chat usage guard parity
This commit is contained in:
@@ -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,
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user