Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 12 additions & 2 deletions backend/internal/service/openai_codex_transform.go
Original file line number Diff line number Diff line change
Expand Up @@ -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>"
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
}
Expand All @@ -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"])) {
Expand All @@ -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"])) {
Expand Down
80 changes: 80 additions & 0 deletions backend/internal/service/openai_codex_transform_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
48 changes: 48 additions & 0 deletions backend/internal/service/openai_image_generation_controls_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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}
Expand Down Expand Up @@ -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 {
Expand All @@ -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())
Expand All @@ -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) {
Expand Down
Loading