From 438510d298a9386864a89942c4d1e84bd5842a4c Mon Sep 17 00:00:00 2001 From: Heatherm Huang Date: Mon, 29 Jun 2026 21:37:30 +0800 Subject: [PATCH] fix: sanitize grok codex responses payloads --- .../internal/service/openai_gateway_grok.go | 159 ++++++++++++++++++ .../service/openai_gateway_grok_test.go | 65 +++++++ 2 files changed, 224 insertions(+) diff --git a/backend/internal/service/openai_gateway_grok.go b/backend/internal/service/openai_gateway_grok.go index eb3fbb66e6..4961b9c589 100644 --- a/backend/internal/service/openai_gateway_grok.go +++ b/backend/internal/service/openai_gateway_grok.go @@ -147,9 +147,168 @@ func patchGrokResponsesBody(body []byte, upstreamModel string) ([]byte, error) { } } } + out, err = sanitizeGrokResponsesUnsupportedFields(out) + if err != nil { + return nil, err + } + out, err = sanitizeGrokResponsesTools(out) + if err != nil { + return nil, err + } return out, nil } +var grokResponsesUnsupportedRecursiveFields = map[string]struct{}{ + "external_web_access": {}, +} + +func sanitizeGrokResponsesUnsupportedFields(body []byte) ([]byte, error) { + if !bytes.Contains(body, []byte(`"external_web_access"`)) { + return body, nil + } + + var payload any + if err := json.Unmarshal(body, &payload); err != nil { + return nil, err + } + if !deleteJSONFields(payload, grokResponsesUnsupportedRecursiveFields) { + return body, nil + } + return json.Marshal(payload) +} + +func deleteJSONFields(value any, fields map[string]struct{}) bool { + switch typed := value.(type) { + case map[string]any: + changed := false + for field := range fields { + if _, ok := typed[field]; ok { + delete(typed, field) + changed = true + } + } + for _, child := range typed { + if deleteJSONFields(child, fields) { + changed = true + } + } + return changed + case []any: + changed := false + for _, child := range typed { + if deleteJSONFields(child, fields) { + changed = true + } + } + return changed + default: + return false + } +} + +var grokResponsesSupportedToolTypes = map[string]struct{}{ + "code_execution": {}, + "code_interpreter": {}, + "collections_search": {}, + "file_search": {}, + "function": {}, + "mcp": {}, + "shell": {}, + "web_search": {}, + "x_search": {}, +} + +func sanitizeGrokResponsesTools(body []byte) ([]byte, error) { + tools := gjson.GetBytes(body, "tools") + if !tools.Exists() || !tools.IsArray() { + return body, nil + } + + rawTools := tools.Array() + filteredTools := make([]json.RawMessage, 0, len(rawTools)) + for _, tool := range rawTools { + toolType := strings.TrimSpace(tool.Get("type").String()) + if _, ok := grokResponsesSupportedToolTypes[toolType]; ok { + filteredTools = append(filteredTools, json.RawMessage(tool.Raw)) + } + } + + var err error + if len(filteredTools) != len(rawTools) { + if len(filteredTools) == 0 { + body, err = sjson.DeleteBytes(body, "tools") + } else { + var encoded []byte + encoded, err = json.Marshal(filteredTools) + if err != nil { + return nil, err + } + body, err = sjson.SetRawBytes(body, "tools", encoded) + } + if err != nil { + return nil, err + } + } + + toolChoice := gjson.GetBytes(body, "tool_choice") + if !toolChoice.Exists() { + return body, nil + } + if shouldDropGrokToolChoice(toolChoice, filteredTools) { + body, err = sjson.DeleteBytes(body, "tool_choice") + if err != nil { + return nil, err + } + } + return body, nil +} + +func shouldDropGrokToolChoice(toolChoice gjson.Result, tools []json.RawMessage) bool { + if len(tools) == 0 { + return true + } + if !toolChoice.IsObject() { + return false + } + choiceType := strings.TrimSpace(toolChoice.Get("type").String()) + if choiceType == "" { + return false + } + if _, ok := grokResponsesSupportedToolTypes[choiceType]; !ok { + return true + } + if choiceType == "function" { + choiceName := strings.TrimSpace(toolChoice.Get("name").String()) + if choiceName == "" { + choiceName = strings.TrimSpace(toolChoice.Get("function.name").String()) + } + if choiceName == "" { + return false + } + for _, tool := range tools { + var item struct { + Type string `json:"type"` + Name string `json:"name"` + Function struct { + Name string `json:"name"` + } `json:"function"` + } + if err := json.Unmarshal(tool, &item); err != nil { + continue + } + name := strings.TrimSpace(item.Name) + if name == "" { + name = strings.TrimSpace(item.Function.Name) + } + if strings.TrimSpace(item.Type) == "function" && name == choiceName { + return false + } + } + return true + } + return false +} + func buildGrokResponsesRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token string) (*http.Request, error) { targetURL, err := xai.BuildResponsesURL(account.GetGrokBaseURL()) if err != nil { diff --git a/backend/internal/service/openai_gateway_grok_test.go b/backend/internal/service/openai_gateway_grok_test.go index 168cddf514..4f6601e843 100644 --- a/backend/internal/service/openai_gateway_grok_test.go +++ b/backend/internal/service/openai_gateway_grok_test.go @@ -39,6 +39,71 @@ func TestPatchGrokResponsesBodySetsMappedModelAndDropsUnsupportedFields(t *testi require.Equal(t, "high", gjson.GetBytes(patched, "reasoning.effort").String()) } +func TestPatchGrokResponsesBodyDropsNestedUnsupportedFields(t *testing.T) { + t.Parallel() + + body := []byte(`{ + "model": "grok", + "input": "hello", + "external_web_access": true, + "tools": [ + {"type": "function", "name": "kept_fn", "external_web_access": true, "parameters": {"type": "object", "properties": {"q": {"type": "string", "external_web_access": true}}}} + ], + "metadata": {"external_web_access": false} + }`) + + patched, err := patchGrokResponsesBody(body, "grok-4.3") + require.NoError(t, err) + 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()) +} + +func TestPatchGrokResponsesBodyDropsUnsupportedNamespaceTools(t *testing.T) { + t.Parallel() + + body := []byte(`{ + "model": "grok", + "input": "hello", + "tools": [ + {"type": "namespace", "namespace": "functions", "tools": [{"type": "function", "name": "inner"}]}, + {"type": "function", "name": "kept_fn", "parameters": {"type": "object"}}, + {"type": "shell", "name": "kept_shell"} + ], + "tool_choice": {"type": "function", "name": "kept_fn"} + }`) + + patched, err := patchGrokResponsesBody(body, "grok-4.3") + require.NoError(t, err) + require.True(t, json.Valid(patched)) + require.Equal(t, "grok-4.3", gjson.GetBytes(patched, "model").String()) + require.Len(t, gjson.GetBytes(patched, "tools").Array(), 2) + require.False(t, gjson.GetBytes(patched, `tools.#(type=="namespace")`).Exists()) + require.True(t, gjson.GetBytes(patched, `tools.#(type=="function")`).Exists()) + require.True(t, gjson.GetBytes(patched, `tools.#(type=="shell")`).Exists()) + require.Equal(t, "kept_fn", gjson.GetBytes(patched, "tool_choice.name").String()) +} + +func TestPatchGrokResponsesBodyDropsToolChoiceWhenNoSupportedToolsRemain(t *testing.T) { + t.Parallel() + + body := []byte(`{ + "model": "grok", + "input": "hello", + "tools": [ + {"type": "namespace", "namespace": "functions"}, + {"type": "image_generation", "model": "gpt-image-2"} + ], + "tool_choice": {"type": "namespace", "namespace": "functions"} + }`) + + patched, err := patchGrokResponsesBody(body, "grok-4.3") + require.NoError(t, err) + require.True(t, json.Valid(patched)) + require.False(t, gjson.GetBytes(patched, "tools").Exists()) + require.False(t, gjson.GetBytes(patched, "tool_choice").Exists()) +} + func TestBuildGrokResponsesRequestUsesAccountBaseURLAndBearerToken(t *testing.T) { t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true")