mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 13:28:39 +08:00
合并上游 main,保留 5925 的 Grok 重试上限与 5888 的协议兼容修复
同账号重试采用次数上限加 deadline;畸形 tools 在出站前删除;compaction 422 与结构化错误扫描一并保留。
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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"))
|
||||
})
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"))
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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).
|
||||
|
||||
@@ -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"} {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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"))
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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 行)。
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 缓存,热路径不触库。
|
||||
|
||||
@@ -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'
|
||||
]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user