fix(responses): restore inject input items and echo the client's inject id

A response.inject frame could carry WebSocket output items whose ids and
encrypted_content LiteLLM wrapped under encrypted-content affinity, but only
response_id was decoded, so the provider got litellm_enc: content it cannot
read. Restore inject input items with the same helper response.create uses.

The inject result frame always wrapped its response_id, so a client that
injected with a raw upstream id got back an id it never sent. Remember the
id each inject was sent with, per connection, and echo that id instead.
This commit is contained in:
Darsh Joshi 2026-10-03 11:51:16 -04:00
parent 5d01530c43
commit b5081877a5
2 changed files with 73 additions and 28 deletions

View file

@ -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

View file

@ -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):