mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
5d01530c43
commit
b5081877a5
2 changed files with 73 additions and 28 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue