feat(grok): enable free OAuth prompt caching

This commit is contained in:
superman2003
2026-07-11 19:34:01 +08:00
parent e316ebf528
commit 0478fd3668
19 changed files with 846 additions and 60 deletions
@@ -710,7 +710,7 @@ func (s *AccountTestService) testGrokAccountConnection(c *gin.Context, account *
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json, text/event-stream")
req.Header.Set("Authorization", "Bearer "+authToken)
req.Header.Set("User-Agent", "sub2api-grok/1.0")
applyGrokCLIHeaders(req.Header)
proxyURL := ""
if account.ProxyID != nil && account.Proxy != nil {
@@ -60,6 +60,7 @@ func TestAccountTestService_TestAccountConnection_GrokUsesXAIResponses(t *testin
require.Equal(t, "https://api.x.ai/v1/responses", upstream.lastReq.URL.String())
require.Equal(t, "Bearer grok-access-token", upstream.lastReq.Header.Get("Authorization"))
require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version"))
require.Equal(t, "grok-4.3", gjson.GetBytes(upstream.lastBody, "model").String())
require.NotContains(t, rec.Body.String(), "claude")
require.Contains(t, rec.Body.String(), `"model":"grok-4.3"`)
+1 -1
View File
@@ -333,7 +333,7 @@ func (s *OpenAIGatewayService) ForwardGrokMedia(
}
upstreamReq.Header.Set("Authorization", "Bearer "+token)
upstreamReq.Header.Set("Accept", "application/json")
upstreamReq.Header.Set("User-Agent", "sub2api-grok/1.0")
applyGrokCLIHeaders(upstreamReq.Header)
if endpoint.RequiresRequestBody() {
contentType = strings.TrimSpace(contentType)
if contentType == "" {
@@ -82,7 +82,7 @@ func (s *GrokQuotaService) ProbeUsage(ctx context.Context, accountID int64) (*Gr
req.Header.Set("Authorization", "Bearer "+token)
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json")
req.Header.Set("User-Agent", "sub2api-grok-quota-probe/1.0")
applyGrokCLIHeaders(req.Header)
resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, maxInt(account.Concurrency, 1))
if err != nil {
@@ -98,6 +98,7 @@ func TestGrokQuotaServiceProbeUsageStoresHeaders(t *testing.T) {
require.EqualValues(t, 7, *result.Snapshot.Requests.Remaining)
require.Equal(t, "https://api.x.ai/v1/responses", upstream.lastReq.URL.String())
require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization"))
require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version"))
require.Equal(t, "grok-4.3", gjson.GetBytes(upstream.lastBody, "model").String())
require.Contains(t, string(upstream.lastBody), `"max_output_tokens":1`)
require.Contains(t, string(upstream.lastBody), `"store":false`)
@@ -159,6 +159,7 @@ func (s *OpenAIGatewayService) sendCCUpstreamRequest(
stream bool,
bearerToken string,
userAgent string,
grokCacheIdentity string,
) (*http.Response, error) {
upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx)
upstreamReq, err := http.NewRequestWithContext(upstreamCtx, http.MethodPost, targetURL, bytes.NewReader(body))
@@ -190,6 +191,10 @@ func (s *OpenAIGatewayService) sendCCUpstreamRequest(
// 账号级请求头覆写(仅 openai api_key 账号启用时生效)
account.ApplyHeaderOverrides(upstreamReq.Header)
if account.Platform == PlatformGrok {
applyGrokCLIHeaders(upstreamReq.Header)
applyGrokCacheHeaders(upstreamReq.Header, grokCacheIdentity)
}
proxyURL := ""
if account.Proxy != nil {
@@ -76,6 +76,12 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions(
// 2. Resolve model mapping (same as ForwardAsChatCompletions)
billingModel := resolveOpenAIForwardModel(account, originalModel, defaultMappedModel)
upstreamModel := normalizeOpenAIModelForUpstream(account, billingModel)
grokCacheIdentity := ""
if account.Platform == PlatformGrok {
// Resolve before image bridging or other body rewrites so the fallback is
// anchored to the client's stable conversation prefix.
grokCacheIdentity = resolveGrokCacheIdentity(c, body, "", upstreamModel)
}
reasoningEffort := extractOpenAIReasoningEffortFromBody(body, upstreamModel, billingModel, originalModel)
// 国产模型默认 effort 补充:需要 mappedModel 判定,推迟到 billingModel 算出之后。
reasoningEffort = ApplyThinkingEnabledFallback(reasoningEffort, body, billingModel)
@@ -134,6 +140,12 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions(
return nil, fmt.Errorf("enable stream usage: %w", usageErr)
}
}
if account.Platform == PlatformGrok {
upstreamBody, err = stripGrokChatPromptCacheKey(upstreamBody)
if err != nil {
return nil, fmt.Errorf("remove Responses-only Grok prompt cache key: %w", err)
}
}
logger.L().Debug("openai chat_completions raw: forwarding without protocol conversion",
zap.Int64("account_id", account.ID),
@@ -152,7 +164,7 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions(
if customUA == "" && account.Platform == PlatformGrok {
customUA = "sub2api-grok/1.0"
}
resp, err := s.sendCCUpstreamRequest(ctx, c, account, targetURL, upstreamBody, clientStream, token, customUA)
resp, err := s.sendCCUpstreamRequest(ctx, c, account, targetURL, upstreamBody, clientStream, token, customUA, grokCacheIdentity)
if err != nil {
return nil, err
}
@@ -49,7 +49,6 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
originalModel := reqModel
if account.Platform == PlatformGrok {
_ = promptCacheKey
return s.forwardGrokResponses(ctx, c, account, body, originalModel, reqStream, startTime)
}
@@ -20,6 +20,8 @@ import (
const (
grokComposerImageBridgeVisionModel = "grok-build-0.1"
grokComposerImageBridgeMaxOutputTokens = 512
grokUpstreamUserAgent = "sub2api-grok/1.0"
grokCLIVersion = "0.2.93"
)
func (s *OpenAIGatewayService) forwardGrokResponses(
@@ -39,10 +41,15 @@ func (s *OpenAIGatewayService) forwardGrokResponses(
if strings.TrimSpace(upstreamModel) == "" {
upstreamModel = "grok-4.3"
}
cacheIdentity := resolveGrokCacheIdentity(c, body, "", upstreamModel)
patchedBody, err := patchGrokResponsesBody(body, upstreamModel)
if err != nil {
return nil, err
}
patchedBody, err = applyGrokResponsesCacheIdentity(patchedBody, body, cacheIdentity, account.IsGrokOAuth())
if err != nil {
return nil, fmt.Errorf("apply grok prompt cache identity: %w", err)
}
token, _, err := s.GetAccessToken(ctx, account)
if err != nil {
@@ -51,7 +58,7 @@ func (s *OpenAIGatewayService) forwardGrokResponses(
upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx)
defer releaseUpstreamCtx()
upstreamReq, err := buildGrokResponsesRequest(upstreamCtx, c, account, patchedBody, token)
upstreamReq, err := buildGrokResponsesRequest(upstreamCtx, c, account, patchedBody, token, cacheIdentity)
if err != nil {
return nil, err
}
@@ -457,7 +464,9 @@ func (s *OpenAIGatewayService) describeGrokComposerImage(
}
upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx)
upstreamReq, err := buildGrokResponsesRequest(upstreamCtx, c, account, body, token)
// Image-description probes are auxiliary requests, not conversation turns.
// Do not bind them to the caller's Grok prompt-cache identity.
upstreamReq, err := buildGrokResponsesRequest(upstreamCtx, c, account, body, token, "")
releaseUpstreamCtx()
if err != nil {
return "", OpenAIUsage{}, fmt.Errorf("build grok composer image bridge request: %w", err)
@@ -623,7 +632,7 @@ func addOpenAIUsage(dst *OpenAIUsage, usage OpenAIUsage) {
dst.ImageOutputTokens += usage.ImageOutputTokens
}
func buildGrokResponsesRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token string) (*http.Request, error) {
func buildGrokResponsesRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token, cacheIdentity string) (*http.Request, error) {
targetURL, err := xai.BuildResponsesURL(account.GetGrokBaseURL())
if err != nil {
return nil, err
@@ -635,7 +644,8 @@ func buildGrokResponsesRequest(ctx context.Context, c *gin.Context, account *Acc
req.Header.Set("Authorization", "Bearer "+token)
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json, text/event-stream")
req.Header.Set("User-Agent", "sub2api-grok/1.0")
applyGrokCLIHeaders(req.Header)
applyGrokCacheHeaders(req.Header, cacheIdentity)
if c != nil {
if v := c.GetHeader("OpenAI-Beta"); strings.TrimSpace(v) != "" {
req.Header.Set("OpenAI-Beta", v)
@@ -644,6 +654,16 @@ func buildGrokResponsesRequest(ctx context.Context, c *gin.Context, account *Acc
return req, nil
}
// applyGrokCLIHeaders identifies subscription traffic as a supported Grok CLI
// version. The CLI gateway rejects otherwise valid OAuth requests without it.
func applyGrokCLIHeaders(headers http.Header) {
if headers == nil {
return
}
headers.Set("User-Agent", grokUpstreamUserAgent)
headers.Set("X-Grok-Client-Version", grokCLIVersion)
}
func (s *OpenAIGatewayService) updateGrokUsageSnapshot(ctx context.Context, accountID int64, snapshot *xai.QuotaSnapshot) {
if s == nil || s.accountRepo == nil || accountID <= 0 || snapshot == nil {
return
@@ -0,0 +1,149 @@
package service
import (
"fmt"
"net/http"
"strings"
"github.com/gin-gonic/gin"
"github.com/tidwall/gjson"
"github.com/tidwall/sjson"
)
const (
grokConversationIDHeader = "X-Grok-Conv-Id"
grokFreeCacheNativeToolsJSON = `[{"type":"web_search"},{"type":"x_search"}]`
grokFreeCacheDisabledToolChoice = "none"
)
// resolveGrokCacheIdentity derives one stable, tenant-isolated routing identity
// for xAI's server-side prompt cache. The returned value is safe to expose to
// the upstream: it never contains the client's raw session identifier.
//
// A valid downstream API key is required. This intentionally fails closed on
// internal probes and incomplete request contexts instead of creating a cache
// identity that could be shared by unrelated tenants.
func resolveGrokCacheIdentity(c *gin.Context, body []byte, explicitKey, upstreamModel string) string {
apiKeyID := getAPIKeyIDFromContext(c)
if apiKeyID <= 0 {
return ""
}
// /responses/compact rejects tool_choice and does not represent a normal
// conversation turn. Keep both cache identity and Free-tier routing
// augmentation out of this path.
if isOpenAIResponsesCompactPath(c) {
return ""
}
model := strings.ToLower(strings.TrimSpace(upstreamModel))
if model == "" {
return ""
}
seed := explicitGrokCacheSeed(c, body, explicitKey)
if seed == "" {
seed = deriveOpenAIContentSessionSeed(body)
}
if seed == "" {
return ""
}
// generateSessionUUID hashes the whole seed before formatting it as a UUID.
// Include a versioned namespace so this identity cannot collide with other
// upstream session identifiers derived by sub2api.
isolatedSeed := fmt.Sprintf("grok-prompt-cache:v1:%d:%s:%s", apiKeyID, model, seed)
return generateSessionUUID(isolatedSeed)
}
func explicitGrokCacheSeed(c *gin.Context, body []byte, explicitKey string) string {
seed := ""
if c != nil {
seed = strings.TrimSpace(c.GetHeader("session_id"))
if seed == "" {
seed = strings.TrimSpace(c.GetHeader("conversation_id"))
}
if seed == "" {
seed = strings.TrimSpace(c.GetHeader(grokConversationIDHeader))
}
}
if seed == "" && len(body) > 0 {
seed = strings.TrimSpace(gjson.GetBytes(body, "prompt_cache_key").String())
}
if seed == "" {
seed = strings.TrimSpace(explicitKey)
}
return seed
}
func isGrokRequestContext(c *gin.Context) bool {
if c == nil {
return false
}
v, exists := c.Get("api_key")
if !exists {
return false
}
apiKey, ok := v.(*APIKey)
return ok && apiKey != nil && apiKey.Group != nil && apiKey.Group.Platform == PlatformGrok
}
// applyGrokResponsesCacheIdentity writes the cache routing identity into an
// xAI Responses request. Existing client values are deliberately replaced by
// the tenant-isolated value to prevent collisions on shared OAuth accounts.
//
// Free OAuth requests without native search tools are routed by xAI to the
// non-cacheable build-free model. For otherwise tool-free requests, add the
// native tools with tool_choice=none: this selects the cache-capable tier
// without allowing an actual search. Any explicit client tools or tool_choice
// disable this augmentation so client function-calling semantics stay intact.
func applyGrokResponsesCacheIdentity(body, intentSourceBody []byte, identity string, injectFreeTierTools bool) ([]byte, error) {
identity = strings.TrimSpace(identity)
if identity == "" {
if gjson.GetBytes(body, "prompt_cache_key").Exists() {
return sjson.DeleteBytes(body, "prompt_cache_key")
}
return body, nil
}
out, err := sjson.SetBytes(body, "prompt_cache_key", identity)
if err != nil {
return nil, err
}
if !injectFreeTierTools {
return out, nil
}
// Inspect the pre-sanitization source. patchGrokResponsesBody may remove an
// unsupported client tool and its tool_choice; that must not turn an
// explicit client tool intent into an eligible native-tool request.
if gjson.GetBytes(intentSourceBody, "tools").Exists() || gjson.GetBytes(intentSourceBody, "tool_choice").Exists() {
return out, nil
}
out, err = sjson.SetRawBytes(out, "tools", []byte(grokFreeCacheNativeToolsJSON))
if err != nil {
return nil, err
}
return sjson.SetBytes(out, "tool_choice", grokFreeCacheDisabledToolChoice)
}
// applyGrokCacheHeaders applies the documented Chat Completions conversation
// routing header. The request is built from a fresh header map, so client
// supplied x-grok headers cannot override this server-derived value.
func applyGrokCacheHeaders(headers http.Header, identity string) {
if headers == nil {
return
}
identity = strings.TrimSpace(identity)
if identity == "" {
headers.Del(grokConversationIDHeader)
return
}
headers.Set(grokConversationIDHeader, identity)
}
// stripGrokChatPromptCacheKey removes the Responses-only body field after it
// has been used as an identity seed. Chat Completions routes cache by header.
func stripGrokChatPromptCacheKey(body []byte) ([]byte, error) {
if !gjson.GetBytes(body, "prompt_cache_key").Exists() {
return body, nil
}
return sjson.DeleteBytes(body, "prompt_cache_key")
}
@@ -0,0 +1,273 @@
//go:build unit
package service
import (
"net/http"
"net/http/httptest"
"sync"
"testing"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
"github.com/tidwall/gjson"
)
func newGrokCacheTestContext(apiKeyID int64) *gin.Context {
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil)
if apiKeyID > 0 {
c.Set("api_key", &APIKey{ID: apiKeyID, Group: &Group{Platform: PlatformGrok}})
}
return c
}
func TestResolveGrokCacheIdentityStableAcrossAppendOnlyTurns(t *testing.T) {
gin.SetMode(gin.TestMode)
c := newGrokCacheTestContext(101)
round1 := []byte(`{"model":"grok","instructions":"be concise","tools":[{"type":"function","name":"lookup","parameters":{"type":"object"}}],"input":[{"role":"user","content":"first question"}]}`)
round2 := []byte(`{"model":"grok","instructions":"be concise","tools":[{"type":"function","name":"lookup","parameters":{"type":"object"}}],"input":[{"role":"user","content":"first question"},{"role":"assistant","content":"first answer"},{"role":"user","content":"second question"}]}`)
first := resolveGrokCacheIdentity(c, round1, "", "grok-4.5")
second := resolveGrokCacheIdentity(c, round2, "", "grok-4.5")
require.NotEmpty(t, first)
require.Len(t, first, 36)
require.Equal(t, first, second)
}
func TestResolveGrokCacheIdentityIsolatesAPIKeyAndMappedModel(t *testing.T) {
gin.SetMode(gin.TestMode)
body := []byte(`{"model":"grok","input":"same prompt"}`)
base := resolveGrokCacheIdentity(newGrokCacheTestContext(201), body, "", "grok-4.5")
otherTenant := resolveGrokCacheIdentity(newGrokCacheTestContext(202), body, "", "grok-4.5")
otherModel := resolveGrokCacheIdentity(newGrokCacheTestContext(201), body, "", "grok-4.3")
require.NotEmpty(t, base)
require.NotEqual(t, base, otherTenant)
require.NotEqual(t, base, otherModel)
}
func TestResolveGrokCacheIdentityUsesAndIsolatesNativeConversationHeader(t *testing.T) {
gin.SetMode(gin.TestMode)
c := newGrokCacheTestContext(301)
c.Request.Header.Set(grokConversationIDHeader, "raw-native-conversation")
body1 := []byte(`{"model":"grok","input":"one"}`)
body2 := []byte(`{"model":"grok","input":"different body that must not replace the explicit session"}`)
first := resolveGrokCacheIdentity(c, body1, "body-cache-key", "grok-4.5")
second := resolveGrokCacheIdentity(c, body2, "another-body-cache-key", "grok-4.5")
require.Equal(t, "raw-native-conversation", (&OpenAIGatewayService{}).ExtractSessionID(c, body1))
require.Equal(t, first, second)
require.NotEqual(t, "raw-native-conversation", first)
require.NotContains(t, first, "raw-native-conversation")
}
func TestResolveGrokCacheIdentityExplicitHeaderPriority(t *testing.T) {
gin.SetMode(gin.TestMode)
body := []byte(`{"model":"grok","prompt_cache_key":"body-key","input":"hi"}`)
c := newGrokCacheTestContext(401)
c.Request.Header.Set(grokConversationIDHeader, "grok-key")
c.Request.Header.Set("conversation_id", "conversation-key")
c.Request.Header.Set("session_id", "session-key")
got := resolveGrokCacheIdentity(c, body, "explicit-argument", "grok-4.5")
onlySession := newGrokCacheTestContext(401)
onlySession.Request.Header.Set("session_id", "session-key")
want := resolveGrokCacheIdentity(onlySession, []byte(`{"model":"grok","input":"unrelated"}`), "", "grok-4.5")
require.Equal(t, want, got)
}
func TestResolveGrokCacheIdentityFailsClosedWithoutAPIKeyContext(t *testing.T) {
gin.SetMode(gin.TestMode)
c := newGrokCacheTestContext(0)
c.Request.Header.Set(grokConversationIDHeader, "native-session")
require.Empty(t, resolveGrokCacheIdentity(c, []byte(`{"model":"grok","input":"hi"}`), "", "grok-4.5"))
require.Empty(t, resolveGrokCacheIdentity(nil, []byte(`{"model":"grok","prompt_cache_key":"key"}`), "key", "grok-4.5"))
}
func TestGrokConversationHeaderIsScopedToGrokRequestScheduling(t *testing.T) {
gin.SetMode(gin.TestMode)
body := []byte(`{"model":"grok","prompt_cache_key":"body-session","input":"hi"}`)
grokContext := newGrokCacheTestContext(601)
grokContext.Request.Header.Set(grokConversationIDHeader, "native-grok-session")
require.Equal(t, "native-grok-session", (&OpenAIGatewayService{}).ExtractSessionID(grokContext, body))
openAIContext := newGrokCacheTestContext(601)
openAIContext.Set("api_key", &APIKey{ID: 601, Group: &Group{Platform: PlatformOpenAI}})
openAIContext.Request.Header.Set(grokConversationIDHeader, "must-be-ignored")
require.Equal(t, "body-session", (&OpenAIGatewayService{}).ExtractSessionID(openAIContext, body))
withoutGrokHeader := newGrokCacheTestContext(601)
withoutGrokHeader.Set("api_key", &APIKey{ID: 601, Group: &Group{Platform: PlatformOpenAI}})
require.Equal(t,
(&OpenAIGatewayService{}).GenerateSessionHash(withoutGrokHeader, body),
(&OpenAIGatewayService{}).GenerateSessionHash(openAIContext, body),
)
}
func TestApplyGrokCacheIdentityWritesResponsesBodyAndHeader(t *testing.T) {
sourceBody := []byte(`{"model":"grok-4.5","prompt_cache_key":"raw-client-key"}`)
body, err := applyGrokResponsesCacheIdentity(sourceBody, sourceBody, "isolated-id", true)
require.NoError(t, err)
require.Equal(t, "isolated-id", gjson.GetBytes(body, "prompt_cache_key").String())
require.Equal(t, "web_search", gjson.GetBytes(body, "tools.0.type").String())
require.Equal(t, "x_search", gjson.GetBytes(body, "tools.1.type").String())
require.Equal(t, grokFreeCacheDisabledToolChoice, gjson.GetBytes(body, "tool_choice").String())
headers := make(http.Header)
headers.Set(grokConversationIDHeader, "spoofed-client-value")
applyGrokCacheHeaders(headers, "isolated-id")
require.Equal(t, "isolated-id", headers.Get(grokConversationIDHeader))
applyGrokCacheHeaders(headers, "")
require.Empty(t, headers.Get(grokConversationIDHeader))
chatBody, err := stripGrokChatPromptCacheKey(body)
require.NoError(t, err)
require.False(t, gjson.GetBytes(chatBody, "prompt_cache_key").Exists())
unscopedSourceBody := []byte(`{"model":"grok","prompt_cache_key":"raw-client-key"}`)
unscopedBody, err := applyGrokResponsesCacheIdentity(unscopedSourceBody, unscopedSourceBody, "", true)
require.NoError(t, err)
require.False(t, gjson.GetBytes(unscopedBody, "prompt_cache_key").Exists())
require.False(t, gjson.GetBytes(unscopedBody, "tools").Exists())
require.False(t, gjson.GetBytes(unscopedBody, "tool_choice").Exists())
}
func TestApplyGrokCacheIdentityPreservesExplicitClientToolFields(t *testing.T) {
tests := []struct {
name string
body string
}{
{
name: "tools only",
body: `{"model":"grok","tools":[{"type":"function","name":"lookup","parameters":{"type":"object"}}]}`,
},
{
name: "empty tools array",
body: `{"model":"grok","tools":[]}`,
},
{
name: "null tools",
body: `{"model":"grok","tools":null}`,
},
{
name: "tool choice only",
body: `{"model":"grok","tool_choice":{"type":"function","name":"lookup"}}`,
},
{
name: "null tool choice",
body: `{"model":"grok","tool_choice":null}`,
},
{
name: "both fields",
body: `{"model":"grok","tools":[{"type":"web_search"}],"tool_choice":"auto"}`,
},
{
name: "unsupported tool",
body: `{"model":"grok","tools":[{"type":"namespace","name":"client_tools"}]}`,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
beforeTools := gjson.Get(tt.body, "tools")
beforeChoice := gjson.Get(tt.body, "tool_choice")
body, err := applyGrokResponsesCacheIdentity([]byte(tt.body), []byte(tt.body), "isolated-id", true)
require.NoError(t, err)
require.Equal(t, "isolated-id", gjson.GetBytes(body, "prompt_cache_key").String())
require.Equal(t, beforeTools.Exists(), gjson.GetBytes(body, "tools").Exists())
require.Equal(t, beforeTools.Raw, gjson.GetBytes(body, "tools").Raw)
require.Equal(t, beforeChoice.Exists(), gjson.GetBytes(body, "tool_choice").Exists())
require.Equal(t, beforeChoice.Raw, gjson.GetBytes(body, "tool_choice").Raw)
})
}
}
func TestApplyGrokCacheIdentityUsesPreSanitizationToolIntent(t *testing.T) {
tests := []struct {
name string
intentBody string
}{
{
name: "unsupported tools removed by sanitizer",
intentBody: `{"model":"grok","tools":[{"type":"namespace","name":"client_tools"}]}`,
},
{
name: "tool choice removed with unsupported tool",
intentBody: `{"model":"grok","tools":[{"type":"namespace","name":"client_tools"}],"tool_choice":{"type":"namespace","name":"client_tools"}}`,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
// This is the shape apply receives after patchGrokResponsesBody has
// removed unsupported tools and their associated tool_choice.
patchedBody := []byte(`{"model":"grok-4.5","input":"hello"}`)
body, err := applyGrokResponsesCacheIdentity(patchedBody, []byte(tt.intentBody), "isolated-id", true)
require.NoError(t, err)
require.Equal(t, "isolated-id", gjson.GetBytes(body, "prompt_cache_key").String())
require.False(t, gjson.GetBytes(body, "tools").Exists())
require.False(t, gjson.GetBytes(body, "tool_choice").Exists())
})
}
}
func TestApplyGrokCacheIdentityWithoutFreeTierRoutingOnlyWritesIdentity(t *testing.T) {
sourceBody := []byte(`{"model":"grok-4.5","input":"hello"}`)
body, err := applyGrokResponsesCacheIdentity(sourceBody, sourceBody, "isolated-id", false)
require.NoError(t, err)
require.Equal(t, "isolated-id", gjson.GetBytes(body, "prompt_cache_key").String())
require.False(t, gjson.GetBytes(body, "tools").Exists())
require.False(t, gjson.GetBytes(body, "tool_choice").Exists())
}
func TestGrokCompactRequestSkipsCacheIdentityAndNativeTools(t *testing.T) {
gin.SetMode(gin.TestMode)
c := newGrokCacheTestContext(701)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses/compact", nil)
body := []byte(`{"model":"grok","input":"compact this","prompt_cache_key":"raw-client-key"}`)
identity := resolveGrokCacheIdentity(c, body, "", "grok-4.5")
patched, err := applyGrokResponsesCacheIdentity(body, body, identity, true)
require.NoError(t, err)
require.Empty(t, identity)
require.False(t, gjson.GetBytes(patched, "prompt_cache_key").Exists())
require.False(t, gjson.GetBytes(patched, "tools").Exists())
require.False(t, gjson.GetBytes(patched, "tool_choice").Exists())
}
func TestResolveGrokCacheIdentityConcurrentDeterminism(t *testing.T) {
gin.SetMode(gin.TestMode)
const workers = 50
body := []byte(`{"model":"grok","messages":[{"role":"system","content":"stable"},{"role":"user","content":"hello"}]}`)
identities := make(chan string, workers)
var wg sync.WaitGroup
for i := 0; i < workers; i++ {
wg.Add(1)
go func() {
defer wg.Done()
identities <- resolveGrokCacheIdentity(newGrokCacheTestContext(501), body, "", "grok-4.5")
}()
}
wg.Wait()
close(identities)
var first string
for identity := range identities {
if first == "" {
first = identity
}
require.Equal(t, first, identity)
}
require.NotEmpty(t, first)
}
@@ -6,6 +6,7 @@ import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"mime/multipart"
"net/http"
@@ -172,13 +173,15 @@ func TestBuildGrokResponsesRequestUsesAccountBaseURLAndBearerToken(t *testing.T)
},
}
req, err := buildGrokResponsesRequest(context.Background(), nil, account, []byte(`{"model":"grok-4.3"}`), "access-token")
req, err := buildGrokResponsesRequest(context.Background(), nil, account, []byte(`{"model":"grok-4.3"}`), "access-token", "isolated-cache-id")
require.NoError(t, err)
require.Equal(t, http.MethodPost, req.Method)
require.Equal(t, "https://xai.test/v1/responses", req.URL.String())
require.Equal(t, "Bearer access-token", req.Header.Get("Authorization"))
require.Equal(t, "application/json", req.Header.Get("Content-Type"))
require.Contains(t, req.Header.Get("Accept"), "text/event-stream")
require.Equal(t, grokCLIVersion, req.Header.Get("X-Grok-Client-Version"))
require.Equal(t, "isolated-cache-id", req.Header.Get(grokConversationIDHeader))
data, err := io.ReadAll(req.Body)
require.NoError(t, err)
@@ -196,7 +199,7 @@ func TestBuildGrokResponsesRequestRejectsUnsafeAccountBaseURL(t *testing.T) {
},
}
_, err := buildGrokResponsesRequest(context.Background(), nil, account, []byte(`{"model":"grok-4.3"}`), "access-token")
_, err := buildGrokResponsesRequest(context.Background(), nil, account, []byte(`{"model":"grok-4.3"}`), "access-token", "")
require.Error(t, err)
require.Contains(t, err.Error(), "invalid base url")
}
@@ -324,6 +327,7 @@ func TestForwardGrokMediaImagesGenerationNormalizesImagineAlias(t *testing.T) {
require.Equal(t, http.MethodPost, upstream.lastReq.Method)
require.Equal(t, "Bearer api-key", upstream.lastReq.Header.Get("Authorization"))
require.Equal(t, "application/json", upstream.lastReq.Header.Get("Content-Type"))
require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version"))
require.JSONEq(t, `{"model":"grok-imagine-image-quality","prompt":"draw a cat"}`, string(upstream.lastBody))
require.Equal(t, http.StatusOK, recorder.Code)
require.JSONEq(t, `{"data":[]}`, recorder.Body.String())
@@ -611,8 +615,9 @@ func TestForwardAsChatCompletionsForGrokUsesXAIChatCompletionsAndSnapshots(t *te
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
body := []byte(`{"model":"grok","messages":[{"role":"user","content":"hi"}],"stream":false}`)
body := []byte(`{"model":"grok","messages":[{"role":"user","content":"hi"}],"stream":false,"prompt_cache_key":"raw-client-cache-key"}`)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body))
c.Set("api_key", &APIKey{ID: 5101})
account := &Account{
ID: 51,
@@ -641,7 +646,7 @@ func TestForwardAsChatCompletionsForGrokUsesXAIChatCompletionsAndSnapshots(t *te
"X-Ratelimit-Limit-Tokens": []string{"1000"},
"X-Ratelimit-Remaining-Tokens": []string{"990"},
},
Body: io.NopCloser(strings.NewReader(`{"id":"chatcmpl","object":"chat.completion","model":"grok-4.3","choices":[],"usage":{"prompt_tokens":1,"completion_tokens":2}}`)),
Body: io.NopCloser(strings.NewReader(`{"id":"chatcmpl","object":"chat.completion","model":"grok-4.3","choices":[],"usage":{"prompt_tokens":1,"completion_tokens":2,"prompt_tokens_details":{"cached_tokens":1}}}`)),
}}
svc := &OpenAIGatewayService{
httpUpstream: upstream,
@@ -653,11 +658,15 @@ func TestForwardAsChatCompletionsForGrokUsesXAIChatCompletionsAndSnapshots(t *te
require.NoError(t, err)
require.Equal(t, xai.DefaultCLIBaseURL+"/chat/completions", upstream.lastReq.URL.String())
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.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, 1, result.Usage.InputTokens)
require.Equal(t, 2, result.Usage.OutputTokens)
require.Equal(t, 1, result.Usage.CacheReadInputTokens)
require.NotNil(t, repo.updates[51][grokQuotaSnapshotExtraKey])
require.Equal(t, http.StatusOK, recorder.Code)
}
@@ -671,6 +680,7 @@ func TestForwardGrokResponsesStreamingUsesXAIResponsesAndSnapshots(t *testing.T)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
c.Request.Header.Set("OpenAI-Beta", "responses=experimental")
c.Set("api_key", &APIKey{ID: 5201})
account := &Account{
ID: 52,
@@ -719,6 +729,11 @@ func TestForwardGrokResponsesStreamingUsesXAIResponsesAndSnapshots(t *testing.T)
require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization"))
require.Equal(t, "responses=experimental", upstream.lastReq.Header.Get("OpenAI-Beta"))
require.Equal(t, "grok-4.5", 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())
require.Equal(t, "x_search", gjson.GetBytes(upstream.lastBody, "tools.1.type").String())
require.Equal(t, "none", gjson.GetBytes(upstream.lastBody, "tool_choice").String())
require.Equal(t, "high", gjson.GetBytes(upstream.lastBody, "reasoning_effort").String())
require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool())
require.True(t, result.Stream)
@@ -734,6 +749,133 @@ func TestForwardGrokResponsesStreamingUsesXAIResponsesAndSnapshots(t *testing.T)
require.NotNil(t, repo.updates[52][grokQuotaSnapshotExtraKey])
}
func TestForwardGrokResponsesNonStreamingUsesCacheIdentityAndCachedUsage(t *testing.T) {
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
body := []byte(`{"model":"grok","input":"hi","stream":false,"tools":[{"type":"namespace","name":"client_tools"}],"tool_choice":{"type":"namespace","name":"client_tools"}}`)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
c.Set("api_key", &APIKey{ID: 5202})
account := &Account{
ID: 56,
Name: "grok",
Platform: PlatformGrok,
Type: AccountTypeOAuth,
Concurrency: 1,
Credentials: map[string]any{
"access_token": "access-token",
"expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339),
"base_url": xai.DefaultCLIBaseURL,
},
}
repo := &grokQuotaAccountRepo{
mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{56: account},
},
}
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{
"Content-Type": []string{"application/json"},
"Xai-Request-Id": []string{"xai-non-stream-req"},
},
Body: io.NopCloser(strings.NewReader(`{"id":"resp_grok_non_stream","object":"response","model":"grok-4.3","status":"completed","output":[{"type":"message","role":"assistant","content":[{"type":"output_text","text":"ok"}]}],"usage":{"input_tokens":7,"output_tokens":2,"total_tokens":9,"input_tokens_details":{"cached_tokens":4}}}`)),
}}
svc := &OpenAIGatewayService{
httpUpstream: upstream,
grokTokenProvider: NewGrokTokenProvider(repo, nil),
accountRepo: repo,
}
result, err := svc.forwardGrokResponses(context.Background(), c, account, body, "grok", false, time.Now())
require.NoError(t, err)
require.NotNil(t, result)
require.False(t, result.Stream)
require.Equal(t, "resp_grok_non_stream", result.ResponseID)
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())
identity := gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String()
require.NotEmpty(t, identity)
require.Equal(t, identity, upstream.lastReq.Header.Get(grokConversationIDHeader))
// The sanitizer drops this unsupported client tool, but its explicit intent
// must still prevent native cache-routing tools from being injected.
require.False(t, gjson.GetBytes(upstream.lastBody, "tools").Exists())
require.False(t, gjson.GetBytes(upstream.lastBody, "tool_choice").Exists())
require.Equal(t, "resp_grok_non_stream", gjson.Get(recorder.Body.String(), "id").String())
}
func TestForwardGrokResponsesFailoverKeepsCacheIdentityAcrossAccounts(t *testing.T) {
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
body := []byte(`{"model":"grok","input":[{"role":"user","content":"stable prefix"}],"stream":false}`)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body))
c.Set("api_key", &APIKey{ID: 5203})
newAccount := func(id int64, token string) *Account {
return &Account{
ID: id,
Name: fmt.Sprintf("grok-%d", id),
Platform: PlatformGrok,
Type: AccountTypeOAuth,
Concurrency: 1,
Credentials: map[string]any{
"access_token": token,
"expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339),
"base_url": xai.DefaultCLIBaseURL,
},
}
}
firstAccount := newAccount(58, "access-token-a")
secondAccount := newAccount(59, "access-token-b")
repo := &grokQuotaAccountRepo{
mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{58: firstAccount, 59: secondAccount},
},
}
upstream := &httpUpstreamRecorder{responses: []*http.Response{
{
StatusCode: http.StatusServiceUnavailable,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(`{"error":{"message":"temporary"}}`)),
},
{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(`{"id":"resp_after_failover","object":"response","model":"grok-4.3","status":"completed","output":[],"usage":{"input_tokens":5,"output_tokens":1}}`)),
},
}}
svc := &OpenAIGatewayService{
httpUpstream: upstream,
grokTokenProvider: NewGrokTokenProvider(repo, nil),
accountRepo: repo,
}
_, err := svc.forwardGrokResponses(context.Background(), c, firstAccount, body, "grok", false, time.Now())
var failoverErr *UpstreamFailoverError
require.ErrorAs(t, err, &failoverErr)
result, err := svc.forwardGrokResponses(context.Background(), c, secondAccount, body, "grok", false, time.Now())
require.NoError(t, err)
require.NotNil(t, result)
require.Len(t, upstream.requests, 2)
require.Len(t, upstream.bodies, 2)
firstIdentity := gjson.GetBytes(upstream.bodies[0], "prompt_cache_key").String()
secondIdentity := gjson.GetBytes(upstream.bodies[1], "prompt_cache_key").String()
require.NotEmpty(t, firstIdentity)
require.Equal(t, firstIdentity, secondIdentity)
require.Equal(t, firstIdentity, upstream.requests[0].Header.Get(grokConversationIDHeader))
require.Equal(t, secondIdentity, upstream.requests[1].Header.Get(grokConversationIDHeader))
require.Equal(t, "Bearer access-token-a", upstream.requests[0].Header.Get("Authorization"))
require.Equal(t, "Bearer access-token-b", upstream.requests[1].Header.Get("Authorization"))
}
func TestForwardAsChatCompletionsForGrokStreamingUsesRawXAIChatCompletions(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -742,6 +884,8 @@ func TestForwardAsChatCompletionsForGrokStreamingUsesRawXAIChatCompletions(t *te
body := []byte(`{"model":"grok","messages":[{"role":"user","content":"hi"}],"stream":true}`)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
c.Request.Header.Set(grokConversationIDHeader, "native-client-conversation")
c.Set("api_key", &APIKey{ID: 5301})
account := &Account{
ID: 53,
@@ -791,6 +935,9 @@ 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, "sub2api-grok/1.0", upstream.lastReq.Header.Get("User-Agent"))
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.True(t, gjson.GetBytes(upstream.lastBody, "stream_options.include_usage").Bool())
require.True(t, result.Stream)
@@ -809,6 +956,7 @@ func TestForwardAsChatCompletionsForGrokComposerBridgesImageInput(t *testing.T)
body := []byte(`{"model":"grok-composer-2.5-fast","messages":[{"role":"system","content":"You are concise."},{"role":"user","content":[{"type":"text","text":"What is shown?"},{"type":"image_url","image_url":{"url":"data:image/png;base64,QUJD"}}]}],"stream":false}`)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
c.Set("api_key", &APIKey{ID: 5501})
account := &Account{
ID: 55,
@@ -858,9 +1006,11 @@ func TestForwardAsChatCompletionsForGrokComposerBridgesImageInput(t *testing.T)
require.NotNil(t, result)
require.Len(t, upstream.requests, 2)
require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.requests[0].URL.String())
require.Empty(t, upstream.requests[0].Header.Get(grokConversationIDHeader))
require.Equal(t, "grok-build-0.1", gjson.GetBytes(upstream.bodies[0], "model").String())
require.Equal(t, "input_image", gjson.GetBytes(upstream.bodies[0], "input.0.content.1.type").String())
require.Equal(t, xai.DefaultCLIBaseURL+"/chat/completions", upstream.requests[1].URL.String())
require.NotEmpty(t, upstream.requests[1].Header.Get(grokConversationIDHeader))
require.Equal(t, "grok-composer-2.5-fast", gjson.GetBytes(upstream.bodies[1], "model").String())
require.False(t, strings.Contains(string(upstream.bodies[1]), "image_url"))
require.Contains(t, gjson.GetBytes(upstream.bodies[1], "messages.1.content").String(), "Image 1 description")
@@ -878,6 +1028,7 @@ func TestForwardAsAnthropicForGrokUsesXAIResponses(t *testing.T) {
c, _ := gin.CreateTestContext(recorder)
body := []byte(`{"model":"grok","max_tokens":32,"stream":false,"messages":[{"role":"user","content":"hi"}]}`)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(body))
c.Set("api_key", &APIKey{ID: 5401})
account := &Account{
ID: 54,
@@ -896,7 +1047,7 @@ func TestForwardAsAnthropicForGrokUsesXAIResponses(t *testing.T) {
accountsByID: map[int64]*Account{54: account},
},
}
upstream := &httpUpstreamRecorder{resp: openAICompatSSECompletedResponse("resp_grok_messages", "grok-4.3")}
upstream := &httpUpstreamRecorder{resp: grokMessagesSSECompletedResponse("resp_grok_messages", 3)}
svc := &OpenAIGatewayService{
httpUpstream: upstream,
grokTokenProvider: NewGrokTokenProvider(repo, nil),
@@ -908,17 +1059,84 @@ func TestForwardAsAnthropicForGrokUsesXAIResponses(t *testing.T) {
require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String())
require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization"))
require.Equal(t, "sub2api-grok/1.0", 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.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())
require.Equal(t, "x_search", gjson.GetBytes(upstream.lastBody, "tools.1.type").String())
require.Equal(t, "none", gjson.GetBytes(upstream.lastBody, "tool_choice").String())
require.Empty(t, upstream.lastReq.Header.Get("session_id"))
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, 5, result.Usage.InputTokens)
require.Equal(t, 2, result.Usage.OutputTokens)
require.Equal(t, 3, result.Usage.CacheReadInputTokens)
require.Contains(t, recorder.Body.String(), `"type":"message"`)
require.Equal(t, int64(3), gjson.Get(recorder.Body.String(), "usage.cache_read_input_tokens").Int())
require.Contains(t, recorder.Body.String(), "ok")
}
func TestForwardAsAnthropicForGrokStreamingPreservesCacheUsage(t *testing.T) {
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
body := []byte(`{"model":"grok","max_tokens":32,"stream":true,"messages":[{"role":"user","content":"hi"}]}`)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(body))
c.Set("api_key", &APIKey{ID: 5402})
account := &Account{
ID: 57,
Name: "grok",
Platform: PlatformGrok,
Type: AccountTypeOAuth,
Concurrency: 1,
Credentials: map[string]any{
"access_token": "access-token",
"expires_at": time.Now().Add(time.Hour).UTC().Format(time.RFC3339),
"base_url": xai.DefaultCLIBaseURL,
},
}
repo := &grokQuotaAccountRepo{
mockAccountRepoForPlatform: &mockAccountRepoForPlatform{
accountsByID: map[int64]*Account{57: account},
},
}
upstream := &httpUpstreamRecorder{resp: grokMessagesSSECompletedResponse("resp_grok_messages_stream", 2)}
svc := &OpenAIGatewayService{
httpUpstream: upstream,
grokTokenProvider: NewGrokTokenProvider(repo, nil),
accountRepo: repo,
}
result, err := svc.ForwardAsAnthropic(context.Background(), c, account, body, "", "")
require.NoError(t, err)
require.NotNil(t, result)
require.Equal(t, 2, result.Usage.CacheReadInputTokens)
identity := gjson.GetBytes(upstream.lastBody, "prompt_cache_key").String()
require.NotEmpty(t, identity)
require.Equal(t, identity, upstream.lastReq.Header.Get(grokConversationIDHeader))
require.Contains(t, recorder.Header().Get("Content-Type"), "text/event-stream")
require.Contains(t, recorder.Body.String(), `"cache_read_input_tokens":2`)
}
func grokMessagesSSECompletedResponse(responseID string, cachedTokens int) *http.Response {
body := strings.Join([]string{
fmt.Sprintf(`data: {"type":"response.completed","response":{"id":%q,"object":"response","model":"grok-4.3","status":"completed","output":[{"type":"message","id":"msg_1","role":"assistant","status":"completed","content":[{"type":"output_text","text":"ok"}]}],"usage":{"input_tokens":5,"output_tokens":2,"total_tokens":7,"input_tokens_details":{"cached_tokens":%d}}}}`, responseID, cachedTokens),
"",
"data: [DONE]",
"",
}, "\n")
return &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"text/event-stream"}},
Body: io.NopCloser(strings.NewReader(body)),
}
}
func TestHandleGrokAccountUpstreamErrorTempUnschedulesReadinessStates(t *testing.T) {
tests := []struct {
name string
@@ -244,12 +244,17 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic(
return nil, policyErr
}
responsesBody = updatedBody
grokCacheIdentity := ""
if account.Platform == PlatformGrok {
grokCacheIdentity = resolveGrokCacheIdentity(c, responsesBody, promptCacheKey, upstreamModel)
patchedBody, patchErr := patchGrokResponsesBody(responsesBody, upstreamModel)
if patchErr != nil {
return nil, patchErr
}
responsesBody = patchedBody
responsesBody, patchErr = applyGrokResponsesCacheIdentity(patchedBody, responsesBody, grokCacheIdentity, account.IsGrokOAuth())
if patchErr != nil {
return nil, fmt.Errorf("apply grok prompt cache identity: %w", patchErr)
}
}
// 5. Get access token
@@ -269,7 +274,7 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic(
upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx)
var upstreamReq *http.Request
if account.Platform == PlatformGrok {
upstreamReq, err = buildGrokResponsesRequest(upstreamCtx, c, account, responsesBody, token)
upstreamReq, err = buildGrokResponsesRequest(upstreamCtx, c, account, responsesBody, token, grokCacheIdentity)
} else {
upstreamReq, err = s.buildUpstreamRequest(upstreamCtx, c, account, responsesBody, token, isStream, promptCacheKey, false)
}
@@ -280,7 +285,7 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic(
// Override session_id with a deterministic UUID derived from the isolated
// session key, ensuring different API keys produce different upstream sessions.
if promptCacheKey != "" {
if account.Platform != PlatformGrok && promptCacheKey != "" {
isolatedSessionID := generateSessionUUID(isolateOpenAISessionID(apiKeyID, promptCacheKey))
upstreamReq.Header.Set("session_id", isolatedSessionID)
if upstreamReq.Header.Get("conversation_id") != "" {
@@ -101,7 +101,7 @@ func (s *OpenAIGatewayService) forwardAnthropicViaRawChatCompletions(
if err != nil {
return nil, err
}
resp, err := s.sendCCUpstreamRequest(ctx, c, account, targetURL, chatBody, clientStream, apiKey, account.GetOpenAIUserAgent())
resp, err := s.sendCCUpstreamRequest(ctx, c, account, targetURL, chatBody, clientStream, apiKey, account.GetOpenAIUserAgent(), "")
if err != nil {
return nil, err
}
@@ -92,7 +92,7 @@ func (s *OpenAIGatewayService) forwardResponsesViaRawChatCompletions(
if err != nil {
return nil, err
}
resp, err := s.sendCCUpstreamRequest(ctx, c, account, targetURL, chatBody, clientStream, apiKey, account.GetOpenAIUserAgent())
resp, err := s.sendCCUpstreamRequest(ctx, c, account, targetURL, chatBody, clientStream, apiKey, account.GetOpenAIUserAgent(), "")
if err != nil {
return nil, err
}
@@ -23,17 +23,7 @@ import (
// ExtractSessionID extracts the raw session ID from headers or body without hashing.
// Used by ForwardAsAnthropic to pass as prompt_cache_key for upstream cache.
func (s *OpenAIGatewayService) ExtractSessionID(c *gin.Context, body []byte) string {
if c == nil {
return ""
}
sessionID := strings.TrimSpace(c.GetHeader("session_id"))
if sessionID == "" {
sessionID = strings.TrimSpace(c.GetHeader("conversation_id"))
}
if sessionID == "" && len(body) > 0 {
sessionID = strings.TrimSpace(gjson.GetBytes(body, "prompt_cache_key").String())
}
return sessionID
return explicitOpenAIRequestSessionID(c, body)
}
func explicitOpenAISessionID(c *gin.Context, body []byte) string {
@@ -51,11 +41,33 @@ func explicitOpenAISessionID(c *gin.Context, body []byte) string {
return sessionID
}
// explicitOpenAIRequestSessionID extends the common OpenAI session signals
// with Grok's native conversation header only for requests authenticated to a
// Grok group. This keeps an unrelated x-grok-conv-id header from changing
// scheduling or upstream session behavior for non-Grok groups.
func explicitOpenAIRequestSessionID(c *gin.Context, body []byte) string {
if c == nil {
return ""
}
sessionID := strings.TrimSpace(c.GetHeader("session_id"))
if sessionID == "" {
sessionID = strings.TrimSpace(c.GetHeader("conversation_id"))
}
if sessionID == "" && isGrokRequestContext(c) {
sessionID = strings.TrimSpace(c.GetHeader(grokConversationIDHeader))
}
if sessionID == "" && len(body) > 0 {
sessionID = strings.TrimSpace(gjson.GetBytes(body, "prompt_cache_key").String())
}
return sessionID
}
// GenerateExplicitSessionHash generates a sticky-session hash only from explicit
// client session signals. It intentionally skips content-derived fallback and is
// used by stateless endpoints such as /v1/images.
func (s *OpenAIGatewayService) GenerateExplicitSessionHash(c *gin.Context, body []byte) string {
sessionID := explicitOpenAISessionID(c, body)
sessionID := explicitOpenAIRequestSessionID(c, body)
if sessionID == "" {
return ""
}
@@ -70,14 +82,15 @@ func (s *OpenAIGatewayService) GenerateExplicitSessionHash(c *gin.Context, body
// Priority:
// 1. Header: session_id
// 2. Header: conversation_id
// 3. Body: prompt_cache_key (opencode)
// 4. Body: content-based fallback (model + system + tools + first user message)
// 3. Header: x-grok-conv-id (Grok groups only)
// 4. Body: prompt_cache_key (opencode)
// 5. Body: content-based fallback (model + system + tools + first user message)
func (s *OpenAIGatewayService) GenerateSessionHash(c *gin.Context, body []byte) string {
if c == nil {
return ""
}
sessionID := explicitOpenAISessionID(c, body)
sessionID := explicitOpenAIRequestSessionID(c, body)
if sessionID == "" && len(body) > 0 {
sessionID = deriveOpenAIContentSessionSeed(body)
}
@@ -421,6 +421,10 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
storeDisabled,
)
currentBridgePayload := firstPayload
// Keep the first turn as the stable conversation seed. The mapped model
// is resolved again for each turn below so an in-connection model switch
// cannot reuse another model's upstream cache identity.
grokCacheSeedPayload := firstPayload.payloadRaw
var bridgeReplayInput []json.RawMessage
bridgeReplayInputExists := false
for turn := 1; ; turn++ {
@@ -469,6 +473,13 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
openAIWSRawPayloadHasToolCallOutput(currentBridgePayload.payloadRaw),
)
}
grokCacheIdentity := ""
if account.Platform == PlatformGrok {
grokCacheIdentity, err = resolveGrokWSCacheIdentity(c, account, grokCacheSeedPayload, currentBridgePayload.originalModel)
if err != nil {
return fmt.Errorf("resolve Grok websocket cache identity: %w", err)
}
}
result, bridgeErr := s.proxyOpenAIWSHTTPBridgeTurn(
ctx,
c,
@@ -480,6 +491,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
currentBridgePayload.imageBillingModel,
currentBridgePayload.imageSizeTier,
currentBridgePayload.imageInputSize,
grokCacheIdentity,
turn,
writeClientMessage,
)
@@ -155,6 +155,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
imageBillingModel string,
imageSizeTier string,
imageInputSize string,
grokCacheIdentity string,
turn int,
writeClientMessage func([]byte) error,
) (*OpenAIForwardResult, error) {
@@ -179,21 +180,19 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx)
var upstreamReq *http.Request
if account.Platform == PlatformGrok {
upstreamModel := strings.TrimSpace(gjson.GetBytes(body, "model").String())
if originalModel != "" {
if mappedModel := normalizeOpenAIModelForUpstream(account, account.GetMappedModel(originalModel)); mappedModel != "" {
upstreamModel = mappedModel
}
}
if upstreamModel == "" {
upstreamModel = "grok-4.3"
}
upstreamModel := resolveGrokWSUpstreamModel(account, body, originalModel)
grokIntentSourceBody := body
body, err = patchGrokResponsesBody(body, upstreamModel)
if err != nil {
releaseUpstreamCtx()
return nil, err
}
upstreamReq, err = buildGrokResponsesRequest(upstreamCtx, c, account, body, token)
body, err = applyGrokResponsesCacheIdentity(body, grokIntentSourceBody, grokCacheIdentity, account.IsGrokOAuth())
if err != nil {
releaseUpstreamCtx()
return nil, fmt.Errorf("apply grok prompt cache identity: %w", err)
}
upstreamReq, err = buildGrokResponsesRequest(upstreamCtx, c, account, body, token, grokCacheIdentity)
} else {
upstreamReq, err = s.buildUpstreamRequestOpenAIPassthrough(upstreamCtx, c, account, body, token)
}
@@ -407,3 +406,25 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
}
return resultWithUsage(), errors.New("upstream http bridge stream ended before terminal event")
}
func resolveGrokWSCacheIdentity(c *gin.Context, account *Account, payload []byte, originalModel string) (string, error) {
body, err := prepareOpenAIWSHTTPBridgeBody(payload)
if err != nil {
return "", err
}
upstreamModel := resolveGrokWSUpstreamModel(account, body, originalModel)
return resolveGrokCacheIdentity(c, body, "", upstreamModel), nil
}
func resolveGrokWSUpstreamModel(account *Account, body []byte, originalModel string) string {
upstreamModel := strings.TrimSpace(gjson.GetBytes(body, "model").String())
if account != nil && originalModel != "" {
if mappedModel := normalizeOpenAIModelForUpstream(account, account.GetMappedModel(originalModel)); mappedModel != "" {
upstreamModel = mappedModel
}
}
if upstreamModel == "" {
upstreamModel = "grok-4.3"
}
return upstreamModel
}
@@ -3,6 +3,7 @@ package service
import (
"context"
"errors"
"fmt"
"io"
"net/http"
"net/http/httptest"
@@ -127,6 +128,7 @@ func TestOpenAIWSHTTPBridgeRelaysSSEFramesAsWebSocketMessages(t *testing.T) {
"",
"",
"",
"",
1,
writeClient,
)
@@ -179,21 +181,28 @@ func TestOpenAIWSHTTPBridgeRelaysSSEFramesAsWebSocketMessages(t *testing.T) {
func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridge(t *testing.T) {
gin.SetMode(gin.TestMode)
sseBody := strings.Join([]string{
`data: {"type":"response.created","response":{"id":"resp_grok_ws","model":"grok-4.3"}}`,
"",
`data: {"type":"response.output_text.delta","response":{"id":"resp_grok_ws"},"delta":"ok"}`,
"",
`data: {"type":"response.completed","response":{"id":"resp_grok_ws","model":"grok-4.3","usage":{"input_tokens":4,"output_tokens":2}}}`,
"",
}, "\n")
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{
"Content-Type": []string{"text/event-stream"},
"Xai-Request-Id": []string{"xai-ws-req"},
},
Body: io.NopCloser(strings.NewReader(sseBody)),
bridgeResponse := func(responseID, requestID string, cachedTokens int) *http.Response {
sseBody := strings.Join([]string{
`data: {"type":"response.created","response":{"id":"` + responseID + `","model":"grok-4.3"}}`,
"",
`data: {"type":"response.output_text.delta","response":{"id":"` + responseID + `"},"delta":"ok"}`,
"",
`data: {"type":"response.completed","response":{"id":"` + responseID + `","model":"grok-4.3","usage":{"input_tokens":4,"output_tokens":2,"input_tokens_details":{"cached_tokens":` + fmt.Sprintf("%d", cachedTokens) + `}}}}`,
"",
}, "\n")
return &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{
"Content-Type": []string{"text/event-stream"},
"Xai-Request-Id": []string{requestID},
},
Body: io.NopCloser(strings.NewReader(sseBody)),
}
}
upstream := &httpUpstreamRecorder{responses: []*http.Response{
bridgeResponse("resp_grok_ws_1", "xai-ws-req-1", 0),
bridgeResponse("resp_grok_ws_2", "xai-ws-req-2", 3),
bridgeResponse("resp_grok_ws_3", "xai-ws-req-3", 0),
}}
svc := &OpenAIGatewayService{
cfg: &config.Config{
@@ -241,6 +250,7 @@ func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridge(t *testing.T)
req := r.Clone(r.Context())
req.Header = req.Header.Clone()
ginCtx.Request = req
ginCtx.Set("api_key", &APIKey{ID: 7101})
errCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "access-token", firstMessage, nil)
}))
@@ -271,6 +281,33 @@ func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridge(t *testing.T)
require.Equal(t, "response.created", gjson.GetBytes(created, "type").String())
require.Equal(t, "response.output_text.delta", gjson.GetBytes(delta, "type").String())
require.Equal(t, "response.completed", gjson.GetBytes(completed, "type").String())
require.Equal(t, "resp_grok_ws_1", gjson.GetBytes(completed, "response.id").String())
writeCtx, cancelWrite = context.WithTimeout(context.Background(), 3*time.Second)
err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","generate":true,"model":"grok","stream":true,"previous_response_id":"resp_grok_ws_1","input":"second turn"}`))
cancelWrite()
require.NoError(t, err)
created = readEvent()
delta = readEvent()
completed = readEvent()
require.Equal(t, "response.created", gjson.GetBytes(created, "type").String())
require.Equal(t, "response.output_text.delta", gjson.GetBytes(delta, "type").String())
require.Equal(t, "response.completed", gjson.GetBytes(completed, "type").String())
require.Equal(t, "resp_grok_ws_2", gjson.GetBytes(completed, "response.id").String())
writeCtx, cancelWrite = context.WithTimeout(context.Background(), 3*time.Second)
err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","generate":true,"model":"grok-4.3","stream":true,"previous_response_id":"resp_grok_ws_2","input":"third turn with a different model"}`))
cancelWrite()
require.NoError(t, err)
created = readEvent()
delta = readEvent()
completed = readEvent()
require.Equal(t, "response.created", gjson.GetBytes(created, "type").String())
require.Equal(t, "response.output_text.delta", gjson.GetBytes(delta, "type").String())
require.Equal(t, "response.completed", gjson.GetBytes(completed, "type").String())
require.Equal(t, "resp_grok_ws_3", gjson.GetBytes(completed, "response.id").String())
_ = clientConn.Close(coderws.StatusNormalClosure, "done")
select {
@@ -280,10 +317,30 @@ func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridge(t *testing.T)
require.Fail(t, "proxy did not finish after client close")
}
require.Len(t, upstream.requests, 3)
require.Len(t, upstream.bodies, 3)
require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String())
require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization"))
require.Equal(t, "sub2api-grok/1.0", upstream.lastReq.Header.Get("User-Agent"))
require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String())
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.5", 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())
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())
require.Equal(t, "x_search", gjson.GetBytes(upstream.lastBody, "tools.1.type").String())
require.Equal(t, "none", gjson.GetBytes(upstream.lastBody, "tool_choice").String())
firstIdentity := gjson.GetBytes(upstream.bodies[0], "prompt_cache_key").String()
secondIdentity := gjson.GetBytes(upstream.bodies[1], "prompt_cache_key").String()
thirdIdentity := gjson.GetBytes(upstream.bodies[2], "prompt_cache_key").String()
require.NotEmpty(t, firstIdentity)
require.Equal(t, firstIdentity, secondIdentity)
require.NotEmpty(t, thirdIdentity)
require.NotEqual(t, firstIdentity, thirdIdentity)
require.Equal(t, firstIdentity, upstream.requests[0].Header.Get(grokConversationIDHeader))
require.Equal(t, secondIdentity, upstream.requests[1].Header.Get(grokConversationIDHeader))
require.Equal(t, thirdIdentity, upstream.requests[2].Header.Get(grokConversationIDHeader))
require.False(t, gjson.GetBytes(upstream.lastBody, "type").Exists())
require.False(t, gjson.GetBytes(upstream.lastBody, "generate").Exists())
require.False(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_retention").Exists())