Merge pull request #6291 from Whisper-stark/fix/grok-codex-responses

fix(grok): sanitize Codex Responses requests
This commit is contained in:
Wesley Liddick
2026-08-28 15:58:39 +08:00
committed by GitHub
4 changed files with 174 additions and 12 deletions
@@ -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 {
@@ -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()
@@ -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:
@@ -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":[