mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
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.
This commit is contained in:
parent
f4b77e0458
commit
7de85ff7a2
2 changed files with 39 additions and 2 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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": {"<TOKEN_1>": "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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue