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
This commit is contained in:
Darsh Joshi 2026-10-02 20:25:31 -04:00
parent f498176a27
commit 5d01530c43
2 changed files with 101 additions and 1 deletions

View file

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

View file

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