mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-24 00:52:24 +00:00
refactor(responses): build the routed websocket request and relay frames without in-place mutation
This commit is contained in:
parent
9662b2a35c
commit
33223920ca
2 changed files with 50 additions and 47 deletions
|
|
@ -2,7 +2,7 @@ import asyncio
|
|||
import contextlib
|
||||
import json
|
||||
import time
|
||||
from collections.abc import AsyncIterator, Awaitable, Mapping
|
||||
from collections.abc import AsyncIterator, Awaitable, Mapping, Sequence
|
||||
from enum import Enum
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, NamedTuple, Protocol, cast, get_args
|
||||
|
|
@ -1376,7 +1376,7 @@ def _extract_model_from_first_ws_event(first_event: Any) -> str | None:
|
|||
class _ResponseCreateRoutingHints(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
|
||||
input: str | list[object] | None = None
|
||||
input: str | Sequence[object] | None = None
|
||||
previous_response_id: str | None = None
|
||||
response: "_ResponseCreateRoutingHints | None" = None
|
||||
|
||||
|
|
@ -1567,12 +1567,13 @@ async def responses_websocket_endpoint(
|
|||
await websocket.close(code=1008, reason="Pre-call error")
|
||||
return
|
||||
|
||||
routed_data: Final = dict(
|
||||
data, user_api_key_dict=user_api_key_dict, **_routing_hints_from_first_ws_frame(first_message)
|
||||
)
|
||||
# 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,
|
||||
data=routed_data,
|
||||
route_type="_aresponses_websocket",
|
||||
llm_router=llm_router,
|
||||
user_model=user_model,
|
||||
|
|
@ -1582,7 +1583,7 @@ async def responses_websocket_endpoint(
|
|||
await proxy_logging_obj.post_call_failure_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
original_exception=failure,
|
||||
request_data=data,
|
||||
request_data=routed_data,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Responses WebSocket error")
|
||||
|
|
@ -1591,6 +1592,6 @@ async def responses_websocket_endpoint(
|
|||
await proxy_logging_obj.post_call_failure_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
original_exception=e,
|
||||
request_data=data,
|
||||
request_data=routed_data,
|
||||
)
|
||||
await websocket.close(code=1011, reason="Internal server error")
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import copy
|
||||
import json
|
||||
import time
|
||||
import traceback
|
||||
|
|
@ -154,7 +155,7 @@ def _load_json_value(payload: str | bytes) -> object:
|
|||
return json.loads(payload)
|
||||
|
||||
|
||||
def _model_id_from_metadata(litellm_metadata: dict[str, object] | None) -> str | None:
|
||||
def _model_id_from_metadata(litellm_metadata: Mapping[str, object] | None) -> str | None:
|
||||
model_info: Final = litellm_metadata.get("model_info") if litellm_metadata else None
|
||||
model_id: Final = model_info.get("id") if _is_json_object(model_info) else None
|
||||
return model_id if isinstance(model_id, str) else None
|
||||
|
|
@ -1701,59 +1702,59 @@ _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:
|
||||
def _ws_event_error(event: Mapping[str, object]) -> 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: Sequence[object]) -> Sequence[object]:
|
||||
return ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input(copy.deepcopy(list(items))) # pyright: ignore[reportPrivateUsage] # same restore the HTTP responses path runs
|
||||
|
||||
|
||||
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:
|
||||
def _restored_container_fields(container: Mapping[str, object]) -> Mapping[str, object]:
|
||||
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
|
||||
restored: Final = {
|
||||
"input": _restore_input_item_ids(input_items) if _is_json_array(input_items) else input_items,
|
||||
"previous_response_id": (
|
||||
ResponsesAPIRequestUtils.decode_previous_response_id_to_original_previous_response_id(previous_response_id)
|
||||
if isinstance(previous_response_id, str)
|
||||
else previous_response_id
|
||||
),
|
||||
}
|
||||
return MappingProxyType({key: value for key, value in restored.items() if value != container.get(key)})
|
||||
|
||||
|
||||
def _restore_wrapped_ids_in_response_create(msg_obj: _MutableJsonObject) -> bool:
|
||||
def _restore_wrapped_ids_in_response_create(msg_obj: Mapping[str, object]) -> dict[str, object] | None:
|
||||
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)
|
||||
nested_fields: Final = _restored_container_fields(nested) if _is_json_object(nested) else EMPTY_MAPPING
|
||||
top_fields: Final = _restored_container_fields(msg_obj)
|
||||
if not nested_fields and not top_fields:
|
||||
return None
|
||||
restored_nested: Final = (
|
||||
{"response": {**nested, **nested_fields}} if _is_json_object(nested) and nested_fields else EMPTY_MAPPING
|
||||
)
|
||||
return {**msg_obj, **top_fields, **restored_nested}
|
||||
|
||||
|
||||
def _wrap_output_item_encrypted_content(event_obj: _MutableJsonObject, litellm_metadata: dict[str, object]) -> bool:
|
||||
def _wrap_output_item_encrypted_content(
|
||||
event_obj: Mapping[str, object], litellm_metadata: Mapping[str, object]
|
||||
) -> dict[str, object] | None:
|
||||
if not litellm_metadata.get("encrypted_content_affinity_enabled"):
|
||||
return False
|
||||
return None
|
||||
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
|
||||
return None
|
||||
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
|
||||
return None
|
||||
wrapped_content: Final = 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
|
||||
return {**event_obj, "item": {**item, "encrypted_content": wrapped_content}}
|
||||
|
||||
|
||||
class ResponsesWebSocketStreaming:
|
||||
|
|
@ -1909,16 +1910,16 @@ class ResponsesWebSocketStreaming:
|
|||
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
|
||||
wrapped_response: Final = 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)
|
||||
return json.dumps({**event_obj, "response": wrapped_response})
|
||||
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
|
||||
wrapped_event: Final = _wrap_output_item_encrypted_content(event_obj, self.litellm_metadata)
|
||||
return response_str if wrapped_event is None else json.dumps(wrapped_event)
|
||||
|
||||
async def backend_to_client(self) -> None:
|
||||
"""Forward events from backend WebSocket to the client."""
|
||||
|
|
@ -2030,13 +2031,14 @@ class ResponsesWebSocketStreaming:
|
|||
if parsed.get("type") != "response.create":
|
||||
return message
|
||||
|
||||
msg_obj: Final = self._with_request_defaults(parsed)
|
||||
defaults_applied: Final = msg_obj != parsed
|
||||
authorized_obj: Final = self._with_request_defaults(parsed)
|
||||
defaults_applied: Final = authorized_obj != parsed
|
||||
|
||||
# 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 or defaults_applied
|
||||
model_modified: Final = self._enforce_authorized_model(authorized_obj)
|
||||
restored_obj: Final = _restore_wrapped_ids_in_response_create(authorized_obj)
|
||||
msg_obj: Final = authorized_obj if restored_obj is None else restored_obj
|
||||
frame_modified: Final = model_modified or restored_obj is not None or defaults_applied
|
||||
|
||||
if not self.guardrail_callbacks:
|
||||
return json.dumps(msg_obj) if frame_modified else message
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue