chore(typing): wip snapshot of basedpyright Any cleanup (streaming_iterator)

Intermediate checkpoint while a background pass fixes LIT001
(mutable-annotation) fallout in streaming_iterator.py introduced by the
Any-cleanup itself. Not yet passing make pre-commit; follow-up commit will
finish this.
This commit is contained in:
mateo-berri 2026-07-31 12:09:34 +00:00
parent 5bd9e12c85
commit 4051c6c4ff
No known key found for this signature in database
3 changed files with 258 additions and 197 deletions

View file

@ -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:

View file

@ -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:

View file

@ -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: