From e5f7836bf369f62549e4e95cc534c181e086c4c4 Mon Sep 17 00:00:00 2001 From: Hao Liu Date: Fri, 26 Jun 2026 19:10:37 +0800 Subject: [PATCH] fix(openai): set tool_choice auto for Codex image bridge --- .../service/openai_codex_transform.go | 14 +++++++++ .../service/openai_gateway_service.go | 4 +++ .../openai_image_generation_controls_test.go | 29 +++++++++++++++++++ .../internal/service/openai_ws_forwarder.go | 4 +++ ...penai_ws_forwarder_ingress_session_test.go | 1 + 5 files changed, 52 insertions(+) diff --git a/backend/internal/service/openai_codex_transform.go b/backend/internal/service/openai_codex_transform.go index a6aa85192f..4be448b2d8 100644 --- a/backend/internal/service/openai_codex_transform.go +++ b/backend/internal/service/openai_codex_transform.go @@ -754,6 +754,20 @@ func ensureOpenAIResponsesImageGenerationTool(reqBody map[string]any) bool { return true } +func ensureOpenAIResponsesImageGenerationToolChoiceAuto(reqBody map[string]any) bool { + if len(reqBody) == 0 || !hasOpenAIImageGenerationTool(reqBody) { + return false + } + if isCodexSparkModel(firstNonEmptyString(reqBody["model"])) { + return false + } + if _, ok := reqBody["tool_choice"]; ok { + return false + } + reqBody["tool_choice"] = "auto" + return true +} + func applyCodexImageGenerationBridgeInstructions(reqBody map[string]any) bool { if len(reqBody) == 0 || !hasOpenAIImageGenerationTool(reqBody) { return false diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index 1def695c91..ee3b6d3f41 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -2696,6 +2696,10 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco markDecodedModified() logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Injected /responses image_generation tool for Codex client") } + if codexImageGenerationBridgeEnabled && ensureOpenAIResponsesImageGenerationToolChoiceAuto(decoded) { + markDecodedModified() + logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Set /responses image_generation tool_choice=auto for Codex client") + } if normalizeOpenAIResponsesImageGenerationTools(decoded) { markDecodedModified() logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Normalized /responses image_generation tool payload") diff --git a/backend/internal/service/openai_image_generation_controls_test.go b/backend/internal/service/openai_image_generation_controls_test.go index ce66d6722b..e0834fa105 100644 --- a/backend/internal/service/openai_image_generation_controls_test.go +++ b/backend/internal/service/openai_image_generation_controls_test.go @@ -116,6 +116,11 @@ func TestOpenAIGatewayServiceForward_CodexImageInjectionRespectsGroupCapability( require.Equal(t, tt.wantInjected, hasImageTool) instructions := gjson.GetBytes(upstream.lastBody, "instructions").String() require.Equal(t, tt.wantInjected, strings.Contains(instructions, "image_generation")) + toolChoice := gjson.GetBytes(upstream.lastBody, "tool_choice") + require.Equal(t, tt.wantInjected, toolChoice.Exists()) + if tt.wantInjected { + require.Equal(t, "auto", toolChoice.String()) + } }) } } @@ -175,10 +180,34 @@ func TestOpenAIGatewayServiceForward_ChannelBridgeOverrideEnablesCodexInjection( require.NotNil(t, result) require.NotNil(t, upstream.lastReq) require.True(t, gjson.GetBytes(upstream.lastBody, `tools.#(type=="image_generation")`).Exists()) + require.Equal(t, "auto", gjson.GetBytes(upstream.lastBody, "tool_choice").String()) instructions := gjson.GetBytes(upstream.lastBody, "instructions").String() require.Contains(t, instructions, "image_generation") } +func TestOpenAIGatewayServiceForward_CodexBridgePreservesExistingToolChoice(t *testing.T) { + gin.SetMode(gin.TestMode) + + upstream := &httpUpstreamRecorder{ + resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"id":"resp_codex_tool_choice","model":"gpt-5.4","usage":{"input_tokens":1,"output_tokens":1}}`)), + }, + } + svc := newOpenAIImageGenerationControlTestService(upstream) + svc.cfg.Gateway.CodexImageGenerationBridgeEnabled = true + c, _ := newOpenAIImageGenerationControlTestContext(true, "codex_cli_rs/0.98.0") + account := newOpenAIImageGenerationControlTestAccount() + + result, err := svc.Forward(context.Background(), c, account, []byte(`{"model":"gpt-5.4","input":"draw","stream":false,"tools":[{"type":"image_generation"}],"tool_choice":{"type":"image_generation"}}`)) + + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, upstream.lastReq) + require.Equal(t, "image_generation", gjson.GetBytes(upstream.lastBody, "tool_choice.type").String()) +} + func TestOpenAIGatewayService_CodexImageGenerationBridgeOverridePrecedence(t *testing.T) { groupID := int64(4242) diff --git a/backend/internal/service/openai_ws_forwarder.go b/backend/internal/service/openai_ws_forwarder.go index 8c34d4defd..a3e1eaf21e 100644 --- a/backend/internal/service/openai_ws_forwarder.go +++ b/backend/internal/service/openai_ws_forwarder.go @@ -2666,6 +2666,10 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( bridgeModified = true logOpenAIWSModeInfo("ingress_ws_codex_image_tool_injected account_id=%d", account.ID) } + if ensureOpenAIResponsesImageGenerationToolChoiceAuto(payloadMap) { + bridgeModified = true + logOpenAIWSModeInfo("ingress_ws_codex_image_tool_choice_auto account_id=%d", account.ID) + } if normalizeOpenAIResponsesImageGenerationTools(payloadMap) { bridgeModified = true } diff --git a/backend/internal/service/openai_ws_forwarder_ingress_session_test.go b/backend/internal/service/openai_ws_forwarder_ingress_session_test.go index dc48b990ea..856e81be02 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress_session_test.go +++ b/backend/internal/service/openai_ws_forwarder_ingress_session_test.go @@ -431,6 +431,7 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_InjectsCodexImag upstreamPayload := requestToJSONString(captureConn.writes[0]) require.True(t, gjson.Get(upstreamPayload, `tools.#(type=="image_generation")`).Exists()) require.Equal(t, "png", gjson.Get(upstreamPayload, `tools.#(type=="image_generation").output_format`).String()) + require.Equal(t, "auto", gjson.Get(upstreamPayload, "tool_choice").String()) require.Contains(t, gjson.Get(upstreamPayload, "instructions").String(), "image_generation") }