diff --git a/backend/internal/service/openai_gateway_grok.go b/backend/internal/service/openai_gateway_grok.go index a99f37ac2a..efcbaa8d44 100644 --- a/backend/internal/service/openai_gateway_grok.go +++ b/backend/internal/service/openai_gateway_grok.go @@ -74,8 +74,14 @@ func (s *OpenAIGatewayService) forwardGrokResponses( } } // Derive the identity from the request xAI will actually see. This makes - // Codex Responses Lite additional_tools part of the stable tool prefix. - cacheIdentity := resolveGrokCacheIdentity(c, patchedBody, "", upstreamModel) + // Codex Responses Lite additional_tools part of the stable tool prefix. If + // the client supplied a Claude Code session only through metadata.user_id, + // keep using that identity even though metadata is stripped before xAI. + cacheIdentityBody := patchedBody + if extractClaudeCodeSessionIDFromPayload(body) != "" { + cacheIdentityBody = body + } + cacheIdentity := resolveGrokCacheIdentity(c, cacheIdentityBody, "", upstreamModel) mixedCacheIntentBody := append([]byte(nil), patchedBody...) patchedBody, err = applyGrokResponsesCacheIdentity(patchedBody, body, cacheIdentity, account.IsGrokOAuth()) if err != nil { @@ -547,7 +553,7 @@ func patchGrokResponsesBodyBase(body []byte, upstreamModel string) ([]byte, erro if err != nil { return nil, err } - for _, unsupportedField := range []string{"prompt_cache_retention", "safety_identifier"} { + for _, unsupportedField := range []string{"prompt_cache_retention", "safety_identifier", "metadata"} { if gjson.GetBytes(out, unsupportedField).Exists() { out, err = sjson.DeleteBytes(out, unsupportedField) if err != nil { @@ -1053,6 +1059,8 @@ var grokResponsesSupportedToolTypes = map[string]struct{}{ "x_search": {}, } +const grokSafeFunctionParameters = `{"type":"object","properties":{},"additionalProperties":true}` + func sanitizeGrokResponsesTools(body []byte) ([]byte, error) { tools := gjson.GetBytes(body, "tools") if !tools.Exists() { @@ -1088,6 +1096,19 @@ func sanitizeGrokResponsesTools(body []byte) ([]byte, error) { } raw = encoded toolsChanged = true + } else if toolType == "function" && grokFunctionParametersHaveInvalidUnionRoot(tool.Get("parameters")) { + var err error + raw, err = sjson.SetRawBytes(raw, "parameters", []byte(grokSafeFunctionParameters)) + if err != nil { + return nil, err + } + if strict := tool.Get("strict"); strict.Exists() && strict.Bool() { + raw, err = sjson.SetBytes(raw, "strict", false) + if err != nil { + return nil, err + } + } + toolsChanged = true } filteredTools = append(filteredTools, raw) } @@ -1139,6 +1160,28 @@ func sanitizeGrokResponsesTools(body []byte) ([]byte, error) { return body, nil } +func grokFunctionParametersHaveInvalidUnionRoot(parameters gjson.Result) bool { + if !parameters.Exists() || !parameters.IsObject() { + return false + } + for _, keyword := range []string{"anyOf", "oneOf"} { + branches := parameters.Get(keyword) + if !branches.IsArray() { + continue + } + values := branches.Array() + if len(values) == 0 { + continue + } + for _, branch := range values { + if !strings.EqualFold(strings.TrimSpace(branch.Get("type").String()), "object") { + return true + } + } + } + return false +} + func grokRawToolsContainType(tools []json.RawMessage, want string) bool { for _, tool := range tools { if strings.TrimSpace(gjson.GetBytes(tool, "type").String()) == want { diff --git a/backend/internal/service/openai_gateway_grok_test.go b/backend/internal/service/openai_gateway_grok_test.go index d1c966ac07..d3256eac93 100644 --- a/backend/internal/service/openai_gateway_grok_test.go +++ b/backend/internal/service/openai_gateway_grok_test.go @@ -325,7 +325,7 @@ func TestPatchGrokResponsesBodyDropsNestedUnsupportedFields(t *testing.T) { require.True(t, json.Valid(patched)) require.False(t, strings.Contains(string(patched), "external_web_access")) require.Equal(t, "kept_fn", gjson.GetBytes(patched, "tools.0.name").String()) - require.Equal(t, "9007199254740993", gjson.GetBytes(patched, "metadata.large_id").Raw) + require.False(t, gjson.GetBytes(patched, "metadata").Exists()) } func TestStripAnthropicThinkingSignaturesPreservesLargeIntegers(t *testing.T) { @@ -395,6 +395,62 @@ func TestSanitizeGrokResponsesToolsRemovesDeferredFlagsWithToolSearch(t *testing require.True(t, gjson.GetBytes(patched, `tools.#(name=="apply_patch")`).Exists()) } +func TestSanitizeGrokResponsesToolsSimplifiesInvalidRootUnion(t *testing.T) { + body := []byte(`{"tools":[ + {"type":"function","name":"mcp__codex_app__automation_update","strict":true,"parameters":{"oneOf":[{"type":"object","properties":{"id":{"type":"string"}}},{"type":"null"}]}}, + {"type":"function","name":"object_only","strict":true,"parameters":{"type":"object","anyOf":[{"type":"object","properties":{"a":{"type":"string"}}},{"type":"object","properties":{"b":{"type":"integer"}}}]}} + ]}`) + + patched, err := sanitizeGrokResponsesTools(body) + require.NoError(t, err) + require.True(t, json.Valid(patched)) + + mixed := gjson.GetBytes(patched, `tools.#(name=="mcp__codex_app__automation_update")`) + require.Equal(t, "object", mixed.Get("parameters.type").String()) + require.True(t, mixed.Get("parameters.properties").IsObject()) + require.True(t, mixed.Get("parameters.additionalProperties").Bool()) + require.False(t, mixed.Get("parameters.oneOf").Exists()) + require.Equal(t, gjson.False, mixed.Get("strict").Type) + + objectOnly := gjson.GetBytes(patched, `tools.#(name=="object_only")`) + require.True(t, objectOnly.Get("parameters.anyOf").Exists()) + require.Equal(t, gjson.True, objectOnly.Get("strict").Type) +} + +func TestPatchGrokResponsesBodySimplifiesTypedInvalidRootUnion(t *testing.T) { + body := []byte(`{ + "model":"grok-4.6", + "metadata":{"session_id":"abc"}, + "tools":[{ + "type":"namespace", + "name":"mcp__codex_app", + "tools":[{ + "type":"function", + "name":"automation_update", + "strict":true, + "parameters":{ + "type":"object", + "oneOf":[{"$ref":"#/$defs/update"},{"type":"null"}], + "$defs":{"update":{"type":"object","properties":{"id":{"type":"string"}}}} + } + }] + }] + }`) + + patched, _, err := patchGrokResponsesBodyWithClientTools(body, "grok-4.6") + require.NoError(t, err) + require.True(t, json.Valid(patched)) + require.False(t, gjson.GetBytes(patched, "metadata").Exists()) + + tool := gjson.GetBytes(patched, `tools.#(name=="mcp__codex_app__automation_update")`) + require.Equal(t, "object", tool.Get("parameters.type").String()) + require.True(t, tool.Get("parameters.properties").IsObject()) + require.True(t, tool.Get("parameters.additionalProperties").Bool()) + require.False(t, tool.Get("parameters.oneOf").Exists()) + require.False(t, tool.Get("parameters.$defs").Exists()) + require.Equal(t, gjson.False, tool.Get("strict").Type) +} + func TestSanitizeGrokResponsesToolsKeepsToolChoiceOnlyWithSupportedTools(t *testing.T) { t.Parallel() @@ -1949,7 +2005,7 @@ func TestForwardGrokResponsesAPIKeyUsesXAIResponses(t *testing.T) { recorder := httptest.NewRecorder() c, _ := gin.CreateTestContext(recorder) - body := []byte(`{"model":"grok","input":"hi","stream":true}`) + body := []byte(`{"model":"grok","input":"hi","metadata":{"session_id":"abc"},"stream":true}`) c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body)) c.Request.Header.Set("Content-Type", "application/json") @@ -1984,11 +2040,63 @@ func TestForwardGrokResponsesAPIKeyUsesXAIResponses(t *testing.T) { require.Empty(t, upstream.lastReq.Header.Get("X-Grok-Client-Version")) require.NotEqual(t, defaultGrokUpstreamUserAgent(), upstream.lastReq.Header.Get("User-Agent")) require.Equal(t, "grok-4.6", gjson.GetBytes(upstream.lastBody, "model").String()) + require.False(t, gjson.GetBytes(upstream.lastBody, "metadata").Exists()) require.Equal(t, "resp_grok_api_key", result.ResponseID) require.Equal(t, 2, result.Usage.InputTokens) require.Equal(t, 1, result.Usage.OutputTokens) } +func TestForwardGrokResponsesUsesMetadataSessionForCacheIdentityWithoutForwardingMetadata(t *testing.T) { + gin.SetMode(gin.TestMode) + + firstBody := []byte(`{"model":"grok","input":"first turn","metadata":{"user_id":"{\"session_id\":\"metadata-session\"}"},"stream":false}`) + secondBody := []byte(`{"model":"grok","input":"different second turn","metadata":{"user_id":"{\"session_id\":\"metadata-session\"}"},"stream":false}`) + account := &Account{ + ID: 5401, + Name: "grok-api-key", + Platform: PlatformGrok, + Type: AccountTypeAPIKey, + Concurrency: 2, + Credentials: map[string]any{ + "api_key": "xai-test-key", + "base_url": "https://api.x.ai/v1", + }, + } + upstream := &httpUpstreamRecorder{responses: []*http.Response{ + { + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"id":"resp_first","object":"response","model":"grok-4.6","status":"completed","output":[],"usage":{"input_tokens":2,"output_tokens":1}}`)), + }, + { + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"id":"resp_second","object":"response","model":"grok-4.6","status":"completed","output":[],"usage":{"input_tokens":3,"output_tokens":1}}`)), + }, + }} + svc := &OpenAIGatewayService{httpUpstream: upstream} + newContext := func(body []byte) *gin.Context { + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body)) + c.Set("api_key", &APIKey{ID: 5401}) + return c + } + + _, err := svc.forwardGrokResponses(context.Background(), newContext(firstBody), account, firstBody, "grok", false, time.Now()) + require.NoError(t, err) + _, err = svc.forwardGrokResponses(context.Background(), newContext(secondBody), account, secondBody, "grok", false, time.Now()) + require.NoError(t, err) + 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.False(t, gjson.GetBytes(upstream.bodies[0], "metadata").Exists()) + require.False(t, gjson.GetBytes(upstream.bodies[1], "metadata").Exists()) +} + func TestForwardGrokResponsesRetriesInvalidEncryptedContentOnce(t *testing.T) { gin.SetMode(gin.TestMode) @@ -2056,8 +2164,8 @@ func TestForwardGrokResponsesRetriesInvalidEncryptedContentOnce(t *testing.T) { require.False(t, gjson.GetBytes(upstream.bodies[1], "input.0.encrypted_content").Exists()) require.Equal(t, "keep this summary", gjson.GetBytes(upstream.bodies[1], "input.0.summary.0.text").String()) require.Equal(t, "message", gjson.GetBytes(upstream.bodies[1], "input.1.type").String()) - require.Equal(t, "9007199254740993", gjson.GetBytes(upstream.bodies[0], "metadata.large_id").Raw) - require.Equal(t, "9007199254740993", gjson.GetBytes(upstream.bodies[1], "metadata.large_id").Raw) + require.False(t, gjson.GetBytes(upstream.bodies[0], "metadata").Exists()) + require.False(t, gjson.GetBytes(upstream.bodies[1], "metadata").Exists()) firstIdentity := gjson.GetBytes(upstream.bodies[0], "prompt_cache_key").String() secondIdentity := gjson.GetBytes(upstream.bodies[1], "prompt_cache_key").String() diff --git a/backend/internal/service/openai_responses_tool_schema.go b/backend/internal/service/openai_responses_tool_schema.go index c77640d0e3..a53eefb982 100644 --- a/backend/internal/service/openai_responses_tool_schema.go +++ b/backend/internal/service/openai_responses_tool_schema.go @@ -28,9 +28,9 @@ var errOpenAIResponsesToolSchemaLimit = errors.New("OpenAI Responses tool schema // shouldRepairOpenAIResponsesNullToolSchemaType reports whether the upstream // path requires a concrete object type at a function tool's parameter root. -// This defect is shared by the OpenAI, Anthropic, and CN-compatible paths. +// This defect is shared by the OpenAI, Anthropic, Grok, and CN-compatible paths. func shouldRepairOpenAIResponsesNullToolSchemaType(platform string) bool { - return platform == PlatformOpenAI || platform == PlatformAnthropic || IsCNProvider(platform) + return platform == PlatformOpenAI || platform == PlatformAnthropic || platform == PlatformGrok || IsCNProvider(platform) } // shouldSanitizeOpenAIResponsesToolSchemaPatterns is intentionally narrower: diff --git a/backend/internal/service/openai_responses_tool_schema_test.go b/backend/internal/service/openai_responses_tool_schema_test.go index e899e17098..b2dfba46bc 100644 --- a/backend/internal/service/openai_responses_tool_schema_test.go +++ b/backend/internal/service/openai_responses_tool_schema_test.go @@ -250,7 +250,7 @@ func TestOpenAIResponsesToolSchemaCapabilities_PlatformBoundary(t *testing.T) { {PlatformKimi, true, false}, {PlatformZhipu, true, false}, {PlatformDeepseek, true, false}, - {PlatformGrok, false, false}, + {PlatformGrok, true, false}, {PlatformGemini, false, false}, {PlatformAntigravity, false, false}, {PlatformComposite, false, false}, @@ -270,7 +270,7 @@ func TestSanitizeOpenAIResponsesToolSchemasForPlatform_ReplayBoundary(t *testing // A malformed tool definition may be replayed after account failover. Every // compatible account must repair it, while non-OpenAI providers retain their // supported regex semantics. - for _, platform := range []string{PlatformAnthropic, PlatformKimi, PlatformZhipu, PlatformDeepseek} { + for _, platform := range []string{PlatformAnthropic, PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek} { t.Run(platform, func(t *testing.T) { for attempt := 0; attempt < 2; attempt++ { normalized, changed, err := sanitizeOpenAIResponsesToolSchemasForPlatform(body, platform) @@ -288,12 +288,23 @@ func TestSanitizeOpenAIResponsesToolSchemasForPlatform_ReplayBoundary(t *testing require.Equal(t, "object", gjson.GetBytes(openAI, "tools.0.parameters.type").String()) require.False(t, gjson.GetBytes(openAI, "tools.0.parameters.properties.query.pattern").Exists()) - unsupported, changed, err := sanitizeOpenAIResponsesToolSchemasForPlatform(body, PlatformGrok) + unsupported, changed, err := sanitizeOpenAIResponsesToolSchemasForPlatform(body, PlatformGemini) require.NoError(t, err) require.False(t, changed) require.Equal(t, string(body), string(unsupported)) } +func TestSanitizeOpenAIResponsesToolSchemasForPlatform_GrokObjectOnlyRootUnion(t *testing.T) { + body := []byte(`{"tools":[{"type":"function","name":"codex_app__automation_update","parameters":{"oneOf":[{"type":"object","properties":{"id":{"type":"string"}}},{"type":"object","properties":{}}]}}]}`) + + sanitized, changed, err := sanitizeOpenAIResponsesToolSchemasForPlatform(body, PlatformGrok) + + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, "object", gjson.GetBytes(sanitized, "tools.0.parameters.type").String()) + require.True(t, gjson.GetBytes(sanitized, "tools.0.parameters.oneOf").Exists()) +} + // 索引映射:只有坏条目被改,前后兄弟条目按原下标保持不变。 func TestSanitizeOpenAIResponsesToolParameterTypes_OnlyOffendingIndexRewritten(t *testing.T) { body := []byte(`{"tools":[