diff --git a/backend/internal/pkg/apicompat/responses_client_tools.go b/backend/internal/pkg/apicompat/responses_client_tools.go index c4b1e0018e..7c529fc8cb 100644 --- a/backend/internal/pkg/apicompat/responses_client_tools.go +++ b/backend/internal/pkg/apicompat/responses_client_tools.go @@ -130,6 +130,42 @@ func AdaptResponsesClientTools(req map[string]any) (ResponsesClientToolMapping, return adapter, changed, nil } +// AdaptResponsesClientToolsWithInheritedMapping lowers client-tool history on +// a follow-up request that omits the session-level tools declaration. An +// explicitly present tools field, including an empty or malformed value, +// always replaces the inherited mapping and is handled by the ordinary +// declaration-driven adapter. +func AdaptResponsesClientToolsWithInheritedMapping( + req map[string]any, + inherited ResponsesClientToolMapping, +) (ResponsesClientToolMapping, bool, error) { + if req == nil { + return ResponsesClientToolMapping{}, false, nil + } + if _, toolsPresent := req["tools"]; toolsPresent { + return AdaptResponsesClientTools(req) + } + if len(inherited.CustomTools) == 0 && !inherited.ToolSearch && len(inherited.NamespaceTools) == 0 { + return ResponsesClientToolMapping{}, false, nil + } + + changed := rewriteClientToolHistory(req["input"], &inherited) + if len(inherited.NamespaceTools) > 0 { + before := changed + rewriteNamespaceQualifiedCalls(req["input"], inherited.NamespaceTools) + // Namespace rewriting does not currently report whether it changed a + // value. A retained namespace mapping is only used for follow-up + // history, so conservatively rebuild the request when input exists. + if _, inputPresent := req["input"]; inputPresent && !before { + changed = true + } + } + if rewriteClientToolChoice(req, &inherited) { + changed = true + } + return inherited, changed, nil +} + func copyClientTool(tool map[string]any) map[string]any { copy := make(map[string]any, len(tool)) for key, value := range tool { @@ -155,10 +191,12 @@ func rewriteClientToolHistory(value any, adapter *ResponsesClientToolMapping) bo typed["type"] = "function_call" typed["arguments"] = customToolCallArguments(stringValue(typed["input"])) delete(typed, "input") + dropInvalidLoweredFunctionItemID(typed) changed = true } case "custom_tool_call_output": typed["type"] = "function_call_output" + dropInvalidLoweredFunctionItemID(typed) normalizeClientToolOutput(typed) changed = true case "tool_search_call": @@ -167,11 +205,13 @@ func rewriteClientToolHistory(value any, adapter *ResponsesClientToolMapping) bo typed["name"] = toolSearchProxyName typed["arguments"] = rawObjectString(typed["arguments"]) delete(typed, "execution") + dropInvalidLoweredFunctionItemID(typed) changed = true } case "tool_search_output": if adapter.ToolSearch { typed["type"] = "function_call_output" + dropInvalidLoweredFunctionItemID(typed) normalizeClientToolOutput(typed) changed = true } @@ -185,6 +225,17 @@ func rewriteClientToolHistory(value any, adapter *ResponsesClientToolMapping) bo return changed } +// dropInvalidLoweredFunctionItemID removes Codex client-only item IDs such as +// ctc_*, ctco_*, tsc_*, and tso_* after their item type is lowered to the +// function protocol. Function upstreams validate these IDs with the fc prefix; +// call_id, which is preserved separately, is the tool call/output pairing key. +func dropInvalidLoweredFunctionItemID(item map[string]any) { + id := strings.TrimSpace(stringValue(item["id"])) + if id != "" && !strings.HasPrefix(id, "fc") { + delete(item, "id") + } +} + func normalizeClientToolOutput(item map[string]any) { output, exists := item["output"] if !exists { diff --git a/backend/internal/pkg/apicompat/responses_client_tools_test.go b/backend/internal/pkg/apicompat/responses_client_tools_test.go index cc1ed467ba..0836e22d64 100644 --- a/backend/internal/pkg/apicompat/responses_client_tools_test.go +++ b/backend/internal/pkg/apicompat/responses_client_tools_test.go @@ -17,10 +17,10 @@ func TestAdaptResponsesClientTools_LowersDeclarationsHistoryChoiceAndNamespaces( }, "tool_choice": map[string]any{"type": "custom", "name": "exec"}, "input": []any{ - map[string]any{"type": "custom_tool_call", "call_id": "c1", "name": "exec", "input": "dir"}, - map[string]any{"type": "custom_tool_call_output", "call_id": "c1", "output": "ok"}, - map[string]any{"type": "tool_search_call", "call_id": "s1", "arguments": map[string]any{"query": "git"}}, - map[string]any{"type": "tool_search_output", "call_id": "s1", "output": map[string]any{"groups": []string{"git"}}}, + map[string]any{"type": "custom_tool_call", "id": "ctc_client", "call_id": "c1", "name": "exec", "input": "dir"}, + map[string]any{"type": "custom_tool_call_output", "id": "ctco_client", "call_id": "c1", "output": "ok"}, + map[string]any{"type": "tool_search_call", "id": "tsc_client", "call_id": "s1", "arguments": map[string]any{"query": "git"}}, + map[string]any{"type": "tool_search_output", "id": "tso_client", "call_id": "s1", "output": map[string]any{"groups": []string{"git"}}}, map[string]any{"type": "function_call", "call_id": "n1", "namespace": "team", "name": "send", "arguments": "{}"}, }, } @@ -48,15 +48,19 @@ func TestAdaptResponsesClientTools_LowersDeclarationsHistoryChoiceAndNamespaces( input := requireResponsesClientToolValue[[]any](t, req["input"]) customCall := requireResponsesClientToolValue[map[string]any](t, input[0]) require.Equal(t, "function_call", customCall["type"]) + require.NotContains(t, customCall, "id") require.JSONEq(t, `{"input":"dir"}`, requireResponsesClientToolValue[string](t, customCall["arguments"])) customOutput := requireResponsesClientToolValue[map[string]any](t, input[1]) require.Equal(t, "function_call_output", customOutput["type"]) + require.NotContains(t, customOutput, "id") searchCall := requireResponsesClientToolValue[map[string]any](t, input[2]) require.Equal(t, "function_call", searchCall["type"]) + require.NotContains(t, searchCall, "id") require.Equal(t, toolSearchProxyName, searchCall["name"]) require.JSONEq(t, `{"query":"git"}`, requireResponsesClientToolValue[string](t, searchCall["arguments"])) searchOutput := requireResponsesClientToolValue[map[string]any](t, input[3]) require.Equal(t, "function_call_output", searchOutput["type"]) + require.NotContains(t, searchOutput, "id") require.JSONEq(t, `{"groups":["git"]}`, requireResponsesClientToolValue[string](t, searchOutput["output"])) namespaceCall := requireResponsesClientToolValue[map[string]any](t, input[4]) require.Equal(t, "team__send", namespaceCall["name"]) @@ -81,6 +85,59 @@ func TestAdaptResponsesClientTools_RejectsAmbiguousNames(t *testing.T) { } } +func TestAdaptResponsesClientToolsWithInheritedMapping_LowersFollowupHistoryWithoutTools(t *testing.T) { + req := map[string]any{ + "input": []any{ + map[string]any{ + "type": "custom_tool_call", "name": "exec", + "call_id": "call_1", "input": "pwd", + }, + map[string]any{ + "type": "custom_tool_call_output", "call_id": "call_1", + "id": "ctco_client_output_1", + "output": []any{map[string]any{"type": "input_text", "text": "ok"}}, + }, + }, + } + inherited := ResponsesClientToolMapping{CustomTools: map[string]bool{"exec": true}} + + mapping, changed, err := AdaptResponsesClientToolsWithInheritedMapping(req, inherited) + + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, inherited, mapping) + items := requireResponsesClientToolValue[[]any](t, req["input"]) + call := requireResponsesClientToolValue[map[string]any](t, items[0]) + require.Equal(t, "function_call", call["type"]) + require.JSONEq(t, `{"input":"pwd"}`, requireResponsesClientToolValue[string](t, call["arguments"])) + require.NotContains(t, call, "input") + output := requireResponsesClientToolValue[map[string]any](t, items[1]) + require.Equal(t, "function_call_output", output["type"]) + require.NotContains(t, output, "id") + require.JSONEq(t, `[{"text":"ok","type":"input_text"}]`, requireResponsesClientToolValue[string](t, output["output"])) +} + +func TestAdaptResponsesClientToolsWithInheritedMapping_ExplicitToolsReplaceInheritedMapping(t *testing.T) { + req := map[string]any{ + "tools": []any{}, + "input": []any{map[string]any{ + "type": "custom_tool_call", "name": "exec", "input": "pwd", + }}, + } + + mapping, changed, err := AdaptResponsesClientToolsWithInheritedMapping( + req, + ResponsesClientToolMapping{CustomTools: map[string]bool{"exec": true}}, + ) + + require.NoError(t, err) + require.False(t, changed) + require.Empty(t, mapping) + items := requireResponsesClientToolValue[[]any](t, req["input"]) + call := requireResponsesClientToolValue[map[string]any](t, items[0]) + require.Equal(t, "custom_tool_call", call["type"]) +} + func TestRestoreResponsesClientToolPayload_RestoresClientAndNamespaceCalls(t *testing.T) { mapping := ResponsesClientToolMapping{ CustomTools: map[string]bool{"exec": true}, ToolSearch: true, diff --git a/backend/internal/service/openai_gateway_grok_tool_protocol.go b/backend/internal/service/openai_gateway_grok_tool_protocol.go index 11122172ae..c22282f8ca 100644 --- a/backend/internal/service/openai_gateway_grok_tool_protocol.go +++ b/backend/internal/service/openai_gateway_grok_tool_protocol.go @@ -16,6 +16,18 @@ import ( const grokResponsesClientToolMappingContextKey = "grok_responses_client_tool_mapping" func adaptResponsesClientToolsForFunctionUpstream(body []byte, upstream string) ([]byte, apicompat.ResponsesClientToolMapping, error) { + return adaptResponsesClientToolsForFunctionUpstreamWithMapping( + body, + upstream, + apicompat.ResponsesClientToolMapping{}, + ) +} + +func adaptResponsesClientToolsForFunctionUpstreamWithMapping( + body []byte, + upstream string, + inherited apicompat.ResponsesClientToolMapping, +) ([]byte, apicompat.ResponsesClientToolMapping, error) { decoder := json.NewDecoder(bytes.NewReader(body)) decoder.UseNumber() var requestBody map[string]any @@ -23,7 +35,7 @@ func adaptResponsesClientToolsForFunctionUpstream(body []byte, upstream string) return body, apicompat.ResponsesClientToolMapping{}, fmt.Errorf("decode %s Responses client tools: %w", upstream, err) } - mapping, changed, err := apicompat.AdaptResponsesClientTools(requestBody) + mapping, changed, err := apicompat.AdaptResponsesClientToolsWithInheritedMapping(requestBody, inherited) if err != nil { return body, apicompat.ResponsesClientToolMapping{}, err } diff --git a/backend/internal/service/openai_ws_http_bridge.go b/backend/internal/service/openai_ws_http_bridge.go index 1575318526..1ea3f2a935 100644 --- a/backend/internal/service/openai_ws_http_bridge.go +++ b/backend/internal/service/openai_ws_http_bridge.go @@ -15,6 +15,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" "github.com/gin-gonic/gin" "github.com/tidwall/gjson" + "github.com/tidwall/sjson" ) const ( @@ -23,6 +24,30 @@ const ( openAIWSHTTPBridgeErrorBodyLimitBytes = 64 * 1024 ) +const openAIWSHTTPBridgeToolStateContextKey = "openai_ws_http_bridge_tool_state" + +type openAIWSHTTPBridgeToolState struct { + ClientMapping apicompat.ResponsesClientToolMapping + LoweredTools json.RawMessage +} + +func openAIWSHTTPBridgeToolStateFromContext(c *gin.Context) (openAIWSHTTPBridgeToolState, bool) { + if c == nil { + return openAIWSHTTPBridgeToolState{}, false + } + value, ok := c.Get(openAIWSHTTPBridgeToolStateContextKey) + state, typed := value.(openAIWSHTTPBridgeToolState) + return state, ok && typed +} + +func setOpenAIWSHTTPBridgeToolState(c *gin.Context, state openAIWSHTTPBridgeToolState) { + if c == nil { + return + } + state.LoweredTools = append(json.RawMessage(nil), state.LoweredTools...) + c.Set(openAIWSHTTPBridgeToolStateContextKey, state) +} + // ResolveOpenAIWSClientFirstMessageTimeout returns the effective client ingress deadline. func ResolveOpenAIWSClientFirstMessageTimeout(cfg *config.Config) time.Duration { seconds := config.DefaultOpenAIWSClientFirstMessageTimeoutSeconds @@ -189,10 +214,29 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( } var clientToolMapping apicompat.ResponsesClientToolMapping if account.Platform == PlatformOpenAI && account.Type == AccountTypeAPIKey { - body, clientToolMapping, err = adaptResponsesClientToolsForFunctionUpstream(body, "OpenAI WS HTTP bridge") + inheritedState, _ := openAIWSHTTPBridgeToolStateFromContext(c) + toolsPresent := gjson.GetBytes(body, "tools").Exists() + body, clientToolMapping, err = adaptResponsesClientToolsForFunctionUpstreamWithMapping( + body, + "OpenAI WS HTTP bridge", + inheritedState.ClientMapping, + ) if err != nil { return nil, fmt.Errorf("adapt OpenAI WS HTTP bridge client tools: %w", err) } + loweredTools := inheritedState.LoweredTools + if toolsPresent { + loweredTools = json.RawMessage(gjson.GetBytes(body, "tools").Raw) + } else if len(loweredTools) > 0 { + body, err = sjson.SetRawBytes(body, "tools", loweredTools) + if err != nil { + return nil, fmt.Errorf("inherit OpenAI WS HTTP bridge tools: %w", err) + } + } + setOpenAIWSHTTPBridgeToolState(c, openAIWSHTTPBridgeToolState{ + ClientMapping: clientToolMapping, + LoweredTools: loweredTools, + }) } upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) diff --git a/backend/internal/service/openai_ws_http_bridge_test.go b/backend/internal/service/openai_ws_http_bridge_test.go index 2f0b78d8ba..b283c36ce6 100644 --- a/backend/internal/service/openai_ws_http_bridge_test.go +++ b/backend/internal/service/openai_ws_http_bridge_test.go @@ -170,6 +170,124 @@ func TestProxyOpenAIWSHTTPBridgeTurnAPIKeyRestoresClientToolsInResponseDone(t *t require.Equal(t, "custom_tool_call", gjson.GetBytes(result.wsReplayInput[0], "type").String()) } +func TestOpenAIWSHTTPBridgeAPIKeyReusesClientToolMappingWhenFollowupOmitsTools(t *testing.T) { + gin.SetMode(gin.TestMode) + + firstSSEBody := strings.Join([]string{ + `data: {"type":"response.completed","response":{"id":"resp_custom_first","model":"gpt-5.6-sol","output":[{"type":"function_call","id":"fc_custom_1","call_id":"call_custom_1","name":"exec","arguments":"{\"input\":\"pwd\"}"}],"usage":{"input_tokens":9,"output_tokens":1}}}`, + "", + }, "\n") + secondSSEBody := strings.Join([]string{ + `data: {"type":"response.completed","response":{"id":"resp_custom_second","model":"gpt-5.6-sol","output":[],"usage":{"input_tokens":1,"output_tokens":1}}}`, + "", + }, "\n") + upstream := &httpUpstreamRecorder{responses: []*http.Response{ + {StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(firstSSEBody))}, + {StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(secondSSEBody))}, + }} + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + cfg.Security.URLAllowlist.AllowInsecureHTTP = true + cfg.Gateway.OpenAIWS.Enabled = true + cfg.Gateway.OpenAIWS.APIKeyEnabled = true + cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true + cfg.Gateway.OpenAIWS.HTTPBridgeEnabled = true + cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes = 1 + cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1 + cfg.Gateway.OpenAIWS.QueueLimitPerConn = 8 + cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3 + cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3 + cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3 + + svc := &OpenAIGatewayService{ + cfg: cfg, httpUpstream: upstream, cache: &stubGatewayCache{}, + openaiWSResolver: NewOpenAIWSProtocolResolver(cfg), toolCorrector: NewCodexToolCorrector(), + } + account := &Account{ + ID: 9001, Name: "api-key-custom-followup", Platform: PlatformOpenAI, Type: AccountTypeAPIKey, + Credentials: map[string]any{"api_key": "sk-upstream"}, Extra: map[string]any{"responses_websockets_v2_enabled": true}, + Concurrency: 1, Status: StatusActive, Schedulable: true, + } + + errCh := make(chan error, 1) + wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := coderws.Accept(w, r, nil) + if err != nil { + errCh <- err + return + } + defer func() { _ = conn.CloseNow() }() + readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second) + _, firstMessage, err := conn.Read(readCtx) + cancelRead() + if err != nil { + errCh <- err + return + } + rec := httptest.NewRecorder() + ginCtx, _ := gin.CreateTestContext(rec) + ginCtx.Request = r.Clone(r.Context()) + errCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "sk-test", firstMessage, nil) + })) + defer wsServer.Close() + + dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second) + clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil) + cancelDial() + require.NoError(t, err) + defer func() { _ = clientConn.CloseNow() }() + + writeMessage := func(payload string) { + writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second) + defer cancelWrite() + require.NoError(t, clientConn.Write(writeCtx, coderws.MessageText, []byte(payload))) + } + readMessage := func() []byte { + readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second) + defer cancelRead() + messageType, event, readErr := clientConn.Read(readCtx) + require.NoError(t, readErr) + require.Equal(t, coderws.MessageText, messageType) + return event + } + + writeMessage(`{"type":"response.create","model":"gpt-5.6-sol","stream":true,"tools":[{"type":"custom","name":"exec"}],"input":"run pwd"}`) + firstEvent := readMessage() + require.Equal(t, "response.completed", gjson.GetBytes(firstEvent, "type").String()) + require.Equal(t, "custom_tool_call", gjson.GetBytes(firstEvent, "response.output.0.type").String()) + require.Equal(t, "pwd", gjson.GetBytes(firstEvent, "response.output.0.input").String()) + + writeMessage(`{"type":"response.create","model":"gpt-5.6-sol","stream":true,"previous_response_id":"resp_custom_first","input":[{"type":"custom_tool_call_output","id":"ctco_client_output_1","call_id":"call_custom_1","output":"ok"}]}`) + secondEvent := readMessage() + require.Equal(t, "response.completed", gjson.GetBytes(secondEvent, "type").String()) + + require.NoError(t, clientConn.Close(coderws.StatusNormalClosure, "done")) + select { + case proxyErr := <-errCh: + require.NoError(t, proxyErr) + case <-time.After(5 * time.Second): + t.Fatal("timed out waiting for websocket bridge proxy to finish") + } + + require.Len(t, upstream.bodies, 2) + firstTools := gjson.GetBytes(upstream.bodies[0], "tools").Array() + require.Len(t, firstTools, 1) + require.Equal(t, "function", firstTools[0].Get("type").String()) + secondTools := gjson.GetBytes(upstream.bodies[1], "tools").Array() + require.Len(t, secondTools, 1) + require.Equal(t, "function", secondTools[0].Get("type").String()) + require.Equal(t, "exec", secondTools[0].Get("name").String()) + secondInput := gjson.GetBytes(upstream.bodies[1], "input").Array() + require.Len(t, secondInput, 3) + require.Equal(t, "run pwd", secondInput[0].String()) + require.Equal(t, "function_call", secondInput[1].Get("type").String()) + require.Equal(t, "fc_custom_1", secondInput[1].Get("id").String()) + require.JSONEq(t, `{"input":"pwd"}`, secondInput[1].Get("arguments").String()) + require.False(t, secondInput[1].Get("input").Exists()) + require.Equal(t, "function_call_output", secondInput[2].Get("type").String()) + require.False(t, secondInput[2].Get("id").Exists()) +} + func TestOpenAIWSHTTPBridgeDecisionKeepsSmallFramesOnWS(t *testing.T) { svc := &OpenAIGatewayService{ cfg: &config.Config{