mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 13:18:18 +08:00
Merge pull request #6291 from Whisper-stark/fix/grok-codex-responses
fix(grok): sanitize Codex Responses requests
This commit is contained in:
@@ -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":[
|
||||
|
||||
Reference in New Issue
Block a user