From 7e579cb28df5f2e7ed9ef84bc8bca6370acfa8f7 Mon Sep 17 00:00:00 2001 From: hansnow Date: Tue, 18 Aug 2026 14:42:30 +0800 Subject: [PATCH] fix(openai): adapt client tools in WS HTTP bridge --- .../openai_gateway_grok_tool_protocol.go | 26 +++++-- .../internal/service/openai_ws_http_bridge.go | 13 +++- .../service/openai_ws_http_bridge_test.go | 75 +++++++++++++++++++ 3 files changed, 108 insertions(+), 6 deletions(-) diff --git a/backend/internal/service/openai_gateway_grok_tool_protocol.go b/backend/internal/service/openai_gateway_grok_tool_protocol.go index 1a8dfc35a6bd..f1683a7adb46 100644 --- a/backend/internal/service/openai_gateway_grok_tool_protocol.go +++ b/backend/internal/service/openai_gateway_grok_tool_protocol.go @@ -15,12 +15,12 @@ import ( const grokResponsesClientToolMappingContextKey = "grok_responses_client_tool_mapping" -func adaptGrokResponsesClientTools(body []byte) ([]byte, apicompat.ResponsesClientToolMapping, error) { +func adaptResponsesClientToolsForFunctionUpstream(body []byte, upstream string) ([]byte, apicompat.ResponsesClientToolMapping, error) { 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 Grok Responses client tools: %w", err) + return body, apicompat.ResponsesClientToolMapping{}, fmt.Errorf("decode %s Responses client tools: %w", upstream, err) } mapping, changed, err := apicompat.AdaptResponsesClientTools(requestBody) @@ -32,15 +32,23 @@ func adaptGrokResponsesClientTools(body []byte) ([]byte, apicompat.ResponsesClie } rebuilt, err := marshalOpenAIUpstreamJSON(requestBody) if err != nil { - return body, apicompat.ResponsesClientToolMapping{}, fmt.Errorf("encode Grok Responses client tools: %w", err) + return body, apicompat.ResponsesClientToolMapping{}, fmt.Errorf("encode %s Responses client tools: %w", upstream, err) } return rebuilt, mapping, nil } -func hasGrokResponsesClientToolMapping(mapping apicompat.ResponsesClientToolMapping) bool { +func adaptGrokResponsesClientTools(body []byte) ([]byte, apicompat.ResponsesClientToolMapping, error) { + return adaptResponsesClientToolsForFunctionUpstream(body, "Grok") +} + +func hasResponsesClientToolMapping(mapping apicompat.ResponsesClientToolMapping) bool { return len(mapping.CustomTools) > 0 || mapping.ToolSearch || len(mapping.NamespaceTools) > 0 } +func hasGrokResponsesClientToolMapping(mapping apicompat.ResponsesClientToolMapping) bool { + return hasResponsesClientToolMapping(mapping) +} + func setGrokResponsesClientToolMapping(c *gin.Context, mapping apicompat.ResponsesClientToolMapping) { if c == nil { return @@ -97,7 +105,7 @@ func (b *grokResponsesClientToolStreamBody) Close() error { return sourceErr } -func newGrokResponsesClientToolStreamBody( +func newResponsesClientToolStreamBody( source io.ReadCloser, mapping apicompat.ResponsesClientToolMapping, maxLineSize int, @@ -108,6 +116,14 @@ func newGrokResponsesClientToolStreamBody( return body } +func newGrokResponsesClientToolStreamBody( + source io.ReadCloser, + mapping apicompat.ResponsesClientToolMapping, + maxLineSize int, +) io.ReadCloser { + return newResponsesClientToolStreamBody(source, mapping, maxLineSize) +} + func transformGrokResponsesClientToolStream( source io.ReadCloser, destination *io.PipeWriter, diff --git a/backend/internal/service/openai_ws_http_bridge.go b/backend/internal/service/openai_ws_http_bridge.go index 87abf1ef04f7..98eee9c7db0c 100644 --- a/backend/internal/service/openai_ws_http_bridge.go +++ b/backend/internal/service/openai_ws_http_bridge.go @@ -12,6 +12,7 @@ import ( "time" "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" "github.com/gin-gonic/gin" "github.com/tidwall/gjson" ) @@ -186,6 +187,13 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( if err != nil { return nil, fmt.Errorf("prepare http bridge body: %w", err) } + var clientToolMapping apicompat.ResponsesClientToolMapping + if account.Platform == PlatformOpenAI && account.Type == AccountTypeAPIKey { + body, clientToolMapping, err = adaptResponsesClientToolsForFunctionUpstream(body, "OpenAI WS HTTP bridge") + if err != nil { + return nil, fmt.Errorf("adapt OpenAI WS HTTP bridge client tools: %w", err) + } + } upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) var upstreamReq *http.Request @@ -329,11 +337,14 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( return result } - scanner := bufio.NewScanner(resp.Body) maxLineSize := defaultMaxLineSize if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 { maxLineSize = s.cfg.Gateway.MaxLineSize } + if hasResponsesClientToolMapping(clientToolMapping) { + resp.Body = newResponsesClientToolStreamBody(resp.Body, clientToolMapping, maxLineSize) + } + scanner := bufio.NewScanner(resp.Body) scanBuf := getSSEScannerBuf64K() scanner.Buffer(scanBuf[:0], maxLineSize) defer putSSEScannerBuf64K(scanBuf) diff --git a/backend/internal/service/openai_ws_http_bridge_test.go b/backend/internal/service/openai_ws_http_bridge_test.go index 1e352f012efb..887d495202e7 100644 --- a/backend/internal/service/openai_ws_http_bridge_test.go +++ b/backend/internal/service/openai_ws_http_bridge_test.go @@ -41,6 +41,81 @@ func TestPrepareOpenAIWSHTTPBridgeBodyStripsWSFields(t *testing.T) { require.Equal(t, "hi", gjson.GetBytes(body, "input").String()) } +func TestProxyOpenAIWSHTTPBridgeTurnAPIKeyAdaptsClientTools(t *testing.T) { + gin.SetMode(gin.TestMode) + + sse := strings.Join([]string{ + `data: {"type":"response.output_item.added","sequence_number":0,"output_index":0,"item":{"type":"function_call","id":"item_exec","call_id":"call_exec","name":"exec","status":"in_progress"}}`, + ``, + `data: {"type":"response.function_call_arguments.done","sequence_number":1,"output_index":0,"item_id":"item_exec","call_id":"call_exec","name":"exec","arguments":"{\"input\":\"pwd\"}"}`, + ``, + `data: {"type":"response.output_item.done","sequence_number":2,"output_index":0,"item":{"type":"function_call","id":"item_exec","call_id":"call_exec","name":"exec","arguments":"{\"input\":\"pwd\"}","status":"completed"}}`, + ``, + `data: {"type":"response.completed","sequence_number":3,"response":{"id":"resp_tools","status":"completed","output":[{"type":"function_call","id":"item_exec","call_id":"call_exec","name":"exec","arguments":"{\"input\":\"pwd\"}","status":"completed"}],"usage":{"input_tokens":1,"output_tokens":1}}}`, + ``, + }, "\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 := &OpenAIGatewayService{ + cfg: &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}}, + httpUpstream: upstream, + } + account := &Account{ID: 5659, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Concurrency: 1} + payload := []byte(`{ + "type":"response.create","model":"gpt-5","stream":true, + "tools":[{"type":"custom","name":"exec","description":"Run a command"}], + "input":[ + {"type":"custom_tool_call","id":"previous_item","call_id":"previous_call","name":"exec","input":"echo ready"}, + {"type":"custom_tool_call_output","call_id":"previous_call","output":"ready"} + ] + }`) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil) + var events [][]byte + + result, err := svc.proxyOpenAIWSHTTPBridgeTurn( + context.Background(), c, account, "test-token", payload, len(payload), + "gpt-5", "", "", "", "", 2, + func(message []byte) error { + events = append(events, append([]byte(nil), message...)) + return nil + }, + ) + + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, "function", gjson.GetBytes(upstream.lastBody, "tools.0.type").String()) + require.Equal(t, "function_call", gjson.GetBytes(upstream.lastBody, "input.0.type").String()) + require.JSONEq(t, `{"input":"echo ready"}`, gjson.GetBytes(upstream.lastBody, "input.0.arguments").String()) + require.False(t, gjson.GetBytes(upstream.lastBody, "input.0.input").Exists()) + require.Equal(t, "function_call_output", gjson.GetBytes(upstream.lastBody, "input.1.type").String()) + + var outputDone, completed []byte + for _, event := range events { + switch gjson.GetBytes(event, "type").String() { + case "response.output_item.done": + outputDone = event + case "response.completed": + completed = event + } + } + require.NotEmpty(t, outputDone) + require.Equal(t, "custom_tool_call", gjson.GetBytes(outputDone, "item.type").String()) + require.Equal(t, "pwd", gjson.GetBytes(outputDone, "item.input").String()) + require.False(t, gjson.GetBytes(outputDone, "item.arguments").Exists()) + require.NotEmpty(t, completed) + require.Equal(t, "custom_tool_call", gjson.GetBytes(completed, "response.output.0.type").String()) + require.Equal(t, "pwd", gjson.GetBytes(completed, "response.output.0.input").String()) + require.True(t, result.wsReplayInputExists) + require.Len(t, result.wsReplayInput, 1) + require.Equal(t, "custom_tool_call", gjson.GetBytes(result.wsReplayInput[0], "type").String()) + require.Equal(t, "pwd", gjson.GetBytes(result.wsReplayInput[0], "input").String()) +} + func TestOpenAIWSHTTPBridgeDecisionKeepsSmallFramesOnWS(t *testing.T) { svc := &OpenAIGatewayService{ cfg: &config.Config{