mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge 58b028ca4c into 635085ac14
This commit is contained in:
commit
89776a7a5d
2 changed files with 234 additions and 1 deletions
|
|
@ -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,34 @@ def _restore_wrapped_ids_in_response_create(msg_obj: Mapping[str, object]) -> di
|
|||
return {**msg_obj, **top_fields, **restored_nested}
|
||||
|
||||
|
||||
def _restore_wrapped_ids_in_response_inject(msg_obj: Mapping[str, object]) -> dict[str, object] | None:
|
||||
response_id: Final = msg_obj.get("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
|
||||
)
|
||||
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 _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:
|
||||
|
|
@ -1930,6 +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
|
||||
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
|
||||
|
|
@ -2030,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)
|
||||
|
|
@ -2043,6 +2087,12 @@ 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:
|
||||
upstream_id: Final = event_obj.get("response_id")
|
||||
if not isinstance(upstream_id, str):
|
||||
return response_str
|
||||
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)
|
||||
|
|
@ -2148,13 +2198,25 @@ 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`` and ``input`` item ids restored to the upstream ones.
|
||||
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_ids_in_response_inject(parsed)
|
||||
inject_id: Final = parsed.get("response_id")
|
||||
if isinstance(inject_id, str):
|
||||
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,
|
||||
)
|
||||
return message if restored_inject is None else json.dumps(restored_inject)
|
||||
if parsed.get("type") != "response.create":
|
||||
return message
|
||||
|
||||
|
|
|
|||
|
|
@ -2940,6 +2940,177 @@ 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_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()
|
||||
|
||||
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"}
|
||||
|
||||
@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