From 44ef88f659a02ca9b725ce11b110e03cd43d814d Mon Sep 17 00:00:00 2001 From: wucm667 Date: Sat, 15 Aug 2026 11:02:59 +0800 Subject: [PATCH] fix(openai): restore API-key custom tools --- .../service/openai_gateway_forward.go | 1 + .../service/openai_gateway_passthrough.go | 102 ++++++++++++ ...nai_gateway_responses_client_tools_test.go | 146 ++++++++++++++++++ 3 files changed, 249 insertions(+) create mode 100644 backend/internal/service/openai_gateway_responses_client_tools_test.go diff --git a/backend/internal/service/openai_gateway_forward.go b/backend/internal/service/openai_gateway_forward.go index f7eb759346..31beb344a0 100644 --- a/backend/internal/service/openai_gateway_forward.go +++ b/backend/internal/service/openai_gateway_forward.go @@ -21,6 +21,7 @@ import ( func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, account *Account, body []byte) (*OpenAIForwardResult, error) { beginUpstreamResponseModelObservation(c) clearGrokResponsesClientToolMapping(c) + clearOpenAIResponsesClientToolMapping(c) clearOpenAIResponsesNamespaceNames(c) startTime := time.Now() // 固定渠道映射后的请求级 canonical body;账号 normalize/strip 不得改写跨 failover hint。 diff --git a/backend/internal/service/openai_gateway_passthrough.go b/backend/internal/service/openai_gateway_passthrough.go index 8c029f0aa8..4fc694f17f 100644 --- a/backend/internal/service/openai_gateway_passthrough.go +++ b/backend/internal/service/openai_gateway_passthrough.go @@ -16,6 +16,7 @@ import ( "strings" "time" + "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" "github.com/Wei-Shaw/sub2api/internal/pkg/logger" "github.com/Wei-Shaw/sub2api/internal/pkg/openai" "github.com/Wei-Shaw/sub2api/internal/util/responseheaders" @@ -25,6 +26,84 @@ import ( "go.uber.org/zap" ) +const openAIResponsesClientToolMappingContextKey = "openai_responses_client_tool_mapping" + +func hasOpenAIResponsesClientToolMapping(mapping apicompat.ResponsesClientToolMapping) bool { + return len(mapping.CustomTools) > 0 || mapping.ToolSearch || len(mapping.NamespaceTools) > 0 +} + +func adaptOpenAIResponsesClientTools(body []byte) ([]byte, apicompat.ResponsesClientToolMapping, error) { + if !needsOpenAIResponsesClientToolAdaptation(body) { + return body, apicompat.ResponsesClientToolMapping{}, nil + } + + decoder := json.NewDecoder(bytes.NewReader(body)) + decoder.UseNumber() + var requestBody map[string]any + if err := decoder.Decode(&requestBody); err != nil { + return body, apicompat.ResponsesClientToolMapping{}, fmt.Errorf("decode OpenAI Responses client tools: %w", err) + } + var trailingValue any + if err := decoder.Decode(&trailingValue); !errors.Is(err, io.EOF) { + if err == nil { + err = errors.New("multiple JSON values") + } + return body, apicompat.ResponsesClientToolMapping{}, fmt.Errorf("decode OpenAI Responses client tools trailing data: %w", err) + } + mapping, changed, err := apicompat.AdaptResponsesClientTools(requestBody) + if err != nil || !changed { + return body, mapping, err + } + rebuilt, err := marshalOpenAIUpstreamJSON(requestBody) + if err != nil { + return body, apicompat.ResponsesClientToolMapping{}, fmt.Errorf("encode OpenAI Responses client tools: %w", err) + } + return rebuilt, mapping, nil +} + +func needsOpenAIResponsesClientToolAdaptation(body []byte) bool { + needsAdaptation := false + var visit func(gjson.Result) bool + visit = func(value gjson.Result) bool { + if value.IsObject() { + switch strings.TrimSpace(value.Get("type").String()) { + case "custom", "custom_tool_call", "custom_tool_call_output", + "tool_search", "tool_search_call", "tool_search_output": + needsAdaptation = true + return false + } + } + if value.IsObject() || value.IsArray() { + value.ForEach(func(_, child gjson.Result) bool { + return visit(child) + }) + } + return !needsAdaptation + } + visit(gjson.ParseBytes(body)) + return needsAdaptation +} + +func openAIResponsesClientToolMapping(c *gin.Context) (apicompat.ResponsesClientToolMapping, bool) { + if c == nil { + return apicompat.ResponsesClientToolMapping{}, false + } + value, ok := c.Get(openAIResponsesClientToolMappingContextKey) + mapping, typed := value.(apicompat.ResponsesClientToolMapping) + return mapping, ok && typed && hasOpenAIResponsesClientToolMapping(mapping) +} + +// clearOpenAIResponsesClientToolMapping removes mapping state from the prior +// forwarding attempt. Forward retries accounts on the same Gin context. +func clearOpenAIResponsesClientToolMapping(c *gin.Context) { + if c == nil { + return + } + if _, exists := c.Get(openAIResponsesClientToolMappingContextKey); exists { + c.Set(openAIResponsesClientToolMappingContextKey, apicompat.ResponsesClientToolMapping{}) + } +} + func (s *OpenAIGatewayService) forwardOpenAIPassthrough( ctx context.Context, c *gin.Context, @@ -82,6 +161,16 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough( reqStream = gjson.GetBytes(body, "stream").Bool() } + if account != nil && account.Platform == PlatformOpenAI && account.Type == AccountTypeAPIKey && + !isOpenAIResponsesCompactPath(c) && needsOpenAIResponsesClientToolAdaptation(body) { + adaptedBody, mapping, adaptErr := adaptOpenAIResponsesClientTools(body) + if adaptErr != nil { + return nil, adaptErr + } + body = adaptedBody + c.Set(openAIResponsesClientToolMappingContextKey, mapping) + } + sanitizedBody, sanitized, err := sanitizeEmptyBase64InputImagesInOpenAIBody(body) if err != nil { return nil, err @@ -236,6 +325,13 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough( return nil, s.handleErrorResponsePassthrough(ctx, resp, c, account, body, probeBody) } defer func() { _ = resp.Body.Close() }() + if mapping, ok := openAIResponsesClientToolMapping(c); ok && isEventStreamResponse(resp.Header) { + maxLineSize := defaultMaxLineSize + if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 { + maxLineSize = s.cfg.Gateway.MaxLineSize + } + resp.Body = newGrokResponsesClientToolStreamBody(resp.Body, mapping, maxLineSize) + } serviceTier := extractOpenAIServiceTierFromBody(body) @@ -1520,6 +1616,12 @@ func (s *OpenAIGatewayService) handleNonStreamingResponsePassthrough( if err != nil { return nil, fmt.Errorf("restore OpenAI passthrough namespace response: %w", err) } + if mapping, ok := openAIResponsesClientToolMapping(c); ok && json.Valid(body) { + body, _, err = apicompat.RestoreResponsesClientToolPayload(body, mapping) + if err != nil { + return nil, fmt.Errorf("restore OpenAI Responses client tools: %w", err) + } + } if !writeOpenAICompactSSEBridge(c, resp.StatusCode, body) { c.Data(resp.StatusCode, contentType, body) } diff --git a/backend/internal/service/openai_gateway_responses_client_tools_test.go b/backend/internal/service/openai_gateway_responses_client_tools_test.go new file mode 100644 index 0000000000..cfb4f6eb62 --- /dev/null +++ b/backend/internal/service/openai_gateway_responses_client_tools_test.go @@ -0,0 +1,146 @@ +package service + +import ( + "bytes" + "context" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +func openAIClientToolsRequest(stream bool) []byte { + streamValue := "false" + if stream { + streamValue = "true" + } + return []byte(`{"model":"gpt-5.4","input":"fix it","stream":` + streamValue + `,"tools":[{"type":"custom","name":"exec"},{"type":"custom","name":"apply_patch"}]}`) +} + +func assertOpenAIClientToolsLowered(t *testing.T, body []byte) { + t.Helper() + for index, name := range []string{"exec", "apply_patch"} { + tool := gjson.GetBytes(body, "tools."+string(rune('0'+index))) + require.Equal(t, "function", tool.Get("type").String()) + require.Equal(t, name, tool.Get("name").String()) + require.Equal(t, "string", tool.Get("parameters.properties.input.type").String()) + } +} + +func openAIClientToolsTestService(upstream *httpUpstreamRecorder) *OpenAIGatewayService { + return &OpenAIGatewayService{ + httpUpstream: upstream, + cfg: &config.Config{Security: config.SecurityConfig{ + URLAllowlist: config.URLAllowlistConfig{Enabled: false}, + }}, + } +} + +func TestAdaptOpenAIResponsesClientToolsLeavesNamespaceOnlyBodyUnchanged(t *testing.T) { + body := []byte(`{ + "model": "gpt-5.5", + "tools": [{"type": "namespace", "name": "code_tools", "tools": [{"type": "function", "name": "run"}]}], + "tool_choice": "auto" + }`) + + adapted, mapping, err := adaptOpenAIResponsesClientTools(body) + + require.NoError(t, err) + require.Equal(t, body, adapted) + require.Empty(t, mapping.CustomTools) + require.Empty(t, mapping.NamespaceTools) + require.False(t, mapping.ToolSearch) +} + +func TestAdaptOpenAIResponsesClientToolsRejectsTrailingData(t *testing.T) { + tests := map[string][]byte{ + "trailing garbage": append(openAIClientToolsRequest(false), []byte(` garbage`)...), + "second JSON document": append(openAIClientToolsRequest(false), []byte(` {"model":"other"}`)...), + } + + for name, body := range tests { + t.Run(name, func(t *testing.T) { + adapted, mapping, err := adaptOpenAIResponsesClientTools(body) + + require.ErrorContains(t, err, "decode OpenAI Responses client tools trailing data") + require.Equal(t, body, adapted) + require.Empty(t, mapping) + }) + } +} + +func TestClearOpenAIResponsesClientToolMappingRemovesStaleContextState(t *testing.T) { + gin.SetMode(gin.TestMode) + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + c.Set(openAIResponsesClientToolMappingContextKey, apicompat.ResponsesClientToolMapping{CustomTools: map[string]bool{"exec": true}}) + + clearOpenAIResponsesClientToolMapping(c) + + _, ok := openAIResponsesClientToolMapping(c) + require.False(t, ok) +} + +func TestOpenAIPassthroughAPIKeyRestoresClientToolsNonStreaming(t *testing.T) { + gin.SetMode(gin.TestMode) + body := openAIClientToolsRequest(false) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body)) + + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"id":"resp_tools","status":"completed","output":[ + {"type":"function_call","id":"i1","call_id":"c1","name":"exec","arguments":"{\"input\":\"pwd\"}"}, + {"type":"function_call","id":"i2","call_id":"c2","name":"apply_patch","arguments":"{\"input\":\"*** Begin Patch\"}"}],"usage":{}}`)), + }} + svc := openAIClientToolsTestService(upstream) + account := &Account{ID: 5659, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Credentials: map[string]any{"api_key": "test-key"}} + + result, err := svc.forwardOpenAIPassthrough(context.Background(), c, account, body, body, "gpt-5.4", false, nil, false, time.Now()) + + require.NoError(t, err) + require.NotNil(t, result) + assertOpenAIClientToolsLowered(t, upstream.lastBody) + require.Equal(t, "custom_tool_call", gjson.Get(recorder.Body.String(), "output.0.type").String()) + require.Equal(t, "pwd", gjson.Get(recorder.Body.String(), "output.0.input").String()) + require.Equal(t, "custom_tool_call", gjson.Get(recorder.Body.String(), "output.1.type").String()) + require.Equal(t, "*** Begin Patch", gjson.Get(recorder.Body.String(), "output.1.input").String()) +} + +func TestOpenAIPassthroughAPIKeyRestoresClientToolsStreaming(t *testing.T) { + gin.SetMode(gin.TestMode) + body := openAIClientToolsRequest(true) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body)) + + sse := strings.Join([]string{ + `data: {"type":"response.output_item.added","sequence_number":0,"output_index":0,"item":{"type":"function_call","id":"i1","call_id":"c1","name":"apply_patch","status":"in_progress"}}`, + `data: {"type":"response.function_call_arguments.done","sequence_number":1,"item_id":"i1","call_id":"c1","name":"apply_patch","arguments":"{\"input\":\"*** Begin Patch\"}"}`, + `data: {"type":"response.output_item.done","sequence_number":2,"output_index":0,"item":{"type":"function_call","id":"i1","call_id":"c1","name":"apply_patch","arguments":"{\"input\":\"*** Begin Patch\"}","status":"completed"}}`, + `data: {"type":"response.completed","sequence_number":3,"response":{"id":"resp_stream_tools","status":"completed","output":[{"type":"function_call","id":"i1","call_id":"c1","name":"apply_patch","arguments":"{\"input\":\"*** Begin Patch\"}"}],"usage":{"input_tokens":1,"output_tokens":1}}}`, + }, "\n\n") + "\n\n" + upstream := &httpUpstreamRecorder{resp: &http.Response{StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(sse))}} + svc := openAIClientToolsTestService(upstream) + account := &Account{ID: 5660, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Credentials: map[string]any{"api_key": "test-key"}} + + result, err := svc.forwardOpenAIPassthrough(context.Background(), c, account, body, body, "gpt-5.4", false, nil, true, time.Now()) + + require.NoError(t, err) + require.NotNil(t, result) + assertOpenAIClientToolsLowered(t, upstream.lastBody) + output := recorder.Body.String() + require.Contains(t, output, `"type":"custom_tool_call"`) + require.Contains(t, output, `"type":"response.custom_tool_call_input.done"`) + require.Contains(t, output, `"input":"*** Begin Patch"`) + require.NotContains(t, output, `"input":{`) +}