fix: sanitize grok codex responses payloads

This commit is contained in:
Heatherm Huang
2026-06-29 21:37:30 +08:00
parent 10e623f674
commit 438510d298
2 changed files with 224 additions and 0 deletions
@@ -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 {
@@ -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")