diff --git a/backend/internal/service/openai_tool_continuation.go b/backend/internal/service/openai_tool_continuation.go index a507213701..3b2356c971 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 569d89eff0..460a659c2e 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 19de862623..b4a8752347 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress.go +++ b/backend/internal/service/openai_ws_forwarder_ingress.go @@ -565,7 +565,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 a13248e393..1f00d0dce6 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, @@ -823,6 +852,91 @@ func TestBuildOpenAIWSReplayInputSequence(t *testing.T) { require.Equal(t, "world", gjson.GetBytes(items[1], "text").String()) }) + t.Run("previous_response_id_filters_orphan_historical_custom_tool_call", func(t *testing.T) { + previousFull := []json.RawMessage{ + json.RawMessage(`{"type":"input_text","text":"hello"}`), + json.RawMessage(`{"type":"custom_tool_call","id":"item_orphan","call_id":"call_orphan","name":"exec","input":"pwd"}`), + } + items, exists, err := buildOpenAIWSReplayInputSequence( + previousFull, + true, + []byte(`{"previous_response_id":"resp_1","input":[{"role":"user","content":"continue"}]}`), + true, + ) + require.NoError(t, err) + require.True(t, exists) + require.Len(t, items, 2) + require.Equal(t, "hello", gjson.GetBytes(items[0], "text").String()) + require.Equal(t, "user", gjson.GetBytes(items[1], "role").String()) + }) + + t.Run("previous_response_id_preserves_paired_historical_function_call", func(t *testing.T) { + previousFull := []json.RawMessage{ + json.RawMessage(`{"type":"function_call","id":"item_1","call_id":"call_1","name":"lookup","arguments":"{}"}`), + json.RawMessage(`{"type":"function_call_output","call_id":"call_1","output":"ok"}`), + } + items, exists, err := buildOpenAIWSReplayInputSequence( + previousFull, + true, + []byte(`{"previous_response_id":"resp_1","input":[{"role":"user","content":"continue"}]}`), + true, + ) + require.NoError(t, err) + require.True(t, exists) + require.Len(t, items, 3) + require.Equal(t, "function_call", gjson.GetBytes(items[0], "type").String()) + require.Equal(t, "function_call_output", gjson.GetBytes(items[1], "type").String()) + }) + + t.Run("previous_response_id_preserves_paired_historical_custom_tool_call", func(t *testing.T) { + previousFull := []json.RawMessage{ + json.RawMessage(`{"type":"custom_tool_call","id":"item_1","call_id":"call_1","name":"exec","input":"pwd"}`), + json.RawMessage(`{"type":"custom_tool_call_output","call_id":"call_1","output":"/tmp"}`), + } + items, exists, err := buildOpenAIWSReplayInputSequence( + previousFull, + true, + []byte(`{"previous_response_id":"resp_1","input":[{"role":"user","content":"continue"}]}`), + true, + ) + 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, "custom_tool_call_output", gjson.GetBytes(items[1], "type").String()) + }) + + t.Run("item_reference_does_not_complete_historical_call", func(t *testing.T) { + previousFull := []json.RawMessage{ + json.RawMessage(`{"type":"custom_tool_call","id":"item_1","call_id":"call_1","name":"exec","input":"pwd"}`), + } + items, exists, err := buildOpenAIWSReplayInputSequence( + previousFull, + true, + []byte(`{"previous_response_id":"resp_1","input":[{"type":"item_reference","id":"call_1"},{"role":"user","content":"continue"}]}`), + true, + ) + require.NoError(t, err) + require.True(t, exists) + require.Len(t, items, 2) + require.Equal(t, "item_reference", gjson.GetBytes(items[0], "type").String()) + require.Equal(t, "user", gjson.GetBytes(items[1], "role").String()) + }) + + t.Run("previous_response_id_preserves_current_orphan_custom_tool_call", func(t *testing.T) { + items, exists, err := buildOpenAIWSReplayInputSequence( + lastFull, + true, + []byte(`{"previous_response_id":"resp_1","input":[{"type":"custom_tool_call","id":"item_live","call_id":"call_live","name":"exec","input":"pwd"}]}`), + true, + ) + require.NoError(t, err) + require.True(t, exists) + require.Len(t, items, 2) + require.Equal(t, "custom_tool_call", gjson.GetBytes(items[1], "type").String()) + require.Equal(t, "call_live", gjson.GetBytes(items[1], "call_id").String()) + }) + t.Run("previous_response_id_full_input_replace", func(t *testing.T) { items, exists, err := buildOpenAIWSReplayInputSequence( lastFull, diff --git a/backend/internal/service/openai_ws_forwarder_payload.go b/backend/internal/service/openai_ws_forwarder_payload.go index aa51a1846d..89d0dfc18c 100644 --- a/backend/internal/service/openai_ws_forwarder_payload.go +++ b/backend/internal/service/openai_ws_forwarder_payload.go @@ -610,6 +610,40 @@ func openAIWSRawItemsHaveToolCallContextForOutputs(items []json.RawMessage) bool return true } +func sanitizeOpenAIWSHistoricalReplayToolCalls( + previousItems []json.RawMessage, + currentItems []json.RawMessage, +) []json.RawMessage { + if len(previousItems) == 0 { + return cloneOpenAIWSRawMessages(previousItems) + } + outputCallIDs := make(map[string]struct{}) + collectOutputCallIDs := func(items []json.RawMessage) { + for _, item := range items { + if !isCodexToolCallOutputItemType(gjson.GetBytes(item, "type").String()) { + continue + } + if callID := strings.TrimSpace(gjson.GetBytes(item, "call_id").String()); callID != "" { + outputCallIDs[callID] = struct{}{} + } + } + } + collectOutputCallIDs(previousItems) + collectOutputCallIDs(currentItems) + + sanitized := make([]json.RawMessage, 0, len(previousItems)) + for _, item := range previousItems { + if isCodexToolCallContextItemType(gjson.GetBytes(item, "type").String()) { + callID := strings.TrimSpace(gjson.GetBytes(item, "call_id").String()) + if _, paired := outputCallIDs[callID]; !paired { + continue + } + } + sanitized = append(sanitized, append(json.RawMessage(nil), item...)) + } + return sanitized +} + func openAIWSRawPayloadHasToolCallOutput(payload []byte) bool { if len(payload) == 0 { return false @@ -648,6 +682,7 @@ func buildOpenAIWSReplayInputSequence( if !previousFullInputExists { return cloneOpenAIWSRawMessages(currentItems), currentExists, nil } + previousFullInput = sanitizeOpenAIWSHistoricalReplayToolCalls(previousFullInput, currentItems) if !currentExists || len(currentItems) == 0 { return cloneOpenAIWSRawMessages(previousFullInput), true, nil } diff --git a/backend/internal/service/openai_ws_http_bridge_test.go b/backend/internal/service/openai_ws_http_bridge_test.go index 4aec9d7d8a..7824591e07 100644 --- a/backend/internal/service/openai_ws_http_bridge_test.go +++ b/backend/internal/service/openai_ws_http_bridge_test.go @@ -416,6 +416,197 @@ 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", `[]`)))}, + {StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(completed("resp_4", `[]`)))}, + }} + 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"}`) + writeAndRead(`{"type":"response.create","model":"gpt-5.1","previous_response_id":"resp_1","input":[{"role":"user","content":"continue without tool output"}]}`) + 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, 4) + orphanInput := gjson.GetBytes(upstream.bodies[1], "input").Array() + require.Len(t, orphanInput, 2) + require.Equal(t, "run pwd", orphanInput[0].String()) + require.Equal(t, "user", orphanInput[1].Get("role").String()) + for _, body := range upstream.bodies[2:] { + 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{