From d3a1835ed76fde8860b9c149613ad8975016615c Mon Sep 17 00:00:00 2001 From: Tassoi Date: Fri, 10 Jul 2026 06:37:06 +0000 Subject: [PATCH] fix(image): strip Codex image_gen namespace declarations --- .../service/image_generation_intent.go | 58 ++++---- .../service/image_generation_intent_test.go | 55 ++++++++ .../service/openai_codex_transform.go | 128 +++++++++++++----- .../service/openai_codex_transform_test.go | 104 ++++++++++++++ .../service/openai_gateway_forward.go | 23 +++- .../openai_image_generation_controls_test.go | 59 ++++++++ .../service/openai_ws_forwarder_ingress.go | 2 +- .../openai_ws_forwarder_ingress_test.go | 87 +++++++++--- .../service/openai_ws_forwarder_v2.go | 20 +-- 9 files changed, 437 insertions(+), 99 deletions(-) diff --git a/backend/internal/service/image_generation_intent.go b/backend/internal/service/image_generation_intent.go index 12590f3f63..5a063a7d71 100644 --- a/backend/internal/service/image_generation_intent.go +++ b/backend/internal/service/image_generation_intent.go @@ -93,7 +93,7 @@ func openAIJSONToolsContainImageGeneration(tools gjson.Result) bool { } found := false tools.ForEach(func(_, item gjson.Result) bool { - if openAIJSONString(item.Get("type")) == "image_generation" { + if isOpenAIImageGenerationType(openAIJSONString(item.Get("type"))) { found = true return false } @@ -106,12 +106,20 @@ func openAIJSONToolsContainImageGeneration(tools gjson.Result) bool { return found } +func isOpenAIImageGenerationType(value string) bool { + return strings.TrimSpace(value) == "image_generation" +} + +func isOpenAIImageGenNamespaceName(value string) bool { + return strings.TrimSpace(value) == "image_gen" +} + // isImageGenNamespaceTool detects the Codex namespace-style image generation // tool declaration: { "type": "namespace", "name": "image_gen", ... }. // Codex /image uses this instead of the flat { "type": "image_generation" }. func isImageGenNamespaceTool(tool gjson.Result) bool { return openAIJSONString(tool.Get("type")) == "namespace" && - openAIJSONString(tool.Get("name")) == "image_gen" + isOpenAIImageGenNamespaceName(openAIJSONString(tool.Get("name"))) } // openAIJSONInputContainsImageGenTool scans Responses input items for @@ -127,27 +135,19 @@ func openAIJSONInputContainsImageGenTool(input gjson.Result) bool { if openAIJSONString(item.Get("type")) != "additional_tools" { return true } - tools := item.Get("tools") - if !tools.IsArray() { - return true - } - tools.ForEach(func(_, tool gjson.Result) bool { - if isImageGenNamespaceTool(tool) { - found = true - return false - } - return true - }) + found = openAIJSONToolsContainImageGeneration(item.Get("tools")) return !found }) return found } -func openAIRequestBodyHasImageGenerationTool(body []byte) bool { +func openAIRequestBodyHasImageGenerationDeclaration(body []byte) bool { if len(body) == 0 || !gjson.ValidBytes(body) { return false } - return openAIJSONToolsContainImageGeneration(gjson.GetBytes(body, "tools")) + return openAIJSONToolsContainImageGeneration(gjson.GetBytes(body, "tools")) || + openAIJSONInputContainsImageGenTool(gjson.GetBytes(body, "input")) || + openAIJSONToolChoiceSelectsImageGeneration(gjson.GetBytes(body, "tool_choice")) } func openAIRequestBodyImageGenerationToolNeedsNormalization(body []byte) bool { @@ -178,18 +178,24 @@ func openAIJSONToolChoiceSelectsImageGeneration(choice gjson.Result) bool { return false } if choice.Type == gjson.String { - return strings.TrimSpace(choice.String()) == "image_generation" + return isOpenAIImageGenerationType(choice.String()) } if !choice.IsObject() { return false } - if strings.TrimSpace(choice.Get("type").String()) == "image_generation" { + choiceType := openAIJSONString(choice.Get("type")) + if isOpenAIImageGenerationType(choiceType) { return true } - if strings.TrimSpace(choice.Get("tool.type").String()) == "image_generation" { + if choiceType == "namespace" && + (isOpenAIImageGenNamespaceName(openAIJSONString(choice.Get("name"))) || + isOpenAIImageGenNamespaceName(openAIJSONString(choice.Get("namespace")))) { return true } - if strings.TrimSpace(choice.Get("function.name").String()) == "image_generation" { + if tool := choice.Get("tool"); tool.IsObject() && openAIJSONToolChoiceSelectsImageGeneration(tool) { + return true + } + if isOpenAIImageGenerationType(openAIJSONString(choice.Get("function.name"))) { return true } return false @@ -198,15 +204,21 @@ func openAIJSONToolChoiceSelectsImageGeneration(choice gjson.Result) bool { func openAIAnyToolChoiceSelectsImageGeneration(choice any) bool { switch v := choice.(type) { case string: - return strings.TrimSpace(v) == "image_generation" + return isOpenAIImageGenerationType(v) case map[string]any: - if strings.TrimSpace(firstNonEmptyString(v["type"])) == "image_generation" { + choiceType := strings.TrimSpace(firstNonEmptyString(v["type"])) + if isOpenAIImageGenerationType(choiceType) { return true } - if tool, ok := v["tool"].(map[string]any); ok && strings.TrimSpace(firstNonEmptyString(tool["type"])) == "image_generation" { + if choiceType == "namespace" && + (isOpenAIImageGenNamespaceName(firstNonEmptyString(v["name"])) || + isOpenAIImageGenNamespaceName(firstNonEmptyString(v["namespace"]))) { return true } - if fn, ok := v["function"].(map[string]any); ok && strings.TrimSpace(firstNonEmptyString(fn["name"])) == "image_generation" { + if tool, ok := v["tool"].(map[string]any); ok && openAIAnyToolChoiceSelectsImageGeneration(tool) { + return true + } + if fn, ok := v["function"].(map[string]any); ok && isOpenAIImageGenerationType(firstNonEmptyString(fn["name"])) { return true } } diff --git a/backend/internal/service/image_generation_intent_test.go b/backend/internal/service/image_generation_intent_test.go index 1a32318cce..7f3ac8a611 100644 --- a/backend/internal/service/image_generation_intent_test.go +++ b/backend/internal/service/image_generation_intent_test.go @@ -41,6 +41,20 @@ func TestIsImageGenerationIntent(t *testing.T) { body: []byte(`{"model":"gpt-5.4","tool_choice":{"type":"image_generation"}}`), want: true, }, + { + name: "namespace image_gen tool choice", + endpoint: "/v1/responses", + model: "gpt-5.5", + body: []byte(`{"model":"gpt-5.5","tool_choice":{"type":"namespace","name":"image_gen"}}`), + want: true, + }, + { + name: "custom imagegen function tool choice is not image intent", + endpoint: "/v1/responses", + model: "gpt-5.5", + body: []byte(`{"model":"gpt-5.5","tool_choice":{"function":{"name":"imagegen"}}}`), + want: false, + }, { name: "required tool choice alone is text", endpoint: "/v1/responses", @@ -62,6 +76,13 @@ func TestIsImageGenerationIntent(t *testing.T) { body: []byte(`{"model":"gpt-5.5","tools":[{"type":"namespace","name":"image_gen","tools":[{"type":"function","name":"imagegen"}]}]}`), want: true, }, + { + name: "custom namespace with nested imagegen function is not image intent", + endpoint: "/v1/responses", + model: "gpt-5.5", + body: []byte(`{"model":"gpt-5.5","tools":[{"type":"namespace","name":"media_tools","tools":[{"type":"function","name":"imagegen"}]}]}`), + want: false, + }, { name: "namespace image_gen in input additional_tools (Responses Lite)", endpoint: "/v1/responses", @@ -118,6 +139,40 @@ func TestIsImageGenerationIntentMap_NamespaceImageGen(t *testing.T) { }, want: true, }, + { + name: "custom namespace with nested imagegen function is not image intent", + reqBody: map[string]any{ + "model": "gpt-5.5", + "tools": []any{ + map[string]any{ + "type": "namespace", + "name": "media_tools", + "tools": []any{ + map[string]any{"type": "function", "name": "imagegen"}, + }, + }, + }, + }, + want: false, + }, + { + name: "namespace image_gen tool choice", + reqBody: map[string]any{ + "model": "gpt-5.5", + "tool_choice": map[string]any{"type": "namespace", "name": "image_gen"}, + }, + want: true, + }, + { + name: "custom imagegen function tool choice is not image intent", + reqBody: map[string]any{ + "model": "gpt-5.5", + "tool_choice": map[string]any{ + "function": map[string]any{"name": "imagegen"}, + }, + }, + want: false, + }, { name: "non-image namespace not flagged", reqBody: map[string]any{ diff --git a/backend/internal/service/openai_codex_transform.go b/backend/internal/service/openai_codex_transform.go index 36426e7377..99355628f2 100644 --- a/backend/internal/service/openai_codex_transform.go +++ b/backend/internal/service/openai_codex_transform.go @@ -596,7 +596,7 @@ func hasOpenAIImageGenerationTool(reqBody map[string]any) bool { if toolsContainImageGeneration(reqBody["tools"]) { return true } - return inputContainsImageGenNamespace(reqBody["input"]) + return inputContainsImageGenerationTool(reqBody["input"]) } func toolsContainImageGeneration(rawTools any) bool { @@ -612,22 +612,24 @@ func toolsContainImageGeneration(rawTools any) bool { if !ok { continue } - if strings.TrimSpace(firstNonEmptyString(toolMap["type"])) == "image_generation" { - return true - } - if isImageGenNamespaceToolMap(toolMap) { + if isOpenAIImageGenerationToolMap(toolMap) { return true } } return false } -func isImageGenNamespaceToolMap(tool map[string]any) bool { - return strings.TrimSpace(firstNonEmptyString(tool["type"])) == "namespace" && - strings.TrimSpace(firstNonEmptyString(tool["name"])) == "image_gen" +func isOpenAIImageGenerationToolMap(tool map[string]any) bool { + return isOpenAIImageGenerationType(firstNonEmptyString(tool["type"])) || + isImageGenNamespaceToolMap(tool) } -func inputContainsImageGenNamespace(rawInput any) bool { +func isImageGenNamespaceToolMap(tool map[string]any) bool { + return strings.TrimSpace(firstNonEmptyString(tool["type"])) == "namespace" && + isOpenAIImageGenNamespaceName(firstNonEmptyString(tool["name"])) +} + +func inputContainsImageGenerationTool(rawInput any) bool { input, ok := rawInput.([]any) if !ok { return false @@ -647,54 +649,110 @@ func inputContainsImageGenNamespace(rawInput any) bool { return false } +// stripOpenAIImageGenerationTools keeps account-level strip policy symmetric +// across standard Responses tools, Responses Lite additional_tools, and tool_choice. func stripOpenAIImageGenerationTools(reqBody map[string]any) bool { - rawTools, ok := reqBody["tools"] + if reqBody == nil { + return false + } + modified := stripOpenAIImageGenerationToolList(reqBody, "tools") + if stripOpenAIImageGenerationToolsFromInput(reqBody) { + modified = true + } + if openAIAnyToolChoiceSelectsImageGeneration(reqBody["tool_choice"]) { + delete(reqBody, "tool_choice") + modified = true + } + return modified +} + +func stripOpenAIImageGenerationToolList(container map[string]any, key string) bool { + rawTools, ok := container[key] if !ok || rawTools == nil { - if openAIAnyToolChoiceSelectsImageGeneration(reqBody["tool_choice"]) { - delete(reqBody, "tool_choice") - return true - } return false } tools, ok := rawTools.([]any) if !ok { - if openAIAnyToolChoiceSelectsImageGeneration(reqBody["tool_choice"]) { - delete(reqBody, "tool_choice") - return true - } return false } filtered := make([]any, 0, len(tools)) removed := false for _, rawTool := range tools { - if toolMap, ok := rawTool.(map[string]any); ok && - strings.TrimSpace(firstNonEmptyString(toolMap["type"])) == "image_generation" { + if toolMap, ok := rawTool.(map[string]any); ok && isOpenAIImageGenerationToolMap(toolMap) { removed = true continue } filtered = append(filtered, rawTool) } - if !removed && !openAIAnyToolChoiceSelectsImageGeneration(reqBody["tool_choice"]) { + if !removed { return false } - if removed { - if len(filtered) == 0 { - delete(reqBody, "tools") - } else { - reqBody["tools"] = filtered - } - } - if openAIAnyToolChoiceSelectsImageGeneration(reqBody["tool_choice"]) { - delete(reqBody, "tool_choice") + if len(filtered) == 0 { + delete(container, key) + } else { + container[key] = filtered } return true } -// stripCodexSparkImageGenerationTools removes image_generation tool entries from -// reqBody["tools"]. gpt-5.3-codex-spark rejects that tool upstream with HTTP 400 -// (invalid_request_error, param=tools), and Codex CLI advertises it by default, so -// it must be dropped for spark. When the tools list becomes empty the key is removed. -// Returns true when the body was modified. +func stripOpenAIImageGenerationToolsFromInput(reqBody map[string]any) bool { + input, ok := reqBody["input"].([]any) + if !ok { + return false + } + + filteredInput := make([]any, 0, len(input)) + modified := false + for _, rawItem := range input { + item, ok := rawItem.(map[string]any) + if !ok || strings.TrimSpace(firstNonEmptyString(item["type"])) != "additional_tools" { + filteredInput = append(filteredInput, rawItem) + continue + } + if !stripOpenAIImageGenerationToolList(item, "tools") { + filteredInput = append(filteredInput, rawItem) + continue + } + modified = true + if _, hasTools := item["tools"]; hasTools { + filteredInput = append(filteredInput, rawItem) + } + // An empty additional_tools carrier is not useful upstream; drop the item + // after its only declared capability has been removed. + } + if modified { + reqBody["input"] = filteredInput + } + return modified +} + +// stripOpenAIImageGenerationToolsFromRawPayload is the shared adapter for paths +// that forward raw HTTP or WebSocket payloads without the normal request map. +func stripOpenAIImageGenerationToolsFromRawPayload(payload []byte) ([]byte, bool, error) { + if !openAIRequestBodyHasImageGenerationDeclaration(payload) { + if json.Valid(payload) { + return payload, false, nil + } + var invalidPayload map[string]any + return payload, false, json.Unmarshal(payload, &invalidPayload) + } + payloadMap := make(map[string]any) + if err := json.Unmarshal(payload, &payloadMap); err != nil { + return payload, false, err + } + if !stripOpenAIImageGenerationTools(payloadMap) { + return payload, false, nil + } + rebuilt, err := json.Marshal(payloadMap) + if err != nil { + return payload, false, err + } + return rebuilt, true, nil +} + +// stripCodexSparkImageGenerationTools removes image tool declarations and choices. +// gpt-5.3-codex-spark rejects those capabilities upstream, while Codex clients may +// advertise them by default. func stripCodexSparkImageGenerationTools(reqBody map[string]any) bool { return stripOpenAIImageGenerationTools(reqBody) } diff --git a/backend/internal/service/openai_codex_transform_test.go b/backend/internal/service/openai_codex_transform_test.go index bb27c81352..b226655eeb 100644 --- a/backend/internal/service/openai_codex_transform_test.go +++ b/backend/internal/service/openai_codex_transform_test.go @@ -796,6 +796,110 @@ func TestApplyCodexOAuthTransform_StripsImageGenerationToolForSparkAlias(t *test require.False(t, hasTools) } +func TestStripOpenAIImageGenerationTools_StripsNamespaceFormats(t *testing.T) { + imageNamespace := func() map[string]any { + return map[string]any{ + "type": "namespace", + "name": "image_gen", + "tools": []any{ + map[string]any{"type": "function", "name": "imagegen"}, + }, + } + } + codeNamespace := func() map[string]any { + return map[string]any{ + "type": "namespace", + "name": "code_tools", + "tools": []any{ + map[string]any{"type": "function", "name": "run"}, + }, + } + } + + reqBody := map[string]any{ + "model": "gpt-5.5", + "tools": []any{ + map[string]any{"type": "function", "name": "shell"}, + imageNamespace(), + codeNamespace(), + }, + "input": []any{ + map[string]any{"type": "message", "role": "user", "content": "hello"}, + map[string]any{ + "type": "additional_tools", + "tools": []any{imageNamespace(), codeNamespace()}, + }, + map[string]any{ + "type": "additional_tools", + "tools": []any{imageNamespace()}, + }, + }, + "tool_choice": map[string]any{"type": "namespace", "name": "image_gen"}, + } + + require.True(t, stripOpenAIImageGenerationTools(reqBody)) + require.False(t, hasOpenAIImageGenerationTool(reqBody)) + require.NotContains(t, reqBody, "tool_choice") + + tools, ok := reqBody["tools"].([]any) + require.True(t, ok) + require.Len(t, tools, 2) + firstTool, ok := tools[0].(map[string]any) + require.True(t, ok) + secondTool, ok := tools[1].(map[string]any) + require.True(t, ok) + require.Equal(t, "shell", firstTool["name"]) + require.Equal(t, "code_tools", secondTool["name"]) + + input, ok := reqBody["input"].([]any) + require.True(t, ok) + require.Len(t, input, 2) + message, ok := input[0].(map[string]any) + require.True(t, ok) + require.Equal(t, "message", message["type"]) + additionalToolsItem, ok := input[1].(map[string]any) + require.True(t, ok) + additionalTools, ok := additionalToolsItem["tools"].([]any) + require.True(t, ok) + require.Len(t, additionalTools, 1) + additionalTool, ok := additionalTools[0].(map[string]any) + require.True(t, ok) + require.Equal(t, "code_tools", additionalTool["name"]) + require.False(t, stripOpenAIImageGenerationTools(reqBody), "stripping should be idempotent") +} + +func TestStripOpenAIImageGenerationTools_KeepsNonImageNamespaces(t *testing.T) { + reqBody := map[string]any{ + "tools": []any{ + map[string]any{"type": "namespace", "name": "code_tools"}, + }, + "input": []any{ + map[string]any{ + "type": "additional_tools", + "tools": []any{ + map[string]any{"type": "namespace", "name": "browser_tools"}, + }, + }, + }, + "tool_choice": "auto", + } + + require.False(t, stripOpenAIImageGenerationTools(reqBody)) + require.Equal(t, "auto", reqBody["tool_choice"]) + require.False(t, hasOpenAIImageGenerationTool(reqBody)) +} + +func TestStripOpenAIImageGenerationTools_KeepsCustomImagegenFunctionChoice(t *testing.T) { + reqBody := map[string]any{ + "tool_choice": map[string]any{ + "function": map[string]any{"name": "imagegen"}, + }, + } + + require.False(t, stripOpenAIImageGenerationTools(reqBody)) + require.Contains(t, reqBody, "tool_choice") +} + // Non-spark Codex models support image_generation; the tool must be preserved. func TestApplyCodexOAuthTransform_KeepsImageGenerationToolForNonSpark(t *testing.T) { reqBody := map[string]any{ diff --git a/backend/internal/service/openai_gateway_forward.go b/backend/internal/service/openai_gateway_forward.go index 974f8eefaa..ee18302056 100644 --- a/backend/internal/service/openai_gateway_forward.go +++ b/backend/internal/service/openai_gateway_forward.go @@ -61,6 +61,10 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco setOpenAICompatMessagesBridgeContext(c, compatMessagesBridge) isCodexCLI := openai.IsCodexOfficialClientByHeaders(c.GetHeader("User-Agent"), c.GetHeader("originator")) || (s.cfg != nil && s.cfg.Gateway.ForceCodexCLI) + codexImageGenerationExplicitToolPolicy := codexImageGenerationExplicitToolPolicyAllow + if isCodexCLI { + codexImageGenerationExplicitToolPolicy = account.CodexImageGenerationExplicitToolPolicy() + } wsDecision := s.getOpenAIWSProtocolResolver().Resolve(account) clientTransport := GetOpenAIClientTransport(c) // 仅允许 WS 入站请求走 WS 上游,避免出现 HTTP -> WS 协议混用。 @@ -95,6 +99,17 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco } passthroughEnabled := account.IsOpenAIPassthroughEnabled() if passthroughEnabled { + if isCodexCLI && codexImageGenerationExplicitToolPolicy == codexImageGenerationExplicitToolPolicyStrip { + strippedBody, changed, stripErr := stripOpenAIImageGenerationToolsFromRawPayload(body) + if stripErr != nil { + return nil, stripErr + } + if changed { + body = strippedBody + originalBody = strippedBody + logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Stripped /responses image_generation tool for Codex client by account policy") + } + } // 透传分支只需要轻量提取字段,避免热路径全量 Unmarshal。 mappedModel := account.GetMappedModel(reqModel) reasoningEffort := extractOpenAIReasoningEffortFromBody(body, mappedModel) @@ -159,10 +174,6 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco if apiKey != nil { imageGenerationAllowed = GroupAllowsImageGeneration(apiKey.Group) } - codexImageGenerationExplicitToolPolicy := codexImageGenerationExplicitToolPolicyAllow - if isCodexCLI { - codexImageGenerationExplicitToolPolicy = account.CodexImageGenerationExplicitToolPolicy() - } codexImageGenerationBridgeEnabled := isCodexCLI && imageGenerationAllowed && codexImageGenerationExplicitToolPolicy != codexImageGenerationExplicitToolPolicyStrip && s.isCodexImageGenerationBridgeEnabled(ctx, account, apiKey) var imageIntent bool if isCodexCLI && codexImageGenerationExplicitToolPolicy == codexImageGenerationExplicitToolPolicyStrip { @@ -272,7 +283,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco markDecodedModified() logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Added Codex image_generation bridge instructions") } - } else if imageGenerationAllowed && imageIntent && openAIRequestBodyHasImageGenerationTool(body) { + } else if imageGenerationAllowed && imageIntent && openAIRequestBodyHasImageGenerationDeclaration(body) { // 完整 image_generation tool 只做 raw 计费读取,校验/桥接/旧字段迁移命中时才展开大 input map。 logger.LegacyPrintf("service.openai_gateway", "[OpenAI] /responses image_generation request inbound_model=%s mapped_model=%s account_type=%s", requestView.Model, upstreamModel, account.Type) } @@ -292,7 +303,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco // gpt-5.3-codex-spark also rejects the image_generation tool (HTTP 400, // param=tools). Strip it here so both APIKey and OAuth /responses paths are // covered regardless of the image-generation feature gate. - if isCodexSparkModel(upstreamModel) && openAIRequestBodyHasImageGenerationTool(body) { + if isCodexSparkModel(upstreamModel) && openAIRequestBodyHasImageGenerationDeclaration(body) { decoded, decodeErr := ensureReqBody() if decodeErr != nil { return nil, decodeErr diff --git a/backend/internal/service/openai_image_generation_controls_test.go b/backend/internal/service/openai_image_generation_controls_test.go index 31edd36097..af0cdf669c 100644 --- a/backend/internal/service/openai_image_generation_controls_test.go +++ b/backend/internal/service/openai_image_generation_controls_test.go @@ -2,6 +2,7 @@ package service import ( "context" + "encoding/json" "io" "net/http" "net/http/httptest" @@ -191,6 +192,64 @@ func TestOpenAIGatewayServiceForward_AccountPolicyStripsExplicitImageTool(t *tes require.NotContains(t, instructions, "image_generation") } +func TestOpenAIGatewayServiceForward_AccountPolicyStripsImageNamespaceTools(t *testing.T) { + gin.SetMode(gin.TestMode) + + tests := []struct { + name string + passthrough bool + }{ + {name: "managed forwarding"}, + {name: "passthrough forwarding", passthrough: true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"id":"resp_stripped_namespace","model":"gpt-5.5","usage":{"input_tokens":2,"output_tokens":1}}`)), + }, + } + svc := newOpenAIImageGenerationControlTestService(upstream) + c, _ := newOpenAIImageGenerationControlTestContext(false, "codex_cli_rs/0.144.1") + account := newOpenAIImageGenerationControlTestAccount() + account.Extra = map[string]any{ + featureKeyCodexImageGenerationExplicitToolPolicy: codexImageGenerationExplicitToolPolicyStrip, + "openai_passthrough": tt.passthrough, + } + body := []byte(`{ + "model":"gpt-5.5", + "stream":false, + "tools":[ + {"type":"function","name":"shell","parameters":{"type":"object"}}, + {"type":"namespace","name":"image_gen","tools":[{"type":"function","name":"imagegen"}]}, + {"type":"namespace","name":"code_tools","tools":[{"type":"function","name":"run"}]} + ], + "input":[ + {"type":"message","role":"user","content":[{"type":"input_text","text":"write code"}]}, + {"type":"additional_tools","tools":[{"type":"namespace","name":"image_gen","tools":[{"type":"function","name":"imagegen"}]}]} + ], + "tool_choice":"auto" + }`) + + result, err := svc.Forward(context.Background(), c, account, body) + + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, upstream.lastReq) + var forwarded map[string]any + require.NoError(t, json.Unmarshal(upstream.lastBody, &forwarded)) + require.False(t, hasOpenAIImageGenerationTool(forwarded)) + require.Equal(t, "auto", forwarded["tool_choice"]) + require.True(t, gjson.GetBytes(upstream.lastBody, `tools.#(name=="shell")`).Exists()) + require.True(t, gjson.GetBytes(upstream.lastBody, `tools.#(name=="code_tools")`).Exists()) + require.Equal(t, "write code", gjson.GetBytes(upstream.lastBody, "input.0.content.0.text").String()) + }) + } +} + func TestOpenAIGatewayServiceForward_ChannelBridgeOverrideEnablesCodexInjection(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/service/openai_ws_forwarder_ingress.go b/backend/internal/service/openai_ws_forwarder_ingress.go index 4a5fa20774..45e004b6d0 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress.go +++ b/backend/internal/service/openai_ws_forwarder_ingress.go @@ -270,7 +270,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( normalized = next } if isCodexCLI && codexImageGenerationExplicitToolPolicy == codexImageGenerationExplicitToolPolicyStrip { - if stripped, changed, stripErr := stripOpenAIImageGenerationToolFromRawPayload(normalized); stripErr != nil { + if stripped, changed, stripErr := stripOpenAIImageGenerationToolsFromRawPayload(normalized); stripErr != nil { return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", stripErr) } else if changed { normalized = stripped diff --git a/backend/internal/service/openai_ws_forwarder_ingress_test.go b/backend/internal/service/openai_ws_forwarder_ingress_test.go index 1d18c46fca..7753ea9598 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress_test.go +++ b/backend/internal/service/openai_ws_forwarder_ingress_test.go @@ -153,6 +153,24 @@ func TestStripCodexSparkImageGenerationToolFromRawPayload(t *testing.T) { require.True(t, gjson.GetBytes(updated, `tools.#(type=="function")`).Exists()) }) + t.Run("strips_namespace_tools_for_spark", func(t *testing.T) { + payload := []byte(`{ + "type":"response.create", + "model":"gpt-5.3-codex-spark", + "input":[ + {"type":"message","role":"user","content":"hello"}, + {"type":"additional_tools","tools":[{"type":"namespace","name":"image_gen"}]} + ], + "tool_choice":{"type":"namespace","name":"image_gen"} + }`) + updated, changed, err := stripCodexSparkImageGenerationToolFromRawPayload(payload, "gpt-5.3-codex-spark") + require.NoError(t, err) + require.True(t, changed) + require.False(t, IsImageGenerationIntent(openAIResponsesEndpoint, "gpt-5.3-codex-spark", updated)) + require.Equal(t, "hello", gjson.GetBytes(updated, "input.0.content").String()) + require.False(t, gjson.GetBytes(updated, "tool_choice").Exists()) + }) + t.Run("keeps_image_generation_for_non_spark", func(t *testing.T) { payload := []byte(`{"type":"response.create","model":"gpt-5.3-codex","tools":[{"type":"image_generation","output_format":"png"}]}`) updated, changed, err := stripCodexSparkImageGenerationToolFromRawPayload(payload, "gpt-5.3-codex") @@ -170,24 +188,61 @@ func TestStripCodexSparkImageGenerationToolFromRawPayload(t *testing.T) { }) } -func TestStripOpenAIImageGenerationToolFromRawPayload(t *testing.T) { - payload := []byte(`{ - "type":"response.create", - "model":"gpt-5.4", - "tools":[ - {"type":"function","name":"shell"}, - {"type":"image_generation","output_format":"png"} - ], - "tool_choice":{"type":"image_generation"} - }`) +func TestStripOpenAIImageGenerationToolsFromRawPayload(t *testing.T) { + t.Run("flat image tool", func(t *testing.T) { + payload := []byte(`{ + "type":"response.create", + "model":"gpt-5.4", + "tools":[ + {"type":"function","name":"shell"}, + {"type":"image_generation","output_format":"png"} + ], + "tool_choice":{"type":"image_generation"} + }`) - updated, changed, err := stripOpenAIImageGenerationToolFromRawPayload(payload) + updated, changed, err := stripOpenAIImageGenerationToolsFromRawPayload(payload) - require.NoError(t, err) - require.True(t, changed) - require.False(t, gjson.GetBytes(updated, `tools.#(type=="image_generation")`).Exists()) - require.True(t, gjson.GetBytes(updated, `tools.#(type=="function")`).Exists()) - require.False(t, gjson.GetBytes(updated, "tool_choice").Exists()) + require.NoError(t, err) + require.True(t, changed) + require.False(t, gjson.GetBytes(updated, `tools.#(type=="image_generation")`).Exists()) + require.True(t, gjson.GetBytes(updated, `tools.#(type=="function")`).Exists()) + require.False(t, gjson.GetBytes(updated, "tool_choice").Exists()) + }) + + t.Run("namespace and Responses Lite tools", func(t *testing.T) { + payload := []byte(`{ + "type":"response.create", + "model":"gpt-5.5", + "tools":[ + {"type":"namespace","name":"image_gen","tools":[{"type":"function","name":"imagegen"}]}, + {"type":"namespace","name":"code_tools","tools":[{"type":"function","name":"run"}]} + ], + "input":[ + {"type":"message","role":"user","content":"hello"}, + {"type":"additional_tools","tools":[{"type":"namespace","name":"image_gen"}]} + ], + "tool_choice":{"type":"namespace","name":"image_gen"} + }`) + + updated, changed, err := stripOpenAIImageGenerationToolsFromRawPayload(payload) + + require.NoError(t, err) + require.True(t, changed) + require.False(t, IsImageGenerationIntent(openAIResponsesEndpoint, "gpt-5.5", updated)) + require.True(t, gjson.GetBytes(updated, `tools.#(name=="code_tools")`).Exists()) + require.Equal(t, "hello", gjson.GetBytes(updated, "input.0.content").String()) + require.False(t, gjson.GetBytes(updated, "tool_choice").Exists()) + }) + + t.Run("non-image namespace is unchanged", func(t *testing.T) { + payload := []byte(`{"type":"response.create","model":"gpt-5.5","tools":[{"type":"namespace","name":"code_tools"}]}`) + + updated, changed, err := stripOpenAIImageGenerationToolsFromRawPayload(payload) + + require.NoError(t, err) + require.False(t, changed) + require.Equal(t, payload, updated) + }) } func TestAlignStoreDisabledPreviousResponseID(t *testing.T) { diff --git a/backend/internal/service/openai_ws_forwarder_v2.go b/backend/internal/service/openai_ws_forwarder_v2.go index 24853d4a6c..f0f71648dd 100644 --- a/backend/internal/service/openai_ws_forwarder_v2.go +++ b/backend/internal/service/openai_ws_forwarder_v2.go @@ -3,7 +3,6 @@ package service import ( "bytes" "context" - "encoding/json" "errors" "fmt" "net/http" @@ -710,23 +709,8 @@ func (s *OpenAIGatewayService) forwardOpenAIWSV2( // Codex clients advertise it by default. Returns the (possibly unchanged) payload, // whether it changed, and any JSON decode error. func stripCodexSparkImageGenerationToolFromRawPayload(payload []byte, model string) ([]byte, bool, error) { - if !isCodexSparkModel(model) || !openAIRequestBodyHasImageGenerationTool(payload) { + if !isCodexSparkModel(model) { return payload, false, nil } - return stripOpenAIImageGenerationToolFromRawPayload(payload) -} - -func stripOpenAIImageGenerationToolFromRawPayload(payload []byte) ([]byte, bool, error) { - payloadMap := make(map[string]any) - if err := json.Unmarshal(payload, &payloadMap); err != nil { - return payload, false, err - } - if !stripOpenAIImageGenerationTools(payloadMap) { - return payload, false, nil - } - rebuilt, err := json.Marshal(payloadMap) - if err != nil { - return payload, false, err - } - return rebuilt, true, nil + return stripOpenAIImageGenerationToolsFromRawPayload(payload) }