mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-21 14:19:18 +08:00
fix(openai): preserve Codex image function tools
This commit is contained in:
@@ -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) {
|
||||
|
||||
Reference in New Issue
Block a user