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