Merge pull request #4489 from fengshao1227/fix/grok-free-cache-tool-injection
fix(grok): 纯客户端函数工具不再注入原生搜索工具
This commit is contained in:
@@ -296,6 +296,13 @@ func appendMissingGrokFreeCacheNativeTools(body []byte) ([]byte, error) {
|
||||
if !hasFunction {
|
||||
return body, nil
|
||||
}
|
||||
// Only complement missing native search tools when the request already contains
|
||||
// at least one search tool (native or function-form). Pure client function tools
|
||||
// (e.g. view_image) must not trigger injection to avoid biasing model tool
|
||||
// selection (#4486).
|
||||
if !present["web_search"] && !present["x_search"] {
|
||||
return body, nil
|
||||
}
|
||||
for _, toolType := range []string{"web_search", "x_search"} {
|
||||
if present[toolType] {
|
||||
continue
|
||||
|
||||
@@ -211,6 +211,7 @@ func TestApplyGrokCacheIdentityAppendsNativeToolsToResponseFunctions(t *testing.
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
// Pure client function tools without search → no native injection (#4486).
|
||||
intentBody := []byte(`{"model":"grok","tools":[{"type":"function","name":"lookup","description":"look up a value","parameters":{"type":"object"}},{"type":"function","name":"save","parameters":{"type":"object"}}]` + tt.toolChoiceJSON + `}`)
|
||||
body, err := applyGrokResponsesCacheIdentity(intentBody, intentBody, "isolated-id", true)
|
||||
require.NoError(t, err)
|
||||
@@ -219,28 +220,35 @@ func TestApplyGrokCacheIdentityAppendsNativeToolsToResponseFunctions(t *testing.
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "isolated-id", gjson.GetBytes(body, "prompt_cache_key").String())
|
||||
tools := gjson.GetBytes(body, "tools").Array()
|
||||
require.Len(t, tools, 4)
|
||||
require.Len(t, tools, 2, "pure client functions should not get native search injected")
|
||||
require.Equal(t, "function", tools[0].Get("type").String())
|
||||
require.Equal(t, "lookup", tools[0].Get("name").String())
|
||||
require.Equal(t, "function", tools[1].Get("type").String())
|
||||
require.Equal(t, "save", tools[1].Get("name").String())
|
||||
require.Equal(t, "web_search", tools[2].Get("type").String())
|
||||
require.Equal(t, "x_search", tools[3].Get("type").String())
|
||||
require.Equal(t, tt.wantChoice, gjson.GetBytes(body, "tool_choice").Exists())
|
||||
if tt.wantChoice {
|
||||
require.Equal(t, "auto", gjson.GetBytes(body, "tool_choice").String())
|
||||
}
|
||||
|
||||
second, err := applyGrokResponsesCacheIdentity(body, intentBody, "isolated-id", true)
|
||||
require.NoError(t, err)
|
||||
second, err = applyGrokFreeMessagesFunctionToolCacheRoute(second, intentBody, account, "isolated-id")
|
||||
require.NoError(t, err)
|
||||
require.JSONEq(t, string(body), string(second), "native tools must not be duplicated")
|
||||
require.Len(t, gjson.GetBytes(second, "tools").Array(), 4)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyGrokCacheIdentityAppendsNativeToolsWhenSearchPresent(t *testing.T) {
|
||||
account := healthyGrokOAuthGatewayTestAccount(901, "access-token")
|
||||
account.Credentials["subscription_tier"] = " FREE "
|
||||
|
||||
// Function tools INCLUDING web_search → convert + complement with x_search.
|
||||
intentBody := []byte(`{"model":"grok","tools":[{"type":"function","name":"lookup","description":"look up a value","parameters":{"type":"object"}},{"type":"function","name":"web_search","description":"search","parameters":{"type":"object"}}]}`)
|
||||
body, err := applyGrokResponsesCacheIdentity(intentBody, intentBody, "isolated-id", true)
|
||||
require.NoError(t, err)
|
||||
body, err = applyGrokFreeMessagesFunctionToolCacheRoute(body, intentBody, account, "isolated-id")
|
||||
require.NoError(t, err)
|
||||
|
||||
tools := gjson.GetBytes(body, "tools").Array()
|
||||
require.Len(t, tools, 3, "lookup(function) + web_search(native) + x_search(native)")
|
||||
require.Equal(t, "function", tools[0].Get("type").String())
|
||||
require.Equal(t, "lookup", tools[0].Get("name").String())
|
||||
require.Equal(t, "web_search", tools[1].Get("type").String())
|
||||
require.Equal(t, "x_search", tools[2].Get("type").String())
|
||||
}
|
||||
|
||||
func TestApplyGrokCacheIdentityRequiresPatchedFunctionTools(t *testing.T) {
|
||||
account := healthyGrokOAuthGatewayTestAccount(902, "access-token")
|
||||
account.Credentials["subscription_tier"] = "free"
|
||||
@@ -272,7 +280,9 @@ func TestApplyGrokCacheIdentityRequiresPatchedFunctionTools(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestGrokFreeMessagesFunctionToolCacheRouteRequiresKnownFreeTier(t *testing.T) {
|
||||
intentBody := []byte(`{"model":"grok","tools":[{"type":"function","name":"lookup"}],"tool_choice":"auto"}`)
|
||||
// Include web_search as function to trigger native tool injection (pure client
|
||||
// functions no longer trigger injection after #4486 fix).
|
||||
intentBody := []byte(`{"model":"grok","tools":[{"type":"function","name":"lookup"},{"type":"function","name":"web_search"}],"tool_choice":"auto"}`)
|
||||
tests := []struct {
|
||||
name string
|
||||
account *Account
|
||||
@@ -382,7 +392,7 @@ func TestGrokFreeMessagesFunctionToolCacheRouteRequiresKnownFreeTier(t *testing.
|
||||
require.Equal(t, "x_search", tools[2].Get("type").String())
|
||||
return
|
||||
}
|
||||
require.Len(t, tools, 1)
|
||||
require.Len(t, tools, 2, "non-free accounts should not get native search injected")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,87 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"github.com/tidwall/gjson"
|
||||
)
|
||||
|
||||
func TestAppendMissingGrokFreeCacheNativeTools_PureClientFunctionNoInject(t *testing.T) {
|
||||
body := []byte(`{
|
||||
"model": "grok-4.5",
|
||||
"tools": [
|
||||
{"type":"function","name":"view_image","description":"View image","parameters":{"type":"object","properties":{"path":{"type":"string"}},"required":["path"]}}
|
||||
],
|
||||
"tool_choice": "auto"
|
||||
}`)
|
||||
|
||||
result, err := appendMissingGrokFreeCacheNativeTools(body)
|
||||
require.NoError(t, err)
|
||||
|
||||
tools := gjson.GetBytes(result, "tools").Array()
|
||||
for _, tool := range tools {
|
||||
toolType := tool.Get("type").String()
|
||||
assert.NotEqual(t, "web_search", toolType, "should not inject web_search for pure client functions")
|
||||
assert.NotEqual(t, "x_search", toolType, "should not inject x_search for pure client functions")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAppendMissingGrokFreeCacheNativeTools_FunctionPlusWebSearchInjects(t *testing.T) {
|
||||
body := []byte(`{
|
||||
"model": "grok-4.5",
|
||||
"tools": [
|
||||
{"type":"function","name":"view_image","description":"View","parameters":{"type":"object","properties":{"path":{"type":"string"}},"required":["path"]}},
|
||||
{"type":"function","name":"web_search","description":"Search","parameters":{"type":"object","properties":{"query":{"type":"string"}},"required":["query"]}}
|
||||
]
|
||||
}`)
|
||||
|
||||
result, err := appendMissingGrokFreeCacheNativeTools(body)
|
||||
require.NoError(t, err)
|
||||
|
||||
tools := gjson.GetBytes(result, "tools").Array()
|
||||
types := make(map[string]bool)
|
||||
for _, tool := range tools {
|
||||
types[tool.Get("type").String()] = true
|
||||
}
|
||||
assert.True(t, types["web_search"], "web_search should be present (converted from function)")
|
||||
assert.True(t, types["x_search"], "x_search should be injected when web_search is present alongside client functions")
|
||||
}
|
||||
|
||||
func TestAppendMissingGrokFreeCacheNativeTools_NativeSearchAlreadyPresent(t *testing.T) {
|
||||
body := []byte(`{
|
||||
"model": "grok-4.5",
|
||||
"tools": [
|
||||
{"type":"function","name":"view_image","description":"View","parameters":{"type":"object","properties":{"path":{"type":"string"}},"required":["path"]}},
|
||||
{"type":"web_search"}
|
||||
]
|
||||
}`)
|
||||
|
||||
result, err := appendMissingGrokFreeCacheNativeTools(body)
|
||||
require.NoError(t, err)
|
||||
|
||||
tools := gjson.GetBytes(result, "tools").Array()
|
||||
types := make(map[string]bool)
|
||||
for _, tool := range tools {
|
||||
types[tool.Get("type").String()] = true
|
||||
}
|
||||
assert.True(t, types["web_search"])
|
||||
assert.True(t, types["x_search"], "x_search should be injected when web_search is already present")
|
||||
}
|
||||
|
||||
func TestAppendMissingGrokFreeCacheNativeTools_MultipleFunctionsNoSearch(t *testing.T) {
|
||||
body := []byte(`{
|
||||
"model": "grok-4.5",
|
||||
"tools": [
|
||||
{"type":"function","name":"view_image","description":"View","parameters":{"type":"object","properties":{"path":{"type":"string"}},"required":["path"]}},
|
||||
{"type":"function","name":"read_file","description":"Read","parameters":{"type":"object","properties":{"path":{"type":"string"}},"required":["path"]}}
|
||||
]
|
||||
}`)
|
||||
|
||||
result, err := appendMissingGrokFreeCacheNativeTools(body)
|
||||
require.NoError(t, err)
|
||||
|
||||
tools := gjson.GetBytes(result, "tools").Array()
|
||||
require.Len(t, tools, 2, "no tools should be injected for pure client functions")
|
||||
}
|
||||
@@ -1604,7 +1604,7 @@ func TestForwardAsAnthropicForGrokFunctionToolUsesCacheCapableMixedRoute(t *test
|
||||
body := []byte(`{
|
||||
"model":"grok","max_tokens":32,"stream":false,
|
||||
"messages":[{"role":"user","content":"look up alpha"}],
|
||||
"tools":[{"name":"lookup","description":"look up a key","input_schema":{"type":"object","properties":{"key":{"type":"string"}},"required":["key"]}}],
|
||||
"tools":[{"name":"lookup","description":"look up a key","input_schema":{"type":"object","properties":{"key":{"type":"string"}},"required":["key"]}},{"name":"web_search","description":"search the web","input_schema":{"type":"object","properties":{"query":{"type":"string"}},"required":["query"]}}],
|
||||
"tool_choice":{"type":"auto"}
|
||||
}`)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(body))
|
||||
|
||||
Reference in New Issue
Block a user