From 58b028ca4c5b771fa9fa6698a4a11f5da0ee580d Mon Sep 17 00:00:00 2001 From: Darsh Joshi Date: Sat, 3 Oct 2026 12:14:13 -0400 Subject: [PATCH] fix(responses): track the client's inject id per in-flight inject The per-connection map from upstream response id to the client's inject id kept one entry per response, so a wrapped and a raw inject for the same response overwrote each other and the first result echoed the wrong form. Entries also stayed until the socket closed. Queue the client's id per upstream response id and release the oldest one when its response.inject.created or response.inject.failed arrives. The inject result frames carry no per-inject identifier, so results for one response are matched in send order. A result with no pending inject gets the same wrap response.created gets. --- litellm/responses/streaming_iterator.py | 46 +++++++++---- .../test_responses_websocket_all_providers.py | 64 +++++++++++++++++++ 2 files changed, 99 insertions(+), 11 deletions(-) diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index dcc4f254905..eb427757a58 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -1882,6 +1882,20 @@ def _restore_wrapped_ids_in_response_inject(msg_obj: Mapping[str, object]) -> di 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: @@ -1946,8 +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 - # Upstream response id -> the id the client sent in response.inject. - self._inject_client_response_ids: dict[str, str] = {} + 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 @@ -2048,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) @@ -2063,12 +2089,10 @@ class ResponsesWebSocketStreaming: 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") - client_id: Final = ( - self._inject_client_response_ids.get(upstream_id) if isinstance(upstream_id, str) else None - ) - if client_id is None or client_id == upstream_id: + if not isinstance(upstream_id, str): return response_str - return json.dumps({**event_obj, "response_id": client_id}) + 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) @@ -2187,11 +2211,11 @@ class ResponsesWebSocketStreaming: 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._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, ) - 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 bf5d8bf7b5d..5b1349f4507 100644 --- a/tests/unit/responses/test_responses_websocket_all_providers.py +++ b/tests/unit/responses/test_responses_websocket_all_providers.py @@ -3047,6 +3047,70 @@ class TestNativeWebSocketEncryptedContentAffinity: 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