From 0478fd36683dbf86e30dbfef0f618a012f7c1daf Mon Sep 17 00:00:00 2001 From: superman2003 <2112076433zcr@gmail.com> Date: Sat, 11 Jul 2026 19:34:01 +0800 Subject: [PATCH] feat(grok): enable free OAuth prompt caching --- .../internal/service/account_test_service.go | 2 +- .../service/account_test_service_grok_test.go | 1 + backend/internal/service/grok_media.go | 2 +- .../internal/service/grok_quota_service.go | 2 +- .../service/grok_quota_service_test.go | 1 + .../service/openai_gateway_cc_pipeline.go | 5 + .../openai_gateway_chat_completions_raw.go | 14 +- .../service/openai_gateway_forward.go | 1 - .../internal/service/openai_gateway_grok.go | 28 +- .../service/openai_gateway_grok_cache.go | 149 ++++++++++ .../service/openai_gateway_grok_cache_test.go | 273 ++++++++++++++++++ .../service/openai_gateway_grok_test.go | 228 ++++++++++++++- .../service/openai_gateway_messages.go | 11 +- .../openai_gateway_messages_chat_fallback.go | 2 +- .../openai_gateway_responses_chat_fallback.go | 2 +- .../service/openai_gateway_scheduling.go | 43 ++- .../service/openai_ws_forwarder_ingress.go | 12 + .../internal/service/openai_ws_http_bridge.go | 41 ++- .../service/openai_ws_http_bridge_test.go | 89 +++++- 19 files changed, 846 insertions(+), 60 deletions(-) create mode 100644 backend/internal/service/openai_gateway_grok_cache.go create mode 100644 backend/internal/service/openai_gateway_grok_cache_test.go diff --git a/backend/internal/service/account_test_service.go b/backend/internal/service/account_test_service.go index b9dae84057..15e8097b79 100644 --- a/backend/internal/service/account_test_service.go +++ b/backend/internal/service/account_test_service.go @@ -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 { diff --git a/backend/internal/service/account_test_service_grok_test.go b/backend/internal/service/account_test_service_grok_test.go index 356fc62fb7..0223070774 100644 --- a/backend/internal/service/account_test_service_grok_test.go +++ b/backend/internal/service/account_test_service_grok_test.go @@ -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"`) diff --git a/backend/internal/service/grok_media.go b/backend/internal/service/grok_media.go index 154e3003ef..0f0e3366ee 100644 --- a/backend/internal/service/grok_media.go +++ b/backend/internal/service/grok_media.go @@ -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 == "" { diff --git a/backend/internal/service/grok_quota_service.go b/backend/internal/service/grok_quota_service.go index 06d90a609e..08eca2e7e2 100644 --- a/backend/internal/service/grok_quota_service.go +++ b/backend/internal/service/grok_quota_service.go @@ -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 { diff --git a/backend/internal/service/grok_quota_service_test.go b/backend/internal/service/grok_quota_service_test.go index fe1a00aa89..5400507827 100644 --- a/backend/internal/service/grok_quota_service_test.go +++ b/backend/internal/service/grok_quota_service_test.go @@ -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`) diff --git a/backend/internal/service/openai_gateway_cc_pipeline.go b/backend/internal/service/openai_gateway_cc_pipeline.go index 816f5a26e4..ea0148fb2a 100644 --- a/backend/internal/service/openai_gateway_cc_pipeline.go +++ b/backend/internal/service/openai_gateway_cc_pipeline.go @@ -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 { diff --git a/backend/internal/service/openai_gateway_chat_completions_raw.go b/backend/internal/service/openai_gateway_chat_completions_raw.go index 6def67c7ec..56130670df 100644 --- a/backend/internal/service/openai_gateway_chat_completions_raw.go +++ b/backend/internal/service/openai_gateway_chat_completions_raw.go @@ -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 } diff --git a/backend/internal/service/openai_gateway_forward.go b/backend/internal/service/openai_gateway_forward.go index 53cc06989e..1c13738020 100644 --- a/backend/internal/service/openai_gateway_forward.go +++ b/backend/internal/service/openai_gateway_forward.go @@ -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) } diff --git a/backend/internal/service/openai_gateway_grok.go b/backend/internal/service/openai_gateway_grok.go index ef823fa2fd..0484d373c1 100644 --- a/backend/internal/service/openai_gateway_grok.go +++ b/backend/internal/service/openai_gateway_grok.go @@ -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 diff --git a/backend/internal/service/openai_gateway_grok_cache.go b/backend/internal/service/openai_gateway_grok_cache.go new file mode 100644 index 0000000000..20934b94c3 --- /dev/null +++ b/backend/internal/service/openai_gateway_grok_cache.go @@ -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") +} diff --git a/backend/internal/service/openai_gateway_grok_cache_test.go b/backend/internal/service/openai_gateway_grok_cache_test.go new file mode 100644 index 0000000000..556f19304f --- /dev/null +++ b/backend/internal/service/openai_gateway_grok_cache_test.go @@ -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) +} diff --git a/backend/internal/service/openai_gateway_grok_test.go b/backend/internal/service/openai_gateway_grok_test.go index 8dbd9ddad0..aa01e43461 100644 --- a/backend/internal/service/openai_gateway_grok_test.go +++ b/backend/internal/service/openai_gateway_grok_test.go @@ -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 diff --git a/backend/internal/service/openai_gateway_messages.go b/backend/internal/service/openai_gateway_messages.go index 219b5e4be4..3615fe4b17 100644 --- a/backend/internal/service/openai_gateway_messages.go +++ b/backend/internal/service/openai_gateway_messages.go @@ -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") != "" { diff --git a/backend/internal/service/openai_gateway_messages_chat_fallback.go b/backend/internal/service/openai_gateway_messages_chat_fallback.go index 3861ae9945..a1a7023be4 100644 --- a/backend/internal/service/openai_gateway_messages_chat_fallback.go +++ b/backend/internal/service/openai_gateway_messages_chat_fallback.go @@ -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 } diff --git a/backend/internal/service/openai_gateway_responses_chat_fallback.go b/backend/internal/service/openai_gateway_responses_chat_fallback.go index 9d3391e4e6..f494429226 100644 --- a/backend/internal/service/openai_gateway_responses_chat_fallback.go +++ b/backend/internal/service/openai_gateway_responses_chat_fallback.go @@ -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 } diff --git a/backend/internal/service/openai_gateway_scheduling.go b/backend/internal/service/openai_gateway_scheduling.go index 6ac0de47b4..f318adc604 100644 --- a/backend/internal/service/openai_gateway_scheduling.go +++ b/backend/internal/service/openai_gateway_scheduling.go @@ -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) } diff --git a/backend/internal/service/openai_ws_forwarder_ingress.go b/backend/internal/service/openai_ws_forwarder_ingress.go index 45e004b6d0..a2af4b760f 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress.go +++ b/backend/internal/service/openai_ws_forwarder_ingress.go @@ -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, ) diff --git a/backend/internal/service/openai_ws_http_bridge.go b/backend/internal/service/openai_ws_http_bridge.go index dce0c7b6db..94bb09e358 100644 --- a/backend/internal/service/openai_ws_http_bridge.go +++ b/backend/internal/service/openai_ws_http_bridge.go @@ -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 +} diff --git a/backend/internal/service/openai_ws_http_bridge_test.go b/backend/internal/service/openai_ws_http_bridge_test.go index 4d1e4a0374..da2ae77917 100644 --- a/backend/internal/service/openai_ws_http_bridge_test.go +++ b/backend/internal/service/openai_ws_http_bridge_test.go @@ -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())