fix(openai): preserve Codex image function tools

This commit is contained in:
xiaohei210509
2026-07-15 15:10:43 +08:00
committed by zxh
parent 44452d4a9d
commit ade5944201
4 changed files with 167 additions and 3 deletions
@@ -83,6 +83,8 @@ type codexOAuthTransformOptions struct {
PreserveToolCallIDs bool
}
const codexImageGenerationFunctionToolName = "image_gen.imagegen"
const (
codexImageGenerationBridgeMarker = "<sub2api-codex-image-generation>"
codexImageGenerationBridgeText = codexImageGenerationBridgeMarker + "\nWhen the user asks for raster image generation or editing, use the OpenAI Responses native `image_generation` tool attached to this request. The local Codex client may not expose an `image_gen` namespace, but that does not mean image generation is unavailable. Do not ask the user to switch to CLI fallback solely because `image_gen` is absent.\n</sub2api-codex-image-generation>"
@@ -616,6 +618,11 @@ func hasOpenAIImageGenerationTool(reqBody map[string]any) bool {
return inputContainsImageGenerationTool(reqBody["input"])
}
func hasCodexImageGenerationFunctionTool(reqBody map[string]any) bool {
return len(reqBody) > 0 &&
codexToolsContainFunctionName(reqBody["tools"], codexImageGenerationFunctionToolName)
}
func toolsContainImageGeneration(rawTools any) bool {
if rawTools == nil {
return false
@@ -855,6 +862,9 @@ func ensureOpenAIResponsesImageGenerationTool(reqBody map[string]any) bool {
if isCodexSparkModel(firstNonEmptyString(reqBody["model"])) {
return false
}
if hasCodexImageGenerationFunctionTool(reqBody) {
return false
}
if hasOpenAIImageGenerationTool(reqBody) {
return false
}
@@ -880,7 +890,7 @@ func ensureOpenAIResponsesImageGenerationTool(reqBody map[string]any) bool {
}
func ensureOpenAIResponsesImageGenerationToolChoiceAuto(reqBody map[string]any) bool {
if len(reqBody) == 0 || !hasOpenAIImageGenerationTool(reqBody) {
if len(reqBody) == 0 || hasCodexImageGenerationFunctionTool(reqBody) || !hasOpenAIImageGenerationTool(reqBody) {
return false
}
if isCodexSparkModel(firstNonEmptyString(reqBody["model"])) {
@@ -894,7 +904,7 @@ func ensureOpenAIResponsesImageGenerationToolChoiceAuto(reqBody map[string]any)
}
func applyCodexImageGenerationBridgeInstructions(reqBody map[string]any) bool {
if len(reqBody) == 0 || !hasOpenAIImageGenerationTool(reqBody) {
if len(reqBody) == 0 || hasCodexImageGenerationFunctionTool(reqBody) || !hasOpenAIImageGenerationTool(reqBody) {
return false
}
if isCodexSparkModel(firstNonEmptyString(reqBody["model"])) {
@@ -676,6 +676,86 @@ func TestEnsureOpenAIResponsesImageGenerationTool_PreservesImageGenNamespace(t *
}
}
func TestCodexImageGenerationBridge_PreservesClientImageFunctionTools(t *testing.T) {
tests := []struct {
name string
reqBody map[string]any
wantClient bool
}{
{
name: "flat image_gen function",
reqBody: map[string]any{
"model": "gpt-5.5",
"input": "draw a cat",
"tools": []any{
map[string]any{"type": "function", "name": "image_gen.imagegen"},
},
},
wantClient: true,
},
{
name: "nested image_gen function",
reqBody: map[string]any{
"model": "gpt-5.5",
"input": "draw a cat",
"tools": []any{
map[string]any{
"type": "function",
"function": map[string]any{
"name": "image_gen.imagegen",
},
},
},
},
wantClient: true,
},
{
name: "similar function name still receives hosted bridge",
reqBody: map[string]any{
"model": "gpt-5.5",
"input": "draw a cat",
"tools": []any{
map[string]any{"type": "function", "name": "image_gen.imagegenerator"},
},
},
wantClient: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
tt.reqBody["instructions"] = "existing instructions"
require.Equal(t, tt.wantClient, hasCodexImageGenerationFunctionTool(tt.reqBody))
toolModified := ensureOpenAIResponsesImageGenerationTool(tt.reqBody)
choiceModified := ensureOpenAIResponsesImageGenerationToolChoiceAuto(tt.reqBody)
instructionsModified := applyCodexImageGenerationBridgeInstructions(tt.reqBody)
require.Equal(t, !tt.wantClient, toolModified)
require.Equal(t, !tt.wantClient, choiceModified)
require.Equal(t, !tt.wantClient, instructionsModified)
hasHostedTool := false
tools, _ := tt.reqBody["tools"].([]any)
for _, rawTool := range tools {
tool, ok := rawTool.(map[string]any)
if ok && firstNonEmptyString(tool["type"]) == "image_generation" {
hasHostedTool = true
}
}
require.Equal(t, !tt.wantClient, hasHostedTool)
if tt.wantClient {
require.NotContains(t, tt.reqBody, "tool_choice")
require.Equal(t, "existing instructions", tt.reqBody["instructions"])
} else {
require.Equal(t, "auto", tt.reqBody["tool_choice"])
require.Contains(t, tt.reqBody["instructions"], codexImageGenerationBridgeMarker)
}
})
}
}
func TestApplyCodexImageGenerationBridgeInstructions_AppendsBridgeOnce(t *testing.T) {
reqBody := map[string]any{
"model": "gpt-5.4",
@@ -353,6 +353,54 @@ func TestOpenAIGatewayServiceForward_CodexBridgeDoesNotInjectHostedToolAlongside
require.Equal(t, "namespace", gjson.GetBytes(upstream.lastBody, `input.#(type=="additional_tools").tools.#(name=="image_gen").type`).String())
}
func TestOpenAIGatewayServiceForward_CodexBridgePreservesImageGenFunction(t *testing.T) {
gin.SetMode(gin.TestMode)
tests := []struct {
name string
tool string
}{
{
name: "flat function",
tool: `{"type":"function","name":"image_gen.imagegen","parameters":{"type":"object"}}`,
},
{
name: "nested function",
tool: `{"type":"function","function":{"name":"image_gen.imagegen","parameters":{"type":"object"}}}`,
},
}
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_function_image","model":"gpt-5.5","usage":{"input_tokens":1,"output_tokens":1}}`)),
},
}
svc := newOpenAIImageGenerationControlTestService(upstream)
svc.cfg.Gateway.CodexImageGenerationBridgeEnabled = true
c, _ := newOpenAIImageGenerationControlTestContext(true, "codex_cli_rs/0.144.1")
account := newOpenAIImageGenerationControlTestAccount()
body := []byte(`{"model":"gpt-5.5","input":"draw a cat","stream":false,"tools":[` + tt.tool + `]}`)
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.True(t, hasCodexImageGenerationFunctionTool(forwarded))
require.False(t, gjson.GetBytes(upstream.lastBody, `tools.#(type=="image_generation")`).Exists())
require.False(t, gjson.GetBytes(upstream.lastBody, "tool_choice").Exists())
require.NotContains(t, gjson.GetBytes(upstream.lastBody, "instructions").String(), codexImageGenerationBridgeMarker)
})
}
}
func TestOpenAIGatewayServiceForward_CodexBridgePreservesExistingToolChoice(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -425,6 +425,7 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_CodexImageBridge
events: [][]byte{
[]byte(`{"type":"response.completed","response":{"id":"resp_codex_image_bridge","model":"gpt-5.5","usage":{"input_tokens":1,"output_tokens":1}}}`),
[]byte(`{"type":"response.completed","response":{"id":"resp_codex_image_lite","model":"gpt-5.5","usage":{"input_tokens":1,"output_tokens":1}}}`),
[]byte(`{"type":"response.completed","response":{"id":"resp_codex_image_function","model":"gpt-5.5","usage":{"input_tokens":1,"output_tokens":1}}}`),
},
}
captureDialer := &openAIWSCaptureDialer{conn: captureConn}
@@ -548,6 +549,25 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_CodexImageBridge
require.Equal(t, coderws.MessageText, msgType)
require.Equal(t, "resp_codex_image_lite", gjson.GetBytes(message, "response.id").String())
writeCtx, cancelWrite = context.WithTimeout(context.Background(), 3*time.Second)
err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{
"type":"response.create",
"model":"gpt-5.5",
"stream":false,
"previous_response_id":"resp_codex_image_lite",
"input":"draw a cat",
"tools":[{"type":"function","name":"image_gen.imagegen","parameters":{"type":"object"}}]
}`))
cancelWrite()
require.NoError(t, err)
readCtx, cancelRead = context.WithTimeout(context.Background(), 3*time.Second)
msgType, message, err = clientConn.Read(readCtx)
cancelRead()
require.NoError(t, err)
require.Equal(t, coderws.MessageText, msgType)
require.Equal(t, "resp_codex_image_function", gjson.GetBytes(message, "response.id").String())
_ = clientConn.Close(coderws.StatusNormalClosure, "done")
select {
@@ -557,7 +577,7 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_CodexImageBridge
t.Fatal("等待 ingress websocket 结束超时")
}
require.Len(t, captureConn.writes, 2)
require.Len(t, captureConn.writes, 3)
nonLitePayload := requestToJSONString(captureConn.writes[0])
require.True(t, gjson.Get(nonLitePayload, `tools.#(type=="image_generation")`).Exists())
require.Equal(t, "png", gjson.Get(nonLitePayload, `tools.#(type=="image_generation").output_format`).String())
@@ -573,6 +593,12 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_CodexImageBridge
require.Equal(t, "collaboration", gjson.Get(litePayload, `input.#(type=="additional_tools").tools.1.name`).String())
require.Equal(t, "namespace", gjson.Get(litePayload, "tool_choice.type").String())
require.Equal(t, "collaboration", gjson.Get(litePayload, "tool_choice.name").String())
functionPayload := requestToJSONString(captureConn.writes[2])
require.True(t, gjson.Get(functionPayload, `tools.#(name=="image_gen.imagegen")`).Exists())
require.False(t, gjson.Get(functionPayload, `tools.#(type=="image_generation")`).Exists())
require.False(t, gjson.Get(functionPayload, "tool_choice").Exists())
require.NotContains(t, gjson.Get(functionPayload, "instructions").String(), codexImageGenerationBridgeMarker)
}
func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_DedicatedModeDoesNotReuseConnAcrossSessions(t *testing.T) {