From 7de85ff7a2340e7534f71313f4b89c9dbe91a907 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 1 Aug 2026 00:49:00 +0000 Subject: [PATCH] fix(responses): guard malformed non-list response.output in WebSocket masking Greptile re-review on PR #35353 flagged that, unlike content and summary, response.output itself was iterated without an isinstance(..., list) check in both _unmask_response_event and _mask_response_completed. A truthy non-list output value raises TypeError, which is caught by the outer exception handler in backend_to_client and terminates the whole forwarding loop for that connection, not just the one malformed event. This predates this branch (same gap exists on litellm_internal_staging) but sits in functions this PR already touches, so fixing it here. Adds a regression test that fails with the guard removed and passes with it in place. --- litellm/responses/streaming_iterator.py | 10 ++++-- .../test_responses_websocket_all_providers.py | 31 +++++++++++++++++++ 2 files changed, 39 insertions(+), 2 deletions(-) diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index dedb4626ac0..6ecebd3d1b8 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -1684,7 +1684,10 @@ class ResponsesWebSocketStreaming: response_obj = evt_obj.get("response") if not isinstance(response_obj, dict): return response_str - for output_item in response_obj.get("output") or []: + output = response_obj.get("output") or [] + if not isinstance(output, list): + return response_str + for output_item in output: if not isinstance(output_item, dict): continue content = output_item.get("content") or [] @@ -1739,7 +1742,10 @@ class ResponsesWebSocketStreaming: response_obj = evt_obj.get("response") if not isinstance(response_obj, dict): continue - for output_item in response_obj.get("output") or []: + output = response_obj.get("output") or [] + if not isinstance(output, list): + continue + for output_item in output: if not isinstance(output_item, dict): continue arguments = output_item.get("arguments") diff --git a/tests/test_litellm/responses/test_responses_websocket_all_providers.py b/tests/test_litellm/responses/test_responses_websocket_all_providers.py index 6153f9adbcc..718a5c277cc 100644 --- a/tests/test_litellm/responses/test_responses_websocket_all_providers.py +++ b/tests/test_litellm/responses/test_responses_websocket_all_providers.py @@ -1122,6 +1122,37 @@ class TestNativeWebSocketGuardrails: assert handler._unmask_response_event(event) == event assert await handler._mask_response_completed(event) == event + @pytest.mark.asyncio + async def test_completed_event_with_non_list_output_passes_through(self): + """Regression test: a "response.output" that is a truthy non-iterable + value (malformed upstream payload) must be skipped, not raise -- an + unhandled exception here terminates the whole backend_to_client loop, + not just this one event.""" + from unittest.mock import MagicMock + + from litellm.responses.streaming_iterator import ResponsesWebSocketStreaming + + class Guardrail: + def get_presidio_settings_from_request_data(self, request_data): + return None + + def _unmask_pii_text(self, text, pii_tokens): + return text + + event = json.dumps({"type": "response.completed", "response": {"output": 42}}) + guardrail = Guardrail() + handler = ResponsesWebSocketStreaming( + websocket=MagicMock(), + backend_ws=MagicMock(), + logging_obj=MagicMock(), + request_data={"metadata": {"pii_tokens": {"": "secret"}}}, + guardrail_callbacks=[guardrail], + output_guardrail_callbacks=[guardrail], + ) + + assert handler._unmask_response_event(event) == event + assert await handler._mask_response_completed(event) == event + @pytest.mark.asyncio async def test_output_masking_suppresses_delta_without_calling_presidio(self): import json