diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 3396f79973a..9061ac39be8 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -14,8 +14,10 @@ from typing import ( Dict, Iterator, List, + Mapping, NoReturn, Optional, + Sequence, Union, cast, ) @@ -104,13 +106,13 @@ def _json_loads_object(raw: Union[str, bytes]) -> object: return json.loads(raw) # any-ok: json.loads is the untyped-JSON boundary; callers narrow via _require_* -def _require_dict(value: object) -> dict[str, object]: +def _require_dict(value: object) -> Mapping[str, object]: if isinstance(value, dict): return value raise ValueError(f"Expected a JSON object, got: {value!r}") -def _require_list(value: object) -> list[object]: +def _require_list(value: object) -> Sequence[object]: if isinstance(value, list): return value raise ValueError(f"Expected a JSON array, got: {value!r}") @@ -697,17 +699,17 @@ class CustomStreamWrapper: except Exception as e: raise e - def model_response_creator(self, chunk: Optional[dict[str, object]] = None, hidden_params: Optional[dict] = None): + def model_response_creator( + self, + chunk: Optional[Mapping[str, object]] = None, + hidden_params: Optional[Mapping[str, object]] = None, + ): _model = self._cached_model_name _logging_obj_llm_provider = self._cached_logging_llm_provider - if chunk is None: - args: dict[str, object] = {"model": _model} - else: - chunk.pop("model", None) - args = {"model": _model} - if chunk: - args.update({k: v for k, v in chunk.items() if k != "stream"}) + args: dict[str, object] = {"model": _model} + if chunk: + args.update({k: v for k, v in chunk.items() if k not in ("model", "stream")}) model_response = ModelResponseStream.model_validate(args) if self.response_id is not None: diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index ec1301e5923..c3c51a8fc50 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -161,6 +161,7 @@ from .http_handler import get_shared_realtime_ssl_context if TYPE_CHECKING: from aiohttp import ClientSession + from starlette.websockets import WebSocket from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( @@ -6137,7 +6138,7 @@ class BaseLLMHTTPHandler: async def async_responses_websocket( self, model: str, - websocket: Any, + websocket: "WebSocket", logging_obj: LiteLLMLoggingObj, responses_api_provider_config: Optional[BaseResponsesAPIConfig], api_base: Optional[str] = None, @@ -6300,7 +6301,10 @@ class BaseLLMHTTPHandler: except websockets.exceptions.InvalidStatusCode as e: # type: ignore verbose_logger.exception(f"Error connecting to responses WS backend: {e}") - await websocket.close(code=e.status_code, reason=_redact_string(str(e))) + await websocket.close( + code=e.status_code if isinstance(e.status_code, int) else 1011, + reason=_redact_string(str(e)), + ) except Exception as e: verbose_logger.exception(f"Error in responses WS: {e}") try: diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index dab666ff0d9..b1766ec0cac 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -8,10 +8,21 @@ import uuid from datetime import datetime from functools import lru_cache from types import MappingProxyType -from typing import Any, Dict, List, Literal, Mapping, Optional +from typing import ( + TYPE_CHECKING, + Any, + Dict, + List, + Literal, + Mapping, + Optional, + Protocol, + Union, +) import httpx from openai._streaming import SSEDecoder +from pydantic import BaseModel, TypeAdapter import litellm from litellm.constants import ( @@ -29,10 +40,56 @@ from litellm.litellm_core_utils.llm_response_utils.response_metadata import ( from litellm.litellm_core_utils.thread_pool_executor import executor from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig from litellm.responses.utils import ResponseAPILoggingUtils, ResponsesAPIRequestUtils -from litellm.types.llms.openai import ResponsesAPIStreamEvents +from litellm.types.guardrails import PresidioPerRequestConfig +from litellm.types.llms.openai import ( + PART_UNION_TYPES, + ResponsesAPIResponse, + ResponsesAPIStreamEvents, + ResponsesAPIStreamingResponse, +) from litellm.types.utils import CallTypes from litellm.utils import async_post_call_success_deployment_hook +if TYPE_CHECKING: + from starlette.websockets import WebSocket + from websockets.asyncio.client import ClientConnection + + from litellm.proxy._types import UserAPIKeyAuth + +_JSON_OBJECT_ADAPTER: TypeAdapter[Dict[str, object]] = TypeAdapter(Dict[str, object]) +_DUMP_ADAPTER: TypeAdapter[Dict[str, Any]] = TypeAdapter(Dict[str, Any]) + + +class PIIMaskingGuardrail(Protocol): + """Structural type for guardrails that expose the Presidio PII masking interface. + + Any guardrail implementing this shape works here (duck typing), not just + the Presidio guardrail hook, to avoid a layering violation (SDK importing + from the proxy-only guardrails package). + """ + + output_parse_pii: bool + apply_to_output: bool + + def get_presidio_settings_from_request_data( + self, data: Dict[str, object] + ) -> Optional[PresidioPerRequestConfig]: ... + + async def check_pii( + self, + text: str, + output_parse_pii: bool, + presidio_config: Optional[PresidioPerRequestConfig], + request_data: Dict[str, object], + ) -> str: ... + + +def _call_unmask_pii_text(guardrail: PIIMaskingGuardrail, text: str, pii_tokens: Dict[str, str]) -> str: + # any-ok: _unmask_pii_text is a private helper on the concrete guardrail class; + # dispatched dynamically so PIIMaskingGuardrail need not declare private methods. + unmask = getattr(guardrail, "_unmask_pii_text") + return unmask(text, pii_tokens) + @lru_cache(maxsize=1) def _get_openai_response_types(): @@ -130,7 +187,7 @@ class BaseResponsesAPIStreamingIterator: self.logging_obj = logging_obj self.finished = False self.responses_api_provider_config = responses_api_provider_config - self.completed_response: Optional[Any] = None + self.completed_response: Optional[ResponsesAPIStreamingResponse] = None self.start_time = getattr(logging_obj, "start_time", datetime.now()) self._failure_handled = False # Track if failure handler has been called self._yielded_first_chunk = False @@ -175,7 +232,7 @@ class BaseResponsesAPIStreamingIterator: llm_provider=self.custom_llm_provider or "", ) - def _process_chunk(self, chunk) -> Optional[Any]: + def _process_chunk(self, chunk: str) -> Optional[ResponsesAPIStreamingResponse]: """Process a single chunk of data from the stream""" if not chunk: return None @@ -298,9 +355,9 @@ class BaseResponsesAPIStreamingIterator: self.completed_response = openai_responses_api_chunk # Add cost to usage object if include_cost_in_streaming_usage is True if litellm.include_cost_in_streaming_usage and self.logging_obj is not None: - response_obj: Optional[Any] = getattr(openai_responses_api_chunk, "response", None) + response_obj = getattr(openai_responses_api_chunk, "response", None) if response_obj: - usage_obj: Optional[Any] = getattr(response_obj, "usage", None) + usage_obj = getattr(response_obj, "usage", None) if usage_obj is not None: try: cost: Optional[float] = self.logging_obj._response_cost_calculator( @@ -403,7 +460,7 @@ class BaseResponsesAPIStreamingIterator: ) self._handle_failure(exception) - def _record_failed_response_usage(self, response_obj: Optional[Any]) -> None: + def _record_failed_response_usage(self, response_obj: Optional[ResponsesAPIResponse]) -> None: if response_obj is None or self.logging_obj is None: return usage_obj = getattr(response_obj, "usage", None) @@ -453,14 +510,13 @@ class BaseResponsesAPIStreamingIterator: is_pre_first_chunk=not self._yielded_first_chunk, ) - def _get_completed_response_object(self) -> Optional[Any]: - openai_types = _get_openai_response_types() + def _get_completed_response_object(self) -> Optional[ResponsesAPIResponse]: completed_response = self.completed_response - if isinstance(completed_response, openai_types.ResponsesAPIResponse): + if isinstance(completed_response, ResponsesAPIResponse): return completed_response response_obj = getattr(completed_response, "response", None) - if isinstance(response_obj, openai_types.ResponsesAPIResponse): + if isinstance(response_obj, ResponsesAPIResponse): return response_obj return None @@ -529,7 +585,9 @@ class BaseResponsesAPIStreamingIterator: self._completed_response_cached = True - async def _call_post_streaming_deployment_hook(self, chunk): + async def _call_post_streaming_deployment_hook( + self, chunk: ResponsesAPIStreamingResponse + ) -> ResponsesAPIStreamingResponse: """ Allow callbacks to modify streaming chunks before returning (parity with chat). """ @@ -547,13 +605,13 @@ class BaseResponsesAPIStreamingIterator: except Exception: typed_call_type = None - request_data = self.request_data or getattr(self.logging_obj, "model_call_details", {}) - callbacks = getattr(litellm, "callbacks", None) or [] + request_data = self.request_data or self.logging_obj.model_call_details hooks_ran = False - for callback in callbacks: - if hasattr(callback, "async_post_call_streaming_deployment_hook"): + for callback in litellm.callbacks: + hook = getattr(callback, "async_post_call_streaming_deployment_hook", None) + if hook is not None: hooks_ran = True - result = await callback.async_post_call_streaming_deployment_hook( + result = await hook( request_data=request_data, response_chunk=chunk, call_type=typed_call_type, @@ -566,7 +624,9 @@ class BaseResponsesAPIStreamingIterator: except Exception: return chunk - async def call_post_streaming_hooks_for_testing(self, chunk): + async def call_post_streaming_hooks_for_testing( + self, chunk: ResponsesAPIStreamingResponse + ) -> ResponsesAPIStreamingResponse: """ Helper to invoke streaming deployment hooks explicitly (used in tests). """ @@ -583,15 +643,12 @@ class BaseResponsesAPIStreamingIterator: if isinstance(self.request_data, dict): request_payload.update(self.request_data) try: - if hasattr(self.logging_obj, "model_call_details"): - request_payload.update(self.logging_obj.model_call_details) + request_payload.update(self.logging_obj.model_call_details) except Exception: pass if "litellm_params" not in request_payload: try: - request_payload["litellm_params"] = getattr(self.logging_obj, "model_call_details", {}).get( - "litellm_params", {} - ) + request_payload["litellm_params"] = self.logging_obj.model_call_details.get("litellm_params", {}) except Exception: request_payload["litellm_params"] = {} @@ -668,14 +725,16 @@ class BaseResponsesAPIStreamingIterator: pass -async def call_post_streaming_hooks_for_testing(iterator, chunk): +async def call_post_streaming_hooks_for_testing( + iterator: object, chunk: ResponsesAPIStreamingResponse +) -> ResponsesAPIStreamingResponse: """ Module-level helper for tests to ensure hooks can be invoked even if the iterator is wrapped. """ hook_fn = getattr(iterator, "_call_post_streaming_deployment_hook", None) if hook_fn is None: return chunk - return await hook_fn(chunk) + return await hook_fn(chunk) # any-ok: test helper must dispatch onto arbitrary wrapped iterator doubles class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): @@ -709,7 +768,7 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): def __aiter__(self): return self - async def __anext__(self) -> Any: + async def __anext__(self) -> ResponsesAPIStreamingResponse: try: self._check_max_streaming_duration() while True: @@ -791,7 +850,7 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): def __iter__(self): return self - def __next__(self): + def __next__(self) -> ResponsesAPIStreamingResponse: try: self._check_max_streaming_duration() while True: @@ -882,10 +941,10 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): def _set_events_from_response( self, - transformed: Any, + transformed: ResponsesAPIResponse, logging_obj: LiteLLMLoggingObj, ) -> None: - self._events = _build_synthetic_response_events( + self._events: List[ResponsesAPIStreamingResponse] = _build_synthetic_response_events( transformed=transformed, logging_obj=logging_obj, chunk_size=self.CHUNK_SIZE, @@ -896,13 +955,12 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): def __aiter__(self): return self - async def __anext__(self) -> Any: + async def __anext__(self) -> ResponsesAPIStreamingResponse: if self._idx >= len(self._events): raise StopAsyncIteration evt = self._events[self._idx] self._idx += 1 - openai_types = _get_openai_response_types() - if getattr(evt, "type", None) == openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED: + if getattr(evt, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED: self.completed_response = evt self._log_completed_response(is_async=True) return evt @@ -910,13 +968,12 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): def __iter__(self): return self - def __next__(self) -> Any: + def __next__(self) -> ResponsesAPIStreamingResponse: if self._idx >= len(self._events): raise StopIteration evt = self._events[self._idx] self._idx += 1 - openai_types = _get_openai_response_types() - if getattr(evt, "type", None) == openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED: + if getattr(evt, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED: self.completed_response = evt self._log_completed_response(is_async=False) return evt @@ -925,7 +982,7 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): class CachedResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): def __init__( self, - response: Any, + response: ResponsesAPIResponse, logging_obj: LiteLLMLoggingObj, request_data: Optional[Dict[str, Any]] = None, call_type: Optional[str] = None, @@ -933,7 +990,7 @@ class CachedResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): BaseResponsesAPIStreamingIterator.__init__( self, response=httpx.Response(200), - model=getattr(response, "model", ""), + model=response.model or "", responses_api_provider_config=None, logging_obj=logging_obj, litellm_metadata=None, @@ -943,13 +1000,13 @@ class CachedResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): ) self._completed_response_cache_hit = True self._persist_completed_response_before_logging = False - self._events: List[Any] = [] + self._events: List[ResponsesAPIStreamingResponse] = [] self._idx = 0 self._set_events_from_response(transformed=response, logging_obj=logging_obj) def _set_events_from_response( self, - transformed: Any, + transformed: ResponsesAPIResponse, logging_obj: LiteLLMLoggingObj, ) -> None: self._events = _build_synthetic_response_events( @@ -963,13 +1020,12 @@ class CachedResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): def __aiter__(self): return self - async def __anext__(self) -> Any: + async def __anext__(self) -> ResponsesAPIStreamingResponse: if self._idx >= len(self._events): raise StopAsyncIteration evt = self._events[self._idx] self._idx += 1 - openai_types = _get_openai_response_types() - if getattr(evt, "type", None) == openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED: + if getattr(evt, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED: self.completed_response = evt self._log_completed_response(is_async=True) return evt @@ -977,23 +1033,29 @@ class CachedResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): def __iter__(self): return self - def __next__(self) -> Any: + def __next__(self) -> ResponsesAPIStreamingResponse: if self._idx >= len(self._events): raise StopIteration evt = self._events[self._idx] self._idx += 1 - openai_types = _get_openai_response_types() - if getattr(evt, "type", None) == openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED: + if getattr(evt, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED: self.completed_response = evt self._log_completed_response(is_async=False) return evt -def _dump_response_object(obj: Any) -> Dict[str, Any]: - if hasattr(obj, "model_dump"): +def _dump_response_object(obj: object) -> Dict[str, Any]: + """Normalize a Responses API output item to a plain dict. + + Returns ``Dict[str, Any]`` (not ``object``) because callers splat this + payload into strongly-typed ``BaseLiteLLMOpenAIResponseObject`` subclass + constructors (e.g. ``logprobs=payload.get("logprobs")``); an ``object`` + value type would fail those field-typed constructor calls. + """ + if isinstance(obj, BaseModel): return obj.model_dump() if isinstance(obj, dict): - return obj + return _DUMP_ADAPTER.validate_python(obj) return {} @@ -1002,8 +1064,8 @@ def _build_response_status_event( "response.created", "response.in_progress", ], - transformed: Any, -) -> Any: + transformed: ResponsesAPIResponse, +) -> ResponsesAPIStreamingResponse: openai_types = _get_openai_response_types() in_progress_response = transformed.model_copy( deep=True, @@ -1020,10 +1082,10 @@ def _build_content_part_done_event( output_index: int, content_index: int, part_payload: Dict[str, Any], -) -> Optional[Any]: +) -> Optional[ResponsesAPIStreamingResponse]: openai_types = _get_openai_response_types() part_type = part_payload.get("type") - part: Any + part: PART_UNION_TYPES if part_type == "output_text": annotations = [ openai_types.BaseLiteLLMOpenAIResponseObject(**annotation) @@ -1059,7 +1121,7 @@ def _build_content_part_done_event( def _add_text_like_part_events( *, - events: List[Any], + events: List[ResponsesAPIStreamingResponse], item_id: str, output_index: int, content_index: int, @@ -1125,13 +1187,13 @@ def _add_text_like_part_events( def _build_synthetic_response_events( *, - transformed: Any, + transformed: ResponsesAPIResponse, logging_obj: LiteLLMLoggingObj, chunk_size: int, -) -> List[Any]: +) -> List[ResponsesAPIStreamingResponse]: openai_types = _get_openai_response_types() if litellm.include_cost_in_streaming_usage and logging_obj is not None: - usage_obj: Optional[Any] = getattr(transformed, "usage", None) + usage_obj = transformed.usage if usage_obj is not None: try: cost: Optional[float] = logging_obj._response_cost_calculator(result=transformed) @@ -1140,13 +1202,13 @@ def _build_synthetic_response_events( except Exception: pass - events: List[Any] = [ + events: List[ResponsesAPIStreamingResponse] = [ _build_response_status_event(openai_types.ResponsesAPIStreamEvents.RESPONSE_CREATED, transformed), _build_response_status_event(openai_types.ResponsesAPIStreamEvents.RESPONSE_IN_PROGRESS, transformed), ] sequence_number = 0 - for output_index, output_item in enumerate(getattr(transformed, "output", []) or []): + for output_index, output_item in enumerate(transformed.output): output_item_payload = _dump_response_object(output_item) item_id = str(output_item_payload.get("id") or transformed.id) item_type = output_item_payload.get("type") @@ -1279,6 +1341,22 @@ RESPONSES_WS_LOGGED_EVENT_TYPES = [ RESPONSES_WS_MASKABLE_TEXT_BLOCK_TYPES = frozenset({"input_text", "output_text", "text"}) +def _parse_json_object(raw: Union[str, bytes]) -> Optional[Dict[str, object]]: + try: + parsed = json.loads(raw) + except (json.JSONDecodeError, TypeError): + return None + return _as_str_object_dict(parsed) + + +def _as_str_object_dict(value: object) -> Optional[Dict[str, object]]: + return _JSON_OBJECT_ADAPTER.validate_python(value) if isinstance(value, dict) else None + + +def _as_object_list(value: object) -> List[object]: + return value if isinstance(value, list) else [] + + class ResponsesWebSocketStreaming: """ Manages bidirectional WebSocket forwarding for the Responses API @@ -1294,55 +1372,47 @@ class ResponsesWebSocketStreaming: def __init__( self, - websocket: Any, - backend_ws: Any, + websocket: "WebSocket", + backend_ws: "ClientConnection", logging_obj: LiteLLMLoggingObj, - user_api_key_dict: Optional[Any] = None, - request_data: Optional[Dict] = None, + user_api_key_dict: Optional["UserAPIKeyAuth"] = None, + request_data: Optional[Dict[str, Any]] = None, first_message: Optional[str] = None, - guardrail_callbacks: Optional[List[Any]] = None, - output_guardrail_callbacks: Optional[List[Any]] = None, + guardrail_callbacks: Optional[List[PIIMaskingGuardrail]] = None, + output_guardrail_callbacks: Optional[List[PIIMaskingGuardrail]] = None, authorized_model: Optional[str] = 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 = request_data or {} - self.messages: list[Dict] = [] + self.request_data: Dict[str, Any] = request_data or {} + self.messages: List[Dict[str, object]] = [] self.input_messages: list[Dict[str, str]] = [] self.first_message = first_message - self.guardrail_callbacks: List[Any] = guardrail_callbacks or [] - self.output_guardrail_callbacks: List[Any] = output_guardrail_callbacks or [] + self.guardrail_callbacks: List[PIIMaskingGuardrail] = guardrail_callbacks or [] + self.output_guardrail_callbacks: List[PIIMaskingGuardrail] = output_guardrail_callbacks or [] # Model name authorized at connection time; enforced on every # response.create frame to prevent deployment-substitution attacks. self.authorized_model: Optional[str] = authorized_model - def _should_store_event(self, event_obj: dict) -> bool: + def _should_store_event(self, event_obj: Dict[str, object]) -> bool: return event_obj.get("type") in RESPONSES_WS_LOGGED_EVENT_TYPES - def _store_event(self, event: Any) -> None: - if isinstance(event, bytes): - event = event.decode("utf-8") - if isinstance(event, str): - try: - event_obj = json.loads(event) - except (json.JSONDecodeError, TypeError): - return - else: - event_obj = event + def _store_event(self, event: Union[str, bytes]) -> None: + decoded = event.decode("utf-8") if isinstance(event, bytes) else event + event_obj = _parse_json_object(decoded) + if event_obj is None: + return if self._should_store_event(event_obj): self.messages.append(event_obj) - def _collect_input_from_client_event(self, message: Any) -> None: + def _collect_input_from_client_event(self, message: Union[str, Dict[str, object]]) -> None: """Extract user input content from response.create for logging.""" try: - if isinstance(message, str): - msg_obj = json.loads(message) - elif isinstance(message, dict): - msg_obj = message - else: + msg_obj = _parse_json_object(message) if isinstance(message, str) else message + if msg_obj is None: return if msg_obj.get("type") != "response.create": @@ -1367,10 +1437,10 @@ class ResponsesWebSocketStreaming: text = c.get("text", "") if text: self.input_messages.append({"role": "user", "content": text}) - except (json.JSONDecodeError, AttributeError, TypeError): + except (AttributeError, TypeError): pass - def _store_input(self, message: Any) -> None: + def _store_input(self, message: str) -> None: self._collect_input_from_client_event(message) if self.logging_obj: self.logging_obj.pre_call(input=message, api_key="") @@ -1390,9 +1460,9 @@ class ResponsesWebSocketStreaming: try: while True: try: - raw_response = await self.backend_ws.recv(decode=False) # type: ignore[union-attr] + raw_response = await self.backend_ws.recv(decode=False) except TypeError: - raw_response = await self.backend_ws.recv() # type: ignore[union-attr, assignment] + raw_response = await self.backend_ws.recv() if isinstance(raw_response, bytes): response_str = raw_response.decode("utf-8") @@ -1408,10 +1478,8 @@ class ResponsesWebSocketStreaming: # before response.completed arrives. The client receives only the # masked response.completed. if self.output_guardrail_callbacks: - try: - _evt_type = json.loads(response_str).get("type") - except (json.JSONDecodeError, TypeError): - _evt_type = None + _parsed_for_type = _parse_json_object(response_str) + _evt_type = _parsed_for_type.get("type") if _parsed_for_type is not None else None if _evt_type in self._DELTA_EVENT_TYPES or _evt_type in self._OUTPUT_DONE_EVENT_TYPES: continue @@ -1431,7 +1499,7 @@ class ResponsesWebSocketStreaming: finally: await self._log_messages() - def _enforce_authorized_model(self, msg_obj: dict) -> bool: + def _enforce_authorized_model(self, msg_obj: Dict[str, object]) -> bool: """ Overwrite any ``model`` field in a ``response.create`` frame with the connection-authorized model to prevent deployment-substitution attacks. @@ -1472,9 +1540,8 @@ class ResponsesWebSocketStreaming: Non-``response.create`` messages are returned unchanged. """ - try: - msg_obj = json.loads(message) - except (json.JSONDecodeError, TypeError): + msg_obj = _parse_json_object(message) + if msg_obj is None: return message if msg_obj.get("type") != "response.create": @@ -1497,10 +1564,10 @@ class ResponsesWebSocketStreaming: # nested: {"type": "response.create", "response": {"input": ..., "instructions": ...}} # Mask "input" and "instructions" in both shapes so PII is never # forwarded unmasked regardless of where the client places it. - nested_response = msg_obj.get("response") if isinstance(msg_obj.get("response"), dict) else None + nested_response = msg_obj.get("response") text_containers: list[tuple[dict, str]] = [] for container in (msg_obj, nested_response): - if container is None: + if not isinstance(container, dict): continue if "input" in container: text_containers.append((container, "input")) @@ -1535,13 +1602,14 @@ class ResponsesWebSocketStreaming: modified = True elif isinstance(value, list): for block in value: - if ( - isinstance(block, dict) - and block.get("type") in RESPONSES_WS_MASKABLE_TEXT_BLOCK_TYPES - and isinstance(block.get("text"), str) - ): + if not isinstance(block, dict): + continue + if block.get("type") not in RESPONSES_WS_MASKABLE_TEXT_BLOCK_TYPES: + continue + block_text = block.get("text") + if isinstance(block_text, str): block["text"] = await cb.check_pii( - text=block["text"], + text=block_text, output_parse_pii=True, presidio_config=presidio_config, request_data=self.request_data, @@ -1592,13 +1660,14 @@ class ResponsesWebSocketStreaming: if not self.guardrail_callbacks: return response_str - pii_tokens: Dict[str, str] = (self.request_data.get("metadata") or {}).get("pii_tokens", {}) + metadata = self.request_data.get("metadata") + raw_pii_tokens = metadata.get("pii_tokens") if isinstance(metadata, dict) else None + pii_tokens: Dict[str, str] = raw_pii_tokens if isinstance(raw_pii_tokens, dict) else {} if not pii_tokens: return response_str - try: - evt_obj = json.loads(response_str) - except (json.JSONDecodeError, TypeError): + evt_obj = _parse_json_object(response_str) + if evt_obj is None: return response_str cb = self.guardrail_callbacks[0] @@ -1606,21 +1675,18 @@ class ResponsesWebSocketStreaming: if event_type == "response.completed": modified = False - response_obj = evt_obj.get("response") or {} + response_obj = evt_obj.get("response") if not isinstance(response_obj, dict): return response_str for output_item in response_obj.get("output") or []: if not isinstance(output_item, dict): continue - content = output_item.get("content") or [] - if not isinstance(content, list): - continue - for content_block in content: + for content_block in output_item.get("content") or []: if not isinstance(content_block, dict): continue text = content_block.get("text") if isinstance(text, str): - unmasked = cb._unmask_pii_text(text, pii_tokens) + unmasked = _call_unmask_pii_text(cb, text, pii_tokens) if unmasked != text: content_block["text"] = unmasked modified = True @@ -1629,7 +1695,7 @@ class ResponsesWebSocketStreaming: if event_type in self._DELTA_EVENT_TYPES: delta = evt_obj.get("delta") if isinstance(delta, str): - unmasked = cb._unmask_pii_text(delta, pii_tokens) + unmasked = _call_unmask_pii_text(cb, delta, pii_tokens) if unmasked != delta: evt_obj["delta"] = unmasked return json.dumps(evt_obj) @@ -1651,9 +1717,8 @@ class ResponsesWebSocketStreaming: if not self.output_guardrail_callbacks: return response_str - try: - evt_obj = json.loads(response_str) - except (json.JSONDecodeError, TypeError): + evt_obj = _parse_json_object(response_str) + if evt_obj is None: return response_str if evt_obj.get("type") != "response.completed": @@ -1662,7 +1727,7 @@ class ResponsesWebSocketStreaming: modified = False for cb in self.output_guardrail_callbacks: presidio_config = cb.get_presidio_settings_from_request_data(self.request_data) - response_obj = evt_obj.get("response") or {} + response_obj = evt_obj.get("response") if not isinstance(response_obj, dict): continue for output_item in response_obj.get("output") or []: @@ -1679,26 +1744,21 @@ class ResponsesWebSocketStreaming: if masked_args != arguments: output_item["arguments"] = masked_args modified = True - summary = output_item.get("summary") or [] - if isinstance(summary, list): - for summary_block in summary: - if not isinstance(summary_block, dict): - continue - summary_text = summary_block.get("text") - if isinstance(summary_text, str): - masked_summary = await cb.check_pii( - text=summary_text, - output_parse_pii=False, - presidio_config=presidio_config, - request_data=self.request_data, - ) - if masked_summary != summary_text: - summary_block["text"] = masked_summary - modified = True - content = output_item.get("content") or [] - if not isinstance(content, list): - continue - for content_block in content: + for summary_block in output_item.get("summary") or []: + if not isinstance(summary_block, dict): + continue + summary_text = summary_block.get("text") + if isinstance(summary_text, str): + masked_summary = await cb.check_pii( + text=summary_text, + output_parse_pii=False, + presidio_config=presidio_config, + request_data=self.request_data, + ) + if masked_summary != summary_text: + summary_block["text"] = masked_summary + modified = True + for content_block in output_item.get("content") or []: if not isinstance(content_block, dict): continue text = content_block.get("text") @@ -1722,14 +1782,14 @@ class ResponsesWebSocketStreaming: masked_first = await self._mask_response_create(self.first_message) self._store_input(masked_first) self._store_event(masked_first) - await self.backend_ws.send(masked_first) # type: ignore[union-attr] + await self.backend_ws.send(masked_first) while True: message = await self.websocket.receive_text() masked = await self._mask_response_create(message) self._store_input(masked) self._store_event(masked) - await self.backend_ws.send(masked) # type: ignore[union-attr] + await self.backend_ws.send(masked) except Exception as e: verbose_logger.debug("Responses WS client_to_backend ended: %s", e) @@ -1795,10 +1855,10 @@ class ManagedResponsesWebSocketHandler: def __init__( self, - websocket: Any, + websocket: "WebSocket", model: str, logging_obj: "LiteLLMLoggingObj", - user_api_key_dict: Optional[Any] = None, + user_api_key_dict: Optional["UserAPIKeyAuth"] = None, litellm_metadata: Optional[Dict[str, Any]] = None, api_key: Optional[str] = None, api_base: Optional[str] = None, @@ -1834,13 +1894,11 @@ class ManagedResponsesWebSocketHandler: # ------------------------------------------------------------------ @staticmethod - def _serialize_chunk(chunk: Any) -> Optional[str]: + def _serialize_chunk(chunk: object) -> Optional[str]: """Serialize a streaming chunk to a JSON string for WebSocket transmission.""" try: - if hasattr(chunk, "model_dump_json"): + if isinstance(chunk, BaseModel): return chunk.model_dump_json(exclude_none=True) - if hasattr(chunk, "model_dump"): - return json.dumps(chunk.model_dump(exclude_none=True), default=str) if isinstance(chunk, dict): return json.dumps(chunk, default=str) return json.dumps(str(chunk)) @@ -1877,41 +1935,41 @@ class ManagedResponsesWebSocketHandler: self._session_history[response_id] = messages @staticmethod - def _extract_response_id(completed_event: Dict[str, Any]) -> Optional[str]: + def _extract_response_id(completed_event: Dict[str, object]) -> Optional[str]: """ Pull the raw (decoded) response ID out of a ``response.completed`` event. Returns *None* if the event doesn't contain a usable ID. """ - resp_obj = completed_event.get("response", {}) - encoded_id: Optional[str] = resp_obj.get("id") if isinstance(resp_obj, dict) else None - if not encoded_id: + resp_obj = _as_str_object_dict(completed_event.get("response")) + encoded_id = resp_obj.get("id") if resp_obj is not None else None + if not isinstance(encoded_id, str) or not encoded_id: return None decoded = ResponsesAPIRequestUtils._decode_responses_api_response_id(encoded_id) return decoded.get("response_id", encoded_id) @staticmethod def _extract_output_messages( - completed_event: Dict[str, Any], + completed_event: Dict[str, object], ) -> List[Dict[str, Any]]: """ Convert the output items in a ``response.completed`` event into Responses API message dicts suitable for the next turn's ``input``. """ - resp_obj = completed_event.get("response", {}) - if not isinstance(resp_obj, dict): + resp_obj = _as_str_object_dict(completed_event.get("response")) + if resp_obj is None: return [] messages: List[Dict[str, Any]] = [] - for item in resp_obj.get("output", []) or []: - if not isinstance(item, dict): + for raw_item in _as_object_list(resp_obj.get("output")): + item = _as_str_object_dict(raw_item) + if item is None: continue item_type = item.get("type") role = item.get("role", "assistant") if item_type == "message": - content_parts = item.get("content") or [] text_parts = [ - p.get("text", "") - for p in content_parts - if isinstance(p, dict) and p.get("type") in ("output_text", "text") + str(p.get("text") or "") + for p in (_as_str_object_dict(raw_p) for raw_p in _as_object_list(item.get("content"))) + if p is not None and p.get("type") in ("output_text", "text") ] text = "".join(text_parts) if text: @@ -1927,7 +1985,7 @@ class ManagedResponsesWebSocketHandler: return messages @staticmethod - def _input_to_messages(input_val: Any) -> List[Dict[str, Any]]: + def _input_to_messages(input_val: object) -> List[Dict[str, object]]: """ Normalise the ``input`` field of a ``response.create`` event to a list of Responses API message dicts. @@ -1940,31 +1998,30 @@ class ManagedResponsesWebSocketHandler: "content": [{"type": "input_text", "text": input_val}], } ] - if isinstance(input_val, list): - return [item for item in input_val if isinstance(item, dict)] - return [] + return [item for item in (_as_str_object_dict(v) for v in _as_object_list(input_val)) if item is not None] # ------------------------------------------------------------------ # _process_response_create sub-methods # ------------------------------------------------------------------ - async def _parse_message(self, raw_message: str) -> Optional[Dict[str, Any]]: + async def _parse_message(self, raw_message: str) -> Optional[Dict[str, object]]: """Parse raw WS text; return the message dict or None (JSON error / ignored type).""" try: - msg_obj = json.loads(raw_message) + parsed = json.loads(raw_message) except json.JSONDecodeError: await self._send_error("Invalid JSON in response.create event", "invalid_request_error") return None - if msg_obj.get("type") != "response.create": + msg_obj = _as_str_object_dict(parsed) + if msg_obj is None or msg_obj.get("type") != "response.create": # Silently ignore non-response.create messages (e.g. warmup pings) return None return msg_obj @staticmethod - def _is_warmup_frame(msg_obj: Dict[str, Any]) -> bool: + def _is_warmup_frame(msg_obj: Dict[str, object]) -> bool: """Return True for a response.create whose generate flag is false.""" - nested = msg_obj.get("response") - source = nested if isinstance(nested, dict) and nested else msg_obj + nested = _as_str_object_dict(msg_obj.get("response")) + source = nested if nested else msg_obj return source.get("generate") is False @staticmethod @@ -1977,13 +2034,13 @@ class ManagedResponsesWebSocketHandler: return str(raw_id).startswith(_WARMUP_RESPONSE_ID_PREFIX) @staticmethod - def _warmup_source_params(msg_obj: Dict[str, Any]) -> Dict[str, Any]: - nested = msg_obj.get("response") - if isinstance(nested, dict) and nested: + def _warmup_source_params(msg_obj: Dict[str, object]) -> Dict[str, object]: + nested = _as_str_object_dict(msg_obj.get("response")) + if nested: return nested return {k: v for k, v in msg_obj.items() if k != "type"} - def _build_warmup_response(self, msg_obj: Dict[str, Any]) -> Dict[str, Any]: + def _build_warmup_response(self, msg_obj: Dict[str, object]) -> Dict[str, Any]: """Build a minimal completed Responses API object for a warmup ack.""" source = self._warmup_source_params(msg_obj) wire_model = source.get("model") or self.model_group or self.model @@ -2001,7 +2058,7 @@ class ManagedResponsesWebSocketHandler: }, } - async def _send_warmup_ack(self, msg_obj: Dict[str, Any]) -> None: + async def _send_warmup_ack(self, msg_obj: Dict[str, object]) -> None: """ Acknowledge a generate=false prewarm without calling the provider. @@ -2024,16 +2081,14 @@ class ManagedResponsesWebSocketHandler: await self.websocket.send_text(serialized) @staticmethod - def _build_base_call_kwargs(msg_obj: Dict[str, Any]) -> Dict[str, Any]: + def _build_base_call_kwargs(msg_obj: Dict[str, object]) -> Dict[str, Any]: """ Extract Responses API params from the event, handling both wire formats: Nested: {"type": "response.create", "response": {"input": [...], ...}} Flat: {"type": "response.create", "input": [...], "model": "...", ...} """ - nested = msg_obj.get("response") - response_params: Dict[str, Any] = ( - nested if isinstance(nested, dict) and nested else {k: v for k, v in msg_obj.items() if k != "type"} - ) + nested = _as_str_object_dict(msg_obj.get("response")) + response_params: Dict[str, object] = nested if nested else {k: v for k, v in msg_obj.items() if k != "type"} return { param: response_params[param] for param in _RESPONSE_CREATE_PARAMS @@ -2133,7 +2188,7 @@ class ManagedResponsesWebSocketHandler: call_kwargs.setdefault("litellm_params", {}) call_kwargs["litellm_params"]["proxy_server_request"] = proxy_server_request - async def _stream_and_forward(self, model: str, call_kwargs: Dict[str, Any]) -> Optional[Dict[str, Any]]: + async def _stream_and_forward(self, model: str, call_kwargs: Dict[str, Any]) -> Optional[Dict[str, object]]: """ Stream ``litellm.aresponses`` and forward every chunk over the WebSocket. @@ -2141,7 +2196,7 @@ class ManagedResponsesWebSocketHandler: directly (before serialization) to avoid a redundant JSON round-trip on every chunk. Returns the completed event dict, or ``None``. """ - completed_event: Optional[Dict[str, Any]] = None + completed_event: Optional[Dict[str, object]] = None stream_response = await litellm.aresponses(model=model, **call_kwargs) async for chunk in stream_response: # type: ignore[union-attr] if chunk is None: @@ -2153,7 +2208,7 @@ class ManagedResponsesWebSocketHandler: continue if chunk_type == "response.completed" and completed_event is None: try: - completed_event = json.loads(serialized) + completed_event = _as_str_object_dict(json.loads(serialized)) except Exception: pass try: @@ -2165,7 +2220,7 @@ class ManagedResponsesWebSocketHandler: def _save_turn_history( self, - completed_event: Optional[Dict[str, Any]], + completed_event: Optional[Dict[str, object]], prior_history: List[Dict[str, Any]], current_messages: List[Dict[str, Any]], ) -> None: