fix(responses): restore encrypted_content and apply affinity on the native WebSocket relay

This commit is contained in:
mateo-berri 2026-09-18 15:44:06 -07:00
parent 6759f28e73
commit 1c15d9f291
8 changed files with 570 additions and 41 deletions

View file

@ -6742,6 +6742,7 @@ class BaseLLMHTTPHandler:
output_guardrail_callbacks=_ws_output_guardrail_callbacks,
quota_callbacks=_ws_quota_callbacks,
authorized_model=model,
custom_llm_provider=custom_llm_provider,
)
await streaming.bidirectional_forward()

View file

@ -19616,7 +19616,7 @@
}
}
},
"description": "\n Unified rate-limit error.\n\n Every rate-limit condition surfaced by litellm \u2014 whether it originated from\n an upstream LLM provider, a vendor batch endpoint, or one of litellm's own\n proxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\n max-iterations, etc.) \u2014 is raised as an instance of this class.\n\n The :attr:`category` attribute lets callers distinguish the source. See\n :class:`RateLimitErrorCategory` for the available values.\n "
"description": "\nUnified rate-limit error.\n\nEvery rate-limit condition surfaced by litellm \u2014 whether it originated from\nan upstream LLM provider, a vendor batch endpoint, or one of litellm's own\nproxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\nmax-iterations, etc.) \u2014 is raised as an instance of this class.\n\nThe :attr:`category` attribute lets callers distinguish the source. See\n:class:`RateLimitErrorCategory` for the available values.\n"
},
"500": {
"content": {

View file

@ -11,10 +11,12 @@ import fastapi
from fastapi import APIRouter, Depends, HTTPException, Request, Response
from fastapi.responses import JSONResponse
from openai.types.responses.response_create_params import ResponseInputParam
from pydantic import BaseModel, ConfigDict, ValidationError
from starlette.websockets import WebSocket, WebSocketDisconnect
from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
from litellm.constants import EMPTY_MAPPING
from litellm.integrations.custom_guardrail import ModifyResponseException
from litellm.llms.base_llm.guardrail_translation.utils import (
blocked_responses_api_usage as _blocked_responses_api_usage,
@ -1289,7 +1291,8 @@ async def cancel_response(
async def _read_ws_model_from_first_frame(
websocket: WebSocket,
) -> tuple | None:
query_model: str | None = None,
) -> tuple[str, str] | None:
"""Read the first WS frame and return (model, raw_message), or None on error.
Sends an appropriate error frame and closes the socket before returning None.
@ -1338,7 +1341,7 @@ async def _read_ws_model_from_first_frame(
await websocket.close(code=1008, reason="Invalid first message")
return None
model: Final = _extract_model_from_first_ws_event(first_event)
model: Final = query_model or _extract_model_from_first_ws_event(first_event)
if not model:
await websocket.send_text(
json.dumps(
@ -1369,6 +1372,29 @@ def _extract_model_from_first_ws_event(first_event: Any) -> str | None:
return (nested.get("model") if isinstance(nested, dict) else None) or first_event.get("model")
class _ResponseCreateRoutingHints(BaseModel):
model_config = ConfigDict(extra="ignore", frozen=True)
input: str | list[object] | None = None
previous_response_id: str | None = None
response: "_ResponseCreateRoutingHints | None" = None
def _routing_hints_from_first_ws_frame(first_message: str) -> Mapping[str, object]:
try:
frame: Final = _ResponseCreateRoutingHints.model_validate_json(first_message)
except ValidationError:
return EMPTY_MAPPING
nested: Final = frame.response or frame
hints: Final = {
"input": frame.input if nested.input is None else nested.input,
"previous_response_id": (
frame.previous_response_id if nested.previous_response_id is None else nested.previous_response_id
),
}
return MappingProxyType({key: value for key, value in hints.items() if value is not None})
async def _enforce_responses_ws_first_frame_model_auth(
request: Request,
model: str,
@ -1455,19 +1481,16 @@ async def responses_websocket_endpoint(
accept_kwargs["subprotocol"] = requested_protocols[0]
await websocket.accept(**accept_kwargs)
first_message: str | None = None
if not model:
result: Final = await _read_ws_model_from_first_frame(websocket)
if result is None:
return
model, first_message = result
result: Final = await _read_ws_model_from_first_frame(websocket, query_model=model)
if result is None:
return
resolved_model, first_message = result
data: dict[str, object] = {
"model": model,
"model": resolved_model,
"websocket": websocket,
"first_message": first_message,
}
if first_message is not None:
data["first_message"] = first_message
# Construct a synthetic Request for pre-call processing
headers_list: Final = list(websocket.scope.get("headers") or [])
@ -1480,7 +1503,7 @@ async def responses_websocket_endpoint(
request: Final = Request(scope=scope)
request._url = websocket.url
_body_bytes: Final = json.dumps({"model": model}).encode()
_body_bytes: Final = json.dumps({"model": resolved_model}).encode()
async def return_body():
return _body_bytes
@ -1490,10 +1513,10 @@ async def responses_websocket_endpoint(
# Phase 1: pre-call processing (auth, guardrails, rate limits)
base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data)
try:
if first_message is not None:
if not model:
await _enforce_responses_ws_first_frame_model_auth(
request=request,
model=model,
model=resolved_model,
user_api_key_dict=user_api_key_dict,
llm_router=llm_router,
)
@ -1512,7 +1535,7 @@ async def responses_websocket_endpoint(
user_request_timeout=user_request_timeout,
user_max_tokens=user_max_tokens,
user_api_base=user_api_base,
model=model,
model=resolved_model,
route_type="_aresponses_websocket",
)
except Exception as e:
@ -1537,6 +1560,7 @@ async def responses_websocket_endpoint(
# Phase 2: route to upstream provider
try:
data["user_api_key_dict"] = user_api_key_dict
data.update(_routing_hints_from_first_ws_frame(first_message))
llm_call: Final = await route_request(
data=data,
route_type="_aresponses_websocket",

View file

@ -2338,6 +2338,8 @@ async def _aresponses_websocket(
"api_base",
"api_key",
"timeout",
"input",
"previous_response_id",
}
remaining_kwargs: Final = {k: v for k, v in kwargs.items() if k not in _explicit_keys}

View file

@ -225,6 +225,29 @@ def _status_code_for_error_fields(error_type: str | None, error_code: str | None
)
def _map_stream_error_to_exception(error_obj: object, model: str, custom_llm_provider: str) -> Exception:
from litellm.llms.base_llm.chat.transformation import BaseLLMException
error_message, error_type, error_code = _error_event_fields(error_obj)
status_code: Final = _status_code_for_error_fields(error_type, error_code)
error_body: Final = {"message": error_message, "type": error_type, "code": error_code}
provider_exception: Final = BaseLLMException(
status_code=status_code,
message=f"Error code: {status_code} - {{'error': {error_body}}}",
body=error_body,
)
try:
return litellm.exception_type(
model=model,
custom_llm_provider=custom_llm_provider,
original_exception=provider_exception,
completion_kwargs={},
extra_kwargs={},
)
except Exception as mapped_exception:
return mapped_exception
def _mid_stream_fallback_eligible(mapped_exception: Exception) -> bool:
if isinstance(mapped_exception, litellm.ContentPolicyViolationError):
return True
@ -588,26 +611,7 @@ class BaseResponsesAPIStreamingIterator:
)
def _map_error_event_exception(self, error_obj: object) -> Exception:
from litellm.llms.base_llm.chat.transformation import BaseLLMException
error_message, error_type, error_code = _error_event_fields(error_obj)
status_code: Final = _status_code_for_error_fields(error_type, error_code)
error_body: Final = {"message": error_message, "type": error_type, "code": error_code}
provider_exception: Final = BaseLLMException(
status_code=status_code,
message=f"Error code: {status_code} - {{'error': {error_body}}}",
body=error_body,
)
try:
return litellm.exception_type(
model=self.model or "",
custom_llm_provider=self.custom_llm_provider or "",
original_exception=provider_exception,
completion_kwargs={},
extra_kwargs={},
)
except Exception as mapped_exception:
return mapped_exception
return _map_stream_error_to_exception(error_obj, self.model or "", self.custom_llm_provider or "")
def _maybe_raise_for_error_event(self, result: object) -> None:
chunk_type: Final = getattr(result, "type", None)
@ -1691,6 +1695,65 @@ RESPONSES_WS_LOGGED_EVENT_TYPES: Final = [
RESPONSES_WS_MASKABLE_TEXT_BLOCK_TYPES: Final = frozenset({"input_text", "output_text", "text"})
_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"})
def _ws_event_error(event: _MutableJsonObject) -> object:
if event.get("type") == "error":
return event.get("error")
response: Final = event.get("response")
return response.get("error") if _is_json_object(response) else None
def _item_id_fields(item: object) -> tuple[object, object]:
return (item.get("id"), item.get("encrypted_content")) if _is_json_object(item) else (None, None)
def _restore_input_item_ids(items: list[object]) -> bool:
before: Final = tuple(_item_id_fields(item) for item in items)
ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input(items) # pyright: ignore[reportPrivateUsage] # same restore the HTTP responses path runs
return before != tuple(_item_id_fields(item) for item in items)
def _restore_wrapped_ids_in_container(container: _MutableJsonObject) -> bool:
input_items: Final = container.get("input")
input_restored: Final = _is_json_array(input_items) and _restore_input_item_ids(input_items)
previous_response_id: Final = container.get("previous_response_id")
if not isinstance(previous_response_id, str):
return input_restored
original_previous_response_id: Final = (
ResponsesAPIRequestUtils.decode_previous_response_id_to_original_previous_response_id(previous_response_id)
)
if original_previous_response_id == previous_response_id:
return input_restored
container["previous_response_id"] = original_previous_response_id
return True
def _restore_wrapped_ids_in_response_create(msg_obj: _MutableJsonObject) -> bool:
nested: Final = msg_obj.get("response")
containers: Final = (msg_obj, nested) if _is_json_object(nested) else (msg_obj,)
restored: Final = tuple(_restore_wrapped_ids_in_container(container) for container in containers)
return any(restored)
def _wrap_output_item_encrypted_content(event_obj: _MutableJsonObject, litellm_metadata: dict[str, object]) -> bool:
if not litellm_metadata.get("encrypted_content_affinity_enabled"):
return False
model_id: Final = _model_id_from_metadata(litellm_metadata)
item: Final = event_obj.get("item")
if model_id is None or not _is_json_object(item):
return False
encrypted_content: Final = item.get("encrypted_content")
if not isinstance(encrypted_content, str) or not encrypted_content:
return False
item["encrypted_content"] = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id( # pyright: ignore[reportPrivateUsage] # same wrap the HTTP streaming path applies
encrypted_content=encrypted_content, model_id=model_id
)
return True
class ResponsesWebSocketStreaming:
"""
@ -1717,12 +1780,16 @@ class ResponsesWebSocketStreaming:
output_guardrail_callbacks: list[PresidioGuardrailCallback] | None = None,
quota_callbacks: Sequence[ProjectQuotaCallback] | None = None,
authorized_model: str | None = None,
custom_llm_provider: str | None = None,
):
self.websocket = websocket
self.backend_ws = backend_ws
self.logging_obj = logging_obj
self.user_api_key_dict = user_api_key_dict
self.request_data: dict[str, object] = request_data or {}
litellm_metadata: Final = self.request_data.get("litellm_metadata")
self.litellm_metadata: dict[str, object] = litellm_metadata if _is_json_object(litellm_metadata) else {}
self.custom_llm_provider: str | None = custom_llm_provider
self.messages: list[_MutableJsonObject] = []
self.input_messages: list[dict[str, object]] = []
self.first_message = first_message
@ -1795,8 +1862,55 @@ class ResponsesWebSocketStreaming:
return
if self.input_messages:
self.logging_obj.model_call_details["messages"] = self.input_messages
if self.messages:
if not self.messages:
return
failed_event: Final = next(
(event for event in self.messages if event.get("type") in _RESPONSES_WS_FAILURE_EVENT_TYPES), None
)
if failed_event is None:
asyncio.create_task(self.logging_obj.dispatch_success_handlers(self.messages, prefer_async_handlers=True))
return
self._record_usage_for_failure()
exception: Final = _map_stream_error_to_exception(
_ws_event_error(failed_event), self.authorized_model or "", self.custom_llm_provider or ""
)
traceback_exception: Final = "".join(traceback.format_exception(exception))
asyncio.create_task(
self.logging_obj.dispatch_failure_handlers(exception, traceback_exception, prefer_async_handlers=True)
)
def _record_usage_for_failure(self) -> None:
from litellm.cost_calculator import ResponsesWebSocketTokenUsageProcessor
from litellm.types.utils import LiteLLMRealtimeStreamLoggingObject
usage: Final = ResponsesWebSocketTokenUsageProcessor.collect_and_combine_usage_from_responses_ws_results(
self.messages
)
tier_partition: Final = ResponsesWebSocketTokenUsageProcessor.partition_results_by_service_tier(self.messages)
service_tier: Final = next(iter(tier_partition)) if len(tier_partition) == 1 else None
logging_result: Final = LiteLLMRealtimeStreamLoggingObject(
usage=usage, results=self.messages, service_tier=service_tier
)
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 _wrap_response_event(self, response_str: str) -> str:
try:
event_obj: Final = _load_json_object(response_str)
except (json.JSONDecodeError, TypeError):
return response_str
response: Final = event_obj.get("response")
if _is_json_object(response):
event_obj["response"] = ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id( # pyright: ignore[reportPrivateUsage] # same wrap the HTTP streaming path applies
responses_api_response=response,
custom_llm_provider=self.custom_llm_provider,
litellm_metadata=self.litellm_metadata,
)
return json.dumps(event_obj)
if event_obj.get("type") not in _RESPONSES_WS_OUTPUT_ITEM_EVENT_TYPES:
return response_str
item_wrapped: Final = _wrap_output_item_encrypted_content(event_obj, self.litellm_metadata)
return json.dumps(event_obj) if item_wrapped else response_str
async def backend_to_client(self) -> None:
"""Forward events from backend WebSocket to the client."""
@ -1833,12 +1947,13 @@ class ResponsesWebSocketStreaming:
unmasked_str = self._unmask_response_event(response_str)
output_masked_str = await self._mask_response_completed(unmasked_str)
wrapped_str = self._wrap_response_event(output_masked_str)
# Log the output-masked form so PII redacted by apply_to_output
# guardrails does not appear in success logs.
self._store_event(output_masked_str)
self._store_event(wrapped_str)
await self.websocket.send_text(output_masked_str)
await self.websocket.send_text(wrapped_str)
except websockets.exceptions.ConnectionClosed as e:
verbose_logger.debug("Responses WS backend connection closed: %s", e)
@ -1898,14 +2013,16 @@ class ResponsesWebSocketStreaming:
# Always enforce the authorized model, even when PII masking is off.
model_modified: Final = self._enforce_authorized_model(msg_obj)
ids_restored: Final = _restore_wrapped_ids_in_response_create(msg_obj)
frame_modified: Final = model_modified or ids_restored
if not self.guardrail_callbacks:
return json.dumps(msg_obj) if model_modified else message
return json.dumps(msg_obj) if frame_modified else message
if "metadata" not in self.request_data:
self.request_data["metadata"] = {}
modified = model_modified
modified = frame_modified
guardrail_cbs: Final[tuple[PresidioGuardrailCallback, ...]] = tuple(self.guardrail_callbacks)
for cb in guardrail_cbs:
presidio_config = cb.get_presidio_settings_from_request_data(self.request_data)

View file

@ -510,6 +510,66 @@ class TestResponsesWSFirstFrameModelAuth:
mock_model_auth.assert_awaited_once()
@pytest.mark.asyncio
@pytest.mark.parametrize("nested", [False, True])
@pytest.mark.parametrize("query_model", [None, "gpt-4o-mini"])
async def test_endpoint_routes_on_first_frame_input_and_previous_response_id(self, nested, query_model):
from litellm.proxy.response_api_endpoints.endpoints import (
responses_websocket_endpoint,
)
replayed_input = [{"type": "reasoning", "id": "encitem_abc", "encrypted_content": "litellm_enc:abc;blob"}]
payload = {"model": "gpt-4o-mini", "input": replayed_input, "previous_response_id": "resp_prev"}
first_frame = {"type": "response.create", "response": payload} if nested else {"type": "response.create", **payload}
raw_first_frame = json.dumps(first_frame)
ws = MagicMock()
ws.headers = {}
ws.query_params = {}
ws.scope = {"headers": []}
ws.url = "ws://testserver/v1/responses"
ws.accept = AsyncMock()
ws.receive_text = AsyncMock(return_value=raw_first_frame)
ws.close = AsyncMock()
processor = MagicMock()
processor.common_processing_pre_call_logic = AsyncMock(
return_value=({"model": "gpt-4o-mini", "litellm_metadata": {}}, MagicMock())
)
async def fake_llm_call():
return None
with (
patch( # test-quality-ok: first-frame model auth needs a live router and key table and has its own tests below
"litellm.proxy.response_api_endpoints.endpoints._enforce_responses_ws_first_frame_model_auth",
new_callable=AsyncMock,
),
patch( # test-quality-ok: the pre-call processor needs a live proxy; the payload it hands to routing is what is under test
"litellm.proxy.response_api_endpoints.endpoints.ProxyBaseLLMRequestProcessing",
return_value=processor,
),
patch( # test-quality-ok: routing is the seam where the first frame's input and previous_response_id become observable
"litellm.proxy.route_llm_request.route_request",
new_callable=AsyncMock,
return_value=fake_llm_call(),
) as mock_route_request,
):
await responses_websocket_endpoint(
websocket=ws,
model=query_model,
user_api_key_dict=MagicMock(),
)
ws.receive_text.assert_awaited_once()
routed = mock_route_request.await_args.kwargs["data"]
assert routed["model"] == "gpt-4o-mini"
assert routed["input"] == replayed_input
assert routed["previous_response_id"] == "resp_prev"
assert processor.common_processing_pre_call_logic.await_args.kwargs["model"] == "gpt-4o-mini"
assert mock_route_request.await_args.kwargs["route_type"] == "_aresponses_websocket"
ws.close.assert_not_awaited()
@pytest.mark.asyncio
async def test_reruns_model_auth_for_first_frame_model(self):
from starlette.requests import Request
@ -636,6 +696,41 @@ class TestReadWSModelFromFirstFrameErrors:
ws.send_text.assert_not_awaited()
ws.close.assert_not_awaited()
@pytest.mark.asyncio
async def test_query_model_wins_over_first_frame_model(self):
from litellm.proxy.response_api_endpoints.endpoints import (
_read_ws_model_from_first_frame,
)
raw = json.dumps({"type": "response.create", "model": "gpt-4o", "input": []})
ws = MagicMock()
ws.receive_text = AsyncMock(return_value=raw)
ws.send_text = AsyncMock()
ws.close = AsyncMock()
result = await _read_ws_model_from_first_frame(ws, query_model="reasoning-group")
assert result == ("reasoning-group", raw)
ws.close.assert_not_awaited()
@pytest.mark.asyncio
async def test_query_model_satisfies_a_first_frame_without_model(self):
from litellm.proxy.response_api_endpoints.endpoints import (
_read_ws_model_from_first_frame,
)
raw = json.dumps({"type": "response.create", "input": []})
ws = MagicMock()
ws.receive_text = AsyncMock(return_value=raw)
ws.send_text = AsyncMock()
ws.close = AsyncMock()
result = await _read_ws_model_from_first_frame(ws, query_model="reasoning-group")
assert result == ("reasoning-group", raw)
ws.send_text.assert_not_awaited()
ws.close.assert_not_awaited()
class TestManagedResponsesSameProvider:
def _handler(self, model, custom_llm_provider=None):

View file

@ -424,6 +424,30 @@ async def test_aresponses_websocket_strips_responses_routing_prefix_from_openai_
assert mock_ws.call_args.kwargs["custom_llm_provider"] == "openai"
@pytest.mark.asyncio
async def test_aresponses_websocket_keeps_routing_hints_out_of_the_relay_kwargs(): # test-quality-ok: the relay kwargs are the only place a dropped key is observable; the provider socket behind them is the boundary
from unittest.mock import MagicMock
from litellm.responses.main import _aresponses_websocket
with patch.object(
import_module("litellm.responses.main").base_llm_http_handler, "async_responses_websocket",
new_callable=AsyncMock,
) as mock_ws:
await _aresponses_websocket(
model="openai/gpt-5.6",
websocket=MagicMock(),
api_key="sk-test",
litellm_logging_obj=MagicMock(),
input=[{"type": "message", "role": "user", "content": "hi"}],
previous_response_id="resp_prev",
)
mock_ws.assert_awaited_once()
assert "input" not in mock_ws.call_args.kwargs
assert "previous_response_id" not in mock_ws.call_args.kwargs
_INJECTION_POINT_INPUT = [{"role": "system", "content": "You are terse."}, {"role": "user", "content": "hi"}]
_SYSTEM_POINT = {"location": "message", "role": "system"}
_USER_POINT = {"location": "message", "role": "user"}

View file

@ -2628,3 +2628,269 @@ class TestNativeWebSocketUrlConstruction:
mock_config.get_websocket_url.assert_called_once()
_, call_kwargs = mock_config.get_websocket_url.call_args
assert call_kwargs["litellm_params"]["api_version"] == "2025-04-01-preview"
_AFFINITY_METADATA = {
"model_info": {"id": "dep-1"},
"encrypted_content_affinity_enabled": True,
}
def _wrapped_reasoning_item():
from litellm.responses.utils import ResponsesAPIRequestUtils
return {
"type": "reasoning",
"id": ResponsesAPIRequestUtils._build_encrypted_item_id("dep-1", "rs_orig"),
"encrypted_content": ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("gAAAA-blob", "dep-1"),
"summary": [],
}
class TestNativeWebSocketEncryptedContentAffinity:
"""The native relay must restore and wrap ids the same way the HTTP /v1/responses path does."""
@pytest.mark.asyncio
@pytest.mark.parametrize("nested", [False, True])
async def test_client_to_backend_restores_wrapped_ids(self, nested):
from unittest.mock import AsyncMock
from litellm.responses.utils import ResponsesAPIRequestUtils
wrapped_previous = ResponsesAPIRequestUtils._build_responses_api_response_id(
custom_llm_provider="openai", model_id="dep-1", response_id="resp_orig"
)
payload = {
"input": [_wrapped_reasoning_item(), {"type": "message", "role": "user", "content": "hi"}],
"previous_response_id": wrapped_previous,
}
frame = {"type": "response.create", "response": payload} if nested else {"type": "response.create", **payload}
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])
body = sent["response"] if nested else sent
assert body["input"][0]["id"] == "rs_orig"
assert body["input"][0]["encrypted_content"] == "gAAAA-blob"
assert body["input"][1] == {"type": "message", "role": "user", "content": "hi"}
assert body["previous_response_id"] == "resp_orig"
@pytest.mark.asyncio
async def test_client_to_backend_leaves_unwrapped_frames_untouched(self):
from unittest.mock import AsyncMock
frame = json.dumps({"type": "response.create", "input": "hello", "previous_response_id": "resp_raw"})
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
from unittest.mock import AsyncMock
import websockets.exceptions # noqa: F401 (lazy submodule must be importable)
from litellm.responses.utils import ResponsesAPIRequestUtils
reasoning_item = {"type": "reasoning", "id": "rs_1", "encrypted_content": "gAAAA-blob", "summary": []}
websocket = MagicMock()
websocket.send_text = AsyncMock()
backend_ws = MagicMock()
backend_ws.recv = AsyncMock(
side_effect=[
json.dumps({"type": "response.output_item.done", "output_index": 0, "item": dict(reasoning_item)}),
json.dumps(
{
"type": "response.completed",
"response": {"id": "resp_1", "output": [dict(reasoning_item)], "usage": {"total_tokens": 3}},
}
),
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": dict(_AFFINITY_METADATA)},
custom_llm_provider="openai",
)
await handler.backend_to_client()
wrapped_content = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("gAAAA-blob", "dep-1")
item_done = json.loads(websocket.send_text.await_args_list[0][0][0])
assert item_done["item"]["encrypted_content"] == wrapped_content
completed = json.loads(websocket.send_text.await_args_list[1][0][0])
assert completed["response"]["id"] == ResponsesAPIRequestUtils._build_responses_api_response_id(
custom_llm_provider="openai", model_id="dep-1", response_id="resp_1"
)
assert completed["response"]["output"][0]["id"] == ResponsesAPIRequestUtils._build_encrypted_item_id(
"dep-1", "rs_1"
)
assert completed["response"]["output"][0]["encrypted_content"] == wrapped_content
await asyncio.sleep(0)
logged = logging_obj.dispatch_success_handlers.await_args[0][0]
assert logged[0]["response"]["id"] == completed["response"]["id"]
@pytest.mark.asyncio
async def test_backend_to_client_wraps_only_response_id_without_affinity(self):
from unittest.mock import AsyncMock
import websockets.exceptions # noqa: F401 (lazy submodule must be importable)
from litellm.responses.utils import ResponsesAPIRequestUtils
reasoning_item = {"type": "reasoning", "id": "rs_1", "encrypted_content": "gAAAA-blob", "summary": []}
websocket = MagicMock()
websocket.send_text = AsyncMock()
backend_ws = MagicMock()
backend_ws.recv = AsyncMock(
side_effect=[
json.dumps({"type": "response.output_item.done", "output_index": 0, "item": dict(reasoning_item)}),
json.dumps({"type": "response.completed", "response": {"id": "resp_1", "output": [dict(reasoning_item)]}}),
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()
item_done = json.loads(websocket.send_text.await_args_list[0][0][0])
assert item_done["item"] == reasoning_item
completed = json.loads(websocket.send_text.await_args_list[1][0][0])
assert completed["response"]["id"] == ResponsesAPIRequestUtils._build_responses_api_response_id(
custom_llm_provider="openai", model_id="dep-1", response_id="resp_1"
)
assert completed["response"]["output"][0] == reasoning_item
@pytest.mark.asyncio
@pytest.mark.parametrize(
"failure_frame, expected_status",
[
(
{
"type": "error",
"error": {
"type": "invalid_request_error",
"code": "invalid_encrypted_content",
"message": "The encrypted content for item rs_1 could not be verified.",
},
},
400,
),
(
{
"type": "response.failed",
"response": {
"id": "resp_1",
"status": "failed",
"error": {"code": "server_error", "message": "upstream blew up"},
},
},
500,
),
],
)
async def test_backend_to_client_books_failure_frames_as_failures(self, failure_frame, expected_status):
import asyncio
from unittest.mock import AsyncMock
import websockets.exceptions # noqa: F401 (lazy submodule must be importable)
websocket = MagicMock()
websocket.send_text = AsyncMock()
backend_ws = MagicMock()
backend_ws.recv = AsyncMock(
side_effect=[
json.dumps({"type": "response.created", "response": {"id": "resp_1", "status": "in_progress"}}),
json.dumps(failure_frame),
Exception("stop"),
]
)
logging_obj = MagicMock()
logging_obj.dispatch_success_handlers = AsyncMock()
logging_obj.dispatch_failure_handlers = AsyncMock()
logging_obj._response_cost_calculator = MagicMock(return_value=0.0)
handler = _make_streaming(
websocket=websocket,
backend_ws=backend_ws,
logging_obj=logging_obj,
request_data={},
authorized_model="gpt-5.6",
custom_llm_provider="openai",
)
await handler.backend_to_client()
await asyncio.sleep(0)
logging_obj.dispatch_success_handlers.assert_not_awaited()
logging_obj.dispatch_failure_handlers.assert_awaited_once()
exception = logging_obj.dispatch_failure_handlers.await_args[0][0]
assert exception.status_code == expected_status
assert failure_frame.get("error", failure_frame.get("response", {}).get("error"))["message"] in str(exception)
@pytest.mark.asyncio
async def test_backend_to_client_bills_completed_turns_before_a_failure(self):
import asyncio
from unittest.mock import AsyncMock
import websockets.exceptions # noqa: F401 (lazy submodule must be importable)
websocket = MagicMock()
websocket.send_text = AsyncMock()
backend_ws = MagicMock()
backend_ws.recv = AsyncMock(
side_effect=[
json.dumps(
{
"type": "response.completed",
"response": {
"id": "resp_1",
"status": "completed",
"output": [],
"usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15},
},
}
),
json.dumps({"type": "error", "error": {"type": "invalid_request_error", "message": "bad turn"}}),
Exception("stop"),
]
)
logging_obj = MagicMock()
logging_obj.dispatch_success_handlers = AsyncMock()
logging_obj.dispatch_failure_handlers = AsyncMock()
logging_obj._response_cost_calculator = MagicMock(return_value=0.01)
handler = _make_streaming(websocket=websocket, backend_ws=backend_ws, logging_obj=logging_obj, request_data={})
await handler.backend_to_client()
await asyncio.sleep(0)
logging_obj.record_partial_usage_for_failure.assert_called_once()
usage, response_cost = logging_obj.record_partial_usage_for_failure.call_args[0]
assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == (10, 5, 15)
assert response_cost == 0.01
logging_obj.dispatch_success_handlers.assert_not_awaited()
logging_obj.dispatch_failure_handlers.assert_awaited_once()