Merge pull request #4719 from superman2003/fix/grok-403-model-sync-4713-4715

fix(grok): sync OAuth models and isolate policy 403s
This commit is contained in:
Wesley Liddick
2026-07-22 14:26:07 +08:00
committed by GitHub
12 changed files with 864 additions and 40 deletions
+17 -1
View File
@@ -850,6 +850,22 @@ func (s *OpenAIGatewayService) handleGrokMediaErrorResponse(
upstreamDetail = truncateString(string(body), maxBytes)
}
setOpsUpstreamError(c, resp.StatusCode, upstreamMsg, upstreamDetail)
if isGrokContentPolicyRejection(resp.StatusCode, body) {
clientMsg := grokContentPolicyClientMessage(body)
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
Platform: account.Platform,
AccountID: account.ID,
AccountName: account.Name,
UpstreamStatusCode: resp.StatusCode,
UpstreamRequestID: requestIDHeader,
Kind: "http_error",
Message: clientMsg,
Detail: upstreamDetail,
})
MarkResponseCommitted(c)
writeGrokMediaErrorResponse(c, http.StatusForbidden, "invalid_request_error", clientMsg)
return nil, fmt.Errorf("grok content policy rejection: %s", clientMsg)
}
if status, errType, errMsg, matched := applyErrorPassthroughRule(
c,
@@ -882,7 +898,7 @@ func (s *OpenAIGatewayService) handleGrokMediaErrorResponse(
}
kind := "http_error"
if s.shouldFailoverUpstreamError(resp.StatusCode) {
if s.shouldFailoverGrokUpstreamError(resp.StatusCode, body) {
kind = "failover"
}
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
@@ -0,0 +1,238 @@
package service
import (
"context"
"encoding/json"
"net/http"
"strings"
"time"
)
// isGrokContentPolicyRejection identifies request-scoped safety refusals from
// xAI. These failures are caused by the prompt or media, so retrying another
// OAuth account cannot change the outcome and would incorrectly drain a pool.
// Keep this matcher deliberately narrow: account entitlement and suspension
// messages may mention policy but must retain the normal account failover path.
func isGrokContentPolicyRejection(statusCode int, responseBody []byte) bool {
if statusCode != http.StatusForbidden || len(responseBody) == 0 {
return false
}
if grokAccountAccessMessage(string(responseBody)) {
return false
}
var payload any
if json.Unmarshal(responseBody, &payload) == nil {
if grokStructuredAccountAccessMarker(payload) {
return false
}
if grokStructuredContentPolicyMarker(payload) {
return true
}
}
return grokContentPolicyMessage(string(responseBody))
}
func grokStructuredAccountAccessMarker(value any) bool {
switch node := value.(type) {
case map[string]any:
for key, child := range node {
normalizedKey := normalizeGrokErrorMarker(key)
switch normalizedKey {
case "code", "error_code", "type", "category", "reason":
if marker, ok := child.(string); ok && isGrokAccountAccessCode(marker) {
return true
}
}
if grokStructuredAccountAccessMarker(child) {
return true
}
}
case []any:
for _, child := range node {
if grokStructuredAccountAccessMarker(child) {
return true
}
}
}
return false
}
func grokStructuredContentPolicyMarker(value any) bool {
switch node := value.(type) {
case map[string]any:
for key, child := range node {
normalizedKey := normalizeGrokErrorMarker(key)
switch normalizedKey {
case "code", "error_code", "type", "category", "reason":
if marker, ok := child.(string); ok && isGrokContentPolicyCode(marker) {
return true
}
}
if grokStructuredContentPolicyMarker(child) {
return true
}
}
case []any:
for _, child := range node {
if grokStructuredContentPolicyMarker(child) {
return true
}
}
}
return false
}
func normalizeGrokErrorMarker(value string) string {
value = strings.ToLower(strings.TrimSpace(value))
value = strings.ReplaceAll(value, "-", "_")
value = strings.ReplaceAll(value, " ", "_")
return value
}
func isGrokContentPolicyCode(value string) bool {
switch normalizeGrokErrorMarker(value) {
case "content_filter",
"content_policy",
"content_policy_violation",
"content_moderation",
"cyber_policy",
"new_sensitive":
return true
default:
return false
}
}
func isGrokAccountAccessCode(value string) bool {
switch normalizeGrokErrorMarker(value) {
case "account_suspended",
"account_disabled",
"user_suspended",
"user_disabled",
"subscription_required",
"entitlement_required",
"not_entitled",
"plan_required",
"permission_denied":
return true
default:
return false
}
}
func grokAccountAccessMessage(value string) bool {
lower := strings.ToLower(strings.TrimSpace(value))
for _, phrase := range []string{
"account suspended",
"account has been suspended",
"account disabled",
"account has been disabled",
"user suspended",
"user has been suspended",
"subscription required",
"entitlement required",
"not entitled",
} {
if strings.Contains(lower, phrase) {
return true
}
}
return false
}
func grokContentPolicyMessage(value string) bool {
lower := strings.ToLower(strings.TrimSpace(value))
if lower == "" {
return false
}
// xAI's media safety responses use these exact phrases. They are specific
// enough not to classify a generic account-policy or entitlement message.
for _, phrase := range []string{
"the moderation feature is not available",
"image is sensitive",
"text is sensitive",
"prohibited content",
"forbidden content",
"content policy violation",
"content policy rejection",
"content policy rejected",
"content moderation rejection",
"content moderation rejected",
"content moderation blocked",
"request blocked by content moderation",
"request rejected by content moderation",
"request blocked by policy",
"request rejected by policy",
"request violates policy",
"prompt violates content policy",
"prompt violates policy",
"input violates content policy",
"input violates policy",
} {
if strings.Contains(lower, phrase) {
return true
}
}
return false
}
func grokContentPolicyClientMessage(responseBody []byte) string {
message := sanitizeUpstreamErrorMessage(strings.TrimSpace(extractUpstreamErrorMessage(responseBody)))
if message == "" {
return "Request blocked by upstream content policy"
}
return message
}
// shouldFailoverGrokUpstreamError is the body-aware counterpart of the
// status-only failover helper. Grok content refusals must stay on the current
// account and be returned to the caller instead of consuming the account pool.
func (s *OpenAIGatewayService) shouldFailoverGrokUpstreamError(statusCode int, responseBody []byte) bool {
if isGrokContentPolicyRejection(statusCode, responseBody) {
return false
}
return s.shouldFailoverUpstreamError(statusCode)
}
// applyGrokForbiddenPolicy applies an administrator's existing temporary
// unschedulable rules to a non-content 403. It reports true only when a rule
// matched; unmatched responses retain the legacy entitlement cooldown.
func (s *OpenAIGatewayService) applyGrokForbiddenPolicy(ctx context.Context, account *Account, responseBody []byte) bool {
if account == nil || !account.IsTempUnschedulableEnabled() {
return false
}
matches := matchTempUnschedulableRules(account, http.StatusForbidden, responseBody)
if len(matches) == 0 {
return false
}
match := matches[0]
// Reuse the central policy implementation when it has a repository. This
// preserves the existing reason/cache format and avoids duplicating writes.
if s != nil && s.rateLimitService != nil && s.rateLimitService.accountRepo != nil {
stateCtx, cancel := openAIAccountStateContext(ctx)
handled := s.rateLimitService.tryTempUnschedulable(
stateCtx,
account,
http.StatusForbidden,
responseBody,
)
cancel()
if handled {
return true
}
}
// A partially constructed service (for example a unit-test gateway) still
// honors the configured duration instead of silently falling back to 30m.
cooldown := time.Duration(match.rule.DurationMinutes) * time.Minute
if cooldown > 0 {
s.tempUnscheduleGrok(ctx, account, cooldown, "grok configured forbidden rule")
}
return true
}
@@ -0,0 +1,319 @@
//go:build unit
package service
import (
"context"
"errors"
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
func TestIsGrokContentPolicyRejection(t *testing.T) {
tests := []struct {
name string
status int
body string
want bool
}{
{
name: "new sensitive code",
status: http.StatusForbidden,
body: `{"error":{"code":"new_sensitive","message":"image is sensitive"}}`,
want: true,
},
{
name: "content policy violation code",
status: http.StatusForbidden,
body: `{"response":{"error":{"code":"content_policy_violation"}}}`,
want: true,
},
{
name: "cyber policy code",
status: http.StatusForbidden,
body: `{"error":{"code":"cyber_policy","message":"request rejected"}}`,
want: true,
},
{
name: "moderation feature unavailable",
status: http.StatusForbidden,
body: `{"error":{"message":"The moderation feature is not available for this request"}}`,
want: true,
},
{
name: "explicit prompt moderation rejection",
status: http.StatusForbidden,
body: `{"error":{"message":"request rejected by content moderation"}}`,
want: true,
},
{
name: "entitlement forbidden",
status: http.StatusForbidden,
body: `{"error":{"message":"subscription required"}}`,
want: false,
},
{
name: "account policy suspension is not request policy",
status: http.StatusForbidden,
body: `{"error":{"message":"account suspended due to policy violation"}}`,
want: false,
},
{
name: "structured account suspension overrides policy reason",
status: http.StatusForbidden,
body: `{"error":{"code":"account_suspended","reason":"policy_violation","message":"account suspended due to policy violation"}}`,
want: false,
},
{
name: "ambiguous policy violation code is not enough",
status: http.StatusForbidden,
body: `{"error":{"code":"policy_violation","message":"policy violation"}}`,
want: false,
},
{
name: "policy violation with request scoped message",
status: http.StatusForbidden,
body: `{"error":{"code":"policy_violation","message":"request blocked by policy"}}`,
want: true,
},
{
name: "wrong status",
status: http.StatusBadRequest,
body: `{"error":{"code":"new_sensitive"}}`,
want: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
require.Equal(t, tt.want, isGrokContentPolicyRejection(tt.status, []byte(tt.body)))
})
}
}
func TestGrokContentPolicy403DoesNotMutateOrFailover(t *testing.T) {
repo := &grokQuotaAccountRepo{}
svc := &OpenAIGatewayService{accountRepo: repo}
account := &Account{ID: 4715, Platform: PlatformGrok, Type: AccountTypeOAuth}
body := []byte(`{"error":{"code":"new_sensitive","message":"text is sensitive"}}`)
svc.handleGrokAccountUpstreamError(context.Background(), account, http.StatusForbidden, nil, body)
require.Zero(t, repo.tempUnschedCalls)
require.Zero(t, repo.rateLimitedCalls)
require.Zero(t, repo.updateCalls)
require.False(t, svc.isOpenAIAccountRuntimeBlocked(account))
require.False(t, svc.shouldFailoverGrokUpstreamError(http.StatusForbidden, body))
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
resp := &http.Response{StatusCode: http.StatusForbidden, Header: http.Header{}}
got := svc.failoverOpenAIUpstreamHTTPError(context.Background(), c, account, resp, body, "text is sensitive", "grok-4.5")
require.Nil(t, got)
require.Zero(t, repo.tempUnschedCalls)
}
func TestGrokContentPolicy403SharedErrorFallbackDoesNotMutate(t *testing.T) {
gin.SetMode(gin.TestMode)
body := []byte(`{"error":{"code":"content_filter","message":"prohibited content"}}`)
repo := &grokQuotaAccountRepo{}
svc := &OpenAIGatewayService{accountRepo: repo}
account := &Account{
ID: 4719,
Platform: PlatformGrok,
Type: AccountTypeOAuth,
Credentials: map[string]any{
"custom_error_codes_enabled": true,
"custom_error_codes": []any{float64(http.StatusTooManyRequests)},
},
}
newContext := func() (*gin.Context, *httptest.ResponseRecorder) {
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil)
return c, recorder
}
c, recorder := newContext()
resp := &http.Response{
StatusCode: http.StatusForbidden,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(string(body))),
}
_, err := svc.handleErrorResponse(context.Background(), resp, c, account, nil, "grok-4.5")
require.Error(t, err)
require.Equal(t, http.StatusForbidden, recorder.Code)
require.Contains(t, recorder.Body.String(), "invalid_request_error")
c, recorder = newContext()
resp = &http.Response{
StatusCode: http.StatusForbidden,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(string(body))),
}
_, err = svc.handleCompatErrorResponse(resp, c, account, writeChatCompletionsError, "grok-4.5")
require.Error(t, err)
require.Equal(t, http.StatusForbidden, recorder.Code)
require.Contains(t, recorder.Body.String(), "invalid_request_error")
require.Zero(t, repo.tempUnschedCalls)
require.Zero(t, repo.rateLimitedCalls)
require.Zero(t, repo.updateCalls)
}
func TestGrokContentPolicy403MediaResponseBypassesCustomErrorCodes(t *testing.T) {
gin.SetMode(gin.TestMode)
body := `{"error":{"code":"new_sensitive","message":"image is sensitive"}}`
repo := &grokQuotaAccountRepo{}
svc := &OpenAIGatewayService{accountRepo: repo}
account := &Account{
ID: 4720,
Platform: PlatformGrok,
Type: AccountTypeOAuth,
Credentials: map[string]any{
"custom_error_codes_enabled": true,
"custom_error_codes": []any{float64(http.StatusTooManyRequests)},
},
}
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/images/generations", nil)
resp := &http.Response{
StatusCode: http.StatusForbidden,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(body)),
}
_, err := svc.handleGrokMediaErrorResponse(context.Background(), resp, c, account, "request-id", "grok-imagine")
require.Error(t, err)
require.Equal(t, http.StatusForbidden, recorder.Code)
require.Contains(t, recorder.Body.String(), "invalid_request_error")
require.Zero(t, repo.tempUnschedCalls)
require.Zero(t, repo.rateLimitedCalls)
require.Zero(t, repo.updateCalls)
}
func TestGrokContentPolicySSEErrorDoesNotMutateOrFailover(t *testing.T) {
gin.SetMode(gin.TestMode)
repo := &grokQuotaAccountRepo{}
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"text/event-stream"}},
Body: io.NopCloser(strings.NewReader(
"data: {\"type\":\"error\",\"error\":{\"code\":\"new_sensitive\",\"message\":\"text is sensitive\"}}\n\n",
)),
}}
svc := &OpenAIGatewayService{accountRepo: repo, httpUpstream: upstream}
account := &Account{ID: 4721, Platform: PlatformGrok, Type: AccountTypeOAuth, Concurrency: 1}
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil)
payload := []byte(`{"type":"response.create","model":"grok-4.5","input":"hi"}`)
var writes [][]byte
result, err := svc.proxyOpenAIWSHTTPBridgeTurn(
context.Background(), c, account, "access-token", payload, len(payload),
"grok-4.5", "", "", "", "cache-id", 1,
func(message []byte) error {
writes = append(writes, append([]byte(nil), message...))
return nil
},
)
require.Error(t, err)
require.NotNil(t, result)
var failoverErr *UpstreamFailoverError
require.False(t, errors.As(err, &failoverErr))
require.Len(t, writes, 1)
require.Contains(t, string(writes[0]), "new_sensitive")
require.Zero(t, repo.tempUnschedCalls)
require.Zero(t, repo.rateLimitedCalls)
require.Zero(t, repo.updateCalls)
require.False(t, svc.isOpenAIAccountRuntimeBlocked(account))
}
func TestHandleGrokAccountUpstreamErrorEntitlement403KeepsDefaultCooldown(t *testing.T) {
repo := &grokQuotaAccountRepo{}
svc := &OpenAIGatewayService{accountRepo: repo}
account := &Account{ID: 4716, Platform: PlatformGrok, Type: AccountTypeOAuth}
before := time.Now()
svc.handleGrokAccountUpstreamError(
context.Background(), account, http.StatusForbidden, nil,
[]byte(`{"error":{"message":"subscription required"}}`),
)
require.Equal(t, 1, repo.tempUnschedCalls)
require.Equal(t, "grok access or entitlement denied", repo.lastTempUnschedReason)
require.Greater(t, repo.lastTempUnschedUntil, before.Add(29*time.Minute))
require.Less(t, repo.lastTempUnschedUntil, before.Add(31*time.Minute))
}
func TestHandleGrokAccountUpstreamError403UsesConfiguredRule(t *testing.T) {
repo := &grokQuotaAccountRepo{}
svc := &OpenAIGatewayService{accountRepo: repo}
account := &Account{
ID: 4717,
Platform: PlatformGrok,
Type: AccountTypeOAuth,
Credentials: map[string]any{
"temp_unschedulable_enabled": true,
"temp_unschedulable_rules": []any{
map[string]any{
"error_code": float64(http.StatusForbidden),
"keywords": []any{"subscription"},
"duration_minutes": float64(7),
},
},
},
}
before := time.Now()
svc.handleGrokAccountUpstreamError(
context.Background(), account, http.StatusForbidden, nil,
[]byte(`{"error":{"message":"subscription required"}}`),
)
require.Equal(t, 1, repo.tempUnschedCalls)
require.Greater(t, repo.lastTempUnschedUntil, before.Add(6*time.Minute))
require.Less(t, repo.lastTempUnschedUntil, before.Add(8*time.Minute))
}
func TestHandleGrokAccountUpstreamError403ConfiguredUnmatchedKeepsDefaultCooldown(t *testing.T) {
repo := &grokQuotaAccountRepo{}
svc := &OpenAIGatewayService{accountRepo: repo}
account := &Account{
ID: 4718,
Platform: PlatformGrok,
Type: AccountTypeOAuth,
Credentials: map[string]any{
"temp_unschedulable_enabled": true,
"temp_unschedulable_rules": []any{
map[string]any{
"error_code": float64(http.StatusForbidden),
"keywords": []any{"different failure"},
"duration_minutes": float64(7),
},
},
},
}
svc.handleGrokAccountUpstreamError(
context.Background(), account, http.StatusForbidden, nil,
[]byte(`{"error":{"message":"subscription required"}}`),
)
require.Equal(t, 1, repo.tempUnschedCalls)
require.Equal(t, "grok access or entitlement denied", repo.lastTempUnschedReason)
require.True(t, svc.isOpenAIAccountRuntimeBlocked(account))
}
@@ -48,6 +48,9 @@ func isOpenAIAccount(account *Account) bool {
// handleOpenAIAccountUpstreamError expects canonicalModel to be the model used
// for scheduling after applying account mapping exactly once.
func (s *OpenAIGatewayService) handleOpenAIAccountUpstreamError(ctx context.Context, account *Account, statusCode int, headers http.Header, responseBody []byte, canonicalModel ...string) bool {
if account != nil && account.Platform == PlatformGrok && isGrokContentPolicyRejection(statusCode, responseBody) {
return false
}
stateCtx, cancel := openAIAccountStateContext(ctx)
defer cancel()
@@ -88,10 +88,14 @@ func (s *OpenAIGatewayService) failoverOpenAIUpstreamHTTPError(
upstreamMsg string,
upstreamModel string,
) *UpstreamFailoverError {
shouldFailover := s.shouldFailoverOpenAIUpstreamResponse(resp.StatusCode, upstreamMsg, respBody)
if account != nil && account.Platform == PlatformGrok {
shouldFailover = s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody)
}
if account != nil && account.Platform == PlatformGrok {
s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
}
if !s.shouldFailoverOpenAIUpstreamResponse(resp.StatusCode, upstreamMsg, respBody) {
if !shouldFailover {
return nil
}
upstreamDetail := ""
@@ -174,17 +174,21 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions(
if resp.StatusCode >= 400 {
respBody, upstreamMsg := s.readOpenAIUpstreamError(resp)
if account.Platform == PlatformGrok {
kind := "http_error"
if s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) {
kind = "failover"
}
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
Platform: account.Platform,
AccountID: account.ID,
AccountName: account.Name,
UpstreamStatusCode: resp.StatusCode,
UpstreamRequestID: firstNonEmpty(resp.Header.Get("x-request-id"), resp.Header.Get("xai-request-id")),
Kind: "failover",
Kind: kind,
Message: upstreamMsg,
})
s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
if s.shouldFailoverUpstreamError(resp.StatusCode) {
if s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) {
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: respBody,
@@ -148,17 +148,21 @@ func (s *OpenAIGatewayService) forwardGrokResponses(
if upstreamMsg == "" {
upstreamMsg = fmt.Sprintf("xAI upstream returned status %d", resp.StatusCode)
}
kind := "http_error"
if s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) {
kind = "failover"
}
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
Platform: account.Platform,
AccountID: account.ID,
AccountName: account.Name,
UpstreamStatusCode: resp.StatusCode,
UpstreamRequestID: firstNonEmpty(resp.Header.Get("x-request-id"), resp.Header.Get("xai-request-id")),
Kind: "failover",
Kind: kind,
Message: upstreamMsg,
})
s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
if s.shouldFailoverUpstreamError(resp.StatusCode) {
if s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) {
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: respBody,
@@ -904,17 +908,21 @@ func (s *OpenAIGatewayService) describeGrokComposerImage(
if upstreamMsg == "" {
upstreamMsg = fmt.Sprintf("xAI image bridge upstream returned status %d", resp.StatusCode)
}
kind := "http_error"
if s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) {
kind = "failover"
}
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
Platform: account.Platform,
AccountID: account.ID,
AccountName: account.Name,
UpstreamStatusCode: resp.StatusCode,
UpstreamRequestID: firstNonEmpty(resp.Header.Get("x-request-id"), resp.Header.Get("xai-request-id")),
Kind: "failover",
Kind: kind,
Message: upstreamMsg,
})
s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
if s.shouldFailoverUpstreamError(resp.StatusCode) {
if s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) {
return "", OpenAIUsage{}, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: respBody,
@@ -1082,6 +1090,7 @@ func applyGrokCLIHeaders(headers http.Header) {
}
headers.Set("User-Agent", grokUpstreamUserAgent)
headers.Set("X-Grok-Client-Version", grokCLIVersion)
headers.Set("X-Grok-Client-Mode", "interactive")
}
func (s *OpenAIGatewayService) updateGrokUsageSnapshot(ctx context.Context, account *Account, snapshot *xai.QuotaSnapshot) {
@@ -1336,12 +1345,18 @@ func (s *OpenAIGatewayService) handleGrokAccountUpstreamError(ctx context.Contex
if s == nil || account == nil {
return
}
if isGrokContentPolicyRejection(statusCode, responseBody) {
return
}
now := time.Now()
s.updateGrokUsageSnapshot(ctx, account, parseGrokQuotaSnapshot(headers, statusCode, now))
switch statusCode {
case http.StatusUnauthorized:
s.tempUnscheduleGrok(ctx, account, 10*time.Minute, "grok credentials unauthorized")
case http.StatusForbidden:
if s.applyGrokForbiddenPolicy(ctx, account, responseBody) {
return
}
s.tempUnscheduleGrok(ctx, account, 30*time.Minute, "grok access or entitlement denied")
case http.StatusTooManyRequests:
// updateGrokUsageSnapshot installs both runtime and durable rate-limit state.
@@ -608,17 +608,21 @@ func (s *OpenAIGatewayService) forwardGrokChatCompletionsViaResponses(
if upstreamMsg == "" {
upstreamMsg = fmt.Sprintf("xAI upstream returned status %d", resp.StatusCode)
}
kind := "http_error"
if s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) {
kind = "failover"
}
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
Platform: account.Platform,
AccountID: account.ID,
AccountName: account.Name,
UpstreamStatusCode: resp.StatusCode,
UpstreamRequestID: firstNonEmpty(resp.Header.Get("x-request-id"), resp.Header.Get("xai-request-id")),
Kind: "failover",
Kind: kind,
Message: upstreamMsg,
})
s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
if s.shouldFailoverUpstreamError(resp.StatusCode) {
if s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) {
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: respBody,
@@ -345,6 +345,19 @@ func (s *OpenAIGatewayService) handleErrorResponse(
}
return nil, fmt.Errorf("openai cyber_policy: %s", cyberMsg)
}
if account != nil && account.Platform == PlatformGrok && isGrokContentPolicyRejection(resp.StatusCode, body) {
clientMsg := grokContentPolicyClientMessage(body)
setOpsUpstreamError(c, resp.StatusCode, clientMsg, truncateString(string(body), 2048))
writeOpenAIPassthroughResponseHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter)
MarkResponseCommitted(c)
c.JSON(http.StatusForbidden, gin.H{
"error": gin.H{
"type": "invalid_request_error",
"message": clientMsg,
},
})
return nil, fmt.Errorf("grok content policy rejection: %s", clientMsg)
}
upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(body))
upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg)
@@ -559,6 +572,13 @@ func (s *OpenAIGatewayService) handleCompatErrorResponse(
}
return nil, fmt.Errorf("openai cyber_policy: %s", cyberMsg)
}
if account != nil && account.Platform == PlatformGrok && isGrokContentPolicyRejection(resp.StatusCode, body) {
clientMsg := grokContentPolicyClientMessage(body)
setOpsUpstreamError(c, resp.StatusCode, clientMsg, truncateString(string(body), 2048))
MarkResponseCommitted(c)
writeError(c, http.StatusForbidden, "invalid_request_error", clientMsg)
return nil, fmt.Errorf("grok content policy rejection: %s", clientMsg)
}
upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(body))
if upstreamMsg == "" {
@@ -248,6 +248,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
}
shouldFailover := s.shouldFailoverOpenAIUpstreamResponse(resp.StatusCode, upstreamMsg, respBody)
if account.Platform == PlatformGrok {
shouldFailover = s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody)
s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
if turn == 1 && shouldFailover {
return nil, newOpenAIUpstreamFailoverError(resp.StatusCode, resp.Header, respBody, upstreamMsg, false)
@@ -393,7 +394,16 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
statusCode := openAIWSErrorHTTPStatusFromRaw(errCodeRaw, errTypeRaw)
shouldFailover := s.shouldFailoverOpenAIUpstreamResponse(statusCode, errMessage, upstreamMessage)
if account.Platform == PlatformGrok {
s.handleGrokAccountUpstreamError(ctx, account, statusCode, resp.Header, upstreamMessage)
// SSE error events do not carry an HTTP status. The local status
// mapper therefore defaults unknown xAI codes (for example
// new_sensitive) to 502; classify the body as a request-scoped
// 403 before applying status-based failover or account state.
if isGrokContentPolicyRejection(http.StatusForbidden, upstreamMessage) {
shouldFailover = false
} else {
shouldFailover = s.shouldFailoverGrokUpstreamError(statusCode, upstreamMessage)
s.handleGrokAccountUpstreamError(ctx, account, statusCode, resp.Header, upstreamMessage)
}
} else if shouldFailover {
accountStatus := statusCode
if transientStatus := openAIWSPayloadTransientStatus(upstreamMessage); transientStatus != 0 {
+125 -22
View File
@@ -116,7 +116,11 @@ func (s *AccountTestService) FetchUpstreamSupportedModels(ctx context.Context, a
)
}
models, err := extractUpstreamModelIDs(body)
extractModels := extractUpstreamModelIDs
if account.IsGrok() {
extractModels = extractGrokUpstreamModelIDs
}
models, err := extractModels(body)
if err != nil {
return nil, newUpstreamModelSyncUpstreamError("Upstream model list response was not valid JSON", err)
}
@@ -147,31 +151,79 @@ func (s *AccountTestService) buildUpstreamModelsRequest(ctx context.Context, acc
}
func (s *AccountTestService) buildGrokUpstreamModelsRequest(ctx context.Context, account *Account) (*http.Request, error) {
if account.Type != AccountTypeAPIKey {
if account == nil {
return nil, newUpstreamModelSyncConfigError("Account is required", nil)
}
var (
authToken string
normalizedBaseURL string
isOAuth = account.IsGrokOAuth()
)
switch account.Type {
case AccountTypeAPIKey:
authToken = strings.TrimSpace(account.GetCredential("api_key"))
if authToken == "" {
return nil, newUpstreamModelSyncConfigError("No Grok API key is available", nil)
}
baseURL := strings.TrimSpace(account.GetCredential("base_url"))
if baseURL == "" {
baseURL = "https://api.x.ai"
}
validatedBaseURL, err := s.validateUpstreamBaseURL(baseURL)
if err != nil {
return nil, newUpstreamModelSyncConfigError("Invalid Grok base URL", err)
}
normalizedBaseURL = validatedBaseURL
case AccountTypeOAuth:
if s.grokTokenProvider == nil {
return nil, newUpstreamModelSyncConfigError("Grok token provider is not configured", nil)
}
accessToken, err := s.grokTokenProvider.GetAccessTokenForManualTest(ctx, account)
if err != nil {
return nil, newUpstreamModelSyncUpstreamError("Failed to get Grok access token", err)
}
authToken = strings.TrimSpace(accessToken)
if authToken == "" {
return nil, newUpstreamModelSyncConfigError("No Grok access token is available", nil)
}
validator, err := grokBaseURLValidator(account, s.cfg)
if err != nil {
return nil, newUpstreamModelSyncConfigError("Invalid Grok base URL", err)
}
validatedBaseURL, err := validator(account.GetGrokBaseURL())
if err != nil {
return nil, newUpstreamModelSyncConfigError("Invalid Grok base URL", err)
}
normalizedBaseURL = validatedBaseURL
default:
return nil, newUpstreamModelSyncUnsupportedError(
fmt.Sprintf("Unsupported Grok account type for upstream model sync: %s", account.Type), nil,
)
}
apiKey := strings.TrimSpace(account.GetCredential("api_key"))
if apiKey == "" {
return nil, newUpstreamModelSyncConfigError("No Grok API key is available", nil)
}
baseURL := strings.TrimSpace(account.GetCredential("base_url"))
if baseURL == "" {
baseURL = "https://api.x.ai"
}
normalizedBaseURL, err := s.validateUpstreamBaseURL(baseURL)
if err != nil {
return nil, newUpstreamModelSyncConfigError("Invalid Grok base URL", err)
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, buildOpenAIModelsURL(normalizedBaseURL), nil)
if err != nil {
return nil, newUpstreamModelSyncConfigError("Invalid Grok model list URL", err)
}
req.Header.Set("Accept", "application/json")
req.Header.Set("Authorization", "Bearer "+apiKey)
req.Header.Set("Authorization", "Bearer "+authToken)
if isOAuth {
// The shared HTTP transport adds the official CLI marker/version for the
// exact proxy host. Keep the request builder aligned with the other Grok
// probes and only forward account identity headers to that trusted host.
applyGrokCLIHeaders(req.Header)
if isGrokCLIProxyTarget(req.URL.String()) {
if userID := strings.TrimSpace(account.GetCredential("sub")); userID != "" {
req.Header.Set("X-UserID", userID)
}
if email := strings.TrimSpace(account.GetCredential("email")); email != "" {
req.Header.Set("X-Email", email)
}
}
}
account.ApplyHeaderOverrides(req.Header)
return req, nil
}
@@ -438,11 +490,31 @@ func buildGeminiModelsURL(base string) string {
}
type upstreamModelEntry struct {
ID string `json:"id"`
Name string `json:"name"`
ID string `json:"id"`
Model string `json:"model"`
ModelID string `json:"modelId"`
ModelIDSnake string `json:"model_id"`
Name string `json:"name"`
Meta json.RawMessage `json:"_meta"`
}
type upstreamModelEntryMetadata struct {
ID string `json:"id"`
Model string `json:"model"`
ModelID string `json:"modelId"`
ModelIDSnake string `json:"model_id"`
Name string `json:"name"`
}
func extractUpstreamModelIDs(body []byte) ([]string, error) {
return extractUpstreamModelIDsWithSelector(body, upstreamModelEntryID)
}
func extractGrokUpstreamModelIDs(body []byte) ([]string, error) {
return extractUpstreamModelIDsWithSelector(body, grokUpstreamModelEntryID)
}
func extractUpstreamModelIDsWithSelector(body []byte, selectID func(upstreamModelEntry) string) ([]string, error) {
var response struct {
Data []upstreamModelEntry `json:"data"`
Models []upstreamModelEntry `json:"models"`
@@ -455,24 +527,24 @@ func extractUpstreamModelIDs(body []byte) ([]string, error) {
models := make([]string, 0, len(arrayResponse))
for _, entry := range arrayResponse {
models = append(models, upstreamModelEntryID(entry))
models = append(models, selectID(entry))
}
return dedupeAndSortModelIDs(models), nil
}
models := make([]string, 0, len(response.Data)+len(response.Models))
for _, entry := range response.Data {
models = append(models, upstreamModelEntryID(entry))
models = append(models, selectID(entry))
}
for _, entry := range response.Models {
models = append(models, upstreamModelEntryID(entry))
models = append(models, selectID(entry))
}
if len(models) == 0 {
var arrayResponse []upstreamModelEntry
if err := json.Unmarshal(body, &arrayResponse); err == nil {
for _, entry := range arrayResponse {
models = append(models, upstreamModelEntryID(entry))
models = append(models, selectID(entry))
}
}
}
@@ -488,6 +560,37 @@ func upstreamModelEntryID(entry upstreamModelEntry) string {
return strings.TrimPrefix(modelID, "models/")
}
func grokUpstreamModelEntryID(entry upstreamModelEntry) string {
candidates := []string{
entry.Model,
entry.ModelID,
entry.ModelIDSnake,
entry.ID,
}
if len(entry.Meta) > 0 {
var meta upstreamModelEntryMetadata
if err := json.Unmarshal(entry.Meta, &meta); err == nil {
candidates = append(candidates,
meta.Model,
meta.ModelID,
meta.ModelIDSnake,
meta.ID,
meta.Name,
)
}
}
// `name` is a display label in the Grok catalog, so keep it as the final
// compatibility fallback rather than preferring it over protocol model IDs.
candidates = append(candidates, entry.Name)
for _, candidate := range candidates {
modelID := strings.TrimSpace(candidate)
if modelID != "" {
return strings.TrimPrefix(modelID, "models/")
}
}
return ""
}
func dedupeAndSortModelIDs(models []string) []string {
seen := make(map[string]struct{}, len(models))
result := make([]string, 0, len(models))
@@ -7,6 +7,7 @@ import (
"net/http"
"strings"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/stretchr/testify/require"
@@ -20,6 +21,25 @@ func upstreamModelSyncTestConfig() *config.Config {
}
}
func grokOAuthModelSyncTestAccount(baseURL string) *Account {
credentials := map[string]any{
"access_token": "oauth-access-token",
"refresh_token": "oauth-refresh-token",
"expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339),
"sub": "grok-user-id",
"email": "grok-user@example.com",
}
if strings.TrimSpace(baseURL) != "" {
credentials["base_url"] = baseURL
}
return &Account{
ID: 10,
Platform: PlatformGrok,
Type: AccountTypeOAuth,
Credentials: credentials,
}
}
func TestBuildV1ModelsURL(t *testing.T) {
t.Parallel()
@@ -115,6 +135,11 @@ func TestExtractUpstreamModelIDs(t *testing.T) {
body: `[{"id":"z-model"},{"name":"models/a-model"}]`,
want: []string{"a-model", "z-model"},
},
{
name: "standard id wins over provider-specific model field",
body: `{"data":[{"id":"canonical-id","model":"display-model"}]}`,
want: []string{"canonical-id"},
},
}
for _, tt := range tests {
@@ -129,6 +154,14 @@ func TestExtractUpstreamModelIDs(t *testing.T) {
}
}
func TestExtractGrokUpstreamModelIDs(t *testing.T) {
t.Parallel()
models, err := extractGrokUpstreamModelIDs([]byte(`{"data":[{"id":"display-id","model":"grok-4.5"},{"modelId":"grok-build-0.1"},{"model_id":"grok-composer-2.5-fast"},{"name":"Grok Meta Display Name","_meta":{"model":"grok-meta"}},{"name":"grok-name"},{"id":"grok-safe","_meta":"not-an-object"}]}`))
require.NoError(t, err)
require.Equal(t, []string{"grok-4.5", "grok-build-0.1", "grok-composer-2.5-fast", "grok-meta", "grok-name", "grok-safe"}, models)
}
func TestBuildUpstreamModelsRequestsForAPIKeyAccounts(t *testing.T) {
t.Parallel()
@@ -214,20 +247,36 @@ func TestBuildUpstreamModelsRequestsForAPIKeyAccounts(t *testing.T) {
require.Equal(t, "antigravity-key", antigravityReq.Header.Get("x-api-key"))
}
func TestBuildUpstreamModelsRequestRejectsGrokOAuth(t *testing.T) {
func TestBuildUpstreamModelsRequestSupportsGrokOAuth(t *testing.T) {
t.Parallel()
svc := &AccountTestService{
cfg: upstreamModelSyncTestConfig(),
grokTokenProvider: NewGrokTokenProvider(nil, nil),
}
req, err := svc.buildUpstreamModelsRequest(context.Background(), grokOAuthModelSyncTestAccount(""))
require.NoError(t, err)
require.Equal(t, "https://cli-chat-proxy.grok.com/v1/models", req.URL.String())
require.Equal(t, "Bearer oauth-access-token", req.Header.Get("Authorization"))
require.Equal(t, grokCLIVersion, req.Header.Get("X-Grok-Client-Version"))
require.Equal(t, "interactive", req.Header.Get("X-Grok-Client-Mode"))
require.Equal(t, grokUpstreamUserAgent, req.Header.Get("User-Agent"))
require.Equal(t, "grok-user-id", req.Header.Get("X-UserID"))
require.Equal(t, "grok-user@example.com", req.Header.Get("X-Email"))
require.NotContains(t, req.Header.Get("Authorization"), "oauth-refresh-token")
}
func TestBuildUpstreamModelsRequestGrokOAuthRequiresTokenProvider(t *testing.T) {
t.Parallel()
svc := &AccountTestService{cfg: upstreamModelSyncTestConfig()}
_, err := svc.buildUpstreamModelsRequest(context.Background(), &Account{
Platform: PlatformGrok,
Type: AccountTypeOAuth,
})
_, err := svc.buildUpstreamModelsRequest(context.Background(), grokOAuthModelSyncTestAccount(""))
require.Error(t, err)
var syncErr *UpstreamModelSyncError
require.True(t, errors.As(err, &syncErr))
require.Equal(t, UpstreamModelSyncErrorUnsupported, syncErr.Kind)
require.Contains(t, syncErr.SafeMessage(), "Unsupported Grok account type")
require.Equal(t, UpstreamModelSyncErrorConfiguration, syncErr.Kind)
require.Contains(t, syncErr.SafeMessage(), "token provider")
}
func TestBuildAntigravityAPIKeyModelsRequestRejectsOfficialCloudCodeBase(t *testing.T) {
@@ -321,6 +370,45 @@ func TestFetchUpstreamSupportedModelsParsesGrokAPIKeyResponse(t *testing.T) {
require.Equal(t, "Bearer xai-key", upstream.lastReq.Header.Get("Authorization"))
}
func TestFetchUpstreamSupportedModelsParsesGrokOAuthResponse(t *testing.T) {
t.Parallel()
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(`{"data":[{"model":"grok-4.5"},{"model":"grok-4.5"},{"modelId":"grok-build-0.1"}]}`)),
}}
svc := &AccountTestService{
httpUpstream: upstream,
cfg: upstreamModelSyncTestConfig(),
grokTokenProvider: NewGrokTokenProvider(nil, nil),
}
models, err := svc.FetchUpstreamSupportedModels(context.Background(), grokOAuthModelSyncTestAccount(""))
require.NoError(t, err)
require.Equal(t, []string{"grok-4.5", "grok-build-0.1"}, models)
require.Equal(t, "https://cli-chat-proxy.grok.com/v1/models", upstream.lastReq.URL.String())
require.Equal(t, "Bearer oauth-access-token", upstream.lastReq.Header.Get("Authorization"))
require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version"))
require.Equal(t, "interactive", upstream.lastReq.Header.Get("X-Grok-Client-Mode"))
require.Equal(t, "grok-user-id", upstream.lastReq.Header.Get("X-UserID"))
require.Equal(t, "grok-user@example.com", upstream.lastReq.Header.Get("X-Email"))
}
func TestBuildUpstreamModelsRequestGrokOAuthDoesNotSendIdentityToCustomBase(t *testing.T) {
t.Parallel()
svc := &AccountTestService{
cfg: upstreamModelSyncTestConfig(),
grokTokenProvider: NewGrokTokenProvider(nil, nil),
}
req, err := svc.buildUpstreamModelsRequest(context.Background(), grokOAuthModelSyncTestAccount("https://relay.example/v1"))
require.NoError(t, err)
require.Equal(t, "https://relay.example/v1/models", req.URL.String())
require.Empty(t, req.Header.Get("X-UserID"))
require.Empty(t, req.Header.Get("X-Email"))
}
func TestFetchUpstreamSupportedModelsDoesNotExposeUpstreamBody(t *testing.T) {
t.Parallel()