mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 15:28:03 +08:00
fix(grok): simplify invalid tool union roots
This commit is contained in:
@@ -1057,6 +1057,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() {
|
||||
@@ -1092,6 +1094,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)
|
||||
}
|
||||
@@ -1143,6 +1158,29 @@ func sanitizeGrokResponsesTools(body []byte) ([]byte, error) {
|
||||
return body, nil
|
||||
}
|
||||
|
||||
func grokFunctionParametersHaveInvalidUnionRoot(parameters gjson.Result) bool {
|
||||
if !parameters.Exists() || !parameters.IsObject() ||
|
||||
strings.EqualFold(strings.TrimSpace(parameters.Get("type").String()), "object") {
|
||||
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 {
|
||||
|
||||
@@ -395,6 +395,28 @@ 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":{"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 TestSanitizeGrokResponsesToolsKeepsToolChoiceOnlyWithSupportedTools(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user