合并上游 main,保留 5925 的 Grok 重试上限与 5888 的协议兼容修复

同账号重试采用次数上限加 deadline;畸形 tools 在出站前删除;compaction 422 与结构化错误扫描一并保留。
This commit is contained in:
IanShaw027
2026-08-21 08:49:44 +08:00
51 changed files with 1414 additions and 240 deletions
+1 -1
View File
@@ -718,7 +718,7 @@ Sub2API supports both Grok subscription accounts through xAI OAuth and standard
- Codex CLI style Responses WebSocket ingress is accepted on the Responses targets and bridged to xAI HTTP/SSE Responses upstream
- Text models: `grok-4.5`, `grok-4.3`, `grok-build-0.1`, `grok-composer-2.5-fast`, `grok-4.20-0309-reasoning`, `grok-4.20-0309-non-reasoning`, and `grok-4.20-multi-agent-0309`
- Media targets for Grok groups: `/v1/images/generations`, `/images/generations`, `/v1/images/edits`, `/images/edits`, `/v1/videos/generations`, `/videos/generations`, `/v1/videos/edits`, `/videos/edits`, `/v1/videos/extensions`, `/videos/extensions`, `/v1/videos/{request_id}`, and `/videos/{request_id}`. Generation, editing, and extension requests require the group image-generation permission.
- Media models: `grok-imagine`, `grok-imagine-image-quality`, `grok-imagine-image`, `grok-imagine-edit`, `grok-imagine-video`, and `grok-imagine-video-1.5`
- Media models: `grok-imagine`, `grok-imagine-image-quality`, `grok-imagine-image`, `grok-imagine-image-2.0`, `grok-imagine-edit`, `grok-imagine-video`, and `grok-imagine-video-1.5`
- JSON image-edit and video-generation requests accept image references in `image`, `images`, `reference_images`, and `mask` objects. Use `url` for xAI-compatible payloads; the legacy `image_url` field remains accepted and is normalized to `url` before forwarding.
- Out of scope for this provider: TTS, transcription, browser automation, cookies, and Grok web scraping
+7
View File
@@ -934,6 +934,9 @@ type GatewayConfig struct {
// OpenAIResponseHeaderTimeout: OpenAI/Codex 上游等待响应头的超时时间(秒),0表示无超时
// OpenAI/Codex 请求可能在上游排队较久;默认不使用通用响应头超时截断。
OpenAIResponseHeaderTimeout int `mapstructure:"openai_response_header_timeout"`
// GrokResponseHeaderTimeout bounds the pre-first-byte wait for xAI/Grok.
// A zero value uses the provider-safe default instead of the generic gateway timeout.
GrokResponseHeaderTimeout int `mapstructure:"grok_response_header_timeout"`
// OpenAIFirstOutputTimeoutSeconds: native HTTP Responses 首个语义输出超时(秒),0表示禁用。
OpenAIFirstOutputTimeoutSeconds int `mapstructure:"openai_first_output_timeout_seconds"`
// OpenAIHighEffortFirstOutputTimeoutSeconds: high/xhigh/max 推理的首个语义输出超时(秒)。
@@ -2327,6 +2330,7 @@ func setDefaults() {
// Gateway
viper.SetDefault("gateway.response_header_timeout", 600) // 600秒(10分钟)等待上游响应头,LLM高负载时可能排队较久
viper.SetDefault("gateway.openai_response_header_timeout", 0)
viper.SetDefault("gateway.grok_response_header_timeout", 120)
viper.SetDefault("gateway.openai_first_output_timeout_seconds", 0)
viper.SetDefault("gateway.openai_high_effort_first_output_timeout_seconds", 0)
viper.SetDefault("gateway.log_upstream_error_body", true)
@@ -3241,6 +3245,9 @@ func (c *Config) Validate() error {
if c.Gateway.OpenAIResponseHeaderTimeout < 0 {
return fmt.Errorf("gateway.openai_response_header_timeout must be non-negative")
}
if c.Gateway.GrokResponseHeaderTimeout < 0 || c.Gateway.GrokResponseHeaderTimeout > 1800 {
return fmt.Errorf("gateway.grok_response_header_timeout must be between 0-1800 seconds")
}
if c.Gateway.OpenAIFirstOutputTimeoutSeconds < 0 || c.Gateway.OpenAIFirstOutputTimeoutSeconds > 600 ||
(c.Gateway.OpenAIFirstOutputTimeoutSeconds > 0 && c.Gateway.OpenAIFirstOutputTimeoutSeconds < 30) {
return fmt.Errorf("gateway.openai_first_output_timeout_seconds must be 0 or between 30-600 seconds")
+41 -7
View File
@@ -56,10 +56,13 @@ const (
const profitVetoExhaustedMessage = "No available accounts: all candidates rejected by group profit control"
func sameAccountRetryDelayFor(failoverErr *service.UpstreamFailoverError, retryCount int) time.Duration {
if failoverErr != nil && failoverErr.SameAccountRetryDelay > 0 {
if failoverErr == nil {
return sameAccountRetryDelay
}
if failoverErr.SameAccountRetryDelay > 0 {
return failoverErr.SameAccountRetryDelay
}
if failoverErr == nil || !failoverErr.RequestScopedTransient || retryCount <= 1 {
if !failoverErr.RequestScopedTransient || retryCount <= 1 {
return sameAccountRetryDelay
}
@@ -77,10 +80,37 @@ func sameAccountRetryAllowed(failoverErr *service.UpstreamFailoverError, retryCo
if failoverErr == nil || !failoverErr.RetryableOnSameAccount {
return false
}
if !failoverErr.SameAccountRetryDeadline.IsZero() {
return time.Now().Before(failoverErr.SameAccountRetryDeadline)
if !sameAccountRetryDeadlineAllows(failoverErr) {
return false
}
return retryCount < retryLimit
// Deadline-window retries (OAuth 429) may pass retryLimit=0 and are not
// bound to pool_mode_retry_count. Pool-mode callers pass a positive limit.
if !failoverErr.SameAccountRetryDeadline.IsZero() && retryLimit <= 0 {
return true
}
if failoverErr.SameAccountRetryMax > 0 && (retryLimit <= 0 || failoverErr.SameAccountRetryMax < retryLimit) {
retryLimit = failoverErr.SameAccountRetryMax
}
return retryLimit > 0 && retryCount < retryLimit
}
// sameAccountRetryDeadlineAllows prevents a retry from starting after the
// service-provided same-account retry window has elapsed.
func sameAccountRetryDeadlineAllows(failoverErr *service.UpstreamFailoverError) bool {
return failoverErr == nil || failoverErr.SameAccountRetryDeadline.IsZero() || time.Now().Before(failoverErr.SameAccountRetryDeadline)
}
// effectiveSameAccountRetryLimit applies an error-specific cap without
// overriding an explicit account setting of zero (which disables retries).
func effectiveSameAccountRetryLimit(failoverErr *service.UpstreamFailoverError, account *service.Account) int {
if account == nil {
return 0
}
limit := account.GetPoolModeRetryCount()
if limit > 0 && failoverErr != nil && failoverErr.SameAccountRetryMax > 0 && failoverErr.SameAccountRetryMax < limit {
return failoverErr.SameAccountRetryMax
}
return limit
}
// FailoverState 跨循环迭代共享的 failover 状态
@@ -171,14 +201,18 @@ func (s *FailoverState) HandleFailoverError(
}
// 同账号重试不算切换账号,粘性会话仅在实际切换时强制缓存计费。
sameAccountRetry := sameAccountRetryAllowed(failoverErr, s.SameAccountRetryCount[accountID], retryLimit)
retryCount := s.SameAccountRetryCount[accountID]
if failoverErr.SameAccountRetryMax > 0 && failoverErr.SameAccountRetryMax < retryLimit {
retryLimit = failoverErr.SameAccountRetryMax
}
sameAccountRetry := failoverErr.RetryableOnSameAccount && retryLimit > 0 && retryCount < retryLimit && sameAccountRetryDeadlineAllows(failoverErr)
if needForceCacheBilling(s.hasBoundSession, failoverErr, sameAccountRetry) {
s.ForceCacheBilling = true
}
// 同账号重试:对 RetryableOnSameAccount 的临时性错误,先在同一账号上重试。
// 重试次数上限 retryLimit 由调用方传入(账号级 pool_mode_retry_count 配置)。
if sameAccountRetryAllowed(failoverErr, s.SameAccountRetryCount[accountID], retryLimit) {
if sameAccountRetry {
s.SameAccountRetryCount[accountID]++
retryDelay := sameAccountRetryDelayFor(failoverErr, s.SameAccountRetryCount[accountID])
logger.FromContext(ctx).Warn("gateway.failover_same_account_retry",
+20 -3
View File
@@ -75,6 +75,23 @@ func TestSameAccountRetryAllowedUsesDeadlineInsteadOfPoolCount(t *testing.T) {
require.False(t, sameAccountRetryAllowed(err, 0, 100))
}
func TestSameAccountRetryDeadlineAllows(t *testing.T) {
require.True(t, sameAccountRetryDeadlineAllows(&service.UpstreamFailoverError{}))
require.True(t, sameAccountRetryDeadlineAllows(&service.UpstreamFailoverError{
SameAccountRetryDeadline: time.Now().Add(time.Second),
}))
require.False(t, sameAccountRetryDeadlineAllows(&service.UpstreamFailoverError{
SameAccountRetryDeadline: time.Now().Add(-time.Second),
}))
}
func TestEffectiveSameAccountRetryLimitHonorsErrorCapAndDisabledAccount(t *testing.T) {
account := &service.Account{Type: service.AccountTypeAPIKey, Credentials: map[string]any{"pool_mode": true, "pool_mode_retry_count": float64(3)}}
require.Equal(t, 1, effectiveSameAccountRetryLimit(&service.UpstreamFailoverError{SameAccountRetryMax: 1}, account))
account.Credentials["pool_mode_retry_count"] = float64(0)
require.Equal(t, 0, effectiveSameAccountRetryLimit(&service.UpstreamFailoverError{SameAccountRetryMax: 1}, account))
}
// ---------------------------------------------------------------------------
// Helper
// ---------------------------------------------------------------------------
@@ -335,7 +352,7 @@ func TestHandleFailoverError_CacheBilling(t *testing.T) {
require.Zero(t, fs.SwitchCount)
})
t.Run("deadline允许超过计数上限时仍不强制缓存计费", func(t *testing.T) {
t.Run("deadline存在但计数已耗尽时按切换处理并强制缓存计费", func(t *testing.T) {
mock := &mockTempUnscheduler{}
fs := NewFailoverState(3, true)
fs.SameAccountRetryCount[100] = maxSameAccountRetries
@@ -345,8 +362,8 @@ func TestHandleFailoverError_CacheBilling(t *testing.T) {
fs.HandleFailoverError(context.Background(), mock, 100, "openai", maxSameAccountRetries, err)
require.False(t, fs.ForceCacheBilling)
require.Zero(t, fs.SwitchCount)
require.True(t, fs.ForceCacheBilling)
require.Equal(t, 1, fs.SwitchCount)
})
t.Run("同账号重试耗尽并实际切换时设置ForceCacheBilling", func(t *testing.T) {
+75 -34
View File
@@ -44,39 +44,84 @@ func (h *OpenAIGatewayHandler) GrokRealtime(c *gin.Context) {
return
}
selection, _, err := h.gatewayService.SelectAccountWithSchedulerForCapability(
c.Request.Context(),
apiKey.GroupID,
"",
"",
"grok-4.5",
nil,
service.OpenAIUpstreamTransportHTTPSSE,
// Grok only advertises chat_completions + media capabilities on HEAD.
service.OpenAIEndpointCapabilityChatCompletions,
false,
false,
false,
service.PlatformGrok,
)
if err != nil || selection == nil || selection.Account == nil {
h.errorResponse(c, http.StatusServiceUnavailable, "api_error", "No available Grok accounts")
return
}
var streamStarted bool
reqLog := requestLogger(c, "handler.openai_gateway.grok_realtime")
release, slotStatus := h.acquireResponsesAccountSlot(c, apiKey.GroupID, "", selection, true, &streamStarted, reqLog)
if slotStatus != openAISlotAcquireOK {
model := c.Query("model")
if strings.TrimSpace(model) == "" {
model = "grok-voice-latest"
}
// Keep the HTTP response uncommitted while selecting and probing an account.
// Realtime is not an HTTP streaming response; using reqStream=true here would
// let the wait queue flush an SSE ping before the WebSocket handshake succeeds.
failed := map[int64]struct{}{}
var selection *service.AccountSelectionResult
var release func()
var token string
var upstream *service.GrokRealtimeUpstream
var candidateSeen bool
for attempts := 0; attempts < 4; attempts++ {
// Realtime's voice model is not a text-model capability. Passing a
// concrete text model here would reject accounts mapped only to an
// older/default text model before the upstream handshake can decide.
// An empty requested model keeps account selection capability-based;
// the actual voice model remains in the upstream WS query below.
candidate, _, selectErr := h.gatewayService.SelectAccountWithSchedulerForCapability(
c.Request.Context(), apiKey.GroupID, "", "", "", failed,
service.OpenAIUpstreamTransportHTTPSSE,
service.OpenAIEndpointCapabilityChatCompletions,
false, false, false, service.PlatformGrok,
)
if selectErr != nil || candidate == nil || candidate.Account == nil {
break
}
candidateSeen = true
account := candidate.Account
var streamStarted bool
var slotStatus openAISlotAcquireResult
release, slotStatus = h.acquireResponsesAccountSlot(c, apiKey.GroupID, "", candidate, false, &streamStarted, reqLog)
if slotStatus != openAISlotAcquireOK {
if slotStatus == openAISlotAcquireFailed {
return
}
failed[account.ID] = struct{}{}
continue
}
var credErr error
token, _, credErr = h.gatewayService.GetRequestCredential(c.Request.Context(), c, account)
if credErr != nil {
release()
release = nil
failed[account.ID] = struct{}{}
continue
}
probeCtx, cancelProbe := context.WithTimeout(c.Request.Context(), service.DefaultGrokRealtimeDialTimeout)
candidateUpstream, openErr := h.gatewayService.OpenGrokRealtime(probeCtx, account, token, model)
cancelProbe()
if openErr != nil {
reqLog.Warn("grok_realtime.pre_accept_failed", zap.Int64("account_id", account.ID), zap.Error(openErr))
statusCode := http.StatusBadGateway
var dialErr *service.GrokRealtimeDialError
if errors.As(openErr, &dialErr) && dialErr.StatusCode > 0 {
statusCode = dialErr.StatusCode
}
h.gatewayService.HandleGrokRealtimeUpstreamError(c.Request.Context(), account, statusCode, []byte(openErr.Error()))
release()
release = nil
failed[account.ID] = struct{}{}
continue
}
selection, upstream = candidate, candidateUpstream
break
}
if selection == nil || selection.Account == nil || release == nil || upstream == nil {
if !candidateSeen {
h.errorResponse(c, http.StatusServiceUnavailable, "api_error", "No available Grok accounts")
} else {
h.errorResponse(c, http.StatusBadGateway, "upstream_error", "Grok realtime upstream unavailable")
}
return
}
defer release()
token, _, err := h.gatewayService.GetRequestCredential(c.Request.Context(), c, selection.Account)
if err != nil {
h.errorResponse(c, http.StatusBadGateway, "upstream_error", "Grok credential unavailable")
return
}
defer func() { _ = upstream.Close() }()
conn, err := coderws.Accept(c.Writer, c.Request, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover})
if err != nil {
@@ -84,12 +129,8 @@ func (h *OpenAIGatewayHandler) GrokRealtime(c *gin.Context) {
}
defer func() { _ = conn.CloseNow() }()
model := c.Query("model")
if strings.TrimSpace(model) == "" {
model = "grok-voice-latest"
}
started := time.Now()
audioObserved, proxyErr := h.gatewayService.ProxyGrokRealtime(c.Request.Context(), c, conn, selection.Account, token, model)
audioObserved, proxyErr := h.gatewayService.ProxyGrokRealtimeConn(c.Request.Context(), c, conn, upstream)
elapsed := time.Since(started)
if proxyErr != nil {
reqLog.Info("grok_realtime.proxy_failed", zap.Error(proxyErr))
+5 -3
View File
@@ -361,19 +361,21 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
return
}
if failoverErr.RetryableOnSameAccount {
retryLimit := account.GetPoolModeRetryCount()
if sameAccountRetryCount[account.ID] < retryLimit {
retryLimit := effectiveSameAccountRetryLimit(failoverErr, account)
if sameAccountRetryCount[account.ID] < retryLimit && sameAccountRetryDeadlineAllows(failoverErr) {
sameAccountRetryCount[account.ID]++
retryDelay := sameAccountRetryDelayFor(failoverErr, sameAccountRetryCount[account.ID])
reqLog.Warn("grok_media.pool_mode_same_account_retry",
zap.Int64("account_id", account.ID),
zap.Int("upstream_status", failoverErr.StatusCode),
zap.Int("retry_limit", retryLimit),
zap.Int("retry_count", sameAccountRetryCount[account.ID]),
zap.Duration("retry_delay", retryDelay),
)
select {
case <-requestCtx.Done():
return
case <-time.After(sameAccountRetryDelay):
case <-time.After(retryDelay):
}
continue
}
@@ -330,8 +330,8 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
}
// Pool mode: retry on the same account
if failoverErr.RetryableOnSameAccount {
retryLimit := account.GetPoolModeRetryCount()
if sameAccountRetryAllowed(failoverErr, sameAccountRetryCount[account.ID], retryLimit) {
retryLimit := effectiveSameAccountRetryLimit(failoverErr, account)
if sameAccountRetryCount[account.ID] < retryLimit && sameAccountRetryDeadlineAllows(failoverErr) {
sameAccountRetryCount[account.ID]++
retryDelay := sameAccountRetryDelayFor(failoverErr, sameAccountRetryCount[account.ID])
reqLog.Warn("openai_chat_completions.pool_mode_same_account_retry",
@@ -754,8 +754,8 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
}
// 池模式:同账号重试
if failoverErr.RetryableOnSameAccount {
retryLimit := account.GetPoolModeRetryCount()
if sameAccountRetryAllowed(failoverErr, sameAccountRetryCount[account.ID], retryLimit) {
retryLimit := effectiveSameAccountRetryLimit(failoverErr, account)
if sameAccountRetryCount[account.ID] < retryLimit && sameAccountRetryDeadlineAllows(failoverErr) {
sameAccountRetryCount[account.ID]++
retryDelay := sameAccountRetryDelayFor(failoverErr, sameAccountRetryCount[account.ID])
reqLog.Warn("openai.pool_mode_same_account_retry",
@@ -1309,8 +1309,8 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
}
// 池模式:同账号重试
if failoverErr.RetryableOnSameAccount {
retryLimit := account.GetPoolModeRetryCount()
if sameAccountRetryAllowed(failoverErr, sameAccountRetryCount[account.ID], retryLimit) {
retryLimit := effectiveSameAccountRetryLimit(failoverErr, account)
if sameAccountRetryCount[account.ID] < retryLimit && sameAccountRetryDeadlineAllows(failoverErr) {
sameAccountRetryCount[account.ID]++
retryDelay := sameAccountRetryDelayFor(failoverErr, sameAccountRetryCount[account.ID])
reqLog.Warn("openai_messages.pool_mode_same_account_retry",
@@ -672,7 +672,7 @@ func TestResolveOpenAIMessagesDispatchMappedModel(t *testing.T) {
Platform: service.PlatformGrok,
},
}
require.Equal(t, "grok-4.5", resolveOpenAIMessagesDispatchMappedModel(nil, apiKey, "claude-sonnet-4-5"))
require.Equal(t, "grok-4.6", resolveOpenAIMessagesDispatchMappedModel(nil, apiKey, "claude-sonnet-4-5"))
require.Empty(t, resolveOpenAIMessagesDispatchMappedModel(nil, apiKey, "grok"))
})
+2 -2
View File
@@ -308,8 +308,8 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
return
}
if failoverErr.RetryableOnSameAccount {
retryLimit := account.GetPoolModeRetryCount()
if sameAccountRetryAllowed(failoverErr, sameAccountRetryCount[account.ID], retryLimit) {
retryLimit := effectiveSameAccountRetryLimit(failoverErr, account)
if sameAccountRetryCount[account.ID] < retryLimit && sameAccountRetryDeadlineAllows(failoverErr) {
sameAccountRetryCount[account.ID]++
retryDelay := sameAccountRetryDelayFor(failoverErr, sameAccountRetryCount[account.ID])
reqLog.Warn("openai.images.pool_mode_same_account_retry",
+14 -15
View File
@@ -51,24 +51,25 @@ type Model struct {
// DefaultTextModel is the built-in fallback for empty model fields and Grok
// text aliases (e.g. "grok", "grok-latest"). Operators may override the runtime
// default via settings key grok_default_text_model.
const DefaultTextModel = "grok-4.5"
const DefaultTextModel = "grok-4.6"
// Official Imagine model IDs (https://docs.x.ai/docs/models).
const (
DefaultImagineImageQualityModel = "grok-imagine-image-quality"
DefaultImagineImageFastModel = "grok-imagine-image"
DefaultImagineImage20Model = "grok-imagine-image-2.0"
DefaultImagineVideoModel = "grok-imagine-video"
DefaultImagineVideo15LegacyModel = "grok-imagine-video-1.5"
DefaultImagineVideo15Model = "grok-imagine-video-1.5-preview"
DefaultImagineVideo15Model = "grok-imagine-video-1.5"
DefaultImagineVideo15LegacyModel = "grok-imagine-video-1.5-preview"
)
// ModelMappingOptions controls optional expansions of the default mapping.
// Cross-client wildcards (gpt-*/claude-*) default ON via settings
// grok_cross_client_model_map_enabled so Codex/Claude clients keep working
// against Grok groups (map to DefaultText / grok-4.5). Operators may disable.
// against Grok groups (map to DefaultText / grok-4.6). Operators may disable.
type ModelMappingOptions struct {
// DefaultText is the target for empty models and optional cross-client maps.
// Empty → DefaultTextModel (grok-4.5).
// Empty → DefaultTextModel (grok-4.6).
DefaultText string
// EnableCrossClientMap merges gpt-*/codex-*/o*/claude-* → DefaultText.
EnableCrossClientMap bool
@@ -86,8 +87,6 @@ var defaultModels = []Model{
{ID: "grok-4.6", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 4.6"},
{ID: "grok-4.5", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 4.5"},
{ID: "grok-4.3", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 4.3"},
{ID: "grok-3-mini", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 3 Mini"},
{ID: "grok-3-mini-fast", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 3 Mini Fast"},
{ID: "grok-build-0.1", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Build 0.1"},
{ID: "grok-composer-2.5-fast", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Composer 2.5 Fast"},
{ID: "grok-4.20-0309-reasoning", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 4.20 Reasoning"},
@@ -96,9 +95,9 @@ var defaultModels = []Model{
// Imagine
{ID: DefaultImagineImageQualityModel, Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Image Quality"},
{ID: DefaultImagineImageFastModel, Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Image"},
{ID: DefaultImagineImage20Model, Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Image 2.0"},
{ID: DefaultImagineVideoModel, Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Video"},
{ID: DefaultImagineVideo15Model, Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Video 1.5 Preview"},
{ID: DefaultImagineVideo15LegacyModel, Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Video 1.5 Legacy"},
{ID: DefaultImagineVideo15Model, Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Video 1.5"},
}
// grokTextResponsesModelAliases is the source of truth for Grok text models
@@ -109,14 +108,14 @@ var grokTextResponsesModelAliases = map[string]string{
"grok-latest": DefaultTextModel,
"grok-4.6": "grok-4.6",
"grok-4.6-latest": "grok-4.6",
"grok-4.5": DefaultTextModel,
"grok-4.5-latest": DefaultTextModel,
"grok-4.5": "grok-4.5",
"grok-4.5-latest": "grok-4.5",
"grok-4.3": "grok-4.3",
"grok-4.3-latest": "grok-4.3",
"grok-3-mini": "grok-3-mini",
"grok-3-mini-fast": "grok-3-mini-fast",
"grok-build": "grok-build-0.1",
"grok-build-latest": DefaultTextModel,
"grok-build-latest": "grok-build-0.1",
"grok-build-0.1": "grok-build-0.1",
"grok-composer-2.5-fast": "grok-composer-2.5-fast",
"grok-composer": "grok-composer-2.5-fast",
@@ -162,7 +161,7 @@ func ModelMappingWithOptions(opts ModelMappingOptions) map[string]string {
}
for alias, canonical := range grokTextResponsesModelAliases {
// Remap aliases that pointed at DefaultTextModel constant to runtime default.
if canonical == DefaultTextModel {
if (alias == "grok" || alias == "grok-latest") && canonical == DefaultTextModel {
mapping[alias] = defaultText
} else {
mapping[alias] = canonical
@@ -179,7 +178,7 @@ func ModelMappingWithOptions(opts ModelMappingOptions) map[string]string {
// Keep official IDs as identity so client-requested model strings are not
// rewritten on the wire (pricing still canonicalizes 1.5* via CanonicalImagineVideoModel).
mapping["grok-imagine-video"] = DefaultImagineVideoModel
mapping["grok-imagine-video-1.5"] = DefaultImagineVideo15LegacyModel
mapping["grok-imagine-video-1.5"] = DefaultImagineVideo15Model
mapping["grok-imagine-video-1.5-preview"] = DefaultImagineVideo15Model
// Informal alias only:
mapping["grok-video-1.5"] = DefaultImagineVideo15Model
@@ -273,7 +272,7 @@ func ResolveGrokTextResponsesModelID(model string, defaultText ...string) string
}
normalized := strings.ToLower(StripGrokProviderPrefix(trimmed))
if canonical, ok := grokTextResponsesModelAliases[normalized]; ok {
if canonical == DefaultTextModel {
if (normalized == "grok" || normalized == "grok-latest") && canonical == DefaultTextModel {
return fallback
}
return canonical
+18 -6
View File
@@ -12,14 +12,14 @@ func TestDefaultModelMappingExcludesCrossClientWildcards(t *testing.T) {
SetRuntimeModelMappingOptions(ModelMappingOptions{})
mapping := DefaultModelMapping()
require.Equal(t, "grok-4.5", mapping["grok"])
require.Equal(t, "grok-4.5", mapping["grok-latest"])
require.Equal(t, "grok-4.6", mapping["grok"])
require.Equal(t, "grok-4.6", mapping["grok-latest"])
require.Equal(t, "grok-build-0.1", mapping["grok-build"])
require.Equal(t, DefaultTextModel, mapping["grok-build-latest"])
require.Equal(t, "grok-build-0.1", mapping["grok-build-latest"])
require.Equal(t, DefaultImagineImageQualityModel, mapping["grok-imagine-edit"])
require.Equal(t, DefaultImagineVideo15LegacyModel, mapping["grok-imagine-video-1.5"])
require.Equal(t, DefaultImagineVideo15Model, mapping["grok-imagine-video-1.5"])
require.Equal(t, DefaultImagineVideo15Model, mapping["grok-imagine-video-1.5-preview"])
require.Equal(t, "grok-4.5", mapping["xai/grok"])
require.Equal(t, "grok-4.6", mapping["xai/grok"])
// Cross-vendor wildcards must stay opt-in.
_, hasGPT := mapping["gpt-*"]
@@ -68,7 +68,19 @@ func TestDefaultModelsIncludesGrok46(t *testing.T) {
func TestResolveGrokTextResponsesModelID(t *testing.T) {
t.Parallel()
require.Equal(t, "grok-4.5", ResolveGrokTextResponsesModelID(""))
require.Equal(t, "grok-4.6", ResolveGrokTextResponsesModelID(""))
require.Equal(t, "grok-4.3", ResolveGrokTextResponsesModelID("grok", "grok-4.3"))
require.Equal(t, "grok-4.20-multi-agent-0309", ResolveGrokTextResponsesModelID("grok-4.20-multi-agent"))
}
func TestExplicitGrok45DoesNotFollowRuntimeDefault(t *testing.T) {
require.Equal(t, "grok-4.5", ResolveGrokTextResponsesModelID("grok-4.5", "grok-4.6"))
require.Equal(t, "grok-4.5", ResolveGrokTextResponsesModelID("grok-4.5-latest", "grok-4.6"))
require.Equal(t, "grok-4.6", ResolveGrokTextResponsesModelID("grok", "grok-4.6"))
}
func TestBareGrokAliasesFollowGrok46Default(t *testing.T) {
require.Equal(t, "grok-4.6", ResolveGrokTextResponsesModelID("grok"))
require.Equal(t, "grok-4.6", ResolveGrokTextResponsesModelID("grok-latest"))
require.Equal(t, "grok-build-0.1", ResolveGrokTextResponsesModelID("grok-build-latest"))
}
+4 -4
View File
@@ -350,14 +350,14 @@ func TestDefaultModelMappingIncludesGrokAliases(t *testing.T) {
t.Cleanup(func() { SetRuntimeModelMappingOptions(original) })
SetRuntimeModelMappingOptions(ModelMappingOptions{})
mapping := DefaultModelMapping()
require.Equal(t, "grok-4.5", mapping["grok"])
require.Equal(t, "grok-4.5", mapping["grok-latest"])
require.Equal(t, "grok-4.6", mapping["grok"])
require.Equal(t, "grok-4.6", mapping["grok-latest"])
require.Equal(t, "grok-4.6", mapping["grok-4.6"])
require.Equal(t, "grok-4.6", mapping["grok-4.6-latest"])
require.Equal(t, "grok-4.5", mapping["grok-4.5"])
require.Equal(t, "grok-4.5", mapping["grok-4.5-latest"])
require.Equal(t, "grok-build-0.1", mapping["grok-build"])
require.Equal(t, "grok-4.5", mapping["grok-build-latest"])
require.Equal(t, "grok-build-0.1", mapping["grok-build-latest"])
require.Equal(t, "grok-composer-2.5-fast", mapping["grok-composer"])
require.Equal(t, "grok-composer-2.5-fast", mapping["composer-2.5"])
require.Equal(t, "grok-4.20-0309-reasoning", mapping["grok-4.20-reasoning"])
@@ -368,7 +368,7 @@ func TestDefaultModelMappingIncludesGrokAliases(t *testing.T) {
require.Equal(t, DefaultImagineImageQualityModel, mapping["grok-imagine-image-quality"])
require.Equal(t, DefaultImagineImageQualityModel, mapping["grok-imagine-edit"])
require.Equal(t, DefaultImagineVideoModel, mapping["grok-imagine-video"])
require.Equal(t, DefaultImagineVideo15LegacyModel, mapping["grok-imagine-video-1.5"])
require.Equal(t, DefaultImagineVideo15Model, mapping["grok-imagine-video-1.5"])
require.Equal(t, DefaultImagineVideo15Model, mapping["grok-imagine-video-1.5-preview"])
_, hasGPT := mapping["gpt-*"]
require.False(t, hasGPT, "cross-client wildcards must be opt-in")
+25
View File
@@ -0,0 +1,25 @@
package xai
// IncludeIndependentReasoningTokens adds reasoning tokens to billed output
// only when total_tokens proves they are not already folded into output.
//
// Official xAI Chat Completions example: prompt=32, completion=9,
// reasoning=94, total=135. Responses example: input=32, output=9,
// reasoning=110, total=151. OpenAI's canonical completion_tokens already
// includes reasoning and total equals input+output.
func IncludeIndependentReasoningTokens(input, output, total, reasoning int64) int64 {
if input < 0 || output < 0 || reasoning <= 0 || total <= 0 {
return output
}
if total == input+output {
return output
}
gap := total - input - output
if gap <= 0 {
return output
}
if reasoning < gap {
gap = reasoning
}
return output + gap
}
+30
View File
@@ -0,0 +1,30 @@
package xai
import "testing"
func TestIncludeIndependentReasoningTokens(t *testing.T) {
t.Parallel()
tests := []struct {
name string
input, output, total, reason int64
want int64
}{
{name: "xAI chat example", input: 32, output: 9, total: 135, reason: 94, want: 103},
{name: "xAI responses example", input: 32, output: 9, total: 151, reason: 110, want: 119},
{name: "OpenAI inclusive output", input: 32, output: 103, total: 135, reason: 94, want: 103},
{name: "gap smaller than reasoning", input: 10, output: 20, total: 33, reason: 5, want: 23},
{name: "absent total", input: 10, output: 20, total: 0, reason: 5, want: 20},
{name: "total already complete", input: 10, output: 20, total: 30, reason: 5, want: 20},
{name: "inconsistent undershoot", input: 10, output: 20, total: 25, reason: 5, want: 20},
{name: "zero reasoning", input: 10, output: 20, total: 35, reason: 0, want: 20},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
got := IncludeIndependentReasoningTokens(tt.input, tt.output, tt.total, tt.reason)
if got != tt.want {
t.Fatalf("got %d, want %d", got, tt.want)
}
})
}
}
@@ -151,6 +151,7 @@ func TestGroupUsageRollupTriggerSerializesInsertTransactionAcrossMidnight(t *tes
insertTx := beginGroupUsageRollupTriggerTestTx(t, ctx, schema)
defer func() { _ = insertTx.Rollback() }()
require.NoError(t, setGroupUsageRollupTriggerTimeZone(ctx, insertTx, "Asia/Shanghai"))
var insertBackendPID int
require.NoError(t, insertTx.QueryRowContext(ctx, "SELECT pg_backend_pid()").Scan(&insertBackendPID))
@@ -206,6 +207,7 @@ func TestGroupUsageRollupTriggerKeepsWatermarkForTodayInsert(t *testing.T) {
tx := beginGroupUsageRollupTriggerTestTx(t, ctx, schema)
defer func() { _ = tx.Rollback() }()
require.NoError(t, setGroupUsageRollupTriggerTimeZone(ctx, tx, "Asia/Shanghai"))
_, err := tx.ExecContext(ctx, `
INSERT INTO groups (id) VALUES (10);
INSERT INTO users (id) VALUES (1);
@@ -476,6 +478,11 @@ func setGroupUsageRollupTriggerSearchPath(ctx context.Context, tx *sql.Tx, quote
return err
}
func setGroupUsageRollupTriggerTimeZone(ctx context.Context, tx *sql.Tx, name string) error {
_, err := tx.ExecContext(ctx, "SET LOCAL TIME ZONE "+pq.QuoteLiteral(name))
return err
}
func waitForGroupUsageRollupStateLock(
ctx context.Context,
backendPID int,
+18 -6
View File
@@ -99,6 +99,7 @@ const (
upstreamProtocolModeOpenAIH1 = "openai_h1"
upstreamProtocolModeOpenAIH2 = "openai_h2"
upstreamProtocolModeOpenAIH1Fallback = "openai_h1_fallback"
upstreamProtocolModeGrok = "grok"
)
var errUpstreamClientLimitReached = errors.New("upstream client cache limit reached")
@@ -899,12 +900,20 @@ func (s *httpUpstreamService) resolvePoolSettings(isolation string, accountConcu
}
func (s *httpUpstreamService) applyProfilePoolSettings(settings poolSettings, profile service.HTTPUpstreamProfile) poolSettings {
if profile != service.HTTPUpstreamProfileOpenAI {
return settings
}
settings.responseHeaderTimeout = 0
if s != nil && s.cfg != nil && s.cfg.Gateway.OpenAIResponseHeaderTimeout > 0 {
settings.responseHeaderTimeout = time.Duration(s.cfg.Gateway.OpenAIResponseHeaderTimeout) * time.Second
switch profile {
case service.HTTPUpstreamProfileOpenAI:
settings.responseHeaderTimeout = 0
if s != nil && s.cfg != nil && s.cfg.Gateway.OpenAIResponseHeaderTimeout > 0 {
settings.responseHeaderTimeout = time.Duration(s.cfg.Gateway.OpenAIResponseHeaderTimeout) * time.Second
}
case service.HTTPUpstreamProfileGrok:
// Grok can stall before its first byte under capacity pressure. Keep the
// generic 600s gateway timeout from turning one request into a 10-minute
// resource hold; streaming after headers is unaffected.
settings.responseHeaderTimeout = 120 * time.Second
if s != nil && s.cfg != nil {
settings.responseHeaderTimeout = time.Duration(s.cfg.Gateway.GrokResponseHeaderTimeout) * time.Second
}
}
return settings
}
@@ -983,6 +992,9 @@ func (s *httpUpstreamService) resolveOpenAIHTTP2Settings() openAIHTTP2Settings {
}
func (s *httpUpstreamService) resolveProtocolMode(profile service.HTTPUpstreamProfile, proxyKey string, parsedProxy *url.URL) string {
if profile == service.HTTPUpstreamProfileGrok {
return upstreamProtocolModeGrok
}
if profile != service.HTTPUpstreamProfileOpenAI {
return upstreamProtocolModeDefault
}
+2 -2
View File
@@ -886,7 +886,7 @@ func TestAPIContracts(t *testing.T) {
"invitation_code_enabled": false,
"home_content": "",
"hide_ccs_import_button": false,
"grok_default_text_model": "grok-4.5",
"grok_default_text_model": "grok-4.6",
"grok_default_base_url_mode": "cli",
"grok_cross_client_model_map_enabled": true,
"purchase_subscription_enabled": false,
@@ -1169,7 +1169,7 @@ func TestAPIContracts(t *testing.T) {
"doc_url": "",
"home_content": "",
"hide_ccs_import_button": false,
"grok_default_text_model": "grok-4.5",
"grok_default_text_model": "grok-4.6",
"grok_default_base_url_mode": "cli",
"grok_cross_client_model_map_enabled": true,
"purchase_subscription_enabled": false,
@@ -101,7 +101,7 @@ const (
AccountTestModeGrokRealtime = "realtime"
defaultGrokRealtimeTestModel = "grok-voice-latest"
grokRealtimeProbeTimeout = 12 * time.Second
grokRealtimeProbeTimeout = DefaultGrokRealtimeDialTimeout
)
// isOpenAIImageModel checks if the model is an OpenAI image generation model (e.g. gpt-image-2).
+45 -12
View File
@@ -667,7 +667,8 @@ func (s *BillingService) initFallbackPricing() {
SupportsCacheBreakdown: false,
}
// xAI Grok 4.5: $2 input / $0.30 cached input / $6 output below 200k.
// xAI Grok 4.5: $2 input / $0.30 cached input / $6 output below 200k;
// long-context rates are $4 / $0.60 / $12 (>=200k prompt tokens).
s.fallbackPrices["grok-4.5"] = &ModelPricing{
InputPricePerToken: 2e-6,
OutputPricePerToken: 6e-6,
@@ -679,9 +680,8 @@ func (s *BillingService) initFallbackPricing() {
LongContextOutputMultiplier: 2,
}
// xAI Grok 4.6 (docs.x.ai/developers/models: $2 input / $0.50 cached input /
// $6 output per MTok under 200k prompt tokens; ≥200k is 2× on input,
// cached input, and output).
// xAI Grok 4.6: $2 input / $0.50 cached input / $6 output below 200k;
// long-context rates are $4 / $1 / $12 (>=200k prompt tokens).
s.fallbackPrices["grok-4.6"] = &ModelPricing{
InputPricePerToken: 2e-6,
OutputPricePerToken: 6e-6,
@@ -693,7 +693,8 @@ func (s *BillingService) initFallbackPricing() {
LongContextOutputMultiplier: 2,
}
// xAI Grok 4.3: $1.25 input / $0.20 cached / $2.50 output below 200k.
// xAI Grok 4.3: $1.25 input / $0.20 cached / $2.50 output below 200k;
// long-context rates are $2.50 / $0.40 / $5.
s.fallbackPrices["grok-4.3"] = &ModelPricing{
InputPricePerToken: 1.25e-6,
OutputPricePerToken: 2.5e-6,
@@ -704,6 +705,33 @@ func (s *BillingService) initFallbackPricing() {
LongContextInputMultiplier: 2,
LongContextOutputMultiplier: 2,
}
// Grok 4.20 variants share the official $1.25 / $0.20 / $2.50 card
// (and $2.50 / $0.40 / $5 long-context rates) with Grok 4.3.
s.fallbackPrices["grok-4.20"] = &ModelPricing{
InputPricePerToken: 1.25e-6,
OutputPricePerToken: 2.5e-6,
CacheReadPricePerToken: 0.2e-6,
SupportsCacheBreakdown: false,
LongContextInputThreshold: 200000,
LongContextThresholdInclusive: true,
LongContextInputMultiplier: 2,
LongContextOutputMultiplier: 2,
}
// Keep legacy Grok 3 Mini requests on their own historical xAI price card;
// otherwise the generic Grok fallback bills them as Grok 4.5.
s.fallbackPrices["grok-3-mini"] = &ModelPricing{
InputPricePerToken: 0.30e-6,
OutputPricePerToken: 0.50e-6,
CacheReadPricePerToken: 0.075e-6,
SupportsCacheBreakdown: false,
}
s.fallbackPrices["grok-3-mini-fast"] = &ModelPricing{
InputPricePerToken: 0.60e-6,
OutputPricePerToken: 4e-6,
CacheReadPricePerToken: 0.15e-6,
SupportsCacheBreakdown: false,
}
// xAI Grok Build 0.1 (official docs: $1 input / $0.20 cached input /
// $2 output per MTok). Composer is available only through Grok Build and
// has no standalone public API rate card, so its aliases use this coding
@@ -909,17 +937,22 @@ func (s *BillingService) getFallbackPricing(model string) *ModelPricing {
}
switch modelLower {
case "grok", "grok-latest", "grok-4.5", "grok-4.5-latest":
return s.fallbackPrices["grok-4.5"]
case "grok-4.6", "grok-4.6-latest":
case "grok", "grok-latest", "grok-4.6", "grok-4.6-latest":
return s.fallbackPrices["grok-4.6"]
case "grok-4.3",
"grok-4.20-0309-reasoning",
case "grok-4.5", "grok-4.5-latest":
return s.fallbackPrices["grok-4.5"]
case "grok-3-mini":
return s.fallbackPrices["grok-3-mini"]
case "grok-3-mini-fast":
return s.fallbackPrices["grok-3-mini-fast"]
case "grok-4.3":
return s.fallbackPrices["grok-4.3"]
case "grok-4.20-0309-reasoning",
"grok-4.20-0309-non-reasoning",
"grok-4.20-multi-agent-0309",
"grok-4.20-reasoning",
"grok-4.20-non-reasoning":
return s.fallbackPrices["grok-4.3"]
return s.fallbackPrices["grok-4.20"]
case "grok-build", "grok-build-latest", "grok-build-0.1", "grok-composer", "grok-composer-2.5-fast", "composer-2.5":
return s.fallbackPrices["grok-build-0.1"]
}
@@ -937,7 +970,7 @@ func (s *BillingService) grokUnknownTextFamilyFallback(model string) *ModelPrici
if s == nil || !isGrokUnknownTextFamilyModel(model) {
return nil
}
return s.fallbackPrices["grok-4.5"]
return s.fallbackPrices["grok-4.6"]
}
func isGrokUnknownTextFamilyModel(model string) bool {
@@ -1179,7 +1179,7 @@ func TestCalculateCostWithLongContext_PropagatesError(t *testing.T) {
func TestGetModelPricing_Grok45OfficialFallback(t *testing.T) {
svc := newTestBillingService()
for _, model := range []string{"grok", "grok-latest", "grok-4.5", "grok-4.5-latest"} {
for _, model := range []string{"grok-4.5", "grok-4.5-latest"} {
model := model
t.Run(model, func(t *testing.T) {
pricing, err := svc.GetModelPricing(model)
@@ -1192,6 +1192,17 @@ func TestGetModelPricing_Grok45OfficialFallback(t *testing.T) {
}
}
func TestGetModelPricing_GrokBareAliasesUseGrok46(t *testing.T) {
svc := newTestBillingService()
for _, model := range []string{"grok", "grok-latest"} {
pricing, err := svc.GetModelPricing(model)
require.NoError(t, err)
require.InDelta(t, 2e-6, pricing.InputPricePerToken, 1e-12)
require.InDelta(t, 0.5e-6, pricing.CacheReadPricePerToken, 1e-12)
require.InDelta(t, 6e-6, pricing.OutputPricePerToken, 1e-12)
}
}
func TestGetModelPricing_Grok46OfficialFallback(t *testing.T) {
svc := newTestBillingService()
@@ -1211,6 +1222,25 @@ func TestGetModelPricing_Grok46OfficialFallback(t *testing.T) {
}
}
func TestGetModelPricing_GrokOfficialFamilyCards(t *testing.T) {
svc := newTestBillingService()
for _, tc := range []struct {
model string
input, cached, output float64
}{
{"grok-4.3", 1.25e-6, 0.2e-6, 2.5e-6},
{"grok-4.20-0309-reasoning", 1.25e-6, 0.2e-6, 2.5e-6},
{"grok-build-0.1", 1e-6, 0.2e-6, 2e-6},
} {
p, err := svc.GetModelPricing(tc.model)
require.NoError(t, err, tc.model)
require.InDelta(t, tc.input, p.InputPricePerToken, 1e-12)
require.InDelta(t, tc.cached, p.CacheReadPricePerToken, 1e-12)
require.InDelta(t, tc.output, p.OutputPricePerToken, 1e-12)
require.Equal(t, 200000, p.LongContextInputThreshold)
}
}
func TestCalculateCostUnified_GroupLongContextToggleUsesPresetLadder(t *testing.T) {
svc := newTestBillingService()
resolver := NewModelPricingResolver(nil, svc)
@@ -1234,9 +1264,9 @@ func TestCalculateCostUnified_GroupLongContextToggleUsesPresetLadder(t *testing.
require.InDelta(t, disabled.OutputCost*2, enabled.OutputCost, 1e-12)
}
func TestGetModelPricing_UnknownGrokTextFallsBackToGrok45(t *testing.T) {
func TestGetModelPricing_UnknownGrokTextFallsBackToGrok46(t *testing.T) {
svc := newTestBillingService()
baseline, err := svc.GetModelPricing("grok-4.5")
baseline, err := svc.GetModelPricing("grok-4.6")
require.NoError(t, err)
for _, model := range []string{"grok-5", "grok-5-latest", "x-ai/grok-7", "grok-4.7-beta"} {
+9 -8
View File
@@ -671,14 +671,15 @@ type GatewayFailureReason string
// source-compatible and preserves their legacy retry-next-account behavior.
type UpstreamFailoverError struct {
StatusCode int
ResponseBody []byte // 上游响应体,用于错误透传规则匹配
ResponseHeaders http.Header // 上游响应头,用于透传 cf-ray/cf-mitigated/content-type 等诊断信息
ForceCacheBilling bool // Antigravity 粘性会话切换时设为 true
RetryableOnSameAccount bool // 临时性错误(如 Google 间歇性 400、空响应),应在同一账号上重试 N 次再切换
SameAccountRetryDelay time.Duration
SameAccountRetryDeadline time.Time
RequestScopedTransient bool // 故障因素与账号无关(如上游按客户端身份/模型容量降载):可同账号重试,但不得据此对账号做临时封禁
SafeToFailoverAfterWrite bool // 仅写出 SSE 注释等非语义字节时,仍可在同一客户端流中切换账号
ResponseBody []byte // 上游响应体,用于错误透传规则匹配
ResponseHeaders http.Header // 上游响应头,用于透传 cf-ray/cf-mitigated/content-type 等诊断信息
ForceCacheBilling bool // Antigravity 粘性会话切换时设为 true
RetryableOnSameAccount bool // 临时性错误(如 Google 间歇性 400、空响应),应在同一账号上重试 N 次再切换
SameAccountRetryDelay time.Duration // 同账号重试的最小间隔;零值使用 handler 默认值
SameAccountRetryDeadline time.Time // 同账号重试截止时间;零值表示仅受 retryLimit 限制
SameAccountRetryMax int // 可选的错误级同账号重试上限,低于 handler 默认预算时优先采用
RequestScopedTransient bool // 故障因素与账号无关(如上游按客户端身份/模型容量降载):可同账号重试,但不得据此对账号做临时封禁
SafeToFailoverAfterWrite bool // 仅写出 SSE 注释等非语义字节时,仍可在同一客户端流中切换账号
Stage GatewayFailureStage
Scope GatewayFailureScope
Reason GatewayFailureReason
+114 -27
View File
@@ -16,6 +16,10 @@ import (
"github.com/tidwall/gjson"
)
// DefaultGrokRealtimeDialTimeout bounds the pre-accept upstream handshake.
// The timeout only covers dialing; an established session is not interrupted.
const DefaultGrokRealtimeDialTimeout = 12 * time.Second
// supportedGrokVoiceHTTPEndpoints are xAI Voice HTTP paths we forward as-is.
var supportedGrokVoiceHTTPEndpoints = map[string]struct{}{
"tts": {},
@@ -125,35 +129,79 @@ func (s *OpenAIGatewayService) ProxyGrokRealtime(ctx context.Context, c *gin.Con
if account.Platform != PlatformGrok {
return false, fmt.Errorf("account platform %s is not supported for grok realtime", account.Platform)
}
base, err := buildGrokVoiceURL(account, s.cfg, "realtime")
if err != nil {
return false, err
}
u, err := url.Parse(base)
if err != nil {
return false, err
}
u.Scheme = "wss"
u.RawQuery = "model=" + url.QueryEscape(firstNonEmpty(model, "grok-voice-latest"))
headers := http.Header{"Authorization": []string{"Bearer " + token}}
// Match media/voice HTTP: CLI headers only on CLI proxy hosts.
if account.IsGrokOAuth() && isGrokCLIProxyTarget(u.String()) {
applyGrokCLIHeaders(headers)
}
if account != nil {
account.ApplyHeaderOverrides(headers)
}
dialer := s.getOpenAIWSPassthroughDialer()
proxyURL := ""
if account.ProxyID != nil && account.Proxy != nil {
proxyURL = account.Proxy.URL()
}
upstream, _, _, err := dialer.Dial(ctx, u.String(), headers, proxyURL)
upstream, err := s.OpenGrokRealtime(ctx, account, token, model)
if err != nil {
return false, err
}
defer func() { _ = upstream.Close() }()
return s.ProxyGrokRealtimeConn(ctx, c, client, upstream)
}
type GrokRealtimeUpstream struct{ conn openAIWSClientConn }
// GrokRealtimeDialError preserves an HTTP status returned before WebSocket
// upgrade so handlers can apply the normal Grok account policy.
type GrokRealtimeDialError struct {
StatusCode int
Err error
}
func (e *GrokRealtimeDialError) Error() string { return e.Err.Error() }
func (e *GrokRealtimeDialError) Unwrap() error { return e.Err }
func (u *GrokRealtimeUpstream) Close() error {
if u == nil || u.conn == nil {
return nil
}
return u.conn.Close()
}
func (s *OpenAIGatewayService) OpenGrokRealtime(ctx context.Context, account *Account, token, model string) (*GrokRealtimeUpstream, error) {
if s == nil || account == nil || account.Platform != PlatformGrok {
return nil, fmt.Errorf("grok realtime account is required")
}
base, err := buildGrokVoiceURL(account, s.cfg, "realtime")
if err != nil {
return nil, err
}
u, err := url.Parse(base)
if err != nil {
return nil, err
}
u.Scheme = "wss"
q := u.Query()
q.Set("model", firstNonEmpty(model, "grok-voice-latest"))
u.RawQuery = q.Encode()
headers := http.Header{"Authorization": []string{"Bearer " + token}}
if account.IsGrokOAuth() && isGrokCLIProxyTarget(u.String()) {
applyGrokCLIHeaders(headers)
}
account.ApplyHeaderOverrides(headers)
proxyURL := ""
if account.ProxyID != nil && account.Proxy != nil {
proxyURL = account.Proxy.URL()
}
conn, status, _, err := s.getOpenAIWSPassthroughDialer().Dial(ctx, u.String(), headers, proxyURL)
if err != nil {
return nil, &GrokRealtimeDialError{StatusCode: status, Err: err}
}
return &GrokRealtimeUpstream{conn: conn}, nil
}
// HandleGrokRealtimeUpstreamError applies the shared Grok account policy to a
// failed pre-accept WebSocket handshake.
func (s *OpenAIGatewayService) HandleGrokRealtimeUpstreamError(ctx context.Context, account *Account, statusCode int, body []byte) {
if statusCode <= 0 {
statusCode = http.StatusBadGateway
}
s.handleGrokAccountUpstreamError(ctx, account, statusCode, nil, body)
}
func (s *OpenAIGatewayService) ProxyGrokRealtimeConn(ctx context.Context, c *gin.Context, client *coderws.Conn, upstream *GrokRealtimeUpstream) (bool, error) {
if s == nil || client == nil || upstream == nil || upstream.conn == nil {
return false, fmt.Errorf("realtime connection is required")
}
conn := upstream.conn
ctx, cancel := context.WithCancel(ctx)
defer cancel()
@@ -163,7 +211,7 @@ func (s *OpenAIGatewayService) ProxyGrokRealtime(ctx context.Context, c *gin.Con
// Upstream → client
go func() {
for {
msg, readErr := upstream.ReadMessage(ctx)
msg, readErr := conn.ReadMessage(ctx)
if readErr != nil {
errCh <- readErr
return
@@ -197,7 +245,7 @@ func (s *OpenAIGatewayService) ProxyGrokRealtime(ctx context.Context, c *gin.Con
errCh <- fmt.Errorf("invalid realtime event: %w", unmarshalErr)
return
}
if writeErr := upstream.WriteJSON(ctx, raw); writeErr != nil {
if writeErr := conn.WriteJSON(ctx, raw); writeErr != nil {
errCh <- writeErr
return
}
@@ -207,6 +255,45 @@ func (s *OpenAIGatewayService) ProxyGrokRealtime(ctx context.Context, c *gin.Con
return awaitGrokRealtimeAudioObserved(errCh, &audioObserved)
}
// ProbeGrokRealtime performs the upstream WebSocket handshake without sending
// any client-visible events. Handlers use it before accepting the downstream
// upgrade so authentication and endpoint failures remain ordinary HTTP errors.
func (s *OpenAIGatewayService) ProbeGrokRealtime(ctx context.Context, account *Account, token, model string) error {
if s == nil || account == nil {
return fmt.Errorf("realtime service and account are required")
}
if account.Platform != PlatformGrok {
return fmt.Errorf("account platform %s is not supported for grok realtime", account.Platform)
}
base, err := buildGrokVoiceURL(account, s.cfg, "realtime")
if err != nil {
return err
}
u, err := url.Parse(base)
if err != nil {
return err
}
u.Scheme = "wss"
q := u.Query()
q.Set("model", firstNonEmpty(model, "grok-voice-latest"))
u.RawQuery = q.Encode()
headers := http.Header{"Authorization": []string{"Bearer " + token}}
if account.IsGrokOAuth() && isGrokCLIProxyTarget(u.String()) {
applyGrokCLIHeaders(headers)
}
account.ApplyHeaderOverrides(headers)
proxyURL := ""
if account.ProxyID != nil && account.Proxy != nil {
proxyURL = account.Proxy.URL()
}
dialer := s.getOpenAIWSPassthroughDialer()
conn, _, _, err := dialer.Dial(ctx, u.String(), headers, proxyURL)
if err != nil {
return err
}
return conn.Close()
}
func awaitGrokRealtimeAudioObserved(errCh <-chan error, audioObserved *atomic.Bool) (bool, error) {
err := <-errCh
if audioObserved == nil {
+25 -10
View File
@@ -59,6 +59,8 @@ type GrokMediaRequestInfo struct {
N int
Size string
SizeTier string
AspectRatio string
ImageResolution string
Resolution string
DurationSeconds int
InputImageURLs []string
@@ -127,6 +129,8 @@ func ParseGrokMediaRequest(contentType string, body []byte) GrokMediaRequestInfo
info.Prompt = strings.TrimSpace(info.Prompt)
info.Size = strings.TrimSpace(info.Size)
info.SizeTier = NormalizeImageBillingTierOrDefault(info.Size)
info.AspectRatio = strings.TrimSpace(info.AspectRatio)
info.ImageResolution = grokImagineImageResolution(info.ImageResolution)
info.Resolution = NormalizeVideoBillingResolutionOrDefault(info.Resolution)
info.DurationSeconds = NormalizeVideoBillingDurationSecondsOrDefault(info.DurationSeconds)
if info.N <= 0 {
@@ -142,7 +146,8 @@ func parseGrokMediaJSONRequest(body []byte, info *GrokMediaRequestInfo) {
info.Model = strings.TrimSpace(gjson.GetBytes(body, "model").String())
info.Prompt = strings.TrimSpace(gjson.GetBytes(body, "prompt").String())
info.Size = strings.TrimSpace(gjson.GetBytes(body, "size").String())
info.Resolution = strings.TrimSpace(gjson.GetBytes(body, "resolution").String())
info.AspectRatio = strings.TrimSpace(gjson.GetBytes(body, "aspect_ratio").String())
assignGrokMediaResolution(strings.TrimSpace(gjson.GetBytes(body, "resolution").String()), info)
if duration := gjson.GetBytes(body, "duration"); duration.Exists() && duration.Type == gjson.Number {
info.DurationSeconds = int(duration.Int())
}
@@ -255,8 +260,10 @@ func parseGrokMediaMultipartRequest(contentType string, body []byte, info *GrokM
info.Prompt = value
case "size":
info.Size = value
case "aspect_ratio":
info.AspectRatio = value
case "resolution":
info.Resolution = value
assignGrokMediaResolution(value, info)
case "duration":
if duration, err := strconv.Atoi(value); err == nil {
info.DurationSeconds = duration
@@ -934,6 +941,12 @@ func prepareGrokMediaForwardBody(endpoint GrokMediaEndpoint, body []byte, conten
if info.Size != "" {
payload["size"] = info.Size
}
if info.ImageResolution != "" {
payload["resolution"] = info.ImageResolution
}
if info.AspectRatio != "" {
payload["aspect_ratio"] = info.AspectRatio
}
images := make([]map[string]string, 0, len(info.InputImageURLs)+len(info.Uploads))
for _, imageURL := range info.InputImageURLs {
@@ -1106,10 +1119,7 @@ func sanitizeGrokMediaForwardBody(endpoint GrokMediaEndpoint, body []byte, conte
}
switch endpoint {
case GrokMediaEndpointImagesGenerations, GrokMediaEndpointImagesEdits:
if !gjson.GetBytes(body, "size").Exists() {
return body, contentType, nil
}
out, err := sjson.DeleteBytes(body, "size")
out, err := applyGrokImagineImageGeometry(body)
if err != nil {
return nil, "", fmt.Errorf("sanitize grok media size: %w", err)
}
@@ -1288,11 +1298,16 @@ func (s *OpenAIGatewayService) handleGrokMediaErrorResponse(
Detail: upstreamDetail,
})
if kind == "failover" {
retryable, retryDelay, retryDeadline, retryMax := grokSameAccountRetryMetadata(account, resp.StatusCode, body)
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: body,
ResponseHeaders: resp.Header.Clone(),
RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode),
StatusCode: resp.StatusCode,
ResponseBody: body,
ResponseHeaders: resp.Header.Clone(),
RetryableOnSameAccount: retryable,
RequestScopedTransient: retryable && resp.StatusCode == http.StatusTooManyRequests,
SameAccountRetryDelay: retryDelay,
SameAccountRetryDeadline: retryDeadline,
SameAccountRetryMax: retryMax,
}
}
@@ -0,0 +1,149 @@
package service
import (
"math"
"strconv"
"strings"
"github.com/tidwall/gjson"
"github.com/tidwall/sjson"
)
// Official Imagine image geometry: https://docs.x.ai/developers/model-capabilities/images/generation
var grokImagineAspectRatioValues = []struct {
label string
ratio float64
}{
{"1:1", 1},
{"16:9", 16.0 / 9.0},
{"9:16", 9.0 / 16.0},
{"4:3", 4.0 / 3.0},
{"3:4", 3.0 / 4.0},
{"3:2", 1.5},
{"2:3", 2.0 / 3.0},
{"2:1", 2},
{"1:2", 0.5},
{"19.5:9", 19.5 / 9.0},
{"9:19.5", 9.0 / 19.5},
{"20:9", 20.0 / 9.0},
{"9:20", 9.0 / 20.0},
}
func applyGrokImagineImageGeometry(body []byte) ([]byte, error) {
size := strings.TrimSpace(gjson.GetBytes(body, "size").String())
resolution := grokImagineImageResolution(gjson.GetBytes(body, "resolution").String())
aspect := strings.TrimSpace(gjson.GetBytes(body, "aspect_ratio").String())
out := append([]byte(nil), body...)
if resolution == "" {
if derived := grokImagineImageResolutionFromSize(size); derived != "" {
next, err := sjson.SetBytes(out, "resolution", derived)
if err != nil {
return nil, err
}
out = next
}
} else if gjson.GetBytes(body, "resolution").String() != resolution {
next, err := sjson.SetBytes(out, "resolution", resolution)
if err != nil {
return nil, err
}
out = next
}
if aspect == "" {
if derived := grokImagineAspectRatioFromSize(size); derived != "" {
next, err := sjson.SetBytes(out, "aspect_ratio", derived)
if err != nil {
return nil, err
}
out = next
}
}
if !gjson.GetBytes(out, "size").Exists() {
return out, nil
}
return sjson.DeleteBytes(out, "size")
}
func assignGrokMediaResolution(value string, info *GrokMediaRequestInfo) {
if info == nil {
return
}
value = strings.TrimSpace(value)
if value == "" {
return
}
if img := grokImagineImageResolution(value); img != "" {
info.ImageResolution = img
return
}
info.Resolution = value
}
func grokImagineImageResolution(value string) string {
switch strings.ToLower(strings.TrimSpace(value)) {
case "1k":
return "1k"
case "2k":
return "2k"
default:
return ""
}
}
func grokImagineImageResolutionFromSize(size string) string {
if explicit := grokImagineImageResolution(size); explicit != "" {
return explicit
}
tier, ok := ClassifyImageBillingTier(size)
if !ok {
return ""
}
if tier == ImageBillingSize1K {
return "1k"
}
return "2k"
}
func grokImagineAspectRatioFromSize(size string) string {
width, height, ok := parseImageBillingDimensions(strings.TrimSpace(size))
if !ok || width <= 0 || height <= 0 {
return ""
}
div := grokImagineGCD(width, height)
exact := strconv.Itoa(width/div) + ":" + strconv.Itoa(height/div)
for _, candidate := range grokImagineAspectRatioValues {
if candidate.label == exact {
return exact
}
}
ratio := float64(width) / float64(height)
bestLabel := ""
bestDelta := math.MaxFloat64
for _, candidate := range grokImagineAspectRatioValues {
delta := math.Abs(ratio - candidate.ratio)
if delta < bestDelta {
bestDelta = delta
bestLabel = candidate.label
}
}
return bestLabel
}
func grokImagineGCD(a, b int) int {
if a < 0 {
a = -a
}
if b < 0 {
b = -b
}
for b != 0 {
a, b = b, a%b
}
if a == 0 {
return 1
}
return a
}
@@ -0,0 +1,60 @@
package service
import (
"testing"
"github.com/stretchr/testify/require"
"github.com/tidwall/gjson"
)
func TestApplyGrokImagineImageGeometryMapsOpenAISize(t *testing.T) {
t.Parallel()
out, err := applyGrokImagineImageGeometry([]byte(`{"model":"grok-imagine-image-2.0","prompt":"hi","size":"1152x1536"}`))
require.NoError(t, err)
require.False(t, gjson.GetBytes(out, "size").Exists())
require.Equal(t, "2k", gjson.GetBytes(out, "resolution").String())
require.Equal(t, "3:4", gjson.GetBytes(out, "aspect_ratio").String())
}
func TestApplyGrokImagineImageGeometryKeepsClientGeometry(t *testing.T) {
t.Parallel()
out, err := applyGrokImagineImageGeometry([]byte(`{"size":"1024x1024","resolution":"2K","aspect_ratio":"16:9"}`))
require.NoError(t, err)
require.False(t, gjson.GetBytes(out, "size").Exists())
require.Equal(t, "2k", gjson.GetBytes(out, "resolution").String())
require.Equal(t, "16:9", gjson.GetBytes(out, "aspect_ratio").String())
}
func TestSanitizeGrokMediaForwardBodyConvertsImageSize(t *testing.T) {
t.Parallel()
out, contentType, err := sanitizeGrokMediaForwardBody(
GrokMediaEndpointImagesGenerations,
[]byte(`{"model":"grok-imagine-image","prompt":"hi","size":"1024x1024"}`),
"application/json",
)
require.NoError(t, err)
require.Equal(t, "application/json", contentType)
require.False(t, gjson.GetBytes(out, "size").Exists())
require.Equal(t, "1k", gjson.GetBytes(out, "resolution").String())
require.Equal(t, "1:1", gjson.GetBytes(out, "aspect_ratio").String())
}
func TestParseGrokMediaRequestKeepsImageResolutionOutOfVideoNormalize(t *testing.T) {
t.Parallel()
info := ParseGrokMediaRequest("application/json", []byte(`{"model":"grok-imagine-image-2.0","resolution":"2K","aspect_ratio":"16:9"}`))
require.Equal(t, "2k", info.ImageResolution)
require.Equal(t, "16:9", info.AspectRatio)
require.Equal(t, VideoBillingResolution480P, info.Resolution)
}
func TestGrokImagineAspectRatioFromSize(t *testing.T) {
t.Parallel()
require.Equal(t, "1:1", grokImagineAspectRatioFromSize("1024x1024"))
require.Equal(t, "3:4", grokImagineAspectRatioFromSize("1152x1536"))
require.Equal(t, "4:3", grokImagineAspectRatioFromSize("1536x1152"))
require.Equal(t, "16:9", grokImagineAspectRatioFromSize("1792x1024"))
}
+11 -2
View File
@@ -32,7 +32,16 @@ func grokStreamIdleFailoverError(account *Account, idle time.Duration) *Upstream
StatusCode: 502,
ResponseBody: []byte(`{"error":{"code":"empty_upstream","message":"` + strings.ReplaceAll(msg, `"`, `'`) + `"}}`),
SafeToFailoverAfterWrite: true,
// Allow pool-mode retries; normal OAuth switches account via handler.
RetryableOnSameAccount: account != nil && account.IsPoolMode(),
// An idle upstream stream is transient and should get the configured
// same-account retry budget before switching credentials. This applies
// to both pooled and dedicated Grok accounts; the handler still enforces
// the request's retry limit.
RetryableOnSameAccount: account != nil && account.Platform == PlatformGrok,
RequestScopedTransient: true,
SameAccountRetryMax: 1,
// Permit at most one same-account replay after the idle failure. The
// deadline is anchored at failure time, so a hung stream cannot consume
// the normal three-attempt budget before failover.
SameAccountRetryDeadline: time.Now().Add(idle),
}
}
@@ -21,5 +21,16 @@ func TestGrokStreamIdleFailoverError(t *testing.T) {
require.NotNil(t, err)
require.Equal(t, 502, err.StatusCode)
require.True(t, err.SafeToFailoverAfterWrite)
require.True(t, err.RetryableOnSameAccount)
require.True(t, err.RequestScopedTransient)
require.Equal(t, 1, err.SameAccountRetryMax)
require.Contains(t, string(err.ResponseBody), "empty_upstream")
require.WithinDuration(t, time.Now().Add(180*time.Second), err.SameAccountRetryDeadline, 2*time.Second)
}
func TestGrokStreamIdleFailoverErrorRequiresGrokAccount(t *testing.T) {
openAI := &Account{ID: 2, Platform: PlatformOpenAI, Type: AccountTypeOAuth}
err := grokStreamIdleFailoverError(openAI, time.Second)
require.False(t, err.RetryableOnSameAccount)
require.True(t, err.RequestScopedTransient)
}
@@ -116,8 +116,9 @@ func isGrokAccountAccessCode(value string) bool {
"subscription_required",
"entitlement_required",
"not_entitled",
"plan_required",
"permission_denied":
"plan_required":
// permission-denied is omitted: xAI reuses it for both entitlement
// refusals and request-scoped safety blocks, so the message decides.
return true
default:
return false
@@ -173,6 +174,7 @@ func grokContentPolicyMessage(value string) bool {
"prompt violates policy",
"input violates content policy",
"input violates policy",
"violates usage guidelines",
} {
if strings.Contains(lower, phrase) {
return true
@@ -207,7 +209,7 @@ func (s *OpenAIGatewayService) shouldFailoverGrokUpstreamError(statusCode int, r
}
decision := classifyGrokUpstreamFailure(statusCode, responseBody, "")
switch decision.Class {
case GrokFailureFreeUsage, GrokFailureEmptyUpstream, GrokFailureBilling, GrokFailureModelCapacity:
case GrokFailureFreeUsage, GrokFailureEmptyUpstream, GrokFailureBilling, GrokFailureModelCapacity, GrokFailureCompatibility:
return decision.ShouldFailover
}
return s.shouldFailoverUpstreamError(statusCode)
@@ -83,6 +83,24 @@ func TestIsGrokContentPolicyRejection(t *testing.T) {
body: `{"error":{"code":"policy_violation","message":"request blocked by policy"}}`,
want: true,
},
{
name: "permission-denied usage guidelines is request scoped",
status: http.StatusForbidden,
body: `{"code":"permission-denied","error":"Content violates usage guidelines. "}`,
want: true,
},
{
name: "permission-denied entitlement stays on the account path",
status: http.StatusForbidden,
body: `{"code":"permission-denied","error":"Access to the chat endpoint is denied"}`,
want: false,
},
{
name: "structured account code overrides usage guidelines phrase",
status: http.StatusForbidden,
body: `{"error":{"code":"account_suspended","message":"Content violates usage guidelines."}}`,
want: false,
},
{
name: "wrong status",
status: http.StatusBadRequest,
@@ -279,6 +297,21 @@ func TestGrokContentPolicySSEErrorDoesNotMutateOrFailover(t *testing.T) {
require.False(t, svc.isOpenAIAccountRuntimeBlocked(account))
}
func TestGrokPermissionDeniedContentRefusalDoesNotMutateOrFailover(t *testing.T) {
repo := &grokQuotaAccountRepo{}
svc := &OpenAIGatewayService{accountRepo: repo}
account := &Account{ID: 4785, Platform: PlatformGrok, Type: AccountTypeOAuth}
body := []byte(`{"code":"permission-denied","error":"Content violates usage guidelines. "}`)
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))
}
func TestHandleGrokAccountUpstreamErrorEntitlement403KeepsDefaultCooldown(t *testing.T) {
repo := &grokQuotaAccountRepo{}
svc := &OpenAIGatewayService{accountRepo: repo}
@@ -26,6 +26,10 @@ const (
GrokFailureRateLimit GrokUpstreamFailureClass = "rate_limit"
GrokFailureAuth GrokUpstreamFailureClass = "auth_error"
GrokFailureServer GrokUpstreamFailureClass = "server_error"
// GrokFailureCompatibility represents a request-history/body-shape that is
// incompatible with the selected account or upstream replay contract. It
// is account-independent: fail over, but never quarantine the pool.
GrokFailureCompatibility GrokUpstreamFailureClass = "compatibility_error"
)
// GrokUpstreamFailureDecision is a pure classification result. Callers map it
@@ -111,6 +115,20 @@ func classifyGrokUpstreamFailure(statusCode int, responseBody []byte, requestedM
}
}
// Responses replay/compaction payloads can be rejected by one Grok
// deployment while the same request is valid on another account. Treat
// these precise decoder/content-shape failures as account compatibility
// errors, rather than durable account health failures.
if isGrokCompatibilityError(statusCode, low, code) {
return GrokUpstreamFailureDecision{
Class: GrokFailureCompatibility,
Model: model,
ShouldFailover: true,
ShouldCooldown: false,
Reason: firstNonEmpty(text, "grok response compatibility error"),
}
}
// Empty HTTP 200 / empty model output (often rewritten to synthetic 502).
if isGrokEmptyModelOutputText(low) || isGrokEmptyModelOutputCode(code) {
return GrokUpstreamFailureDecision{
@@ -129,7 +147,7 @@ func classifyGrokUpstreamFailure(statusCode int, responseBody []byte, requestedM
return GrokUpstreamFailureDecision{
Class: GrokFailureModelCapacity,
Model: model,
Cooldown: 3 * time.Minute,
Cooldown: time.Minute,
ShouldCooldown: true,
ShouldFailover: true,
BlockModel: false,
@@ -355,7 +373,8 @@ func isGrokBillingQuotaText(low string) bool {
if strings.Contains(low, "payment") && (strings.Contains(low, "required") || strings.Contains(low, "fail")) {
return true
}
if strings.Contains(low, "spending limit") || strings.Contains(low, "run out of credits") || strings.Contains(low, "out of credits") {
if strings.Contains(low, "spending limit") || strings.Contains(low, "run out of credits") || strings.Contains(low, "out of credits") ||
(strings.Contains(low, "need a grok subscription") || strings.Contains(low, "need grok subscription")) {
return true
}
if strings.Contains(low, "余额不足") || strings.Contains(low, "欠费") || strings.Contains(low, "需要付费") {
@@ -364,6 +383,86 @@ func isGrokBillingQuotaText(low string) bool {
return false
}
// grokRetryableOnSameAccount marks transient 429 classes for the shared
// failover loop. Capacity and ordinary throttles are request/model pressure,
// not evidence that the credential is invalid, so a bounded retry on the same
// account is preferable before switching accounts. Free-usage and billing
// exhaustion deliberately skip same-account retry and fail over immediately.
func grokRetryableOnSameAccount(account *Account, statusCode int, responseBody []byte) bool {
if account == nil || !account.IsGrok() {
return false
}
decision := classifyGrokUpstreamFailure(statusCode, responseBody, "")
switch decision.Class {
case GrokFailureFreeUsage, GrokFailureBilling, GrokFailureCompatibility:
// Quota/entitlement exhaustion is account state, not transient
// pressure. Retrying the same account only repeats the failure.
return false
case GrokFailureModelCapacity:
if statusCode == http.StatusTooManyRequests {
return true
}
}
return account.IsPoolMode() && account.IsPoolModeRetryableStatus(statusCode)
}
func grokSameAccountRetryMetadata(account *Account, statusCode int, responseBody []byte) (bool, time.Duration, time.Time, int) {
if !grokRetryableOnSameAccount(account, statusCode, responseBody) {
return false, 0, time.Time{}, 0
}
decision := classifyGrokUpstreamFailure(statusCode, responseBody, "")
if decision.Class != GrokFailureModelCapacity {
return true, 0, time.Time{}, 0
}
// The error is reconstructed after every upstream attempt, so a deadline
// stored on the error cannot provide a request-wide window. Cap capacity
// retries explicitly to one replay; this remains effective even when the
// first attempt itself takes longer than the nominal 30-second window.
return true, 500 * time.Millisecond, time.Now().Add(30 * time.Second), 1
}
// shouldMarkGrokTeamModelRateLimit controls the process-local sibling-account
// overlay. Model-capacity responses are request pressure, not a team quota;
// marking them would hide healthy sibling credentials while the bounded
// same-account retry is still in progress. Ordinary 429s and free-usage
// exhaustion retain the existing quota/team isolation behavior.
func shouldMarkGrokTeamModelRateLimit(statusCode int, responseBody []byte) bool {
decision := classifyGrokUpstreamFailure(statusCode, responseBody, "")
if decision.Class == GrokFailureModelCapacity {
return false
}
return statusCode == http.StatusTooManyRequests || decision.Class == GrokFailureFreeUsage
}
func isGrokCompatibilityError(statusCode int, low, code string) bool {
if statusCode != http.StatusBadRequest && statusCode != http.StatusUnprocessableEntity {
return false
}
combined := strings.ToLower(strings.TrimSpace(low + " " + code))
// Compaction blobs are account/session-bound and frequently fail with 400
// or 422 after a reconnect. Also cover xAI's JSON decoder shape errors.
for _, phrase := range []string{
"could not decode the compaction blob",
"cannot decode the compaction blob",
"decode the compaction blob",
"ensure it is unmodified from the compact response",
"compaction blob",
} {
if strings.Contains(combined, phrase) {
return true
}
}
for _, marker := range []string{
"invalid_compaction",
"compaction_decode_error",
} {
if strings.Contains(combined, marker) {
return true
}
}
return false
}
func isGrokModelCapacityText(low string) bool {
return strings.Contains(low, "capacity") ||
strings.Contains(low, "overloaded") ||
@@ -489,10 +588,11 @@ func (s *OpenAIGatewayService) applyGrokUpstreamFailureDecision(
case GrokFailureEmptyUpstream:
reason = "grok empty model output"
case GrokFailureModelCapacity:
if persistGrokTransientModelCooldown(account, decision) {
return true
}
reason = "grok model capacity"
// Capacity is scoped to the requested model. Never persist an account-wide
// unschedulable state for this transient class; the failover loop performs
// a bounded same-account retry before selecting another account.
_ = persistGrokTransientModelCooldown(account, decision)
return true
case GrokFailureRateLimit:
// Pure 429 without free-usage language keeps the existing rate-limit
// snapshot path (Retry-After / quota headers). Body-only rate-limit
@@ -501,6 +601,11 @@ func (s *OpenAIGatewayService) applyGrokUpstreamFailureDecision(
return false
case GrokFailureServer:
reason = "grok upstream temporary error"
case GrokFailureCompatibility:
// Deliberately no account mutation. The caller uses ShouldFailover to
// retry another account; cooling a pool for a request-shape mismatch
// would remove healthy accounts.
return true
default:
return false
}
@@ -54,6 +54,14 @@ func TestClassifyGrokUpstreamFailure_EmptyUpstream(t *testing.T) {
require.Equal(t, 4*time.Minute, d.Cooldown)
}
func TestClassifyGrokUpstreamFailure_ModelCapacityUsesShortCooldown(t *testing.T) {
d := classifyGrokUpstreamFailure(http.StatusTooManyRequests,
[]byte(`{"error":{"message":"The model is currently at capacity due to high demand"}}`), "grok-4.6")
require.Equal(t, GrokFailureModelCapacity, d.Class)
require.Equal(t, time.Minute, d.Cooldown)
require.False(t, d.BlockModel)
}
func TestClassifyGrokUpstreamFailure_Billing(t *testing.T) {
d := classifyGrokUpstreamFailure(http.StatusForbidden, []byte(`{"code":"personal-team-blocked:spending-limit","error":"spending limit reached"}`), "")
require.Equal(t, GrokFailureBilling, d.Class)
@@ -61,6 +69,62 @@ func TestClassifyGrokUpstreamFailure_Billing(t *testing.T) {
require.True(t, d.ShouldFailover)
}
func TestClassifyGrokUpstreamFailure_GrokSubscriptionRequiredIsBilling(t *testing.T) {
d := classifyGrokUpstreamFailure(http.StatusPaymentRequired,
[]byte(`{"error":{"message":"You have run out of credits or need a Grok subscription"}}`), "grok-4.6")
require.Equal(t, GrokFailureBilling, d.Class)
require.True(t, d.ShouldFailover)
require.True(t, d.ShouldCooldown)
}
func TestGrokRetryableOnSameAccount_CapacityAndRateLimit(t *testing.T) {
account := &Account{ID: 9105, Platform: PlatformGrok, Type: AccountTypeOAuth}
require.True(t, grokRetryableOnSameAccount(account, http.StatusTooManyRequests,
[]byte(`{"error":{"message":"The model is currently at capacity due to high demand"}}`)))
require.False(t, grokRetryableOnSameAccount(account, http.StatusTooManyRequests,
[]byte(`{"error":{"message":"rate limit exceeded"}}`)))
require.False(t, grokRetryableOnSameAccount(account, http.StatusPaymentRequired,
[]byte(`{"error":{"message":"You have run out of credits or need a Grok subscription"}}`)))
poolAccount := &Account{ID: 9108, Platform: PlatformGrok, Type: AccountTypeOAuth,
Credentials: map[string]any{"pool_mode": true}}
require.False(t, grokRetryableOnSameAccount(poolAccount, http.StatusTooManyRequests,
[]byte(`{"error":{"code":"subscription:free-usage-exhausted"}}`)),
"pool free-usage must fail over instead of retrying the exhausted account")
require.False(t, grokRetryableOnSameAccount(account, http.StatusBadRequest,
[]byte(`{"error":{"message":"capacity field is invalid"}}`)))
nonGrok := &Account{ID: 9106, Platform: PlatformOpenAI, Type: AccountTypeOAuth}
require.False(t, grokRetryableOnSameAccount(nonGrok, http.StatusTooManyRequests,
[]byte(`{"error":{"message":"model at capacity"}}`)))
}
func TestShouldMarkGrokTeamModelRateLimit_ExcludesCapacity(t *testing.T) {
require.False(t, shouldMarkGrokTeamModelRateLimit(http.StatusTooManyRequests,
[]byte(`{"error":{"message":"The model is currently at capacity due to high demand"}}`)))
require.True(t, shouldMarkGrokTeamModelRateLimit(http.StatusTooManyRequests,
[]byte(`{"error":{"message":"rate limit exceeded"}}`)))
require.True(t, shouldMarkGrokTeamModelRateLimit(http.StatusBadRequest,
[]byte(`{"error":{"code":"subscription:free-usage-exhausted"}}`)))
require.False(t, shouldMarkGrokTeamModelRateLimit(http.StatusBadRequest,
[]byte(`{"error":{"message":"invalid request"}}`)))
}
func TestGrokSameAccountRetryMetadata_CapacityDeadline(t *testing.T) {
account := &Account{ID: 9107, Platform: PlatformGrok, Type: AccountTypeOAuth}
retryable, delay, deadline, retryMax := grokSameAccountRetryMetadata(account, http.StatusTooManyRequests,
[]byte(`{"error":{"message":"model capacity exceeded"}}`))
require.True(t, retryable)
require.Equal(t, 500*time.Millisecond, delay)
require.WithinDuration(t, time.Now().Add(30*time.Second), deadline, 2*time.Second)
require.Equal(t, 1, retryMax)
retryable, delay, deadline, retryMax = grokSameAccountRetryMetadata(account, http.StatusTooManyRequests,
[]byte(`{"error":{"message":"rate limit exceeded"}}`))
require.False(t, retryable)
require.Zero(t, delay)
require.True(t, deadline.IsZero())
require.Zero(t, retryMax)
}
func TestClassifyGrokUpstreamFailure_ValidationNoCool(t *testing.T) {
d := classifyGrokUpstreamFailure(http.StatusBadRequest, []byte(`{"error":{"message":"invalid tool schema"}}`), "")
require.Equal(t, GrokFailureNone, d.Class)
@@ -75,12 +139,48 @@ func TestClassifyGrokUpstreamFailure_FreeUsageWinsOver5xx(t *testing.T) {
require.NotEqual(t, GrokFailureServer, d.Class)
}
func TestClassifyGrokUpstreamFailure_CompatibilityDoesNotCooldown(t *testing.T) {
cases := []string{
`{"error":{"message":"Could not decode the compaction blob. Ensure it is unmodified from the compact response"}}`,
`{"code":"compaction_decode_error","message":"invalid response history"}`,
}
for _, body := range cases {
d := classifyGrokUpstreamFailure(http.StatusUnprocessableEntity, []byte(body), "grok-4.6")
require.Equal(t, GrokFailureCompatibility, d.Class, body)
require.True(t, d.ShouldFailover, body)
require.False(t, d.ShouldCooldown, body)
require.Zero(t, d.Cooldown, body)
}
}
func TestClassifyGrokUpstreamFailure_CompatibilityRequiresClientError(t *testing.T) {
body := []byte(`{"error":{"message":"upstream failed while handling the compaction blob"}}`)
for _, status := range []int{http.StatusBadGateway, http.StatusInternalServerError} {
d := classifyGrokUpstreamFailure(status, body, "grok-4.6")
require.NotEqual(t, GrokFailureCompatibility, d.Class)
require.True(t, d.ShouldCooldown)
}
}
func TestClassifyGrokUpstreamFailure_GenericShapeErrorDoesNotFailover(t *testing.T) {
d := classifyGrokUpstreamFailure(http.StatusBadRequest,
[]byte(`{"error":{"message":"data did not match any variant of the untagged enum content"}}`), "grok-4.6")
require.NotEqual(t, GrokFailureCompatibility, d.Class)
require.False(t, d.ShouldFailover)
}
func TestShouldFailoverGrokUpstreamError_FreeUsageBody(t *testing.T) {
svc := &OpenAIGatewayService{}
body := []byte(`{"error":{"code":"subscription:free-usage-exhausted","message":"free usage exhausted"}}`)
require.True(t, svc.shouldFailoverGrokUpstreamError(http.StatusBadRequest, body))
}
func TestShouldFailoverGrokUpstreamError_CompatibilityBody(t *testing.T) {
svc := &OpenAIGatewayService{}
body := []byte(`{"error":{"message":"Could not decode the compaction blob"}}`)
require.True(t, svc.shouldFailoverGrokUpstreamError(http.StatusUnprocessableEntity, body))
}
func TestShouldFailoverGrokUpstreamError_ContentPolicyStillNoFailover(t *testing.T) {
svc := &OpenAIGatewayService{}
body := []byte(`{"error":{"code":"new_sensitive","message":"text is sensitive"}}`)
@@ -149,6 +249,19 @@ func TestHandleGrokAccountUpstreamError_MultiAgentCapacityBlocksOnlyThatModel(t
require.False(t, isGrokModelQuotaBlocked(account.ID, "grok-4.5", time.Now()))
}
func TestHandleGrokAccountUpstreamError_CapacityNeverCoolsAccount(t *testing.T) {
repo := &grokQuotaAccountRepo{}
svc := &OpenAIGatewayService{accountRepo: repo}
account := &Account{ID: 9121, Platform: PlatformGrok, Type: AccountTypeOAuth}
ctx := withGrokTeamRateLimitModel(context.Background(), "grok-4.6")
svc.handleGrokAccountUpstreamError(ctx, account, http.StatusTooManyRequests, nil,
[]byte(`{"error":{"message":"The model is currently at capacity due to high demand"}}`))
require.Zero(t, repo.tempUnschedCalls)
require.False(t, svc.isOpenAIAccountRuntimeBlocked(account))
}
func TestHandleGrokAccountUpstreamError_FreeUsageDoesNotCoolPoolMode(t *testing.T) {
repo := &grokQuotaAccountRepo{}
svc := &OpenAIGatewayService{accountRepo: repo}
@@ -9,6 +9,7 @@ type HTTPUpstreamProfile string
const (
HTTPUpstreamProfileDefault HTTPUpstreamProfile = ""
HTTPUpstreamProfileOpenAI HTTPUpstreamProfile = "openai"
HTTPUpstreamProfileGrok HTTPUpstreamProfile = "grok"
)
type httpUpstreamProfileContextKey struct{}
@@ -35,7 +36,7 @@ func HTTPUpstreamProfileFromContext(ctx context.Context) HTTPUpstreamProfile {
return HTTPUpstreamProfileDefault
}
switch profile {
case HTTPUpstreamProfileOpenAI:
case HTTPUpstreamProfileOpenAI, HTTPUpstreamProfileGrok:
return profile
default:
return HTTPUpstreamProfileDefault
@@ -201,11 +201,16 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions(
})
s.handleGrokAccountUpstreamError(withGrokTeamRateLimitModel(ctx, upstreamModel), account, resp.StatusCode, resp.Header, respBody)
if s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) {
retryable, retryDelay, retryDeadline, retryMax := grokSameAccountRetryMetadata(account, resp.StatusCode, respBody)
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: respBody,
ResponseHeaders: resp.Header.Clone(),
RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode),
StatusCode: resp.StatusCode,
ResponseBody: respBody,
ResponseHeaders: resp.Header.Clone(),
RetryableOnSameAccount: retryable,
RequestScopedTransient: retryable && resp.StatusCode == http.StatusTooManyRequests,
SameAccountRetryDelay: retryDelay,
SameAccountRetryDeadline: retryDeadline,
SameAccountRetryMax: retryMax,
}
}
return s.handleChatCompletionsErrorResponse(resp, c, account, billingModel)
+132 -43
View File
@@ -119,7 +119,7 @@ func (s *OpenAIGatewayService) forwardGrokResponses(
// xAI can reject encrypted reasoning or a compaction blob copied from a
// different decoder/cache context. Retry once on the same account after
// preserving visible summaries and removing only opaque replay state.
if attempt > 0 || resp.StatusCode != http.StatusBadRequest {
if attempt > 0 || (resp.StatusCode != http.StatusBadRequest && resp.StatusCode != http.StatusUnprocessableEntity) {
break
}
respBody := s.readUpstreamErrorBody(resp)
@@ -175,17 +175,22 @@ func (s *OpenAIGatewayService) forwardGrokResponses(
})
errCtx := withGrokTeamRateLimitModel(ctx, upstreamModel)
s.handleGrokAccountUpstreamError(errCtx, account, resp.StatusCode, resp.Header, respBody)
// 429 / free-usage: stamp team+model cool so sibling accounts skip this model.
if resp.StatusCode == http.StatusTooManyRequests ||
classifyGrokUpstreamFailure(resp.StatusCode, respBody, upstreamModel).Class == GrokFailureFreeUsage {
// Quota/rate-limit responses stamp the team+model overlay. Capacity is
// request pressure and must not hide sibling accounts.
if shouldMarkGrokTeamModelRateLimit(resp.StatusCode, respBody) {
markGrokTeamModelRateLimit(account, upstreamModel, resolveGrokTeamRateLimitUntil(time.Now().Add(grokTeamRateLimitDefaultTTL), time.Now()))
}
if s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) {
retryable, retryDelay, retryDeadline, retryMax := grokSameAccountRetryMetadata(account, resp.StatusCode, respBody)
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: respBody,
ResponseHeaders: resp.Header.Clone(),
RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode),
StatusCode: resp.StatusCode,
ResponseBody: respBody,
ResponseHeaders: resp.Header.Clone(),
RetryableOnSameAccount: retryable,
RequestScopedTransient: retryable && resp.StatusCode == http.StatusTooManyRequests,
SameAccountRetryDelay: retryDelay,
SameAccountRetryDeadline: retryDeadline,
SameAccountRetryMax: retryMax,
}
}
return s.handleErrorResponse(ctx, resp, c, account, patchedBody, upstreamModel)
@@ -262,7 +267,7 @@ func (s *OpenAIGatewayService) forwardGrokResponses(
}
func isGrokInvalidEncryptedContentResponse(statusCode int, body []byte) bool {
if statusCode != http.StatusBadRequest {
if statusCode != http.StatusBadRequest && statusCode != http.StatusUnprocessableEntity {
return false
}
@@ -275,7 +280,7 @@ func isGrokInvalidEncryptedContentResponse(statusCode int, body []byte) bool {
code = strings.TrimSpace(errNode.Get("code").String())
}
if strings.EqualFold(code, "invalid_encrypted_content") {
if strings.EqualFold(code, "invalid_encrypted_content") || strings.EqualFold(code, "invalid_compaction") || strings.EqualFold(code, "compaction_decode_error") {
return true
}
// Keep the official xAI flat-code gate so unrelated 400s are not retried.
@@ -285,19 +290,22 @@ func isGrokInvalidEncryptedContentResponse(statusCode int, body []byte) bool {
for _, candidate := range grokStructuredErrorMessageCandidates(body) {
normalizedMessage := strings.ToLower(candidate)
// Nested OpenAI-style envelopes may omit top-level code; require decrypt text.
if code == "" && !strings.Contains(normalizedMessage, "decrypt") {
if code == "" && !strings.Contains(normalizedMessage, "decrypt") && !strings.Contains(normalizedMessage, "decode the compaction blob") {
continue
}
if strings.Contains(normalizedMessage, "encrypted_content") &&
(strings.Contains(normalizedMessage, "decrypt") || strings.Contains(normalizedMessage, "unmodified")) {
return true
}
if strings.Contains(normalizedMessage, "decode the compaction blob") {
return true
}
}
return false
}
func isGrokCompactionReplayDecodeError(statusCode int, body []byte) bool {
if statusCode != http.StatusBadRequest || len(body) == 0 {
if (statusCode != http.StatusBadRequest && statusCode != http.StatusUnprocessableEntity) || len(body) == 0 {
return false
}
for _, candidate := range grokStructuredErrorMessageCandidates(body) {
@@ -472,7 +480,8 @@ func trimGrokInvalidEncryptedContentRetryBody(body []byte) ([]byte, bool, error)
hasEncryptedReasoning := false
for _, item := range items {
if strings.TrimSpace(item.Get("type").String()) == "reasoning" && item.Get("encrypted_content").Exists() {
if (strings.TrimSpace(item.Get("type").String()) == "reasoning" && item.Get("encrypted_content").Exists()) ||
(isOpenAICompactionType(strings.TrimSpace(item.Get("type").String())) && item.Get("encrypted_content").Exists()) {
hasEncryptedReasoning = true
break
}
@@ -734,7 +743,7 @@ func grokSupportsXHighReasoningEffort(model string) bool {
func grokSupportsReasoningEffort(model string) bool {
model = strings.ToLower(xai.StripGrokProviderPrefix(strings.TrimSpace(model)))
switch model {
case xai.DefaultTextModel, "grok-4.5-latest", "grok-4.6", "grok-4.6-latest",
case "grok-4.5", "grok-4.5-latest", "grok-4.6", "grok-4.6-latest",
"grok-4.3", "grok-4.3-latest",
"grok-3-mini", "grok-3-mini-fast", "grok-4.20-0309-reasoning",
"grok-4.20-reasoning", "grok-4.20-multi-agent-0309":
@@ -937,33 +946,98 @@ func grokResponsesToolDedupKey(tool gjson.Result) string {
return "json:" + normalizeCompatSeedJSON(json.RawMessage(tool.Raw))
}
// sanitizeGrokReasoningNullContent 删除 reasoning 项中的 "content": null。
// xAI 的 untagged enum 反序列化器拒收该字段,返回 422。
// sanitizeGrokReasoningNullContent drops explicit JSON nulls from Responses
// input items. xAI's untagged ModelInput decoder 422s on those fields.
// Compaction items stay unmodified per the compact contract.
func sanitizeGrokReasoningNullContent(body []byte) ([]byte, error) {
input := gjson.GetBytes(body, "input")
if !input.Exists() || !input.IsArray() {
if !input.Exists() || (!input.IsArray() && !input.IsObject()) {
return body, nil
}
items := input.Array()
var decoded map[string]any
decoder := json.NewDecoder(bytes.NewReader(body))
decoder.UseNumber()
if err := decoder.Decode(&decoded); err != nil {
return body, nil
}
rawInput, ok := decoded["input"]
if !ok {
return body, nil
}
cleaned, changed := stripExplicitNullsFromGrokInput(rawInput)
if !changed {
return body, nil
}
decoded["input"] = cleaned
out, err := marshalOpenAIUpstreamJSON(decoded)
if err != nil {
return nil, err
}
return out, nil
}
func stripExplicitNullsFromGrokInput(value any) (any, bool) {
switch node := value.(type) {
case []any:
changed := false
for i, item := range node {
itemMap, ok := item.(map[string]any)
if !ok {
next, childChanged := stripExplicitNullsFromGrokInput(item)
if childChanged {
node[i] = next
changed = true
}
continue
}
if isOpenAICompactionType(stringValue(itemMap["type"])) {
continue
}
next, childChanged := stripExplicitNullsFromJSONObject(itemMap)
if childChanged {
node[i] = next
changed = true
}
}
return node, changed
case map[string]any:
if isOpenAICompactionType(stringValue(node["type"])) {
return node, false
}
return stripExplicitNullsFromJSONObject(node)
default:
return value, false
}
}
func stripExplicitNullsFromJSONObject(node map[string]any) (map[string]any, bool) {
if node == nil {
return node, false
}
changed := false
for i := len(items) - 1; i >= 0; i-- {
item := items[i]
if strings.TrimSpace(item.Get("type").String()) != "reasoning" {
for key, child := range node {
if child == nil {
delete(node, key)
changed = true
continue
}
contentResult := item.Get("content")
if contentResult.Exists() && contentResult.Type == gjson.Null {
var err error
body, err = sjson.DeleteBytes(body, fmt.Sprintf("input.%d.content", i))
if err != nil {
return nil, err
switch typed := child.(type) {
case map[string]any:
next, childChanged := stripExplicitNullsFromJSONObject(typed)
if childChanged {
node[key] = next
changed = true
}
case []any:
next, childChanged := stripExplicitNullsFromGrokInput(typed)
if childChanged {
node[key] = next
changed = true
}
changed = true
}
}
_ = changed
return body, nil
return node, changed
}
var grokResponsesSupportedToolTypes = map[string]struct{}{
@@ -984,6 +1058,13 @@ func sanitizeGrokResponsesTools(body []byte) ([]byte, error) {
return deleteGrokOrphanToolControls(body)
}
if !tools.IsArray() {
// xAI rejects tool_choice when tools is null/object. Drop the malformed
// collection and any orphan tool controls instead of forwarding a pair
// the Grok Responses endpoint cannot interpret.
body, err := sjson.DeleteBytes(body, "tools")
if err != nil {
return nil, err
}
return deleteGrokOrphanToolControls(body)
}
@@ -1275,11 +1356,16 @@ func (s *OpenAIGatewayService) describeGrokComposerImage(
})
s.handleGrokAccountUpstreamError(withGrokTeamRateLimitModel(ctx, grokComposerImageBridgeVisionModel), account, resp.StatusCode, resp.Header, respBody)
if s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) {
retryable, retryDelay, retryDeadline, retryMax := grokSameAccountRetryMetadata(account, resp.StatusCode, respBody)
return "", OpenAIUsage{}, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: respBody,
ResponseHeaders: resp.Header.Clone(),
RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode),
StatusCode: resp.StatusCode,
ResponseBody: respBody,
ResponseHeaders: resp.Header.Clone(),
RetryableOnSameAccount: retryable,
RequestScopedTransient: retryable && resp.StatusCode == http.StatusTooManyRequests,
SameAccountRetryDelay: retryDelay,
SameAccountRetryDeadline: retryDeadline,
SameAccountRetryMax: retryMax,
}
}
return "", OpenAIUsage{}, fmt.Errorf("grok composer image bridge upstream error: %s", upstreamMsg)
@@ -1416,6 +1502,7 @@ func buildGrokResponsesRequest(ctx context.Context, c *gin.Context, account *Acc
if err != nil {
return nil, err
}
req = req.WithContext(WithHTTPUpstreamProfile(req.Context(), HTTPUpstreamProfileGrok))
req.Header.Set("Authorization", "Bearer "+token)
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json, text/event-stream")
@@ -1452,6 +1539,10 @@ func applyGrokCLIHeaders(headers http.Header) {
}
func (s *OpenAIGatewayService) updateGrokUsageSnapshot(ctx context.Context, account *Account, snapshot *xai.QuotaSnapshot) {
s.updateGrokUsageSnapshotWithRateLimit(ctx, account, snapshot, true)
}
func (s *OpenAIGatewayService) updateGrokUsageSnapshotWithRateLimit(ctx context.Context, account *Account, snapshot *xai.QuotaSnapshot, installRateLimit bool) {
if s == nil || account == nil || account.ID <= 0 || snapshot == nil {
return
}
@@ -1499,7 +1590,7 @@ func (s *OpenAIGatewayService) updateGrokUsageSnapshot(ctx context.Context, acco
// API keys retain the snapshot for observability but leave account health to
// the upstream pool. Other accounts install the immediate runtime and durable
// rate-limit state when the observed window is exhausted.
if hasActiveLimit && !account.IsPoolMode() {
if installRateLimit && hasActiveLimit && !account.IsPoolMode() {
s.rateLimitGrok(stateCtx, account, resetAt)
} else if recovery {
clearGrokRateLimitAfterRecovery(stateCtx, s.accountRepo, account)
@@ -1822,17 +1913,12 @@ func grokRequestedModelFromCtx(ctx context.Context) string {
return strings.TrimSpace(model)
}
func isGrokHeavyTransientModel(requestedModel string) bool {
model := strings.ToLower(strings.TrimSpace(xai.ResolveGrokTextResponsesModelID(requestedModel)))
return strings.Contains(model, "multi-agent")
}
func persistGrokTransientModelCooldown(account *Account, decision GrokUpstreamFailureDecision) bool {
if account == nil {
return false
}
model := strings.TrimSpace(decision.Model)
if model == "" || !isGrokHeavyTransientModel(model) {
if model == "" {
return false
}
cooldown := decision.Cooldown
@@ -1851,14 +1937,17 @@ func (s *OpenAIGatewayService) handleGrokAccountUpstreamError(ctx context.Contex
return
}
now := time.Now()
decision := classifyGrokUpstreamFailure(statusCode, responseBody, grokRequestedModelFromCtx(ctx))
snapshot := parseGrokQuotaSnapshot(headers, statusCode, now)
stampGrokQuotaSnapshotForPlan(account, snapshot, grokRequestedModelFromCtx(ctx))
s.updateGrokUsageSnapshot(ctx, account, snapshot)
// Capacity 429 is model pressure, not account quota exhaustion. Keep the
// snapshot for observability but do not install account-level rate limiting;
// the failover decision below applies a bounded model-scoped block instead.
s.updateGrokUsageSnapshotWithRateLimit(ctx, account, snapshot, decision.Class != GrokFailureModelCapacity)
// Body-first free-usage / empty / billing / capacity must run before the
// status switch so non-429 free-usage bodies still cool the account.
// Pool-mode still skips durable mutation unless an explicit temp rule matches.
decision := classifyGrokUpstreamFailure(statusCode, responseBody, grokRequestedModelFromCtx(ctx))
if decision.ShouldCooldown && decision.Class != GrokFailureNone && decision.Class != GrokFailureRateLimit {
if account.IsPoolMode() {
// Allow configured temp rules (403) below; skip default body cools.
@@ -633,11 +633,16 @@ func (s *OpenAIGatewayService) forwardGrokChatCompletionsViaResponses(
})
s.handleGrokAccountUpstreamError(withGrokTeamRateLimitModel(ctx, upstreamModel), account, resp.StatusCode, resp.Header, respBody)
if s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) {
retryable, retryDelay, retryDeadline, retryMax := grokSameAccountRetryMetadata(account, resp.StatusCode, respBody)
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: respBody,
ResponseHeaders: resp.Header.Clone(),
RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode),
StatusCode: resp.StatusCode,
ResponseBody: respBody,
ResponseHeaders: resp.Header.Clone(),
RetryableOnSameAccount: retryable,
RequestScopedTransient: retryable && resp.StatusCode == http.StatusTooManyRequests,
SameAccountRetryDelay: retryDelay,
SameAccountRetryDeadline: retryDeadline,
SameAccountRetryMax: retryMax,
}
}
return s.handleChatCompletionsErrorResponse(resp, c, account, billingModel)
@@ -251,7 +251,7 @@ func TestForwardGrokChatViaResponsesNonStreamingCachesAndReturnsChat(t *testing.
require.NotNil(t, result)
require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String())
require.Equal(t, grokChatResponsesEndpoint, result.UpstreamEndpoint)
require.Equal(t, "grok-4.5", result.UpstreamModel)
require.Equal(t, "grok-4.6", result.UpstreamModel)
require.Equal(t, 9908, result.Usage.InputTokens)
require.Equal(t, 12, result.Usage.OutputTokens)
require.Equal(t, 9856, result.Usage.CacheReadInputTokens)
@@ -352,7 +352,7 @@ func TestForwardGrokChatViaResponsesCodeBuddyUsesStableConversationHeader(t *tes
c.Request.Header.Set("X-Request-ID", "generic-"+tt.requestID)
c.Set("api_key", &APIKey{ID: 7111})
identity := resolveGrokCacheIdentity(c, tt.body, "", "grok-4.5")
identity := resolveGrokCacheIdentity(c, tt.body, "", "grok-4.6")
require.NotEmpty(t, identity)
if index == 0 {
stableIdentity = identity
@@ -406,8 +406,8 @@ func TestForwardGrokChatViaResponsesTraeToolHistoryKeepsCacheRoute(t *testing.T)
accountRepo: repo,
}
firstTurnIdentity := resolveGrokCacheIdentity(c, firstTurnBody, "", "grok-4.5")
extendedTurnIdentity := resolveGrokCacheIdentity(c, body, "", "grok-4.5")
firstTurnIdentity := resolveGrokCacheIdentity(c, firstTurnBody, "", "grok-4.6")
extendedTurnIdentity := resolveGrokCacheIdentity(c, body, "", "grok-4.6")
require.NotEmpty(t, firstTurnIdentity)
require.Equal(t, firstTurnIdentity, extendedTurnIdentity)
@@ -463,8 +463,8 @@ func TestForwardGrokChatViaResponsesTraeCompatibilityFieldsKeepCacheRoute(t *tes
accountRepo: repo,
}
firstTurnIdentity := resolveGrokCacheIdentity(c, firstTurnBody, "", "grok-4.5")
extendedTurnIdentity := resolveGrokCacheIdentity(c, body, "", "grok-4.5")
firstTurnIdentity := resolveGrokCacheIdentity(c, firstTurnBody, "", "grok-4.6")
extendedTurnIdentity := resolveGrokCacheIdentity(c, body, "", "grok-4.6")
require.NotEmpty(t, firstTurnIdentity)
require.Equal(t, firstTurnIdentity, extendedTurnIdentity)
@@ -534,7 +534,7 @@ func TestForwardGrokChatRuntimeGateFallsBackToRaw(t *testing.T) {
mappedModel string
wantUpstream string
}{
{name: "missing cache identity", wantUpstream: "grok-4.5"},
{name: "missing cache identity", wantUpstream: "grok-4.6"},
{name: "non cache capable mapped model", setAPIKey: true, mappedModel: "grok-4.3", wantUpstream: "grok-4.3"},
}
@@ -418,9 +418,10 @@ func TestSanitizeGrokResponsesToolsKeepsToolChoiceOnlyWithSupportedTools(t *test
wantToolChoice: true,
},
{
name: "malformed non-array tools drop orphan controls",
body: `{"input":"hello","tools":{"type":"function","name":"lookup"},"tool_choice":"auto"}`,
wantTools: true,
name: "malformed non-array tools are removed",
body: `{"input":"hello","tools":{"type":"function","name":"lookup"},"tool_choice":"auto"}`,
wantTools: false,
wantToolChoice: false,
},
}
@@ -1188,7 +1189,7 @@ func TestForwardGrokMediaAppliesAccountModelMappingAfterEndpointNormalization(t
modelMapping: map[string]any{"grok-imagine-image-quality": "vendor-image-model"},
wantRequestModel: "grok-imagine-image-quality",
wantUpstream: "vendor-image-model",
wantBody: `{"model":"vendor-image-model","prompt":"draw"}`,
wantBody: `{"model":"vendor-image-model","prompt":"draw","resolution":"1k","aspect_ratio":"1:1"}`,
responseBody: `{"data":[{"url":"https://images.test/mapped.png"}]}`,
},
{
@@ -1311,7 +1312,8 @@ func TestForwardGrokMediaImagesGenerationStripsUnsupportedSize(t *testing.T) {
result, err := svc.ForwardGrokMedia(context.Background(), c, account, GrokMediaEndpointImagesGenerations, "", body, "application/json")
require.NoError(t, err)
require.JSONEq(t, `{"model":"grok-imagine-image","prompt":"draw a cat"}`, string(upstream.lastBody))
require.JSONEq(t, `{"model":"grok-imagine-image","prompt":"draw a cat","resolution":"1k","aspect_ratio":"1:1"}`, string(upstream.lastBody))
require.False(t, gjson.GetBytes(upstream.lastBody, "size").Exists())
require.Equal(t, ImageBillingSize1K, result.ImageSize)
require.Equal(t, "1024x1024", result.ImageInputSize)
}
@@ -1372,6 +1374,58 @@ func TestForwardGrokMediaImagesEditMultipartConvertsToJSON(t *testing.T) {
require.Equal(t, "vendor-image-edit", result.UpstreamModel)
}
func TestForwardGrokMediaImagesEditMultipartPreservesExplicitGeometry(t *testing.T) {
t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true")
gin.SetMode(gin.TestMode)
var buf bytes.Buffer
writer := multipart.NewWriter(&buf)
require.NoError(t, writer.WriteField("model", "grok-imagine-edit"))
require.NoError(t, writer.WriteField("prompt", "edit this private image"))
require.NoError(t, writer.WriteField("size", "1024x1024"))
require.NoError(t, writer.WriteField("resolution", "2k"))
require.NoError(t, writer.WriteField("aspect_ratio", "16:9"))
partHeader := textproto.MIMEHeader{}
partHeader.Set("Content-Disposition", `form-data; name="image"; filename="input.png"`)
partHeader.Set("Content-Type", "image/png")
part, err := writer.CreatePart(partHeader)
require.NoError(t, err)
_, err = part.Write([]byte{0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a})
require.NoError(t, err)
require.NoError(t, writer.Close())
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/images/edits", bytes.NewReader(buf.Bytes()))
c.Request.Header.Set("Content-Type", writer.FormDataContentType())
account := &Account{
ID: 67,
Name: "grok",
Platform: PlatformGrok,
Type: AccountTypeAPIKey,
Concurrency: 1,
Credentials: map[string]any{
"api_key": "api-key",
"base_url": "https://xai.test/v1",
},
}
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(`{"data":[{"url":"https://images.test/edited.png"}]}`)),
}}
svc := &OpenAIGatewayService{httpUpstream: upstream}
result, err := svc.ForwardGrokMedia(context.Background(), c, account, GrokMediaEndpointImagesEdits, "", buf.Bytes(), writer.FormDataContentType())
require.NoError(t, err)
require.False(t, gjson.GetBytes(upstream.lastBody, "size").Exists())
require.Equal(t, "2k", gjson.GetBytes(upstream.lastBody, "resolution").String())
require.Equal(t, "16:9", gjson.GetBytes(upstream.lastBody, "aspect_ratio").String())
require.Equal(t, ImageBillingSize1K, result.ImageSize)
require.Equal(t, "1024x1024", result.ImageInputSize)
}
func TestForwardGrokMediaVideoGenerationReturnsUsageAndResponseID(t *testing.T) {
t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true")
gin.SetMode(gin.TestMode)
@@ -1802,10 +1856,10 @@ func TestForwardAsChatCompletionsForGrokStopFallsBackToXAIChatCompletions(t *tes
require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization"))
require.NotEmpty(t, upstream.lastReq.Header.Get(grokConversationIDHeader))
require.NotEqual(t, "raw-client-cache-key", upstream.lastReq.Header.Get(grokConversationIDHeader))
require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String())
require.Equal(t, "grok-4.6", gjson.GetBytes(upstream.lastBody, "model").String())
require.False(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_key").Exists())
require.Equal(t, "grok", result.Model)
require.Equal(t, "grok-4.5", result.UpstreamModel)
require.Equal(t, "grok-4.6", result.UpstreamModel)
require.Equal(t, 1, result.Usage.InputTokens)
require.Equal(t, 2, result.Usage.OutputTokens)
require.Equal(t, 1, result.Usage.CacheReadInputTokens)
@@ -1919,7 +1973,7 @@ func TestForwardGrokResponsesAPIKeyUsesXAIResponses(t *testing.T) {
require.Equal(t, "Bearer xai-test-key", upstream.lastReq.Header.Get("Authorization"))
require.Empty(t, upstream.lastReq.Header.Get("X-Grok-Client-Version"))
require.NotEqual(t, grokUpstreamUserAgent, upstream.lastReq.Header.Get("User-Agent"))
require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String())
require.Equal(t, "grok-4.6", gjson.GetBytes(upstream.lastBody, "model").String())
require.Equal(t, "resp_grok_api_key", result.ResponseID)
require.Equal(t, 2, result.Usage.InputTokens)
require.Equal(t, 1, result.Usage.OutputTokens)
@@ -2067,6 +2121,16 @@ func TestForwardGrokResponsesInvalidEncryptedContentRecoveryDoesNotOvermatch(t *
}
}
func TestGrokCompactionBlobRecoveryStripsCompactionItem(t *testing.T) {
body := []byte(`{"model":"grok","input":[{"type":"compaction","id":"cmp_1","encrypted_content":"blob"},{"type":"message","role":"user","content":"hi"}]}`)
require.True(t, isGrokInvalidEncryptedContentResponse(http.StatusUnprocessableEntity, []byte(`{"code":"invalid_compaction","error":"could not decode the compaction blob"}`)))
retry, changed, err := trimGrokInvalidEncryptedContentRetryBody(body)
require.NoError(t, err)
require.True(t, changed)
require.Equal(t, "message", gjson.GetBytes(retry, "input.0.type").String())
require.False(t, gjson.GetBytes(retry, "input.#(type==\"compaction\")").Exists())
}
func TestForwardGrokResponsesInvalidEncryptedContentRecoveryNestedErrorShape(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -2371,7 +2435,7 @@ func TestForwardAsChatCompletionsForGrokStreamingUsesRawXAIChatCompletions(t *te
require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization"))
require.Equal(t, "text/event-stream", upstream.lastReq.Header.Get("Accept"))
require.Equal(t, xai.CLIUserAgent(xai.CLIClientVersion), upstream.lastReq.Header.Get("User-Agent"))
require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String())
require.Equal(t, "grok-4.6", gjson.GetBytes(upstream.lastBody, "model").String())
require.True(t, gjson.GetBytes(upstream.lastBody, "stream_options.include_usage").Bool())
require.True(t, result.Stream)
require.Equal(t, 6, result.Usage.InputTokens)
@@ -2424,7 +2488,7 @@ func TestForwardGrokResponsesNonStreamingUsesCacheIdentityAndCachedUsage(t *test
require.Equal(t, 7, result.Usage.InputTokens)
require.Equal(t, 2, result.Usage.OutputTokens)
require.Equal(t, 4, result.Usage.CacheReadInputTokens)
require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String())
require.Equal(t, "grok-4.6", gjson.GetBytes(upstream.lastBody, "model").String())
identity := gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String()
require.NotEmpty(t, identity)
require.Equal(t, identity, upstream.lastReq.Header.Get(grokConversationIDHeader))
@@ -2547,7 +2611,7 @@ func TestForwardAsChatCompletionsForGrokStreamingStopFallsBackToRawXAIChatComple
require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version"))
require.NotEmpty(t, upstream.lastReq.Header.Get(grokConversationIDHeader))
require.NotEqual(t, "native-client-conversation", upstream.lastReq.Header.Get(grokConversationIDHeader))
require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String())
require.Equal(t, "grok-4.6", gjson.GetBytes(upstream.lastBody, "model").String())
require.True(t, gjson.GetBytes(upstream.lastBody, "stream_options.include_usage").Bool())
require.True(t, result.Stream)
require.Equal(t, 6, result.Usage.InputTokens)
@@ -2653,7 +2717,7 @@ func TestForwardAsAnthropicForGrokUsesXAIResponses(t *testing.T) {
require.Equal(t, "grok-experimental", upstream.lastReq.Header.Get("OpenAI-Beta"))
require.Empty(t, upstream.lastReq.Header.Get("originator"))
require.Empty(t, upstream.lastReq.Header.Get("version"))
require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String())
require.Equal(t, "grok-4.6", gjson.GetBytes(upstream.lastBody, "model").String())
require.NotEmpty(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String())
require.Equal(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String(), upstream.lastReq.Header.Get(grokConversationIDHeader))
require.Equal(t, "web_search", gjson.GetBytes(upstream.lastBody, "tools.0.type").String())
@@ -2663,7 +2727,7 @@ func TestForwardAsAnthropicForGrokUsesXAIResponses(t *testing.T) {
require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool())
require.NotContains(t, string(upstream.lastBody), "chatgpt.com")
require.Equal(t, "grok", result.Model)
require.Equal(t, "grok-4.5", result.UpstreamModel)
require.Equal(t, "grok-4.6", result.UpstreamModel)
require.Equal(t, 5, result.Usage.InputTokens)
require.Equal(t, 2, result.Usage.OutputTokens)
require.Equal(t, 3, result.Usage.CacheReadInputTokens)
@@ -3451,6 +3515,43 @@ func TestPatchGrokResponsesBody_StripsReasoningContentNull(t *testing.T) {
require.False(t, reasoning.Get("content").Exists(), "content: null should be stripped")
}
func TestPatchGrokResponsesBody_StripsNestedInputNulls(t *testing.T) {
t.Parallel()
body := []byte(`{
"model": "grok-4.6",
"input": [
{"type":"message","role":"user","content":[{"type":"input_text","text":"hi","logprobs":null}]}
]
}`)
patched, err := patchGrokResponsesBody(body, "grok-4.6")
require.NoError(t, err)
require.False(t, gjson.GetBytes(patched, "input.0.content.0.logprobs").Exists())
require.Equal(t, "hi", gjson.GetBytes(patched, "input.0.content.0.text").String())
}
func TestStripExplicitNullsFromGrokInputSkipsCompaction(t *testing.T) {
t.Parallel()
input := []any{
map[string]any{"type": "message", "role": "user", "content": []any{map[string]any{"type": "input_text", "text": "hi", "logprobs": nil}}},
map[string]any{"type": "compaction", "id": "cmp_1", "encrypted_content": "blob", "status": nil},
}
out, changed := stripExplicitNullsFromGrokInput(input)
require.True(t, changed)
items, ok := out.([]any)
require.True(t, ok)
require.Len(t, items, 2)
msg := items[0].(map[string]any)
content := msg["content"].([]any)[0].(map[string]any)
_, hasLogprobs := content["logprobs"]
require.False(t, hasLogprobs)
compaction := items[1].(map[string]any)
require.Equal(t, "blob", compaction["encrypted_content"])
_, hasStatus := compaction["status"]
require.True(t, hasStatus, "compaction items must stay unmodified")
}
func TestPatchGrokResponsesBody_KeepsReasoningContentNonNull(t *testing.T) {
t.Parallel()
@@ -15,6 +15,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/pkg/apicompat"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
"github.com/Wei-Shaw/sub2api/internal/pkg/xai"
"github.com/Wei-Shaw/sub2api/internal/util/responseheaders"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
@@ -1396,6 +1397,17 @@ func openAIUsageFromGJSON(value gjson.Result) (OpenAIUsage, bool) {
if outputTokens == 0 {
outputTokens = value.Get("completion_tokens").Int()
}
// xAI reports visible output separately from reasoning_tokens; OpenAI
// folds reasoning into completion/output. Use total_tokens to tell them apart.
reasoningTokens := max(int(firstPositiveGJSONInt(
value.Get("completion_tokens_details.reasoning_tokens"),
value.Get("output_tokens_details.reasoning_tokens"),
)), 0)
if reasoningTokens > 0 {
outputTokens = xai.IncludeIndependentReasoningTokens(
inputTokens, outputTokens, value.Get("total_tokens").Int(), int64(reasoningTokens),
)
}
cacheReadTokens := openAICacheReadTokensFromUsage(value)
cacheCreationTokens := openAICacheCreationTokensFromUsage(value)
imageOutputTokens := value.Get("output_tokens_details.image_tokens").Int()
@@ -3509,6 +3509,30 @@ func TestExtractOpenAIUsageFromJSONBytes_AcceptsResponseAndChatUsageShapes(t *te
usage, ok = extractOpenAIUsageFromJSONBytes([]byte(`{"usage":{"input_tokens":20,"output_tokens":2,"cache_read_input_tokens":19,"input_tokens_details":{"cached_tokens":0}}}`))
require.True(t, ok)
require.Zero(t, usage.CacheReadInputTokens, "官方嵌套缓存读取字段显式为零时仍应优先于兼容顶层别名")
// xAI reports reasoning_tokens outside visible output_tokens. Only the
// arithmetic-consistent shape is independent; OpenAI's canonical shape
// already includes reasoning in completion/output_tokens.
usage, ok = extractOpenAIUsageFromJSONBytes([]byte(`{"usage":{"input_tokens":10000,"output_tokens":500,"total_tokens":10800,"output_tokens_details":{"reasoning_tokens":300}}}`))
require.True(t, ok)
require.Equal(t, 800, usage.OutputTokens)
usage, ok = extractOpenAIUsageFromJSONBytes([]byte(`{"usage":{"input_tokens":10000,"output_tokens":500,"total_tokens":10500,"output_tokens_details":{"reasoning_tokens":300}}}`))
require.True(t, ok)
require.Equal(t, 500, usage.OutputTokens)
}
func TestExtractOpenAIUsageFromJSONBytes_IncludesGrokReasoningTokens(t *testing.T) {
usage, ok := extractOpenAIUsageFromJSONBytes([]byte(`{"usage":{"prompt_tokens":32,"completion_tokens":9,"total_tokens":135,"completion_tokens_details":{"reasoning_tokens":94}}}`))
require.True(t, ok)
require.Equal(t, 103, usage.OutputTokens, "Grok Chat usage bills visible completion plus reasoning tokens")
usage, ok = extractOpenAIUsageFromJSONBytes([]byte(`{"usage":{"input_tokens":32,"output_tokens":103,"total_tokens":135,"output_tokens_details":{"reasoning_tokens":94}}}`))
require.True(t, ok)
require.Equal(t, 103, usage.OutputTokens, "Responses output_tokens already includes reasoning when total confirms it")
usage, ok = extractOpenAIUsageFromJSONBytes([]byte(`{"usage":{"input_tokens":32,"output_tokens":9,"total_tokens":135,"output_tokens_details":{"reasoning_tokens":94}}}`))
require.True(t, ok)
require.Equal(t, 103, usage.OutputTokens, "Responses detail-only shape is normalized when total exposes the full output")
}
func TestExtractCodexFinalResponse_SampleReplay(t *testing.T) {
@@ -1298,7 +1298,7 @@ func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridgeAndPreservesMa
require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization"))
require.Equal(t, xai.CLIUserAgent(xai.CLIClientVersion), upstream.lastReq.Header.Get("User-Agent"))
require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version"))
require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.bodies[0], "model").String())
require.Equal(t, "grok-4.6", gjson.GetBytes(upstream.bodies[0], "model").String())
require.Equal(t, "grok-4.3", gjson.GetBytes(upstream.bodies[1], "model").String())
require.Equal(t, "grok-4.3", gjson.GetBytes(upstream.bodies[2], "model").String())
require.NotEmpty(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String())
@@ -14,6 +14,8 @@ import (
coderws "github.com/coder/websocket"
"github.com/tidwall/gjson"
"github.com/Wei-Shaw/sub2api/internal/pkg/xai"
)
type FrameConn interface {
@@ -1040,6 +1042,15 @@ func parseUsageAndAccumulate(
// 解析失败时不做部分字段累加,避免计费 usage 出现“半有效”状态。
return Usage{}
}
reasoningTokens := usageResult.Get("output_tokens_details.reasoning_tokens").Int()
if reasoningTokens == 0 {
reasoningTokens = usageResult.Get("completion_tokens_details.reasoning_tokens").Int()
}
if reasoningTokens > 0 {
outputTokens = int(xai.IncludeIndependentReasoningTokens(
int64(inputTokens), int64(outputTokens), usageResult.Get("total_tokens").Int(), reasoningTokens,
))
}
parsedUsage := Usage{
InputTokens: inputTokens,
OutputTokens: outputTokens,
@@ -325,6 +325,29 @@ func TestParseUsageAndEnrichCoverage(t *testing.T) {
enrichResult(nil, state, 0)
}
func TestParseUsageAndAccumulateIncludesIndependentReasoningTokens(t *testing.T) {
t.Parallel()
state := &relayState{}
got := parseUsageAndAccumulate(
state,
[]byte(`{"type":"response.completed","response":{"usage":{"input_tokens":32,"output_tokens":9,"total_tokens":151,"output_tokens_details":{"reasoning_tokens":110}}}}`),
"response.completed",
nil,
)
require.Equal(t, 32, got.InputTokens)
require.Equal(t, 119, got.OutputTokens)
state = &relayState{}
got = parseUsageAndAccumulate(
state,
[]byte(`{"type":"response.completed","response":{"usage":{"input_tokens":32,"output_tokens":119,"total_tokens":151,"output_tokens_details":{"reasoning_tokens":110}}}}`),
"response.completed",
nil,
)
require.Equal(t, 119, got.OutputTokens, "inclusive Responses output must not double-count reasoning")
}
func TestParseUsageAndAccumulateAcceptsChatUsageAliases(t *testing.T) {
t.Parallel()
@@ -457,6 +457,37 @@ func (s *SettingService) MigrateOpenAIAllowClaudeCodeCodexPluginSetting(ctx cont
return nil
}
// MigrateGrokDefaultTextModel upgrades the pre-4.6 built-in default for
// existing installations. Explicit operator choices are left untouched.
func (s *SettingService) MigrateGrokDefaultTextModel(ctx context.Context) error {
if s == nil || s.settingRepo == nil {
return nil
}
if ctx == nil {
ctx = context.Background()
}
dbCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), codexRestrictionPolicyDBTimeout)
defer cancel()
value, err := s.settingRepo.GetValue(dbCtx, SettingKeyGrokDefaultTextModel)
if err != nil {
if errors.Is(err, ErrSettingNotFound) {
return nil
}
return fmt.Errorf("get %s setting: %w", SettingKeyGrokDefaultTextModel, err)
}
// Only migrate the value that was previously shipped as the built-in
// default. Any other value is an explicit operator choice or a future
// default and must remain unchanged.
if strings.TrimSpace(value) != "grok-4.5" {
return nil
}
if err := s.settingRepo.Set(dbCtx, SettingKeyGrokDefaultTextModel, "grok-4.6"); err != nil {
return fmt.Errorf("set %s setting: %w", SettingKeyGrokDefaultTextModel, err)
}
return nil
}
// MigrateCodexBodyFingerprintToSignals 把已废弃的 codex_cli_only_allow_body_engine_fingerprint
// 开关并入引擎指纹信号列表。幂等:信号键已存在(非空)则不动;缺失时写默认种子,
// 并把 body 路径行的 Required 设为旧 body 开关的值(旧 true ⇒ 勾上 body 行)。
+2 -2
View File
@@ -193,7 +193,7 @@ func (s *SettingService) InitializeDefaultSettings(ctx context.Context) error {
SettingKeyChannelMonitorShowQuota: "false",
// Grok: safe defaults — no cross-vendor model rewrite unless operators enable it.
SettingKeyGrokDefaultTextModel: "grok-4.5",
SettingKeyGrokDefaultTextModel: "grok-4.6",
SettingKeyGrokCrossClientModelMapEnabled: "true",
SettingKeyGrokDefaultBaseURLMode: GrokDefaultBaseURLModeCLI,
@@ -806,7 +806,7 @@ func (s *SettingService) parseSettings(settings map[string]string) *SystemSettin
// Grok default mapping policy
result.GrokDefaultTextModel = strings.TrimSpace(settings[SettingKeyGrokDefaultTextModel])
if result.GrokDefaultTextModel == "" {
result.GrokDefaultTextModel = "grok-4.5"
result.GrokDefaultTextModel = "grok-4.6"
}
// Default true (missing/empty → enabled) so Claude/Codex→Grok mapping keeps working.
// Operators can set false to disable silent cross-client rewrite.
@@ -202,3 +202,34 @@ func TestMigrateCodexBodyFingerprintToSignals(t *testing.T) {
require.Equal(t, openai.DefaultEngineFingerprintSignalsJSON(), repo.values[SettingKeyCodexCLIOnlyEngineFingerprintSignals])
})
}
func TestMigrateGrokDefaultTextModel(t *testing.T) {
t.Run("upgrades legacy built-in default", func(t *testing.T) {
repo := &codexPolicyMigrationRepoStub{values: map[string]string{
SettingKeyGrokDefaultTextModel: "grok-4.5",
}}
svc := NewSettingService(repo, &config.Config{})
require.NoError(t, svc.MigrateGrokDefaultTextModel(context.Background()))
require.Equal(t, "grok-4.6", repo.values[SettingKeyGrokDefaultTextModel])
require.Equal(t, "grok-4.6", repo.sets[SettingKeyGrokDefaultTextModel])
})
t.Run("does not overwrite an explicit model", func(t *testing.T) {
repo := &codexPolicyMigrationRepoStub{values: map[string]string{
SettingKeyGrokDefaultTextModel: "grok-4.3",
}}
svc := NewSettingService(repo, &config.Config{})
require.NoError(t, svc.MigrateGrokDefaultTextModel(context.Background()))
require.Equal(t, "grok-4.3", repo.values[SettingKeyGrokDefaultTextModel])
_, wrote := repo.sets[SettingKeyGrokDefaultTextModel]
require.False(t, wrote)
})
t.Run("missing setting is left for normal defaults", func(t *testing.T) {
repo := &codexPolicyMigrationRepoStub{values: map[string]string{}}
svc := NewSettingService(repo, &config.Config{})
require.NoError(t, svc.MigrateGrokDefaultTextModel(context.Background()))
_, wrote := repo.sets[SettingKeyGrokDefaultTextModel]
require.False(t, wrote)
})
}
+1 -1
View File
@@ -423,7 +423,7 @@ func (s *SettingService) buildSystemSettingsUpdates(ctx context.Context, setting
if v := strings.TrimSpace(settings.GrokDefaultTextModel); v != "" {
updates[SettingKeyGrokDefaultTextModel] = v
} else {
updates[SettingKeyGrokDefaultTextModel] = "grok-4.5"
updates[SettingKeyGrokDefaultTextModel] = "grok-4.6"
}
updates[SettingKeyGrokCrossClientModelMapEnabled] = strconv.FormatBool(settings.GrokCrossClientModelMapEnabled)
updates[SettingKeyGrokDefaultBaseURLMode] = normalizeGrokDefaultBaseURLMode(settings.GrokDefaultBaseURLMode)
+3
View File
@@ -748,6 +748,9 @@ func ProvideSettingService(settingRepo SettingRepository, groupRepo GroupReposit
if err := svc.MigrateCodexBodyFingerprintToSignals(context.Background()); err != nil {
logger.LegacyPrintf("service.setting", "Warning: migrate codex body fingerprint to signals failed: %v", err)
}
if err := svc.MigrateGrokDefaultTextModel(context.Background()); err != nil {
logger.LegacyPrintf("service.setting", "Warning: migrate Grok default text model failed: %v", err)
}
antigravity.SetUserAgentVersionResolver(svc.GetAntigravityUserAgentVersion)
// enforceCodexIdentityHeaders 是所有 Codex 出站路径共用的纯函数收口点,拿不到 ctx,
// 故注入无参解析器;解析器内部自带 60s TTL 缓存,热路径不触库。
+3
View File
@@ -214,6 +214,9 @@ gateway:
# OpenAI/Codex upstream response header timeout (seconds, 0=disabled)
# OpenAI/Codex 等待上游响应头超时时间(秒,0=禁用本地响应头超时)
openai_response_header_timeout: 0
# Grok text Responses first-byte timeout (seconds, 0=disabled; default 120)
# Grok 文本 Responses 首字节超时时间(秒,0=禁用;默认 120)
grok_response_header_timeout: 120
# Native OpenAI HTTP Responses first semantic output timeout (seconds, 0=disabled)
# Includes response-header wait; does not apply to passthrough or WebSocket transports.
# A timed-out request may already have incurred upstream usage; account failover can therefore duplicate upstream billing.
@@ -51,7 +51,8 @@ describe('useModelWhitelist', () => {
expect(models).toContain('grok-4.5')
expect(models).toContain('grok-4.5-latest')
expect(models).toContain('grok-build-latest')
expect(models).toContain('grok-imagine-video-1.5-preview')
expect(models).toContain('grok-imagine-image-2.0')
expect(models).toContain('grok-imagine-video-1.5')
})
it('combined 模式支持 Grok 4.5 官方别名映射', () => {
@@ -157,8 +157,8 @@ const xaiModels = [
'grok-imagine',
'grok-imagine-image-quality',
'grok-imagine-image',
'grok-imagine-image-2.0',
'grok-imagine-video',
'grok-imagine-video-1.5-preview',
'grok-imagine-video-1.5'
]