diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index a8b62d2607e..dcc4f254905 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -1868,28 +1868,18 @@ 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: +def _restore_wrapped_ids_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 + original_id: Final = ( + ResponsesAPIRequestUtils.decode_previous_response_id_to_original_previous_response_id(response_id) + if isinstance(response_id, str) + else 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} + restored: Final = { + **_restored_container_fields(msg_obj), + **({"response_id": original_id} if original_id != response_id else EMPTY_MAPPING), + } + return {**msg_obj, **restored} if restored else None def _wrap_output_item_encrypted_content( @@ -1956,6 +1946,8 @@ class ResponsesWebSocketStreaming: # response.create frame to prevent deployment-substitution attacks. self.authorized_model: str | None = authorized_model self.request_defaults: ResponsesWebSocketRequestDefaults | None = request_defaults + # Upstream response id -> the id the client sent in response.inject. + self._inject_client_response_ids: dict[str, str] = {} def _should_store_event(self, event_obj: _MutableJsonObject) -> bool: return event_obj.get("type") in RESPONSES_WS_LOGGED_EVENT_TYPES @@ -2070,10 +2062,13 @@ class ResponsesWebSocketStreaming: ) 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 + upstream_id: Final = event_obj.get("response_id") + client_id: Final = ( + self._inject_client_response_ids.get(upstream_id) if isinstance(upstream_id, str) else None ) - return response_str if wrapped_inject is None else json.dumps(wrapped_inject) + if client_id is None or client_id == upstream_id: + return response_str + return json.dumps({**event_obj, "response_id": client_id}) 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) @@ -2180,8 +2175,8 @@ class ResponsesWebSocketStreaming: ``self.request_data["metadata"]`` for later unmasking. A ``response.inject`` message only gets its LiteLLM-wrapped - ``response_id`` restored to the upstream id. Other messages are - returned unchanged. + ``response_id`` and ``input`` item ids restored to the upstream ones. + Other messages are returned unchanged. """ try: parsed: Final = _load_json_object(message) @@ -2189,7 +2184,14 @@ class ResponsesWebSocketStreaming: return message if parsed.get("type") == "response.inject": - restored_inject: Final = _restore_wrapped_id_in_response_inject(parsed) + restored_inject: Final = _restore_wrapped_ids_in_response_inject(parsed) + inject_id: Final = parsed.get("response_id") + if isinstance(inject_id, str): + # Remember the id as the client sent it, wrapped or raw, so the result echoes it back. + upstream_inject_id: Final = ( + ResponsesAPIRequestUtils.decode_previous_response_id_to_original_previous_response_id(inject_id) + ) + self._inject_client_response_ids[upstream_inject_id] = inject_id 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 d7dbfe7ff4d..bf5d8bf7b5d 100644 --- a/tests/unit/responses/test_responses_websocket_all_providers.py +++ b/tests/unit/responses/test_responses_websocket_all_providers.py @@ -2990,19 +2990,62 @@ class TestNativeWebSocketEncryptedContentAffinity: 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): + async def test_response_inject_with_an_unwrapped_id_round_trips_the_raw_id(self): from unittest.mock import AsyncMock + import websockets.exceptions # noqa: F401 (lazy submodule must be importable) + frame = json.dumps({"type": "response.inject", "response_id": "resp_raw", "input": []}) backend_ws = MagicMock() backend_ws.send = AsyncMock() + backend_ws.recv = AsyncMock( + side_effect=[ + json.dumps({"type": "response.inject.created", "sequence_number": 1, "response_id": "resp_raw"}), + Exception("stop"), + ] + ) websocket = MagicMock() + websocket.send_text = AsyncMock() websocket.receive_text = AsyncMock(side_effect=[frame, 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.client_to_backend() + await handler.backend_to_client() + + assert backend_ws.send.await_args_list[0][0][0] == frame + result = json.loads(websocket.send_text.await_args_list[0][0][0]) + assert result == {"type": "response.inject.created", "sequence_number": 1, "response_id": "resp_raw"} + + @pytest.mark.asyncio + async def test_response_inject_restores_wrapped_input_items(self): + from unittest.mock import AsyncMock + + frame = { + "type": "response.inject", + "response_id": "resp_raw", + "input": [_wrapped_reasoning_item(), {"type": "message", "role": "user", "content": "hi"}], + } + backend_ws = MagicMock() + backend_ws.send = AsyncMock() + websocket = MagicMock() + websocket.receive_text = AsyncMock(side_effect=[json.dumps(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 + sent = json.loads(backend_ws.send.await_args_list[0][0][0]) + assert sent["response_id"] == "resp_raw" + assert sent["input"][0]["id"] == "rs_orig" + assert sent["input"][0]["encrypted_content"] == "gAAAA-blob" + assert sent["input"][1] == {"type": "message", "role": "user", "content": "hi"} @pytest.mark.asyncio async def test_backend_to_client_wraps_ids_when_affinity_is_enabled(self):