diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 5e045c3e84f..eb427757a58 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,34 @@ def _restore_wrapped_ids_in_response_create(msg_obj: Mapping[str, object]) -> di return {**msg_obj, **top_fields, **restored_nested} +def _restore_wrapped_ids_in_response_inject(msg_obj: Mapping[str, object]) -> dict[str, object] | None: + response_id: Final = msg_obj.get("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 + ) + 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 _with_pending_inject( + pending: Mapping[str, tuple[str, ...]], upstream_id: str, client_id: str +) -> Mapping[str, tuple[str, ...]]: + return MappingProxyType({**pending, upstream_id: (*pending.get(upstream_id, ()), client_id)}) + + +def _without_oldest_pending_inject( + pending: Mapping[str, tuple[str, ...]], upstream_id: str +) -> Mapping[str, tuple[str, ...]]: + remaining: Final = pending.get(upstream_id, ())[1:] + others: Final = {key: value for key, value in pending.items() if key != upstream_id} + return MappingProxyType({**others, upstream_id: remaining} if remaining else others) + + def _wrap_output_item_encrypted_content( event_obj: Mapping[str, object], litellm_metadata: Mapping[str, object] ) -> dict[str, object] | None: @@ -1930,6 +1960,7 @@ class ResponsesWebSocketStreaming: # response.create frame to prevent deployment-substitution attacks. self.authorized_model: str | None = authorized_model self.request_defaults: ResponsesWebSocketRequestDefaults | None = request_defaults + self._pending_inject_client_ids: Mapping[str, tuple[str, ...]] = EMPTY_MAPPING def _should_store_event(self, event_obj: _MutableJsonObject) -> bool: return event_obj.get("type") in RESPONSES_WS_LOGGED_EVENT_TYPES @@ -2030,6 +2061,19 @@ class ResponsesWebSocketStreaming: response_cost: Final = self.logging_obj._response_cost_calculator(result=logging_result) or 0.0 # pyright: ignore[reportPrivateUsage] # as the HTTP streaming iterator does self.logging_obj.record_partial_usage_for_failure(usage, response_cost) + def _release_inject_client_id(self, upstream_id: str) -> str: + pending: Final = self._pending_inject_client_ids.get(upstream_id, ()) + if not pending: + model_info: Final = self.litellm_metadata.get("model_info") + model_id: Final = model_info.get("id") if _is_json_object(model_info) else None + return ResponsesAPIRequestUtils._build_responses_api_response_id( # pyright: ignore[reportPrivateUsage] # same wrap response.created gets + custom_llm_provider=self.custom_llm_provider, + model_id=model_id if isinstance(model_id, str) else None, + response_id=upstream_id, + ) + self._pending_inject_client_ids = _without_oldest_pending_inject(self._pending_inject_client_ids, upstream_id) + return pending[0] + def _wrap_response_event(self, response_str: str) -> str: try: event_obj: Final = _load_json_object(response_str) @@ -2043,6 +2087,12 @@ 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: + upstream_id: Final = event_obj.get("response_id") + if not isinstance(upstream_id, str): + return response_str + client_id: Final = self._release_inject_client_id(upstream_id) + return response_str if client_id == upstream_id else 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) @@ -2148,13 +2198,25 @@ 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`` and ``input`` item ids restored to the upstream ones. + 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_ids_in_response_inject(parsed) + inject_id: Final = parsed.get("response_id") + if isinstance(inject_id, str): + self._pending_inject_client_ids = _with_pending_inject( + self._pending_inject_client_ids, + ResponsesAPIRequestUtils.decode_previous_response_id_to_original_previous_response_id(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 6f346a25d9c..5b1349f4507 100644 --- a/tests/unit/responses/test_responses_websocket_all_providers.py +++ b/tests/unit/responses/test_responses_websocket_all_providers.py @@ -2940,6 +2940,177 @@ 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_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() + + 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"} + + @staticmethod + async def _inject_session(client_frames, upstream_results): + from unittest.mock import AsyncMock + + 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() + wrapped_id = json.loads(websocket.send_text.await_args_list[0][0][0])["response"]["id"] + websocket.receive_text = AsyncMock( + side_effect=[*(json.dumps(frame(wrapped_id)) for frame in client_frames), Exception("stop")] + ) + await handler.client_to_backend() + backend_ws.recv = AsyncMock(side_effect=[*(json.dumps(event) for event in upstream_results), Exception("stop")]) + await handler.backend_to_client() + echoed = [json.loads(call[0][0])["response_id"] for call in websocket.send_text.await_args_list[1:]] + return wrapped_id, echoed + + @pytest.mark.asyncio + async def test_overlapping_injects_for_one_response_each_echo_their_own_id(self): + import websockets.exceptions # noqa: F401 (lazy submodule must be importable) + + result = {"type": "response.inject.created", "sequence_number": 1, "response_id": "resp_upstream"} + + wrapped_id, echoed = await self._inject_session( + client_frames=[ + lambda wrapped: {"type": "response.inject", "response_id": wrapped, "input": []}, + lambda _: {"type": "response.inject", "response_id": "resp_upstream", "input": []}, + ], + upstream_results=[result, {**result, "sequence_number": 2}], + ) + + assert echoed == [wrapped_id, "resp_upstream"] + + @pytest.mark.parametrize("result_type", ["response.inject.created", "response.inject.failed"]) + @pytest.mark.asyncio + async def test_inject_result_releases_the_client_id_so_a_later_result_gets_the_default_wrap(self, result_type): + import websockets.exceptions # noqa: F401 (lazy submodule must be importable) + + result = {"type": result_type, "sequence_number": 1, "response_id": "resp_upstream"} + + wrapped_id, echoed = await self._inject_session( + client_frames=[lambda _: {"type": "response.inject", "response_id": "resp_upstream", "input": []}], + upstream_results=[result, {**result, "sequence_number": 2}], + ) + + assert echoed == ["resp_upstream", wrapped_id] + @pytest.mark.asyncio async def test_backend_to_client_wraps_ids_when_affinity_is_enabled(self): import asyncio