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.
This commit is contained in:
Darsh Joshi 2026-10-03 12:14:13 -04:00
parent b5081877a5
commit 58b028ca4c
2 changed files with 99 additions and 11 deletions

View file

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

View file

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