diff --git a/backend/internal/service/openai_tool_continuation.go b/backend/internal/service/openai_tool_continuation.go index a50721370..3b2356c97 100644 --- a/backend/internal/service/openai_tool_continuation.go +++ b/backend/internal/service/openai_tool_continuation.go @@ -235,16 +235,16 @@ func AnalyzeToolCallOutputContextCoverageBytes(body []byte) ToolCallOutputContex return coverage } input := parseRawJSONView(body).Get("input") - if !input.IsArray() { + if !input.IsArray() && !input.IsObject() { return coverage } missingCallID := false var outputCallIDs map[string]struct{} var contextIDs map[string]struct{} - input.ForEach(func(_, item gjson.Result) bool { + analyzeItem := func(item gjson.Result) { if !item.IsObject() { - return true + return } itemType := item.Get("type").String() switch { @@ -253,7 +253,7 @@ func AnalyzeToolCallOutputContextCoverageBytes(body []byte) ToolCallOutputContex callID := strings.TrimSpace(item.Get("call_id").String()) if callID == "" { missingCallID = true - return true + return } if outputCallIDs == nil { outputCallIDs = make(map[string]struct{}) @@ -262,7 +262,7 @@ func AnalyzeToolCallOutputContextCoverageBytes(body []byte) ToolCallOutputContex case isCodexToolCallContextItemType(itemType): callID := strings.TrimSpace(item.Get("call_id").String()) if callID == "" { - return true + return } if contextIDs == nil { contextIDs = make(map[string]struct{}) @@ -271,15 +271,22 @@ func AnalyzeToolCallOutputContextCoverageBytes(body []byte) ToolCallOutputContex case itemType == "item_reference": idValue := strings.TrimSpace(item.Get("id").String()) if idValue == "" { - return true + return } if contextIDs == nil { contextIDs = make(map[string]struct{}) } contextIDs[idValue] = struct{}{} } - return true - }) + } + if input.IsArray() { + input.ForEach(func(_, item gjson.Result) bool { + analyzeItem(item) + return true + }) + } else { + analyzeItem(input) + } if !coverage.HasFunctionCallOutput || missingCallID { return coverage diff --git a/backend/internal/service/openai_tool_continuation_test.go b/backend/internal/service/openai_tool_continuation_test.go index 569d89eff..460a659c2 100644 --- a/backend/internal/service/openai_tool_continuation_test.go +++ b/backend/internal/service/openai_tool_continuation_test.go @@ -206,6 +206,14 @@ func TestAnalyzeToolCallOutputContextCoverageBytes(t *testing.T) { hasOutput: false, coversAllIDs: false, }, + { + name: "object_tool_output_requires_context_replay", + body: map[string]any{"input": map[string]any{ + "type": "custom_tool_call_output", "call_id": "call_a", + }}, + hasOutput: true, + coversAllIDs: false, + }, { name: "all_outputs_covered_by_context", body: map[string]any{"input": []any{ diff --git a/backend/internal/service/openai_ws_forwarder_ingress.go b/backend/internal/service/openai_ws_forwarder_ingress.go index 4b0fe5a01..f02ad2513 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress.go +++ b/backend/internal/service/openai_ws_forwarder_ingress.go @@ -511,7 +511,9 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( } bridgePayloadRaw := currentBridgePayload.payloadRaw bridgePayloadBytes := currentBridgePayload.payloadBytes - needsBridgeReplay := currentBridgePayload.previousResponseID != "" || openAIWSRawPayloadHasToolCallOutput(currentBridgePayload.payloadRaw) + toolOutputCoverage := AnalyzeToolCallOutputContextCoverageBytes(currentBridgePayload.payloadRaw) + needsBridgeReplay := currentBridgePayload.previousResponseID != "" || + (toolOutputCoverage.HasFunctionCallOutput && !toolOutputCoverage.ContextCoversAllCallIDs) turnReplayInput, turnReplayInputExists, replayInputErr := buildOpenAIWSReplayInputSequence( bridgeReplayInput, bridgeReplayInputExists, diff --git a/backend/internal/service/openai_ws_forwarder_ingress_test.go b/backend/internal/service/openai_ws_forwarder_ingress_test.go index a13248e39..a7bdf31b7 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress_test.go +++ b/backend/internal/service/openai_ws_forwarder_ingress_test.go @@ -809,6 +809,35 @@ func TestBuildOpenAIWSReplayInputSequence(t *testing.T) { require.Equal(t, "new", gjson.GetBytes(items[0], "text").String()) }) + t.Run("no_previous_response_id_custom_tool_history_does_not_accumulate", func(t *testing.T) { + previousFull := []json.RawMessage{ + json.RawMessage(`{"type":"input_text","text":"stale"}`), + json.RawMessage(`{"type":"custom_tool_call","id":"stale_item","call_id":"stale_call","name":"exec","input":"stale"}`), + } + currentPayload := []byte(`{"input":[ + {"type":"custom_tool_call","id":"item_1","call_id":"call_1","name":"exec","input":"pwd"}, + {"type":"custom_tool_call_output","call_id":"call_1","output":"/tmp"}, + {"type":"input_text","text":"continue"} + ]}`) + + for range 3 { + items, exists, err := buildOpenAIWSReplayInputSequence( + previousFull, + true, + currentPayload, + false, + ) + require.NoError(t, err) + require.True(t, exists) + require.Len(t, items, 3) + require.Equal(t, "custom_tool_call", gjson.GetBytes(items[0], "type").String()) + require.Equal(t, "call_1", gjson.GetBytes(items[0], "call_id").String()) + require.Equal(t, "custom_tool_call_output", gjson.GetBytes(items[1], "type").String()) + require.Equal(t, "call_1", gjson.GetBytes(items[1], "call_id").String()) + previousFull = append(items, json.RawMessage(`{"type":"custom_tool_call","id":"replayed_item","call_id":"replayed_call","name":"exec","input":"ignored"}`)) + } + }) + t.Run("previous_response_id_delta_append", func(t *testing.T) { items, exists, err := buildOpenAIWSReplayInputSequence( lastFull, diff --git a/backend/internal/service/openai_ws_http_bridge_test.go b/backend/internal/service/openai_ws_http_bridge_test.go index 49b4c7073..74838026c 100644 --- a/backend/internal/service/openai_ws_http_bridge_test.go +++ b/backend/internal/service/openai_ws_http_bridge_test.go @@ -413,6 +413,191 @@ func TestOpenAIWSHTTPBridgeAPIKeyReusesClientToolMappingWhenFollowupOmitsTools(t require.False(t, secondInput[2].Get("id").Exists()) } +func TestOpenAIWSHTTPBridgeFullCustomToolHistoryWithoutPreviousResponseIDDoesNotReplay(t *testing.T) { + gin.SetMode(gin.TestMode) + + completed := func(responseID string, output string) string { + return "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"" + responseID + "\",\"model\":\"gpt-5.1\",\"output\":" + output + ",\"usage\":{\"input_tokens\":1,\"output_tokens\":1}}}\n\n" + } + upstream := &httpUpstreamRecorder{responses: []*http.Response{ + {StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(completed("resp_1", `[{"type":"custom_tool_call","id":"item_1","call_id":"call_1","name":"exec","input":"pwd"}]`)))}, + {StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(completed("resp_2", `[]`)))}, + {StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(completed("resp_3", `[]`)))}, + }} + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + cfg.Security.URLAllowlist.AllowInsecureHTTP = true + cfg.Gateway.OpenAIWS.Enabled = true + cfg.Gateway.OpenAIWS.OAuthEnabled = true + cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true + cfg.Gateway.OpenAIWS.HTTPBridgeEnabled = true + cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes = 1 + 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: 9002, Name: "oauth-full-context", Platform: PlatformOpenAI, Type: AccountTypeOAuth, + Credentials: map[string]any{"access_token": "test-token"}, 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, "test-token", 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() }() + + writeAndRead := func(payload string) { + writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second) + require.NoError(t, clientConn.Write(writeCtx, coderws.MessageText, []byte(payload))) + cancelWrite() + readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second) + _, event, readErr := clientConn.Read(readCtx) + cancelRead() + require.NoError(t, readErr) + require.Equal(t, "response.completed", gjson.GetBytes(event, "type").String()) + } + + writeAndRead(`{"type":"response.create","model":"gpt-5.1","input":"run pwd"}`) + fullContext := `{"type":"response.create","model":"gpt-5.1","input":[{"type":"custom_tool_call","id":"item_1","call_id":"call_1","name":"exec","input":"pwd"},{"type":"custom_tool_call_output","call_id":"call_1","output":"/tmp"},{"role":"user","content":"continue"}]}` + writeAndRead(fullContext) + writeAndRead(fullContext) + + 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, 3) + for _, body := range upstream.bodies[1:] { + input := gjson.GetBytes(body, "input").Array() + require.Len(t, input, 3) + require.Equal(t, "custom_tool_call", input[0].Get("type").String()) + require.Equal(t, "call_1", input[0].Get("call_id").String()) + require.Equal(t, "custom_tool_call_output", input[1].Get("type").String()) + require.Equal(t, "call_1", input[1].Get("call_id").String()) + } +} + +func TestOpenAIWSHTTPBridgeObjectToolOutputWithoutPreviousResponseIDReplaysMatchingCall(t *testing.T) { + gin.SetMode(gin.TestMode) + + completed := func(responseID string, output string) string { + return "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"" + responseID + "\",\"model\":\"gpt-5.1\",\"output\":" + output + ",\"usage\":{\"input_tokens\":1,\"output_tokens\":1}}}\n\n" + } + upstream := &httpUpstreamRecorder{responses: []*http.Response{ + {StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(completed("resp_1", `[{"type":"custom_tool_call","id":"item_1","call_id":"call_1","name":"exec","input":"pwd"}]`)))}, + {StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(completed("resp_2", `[]`)))}, + }} + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + cfg.Security.URLAllowlist.AllowInsecureHTTP = true + cfg.Gateway.OpenAIWS.Enabled = true + cfg.Gateway.OpenAIWS.OAuthEnabled = true + cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true + cfg.Gateway.OpenAIWS.HTTPBridgeEnabled = true + cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes = 1 + 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: 9003, Name: "oauth-output-only", Platform: PlatformOpenAI, Type: AccountTypeOAuth, + Credentials: map[string]any{"access_token": "test-token"}, 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, "test-token", 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() }() + + writeAndRead := func(payload string) { + writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second) + require.NoError(t, clientConn.Write(writeCtx, coderws.MessageText, []byte(payload))) + cancelWrite() + readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second) + _, event, readErr := clientConn.Read(readCtx) + cancelRead() + require.NoError(t, readErr) + require.Equal(t, "response.completed", gjson.GetBytes(event, "type").String()) + } + + writeAndRead(`{"type":"response.create","model":"gpt-5.1","input":"run pwd"}`) + writeAndRead(`{"type":"response.create","model":"gpt-5.1","input":{"type":"custom_tool_call_output","call_id":"call_1","output":"/tmp"}}`) + + 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) + secondInput := gjson.GetBytes(upstream.bodies[1], "input").Array() + require.Len(t, secondInput, 3) + require.Equal(t, "custom_tool_call", secondInput[1].Get("type").String()) + require.Equal(t, "call_1", secondInput[1].Get("call_id").String()) + require.Equal(t, "custom_tool_call_output", secondInput[2].Get("type").String()) + require.Equal(t, "call_1", secondInput[2].Get("call_id").String()) +} + func TestOpenAIWSHTTPBridgeDecisionKeepsSmallFramesOnWS(t *testing.T) { svc := &OpenAIGatewayService{ cfg: &config.Config{