From 5d01530c4339e20771963f4b8f0a9bcf6dbbf524 Mon Sep 17 00:00:00 2001 From: Darsh Joshi Date: Fri, 2 Oct 2026 20:25:31 -0400 Subject: [PATCH] fix(responses): restore the upstream id on WebSocket response.inject frames The native Responses WebSocket relay wraps the upstream response id in response.created, but forwarded response.inject frames without decoding response_id, so the upstream got an id it never issued and answered response_not_found. Decode the wrapped id before forwarding and wrap the response_id echoed in response.inject.created / response.inject.failed so the client sees the same id it used. Fixes #44242 --- litellm/responses/streaming_iterator.py | 38 ++++++++++- .../test_responses_websocket_all_providers.py | 64 +++++++++++++++++++ 2 files changed, 101 insertions(+), 1 deletion(-) diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 5e045c3e84f..a8b62d2607e 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -1828,6 +1828,8 @@ _RESPONSES_WS_FAILURE_EVENT_TYPES: Final = frozenset({"error", "response.failed" _RESPONSES_WS_OUTPUT_ITEM_EVENT_TYPES: Final = frozenset({"response.output_item.added", "response.output_item.done"}) +_RESPONSES_WS_INJECT_RESULT_EVENT_TYPES: Final = frozenset({"response.inject.created", "response.inject.failed"}) + def _ws_event_error(event: Mapping[str, object]) -> object: if event.get("type") == "error": @@ -1866,6 +1868,30 @@ def _restore_wrapped_ids_in_response_create(msg_obj: Mapping[str, object]) -> di return {**msg_obj, **top_fields, **restored_nested} +def _restore_wrapped_id_in_response_inject(msg_obj: Mapping[str, object]) -> dict[str, object] | None: + response_id: Final = msg_obj.get("response_id") + if not isinstance(response_id, str): + return None + original_id: Final = ResponsesAPIRequestUtils.decode_previous_response_id_to_original_previous_response_id( + response_id + ) + return None if original_id == response_id else {**msg_obj, "response_id": original_id} + + +def _wrap_inject_result_response_id( + event_obj: Mapping[str, object], custom_llm_provider: str | None, litellm_metadata: Mapping[str, object] +) -> dict[str, object] | None: + response_id: Final = event_obj.get("response_id") + if not isinstance(response_id, str) or ResponsesAPIRequestUtils._is_litellm_encoded_response_id(response_id): # pyright: ignore[reportPrivateUsage] # same check the HTTP response id wrap applies + return None + wrapped_id: Final = ResponsesAPIRequestUtils._build_responses_api_response_id( # pyright: ignore[reportPrivateUsage] # same wrap the HTTP response id path applies + custom_llm_provider=custom_llm_provider, + model_id=_model_id_from_metadata(litellm_metadata), + response_id=response_id, + ) + return {**event_obj, "response_id": wrapped_id} + + def _wrap_output_item_encrypted_content( event_obj: Mapping[str, object], litellm_metadata: Mapping[str, object] ) -> dict[str, object] | None: @@ -2043,6 +2069,11 @@ class ResponsesWebSocketStreaming: litellm_metadata=self.litellm_metadata, ) return json.dumps({**event_obj, "response": wrapped_response}) + if event_obj.get("type") in _RESPONSES_WS_INJECT_RESULT_EVENT_TYPES: + wrapped_inject: Final = _wrap_inject_result_response_id( + event_obj, self.custom_llm_provider, self.litellm_metadata + ) + return response_str if wrapped_inject is None else json.dumps(wrapped_inject) if event_obj.get("type") not in _RESPONSES_WS_OUTPUT_ITEM_EVENT_TYPES: return response_str wrapped_event: Final = _wrap_output_item_encrypted_content(event_obj, self.litellm_metadata) @@ -2148,13 +2179,18 @@ class ResponsesWebSocketStreaming: on every text block, and stores the resulting ``pii_tokens`` map in ``self.request_data["metadata"]`` for later unmasking. - Non-``response.create`` messages are returned unchanged. + A ``response.inject`` message only gets its LiteLLM-wrapped + ``response_id`` restored to the upstream id. Other messages are + returned unchanged. """ try: parsed: Final = _load_json_object(message) except (json.JSONDecodeError, TypeError): return message + if parsed.get("type") == "response.inject": + restored_inject: Final = _restore_wrapped_id_in_response_inject(parsed) + return message if restored_inject is None else json.dumps(restored_inject) if parsed.get("type") != "response.create": return message diff --git a/tests/unit/responses/test_responses_websocket_all_providers.py b/tests/unit/responses/test_responses_websocket_all_providers.py index 6f346a25d9c..d7dbfe7ff4d 100644 --- a/tests/unit/responses/test_responses_websocket_all_providers.py +++ b/tests/unit/responses/test_responses_websocket_all_providers.py @@ -2940,6 +2940,70 @@ class TestNativeWebSocketEncryptedContentAffinity: assert backend_ws.send.await_args_list[0][0][0] == frame + @pytest.mark.asyncio + async def test_response_inject_targets_the_upstream_id_and_its_result_echoes_the_client_id(self): + from unittest.mock import AsyncMock + + import websockets.exceptions # noqa: F401 (lazy submodule must be importable) + + inject_input = [{"type": "function_call_output", "call_id": "call_1", "output": "sunny"}] + websocket = MagicMock() + websocket.send_text = AsyncMock() + backend_ws = MagicMock() + backend_ws.send = AsyncMock() + backend_ws.recv = AsyncMock( + side_effect=[ + json.dumps({"type": "response.created", "response": {"id": "resp_upstream", "output": []}}), + Exception("stop"), + ] + ) + logging_obj = MagicMock() + logging_obj.dispatch_success_handlers = AsyncMock() + handler = _make_streaming( + websocket=websocket, + backend_ws=backend_ws, + logging_obj=logging_obj, + request_data={"litellm_metadata": {"model_info": {"id": "dep-1"}}}, + custom_llm_provider="openai", + ) + await handler.backend_to_client() + client_id = json.loads(websocket.send_text.await_args_list[0][0][0])["response"]["id"] + websocket.receive_text = AsyncMock( + side_effect=[ + json.dumps({"type": "response.inject", "response_id": client_id, "input": inject_input}), + Exception("stop"), + ] + ) + + await handler.client_to_backend() + + sent = json.loads(backend_ws.send.await_args_list[0][0][0]) + assert sent == {"type": "response.inject", "response_id": "resp_upstream", "input": inject_input} + backend_ws.recv = AsyncMock( + side_effect=[ + json.dumps({"type": "response.inject.created", "sequence_number": 1, "response_id": "resp_upstream"}), + Exception("stop"), + ] + ) + await handler.backend_to_client() + result = json.loads(websocket.send_text.await_args_list[1][0][0]) + assert result == {"type": "response.inject.created", "sequence_number": 1, "response_id": client_id} + + @pytest.mark.asyncio + async def test_response_inject_with_an_unwrapped_id_is_forwarded_untouched(self): + from unittest.mock import AsyncMock + + frame = json.dumps({"type": "response.inject", "response_id": "resp_raw", "input": []}) + backend_ws = MagicMock() + backend_ws.send = AsyncMock() + websocket = MagicMock() + websocket.receive_text = AsyncMock(side_effect=[frame, Exception("stop")]) + handler = _make_streaming(websocket=websocket, backend_ws=backend_ws, request_data={}) + + await handler.client_to_backend() + + assert backend_ws.send.await_args_list[0][0][0] == frame + @pytest.mark.asyncio async def test_backend_to_client_wraps_ids_when_affinity_is_enabled(self): import asyncio