From 23bd48f357146abbb7d50d08490ffa5acae08f98 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 10 Aug 2026 13:48:46 +0000 Subject: [PATCH] chore(typing): clear basedpyright Any errors in responses streaming and websearch interception Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../websearch_interception/handler.py | 80 +++++++--- litellm/responses/streaming_iterator.py | 141 ++++++++++++------ 2 files changed, 151 insertions(+), 70 deletions(-) diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 2c5f7484dac..25c83fd7235 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -10,7 +10,10 @@ import asyncio import math import uuid from collections.abc import AsyncIterator, Mapping, Sequence -from typing import TYPE_CHECKING, Any, Final, cast +from typing import TYPE_CHECKING, Any, Final, TypeVar, cast + +from pydantic import TypeAdapter, ValidationError +from typing_extensions import TypeIs import litellm from litellm._logging import verbose_logger @@ -72,6 +75,43 @@ WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY: Final = "_websearch_interception_emit_native_b # ``web_search_tool_result`` blocks to inject into the final response. WEBSEARCH_NATIVE_BLOCKS_METADATA_KEY: Final = "websearch_native_blocks" +_ResponseT = TypeVar("_ResponseT") + +_CONTENT_ATTR: Final = "content" + +_OBJECT_ITEMS_ADAPTER: Final[TypeAdapter[tuple[object, ...]]] = TypeAdapter(tuple[object, ...]) + +_NATIVE_BLOCKS_ADAPTER: Final[TypeAdapter[tuple[Mapping[str, object], ...]]] = TypeAdapter( + tuple[Mapping[str, object], ...] +) + +_WEBSEARCH_CONFIG_ADAPTER: Final[TypeAdapter[WebSearchInterceptionConfig]] = TypeAdapter(WebSearchInterceptionConfig) + + +def _is_json_object(value: object) -> TypeIs[dict[str, object]]: # guard-ok: trivial isinstance; JSON keys are str + return isinstance(value, dict) + + +def _parse_object_items(value: object) -> tuple[object, ...]: + try: + return _OBJECT_ITEMS_ADAPTER.validate_python(value) + except ValidationError: + return () + + +def _parse_native_blocks(value: object) -> tuple[Mapping[str, object], ...]: + try: + return _NATIVE_BLOCKS_ADAPTER.validate_python(value) + except ValidationError: + return () + + +def _parse_websearch_config(value: object) -> WebSearchInterceptionConfig: + try: + return _WEBSEARCH_CONFIG_ADAPTER.validate_python(value) + except ValidationError: + return {} + class WebSearchInterceptionLogger(CustomLogger): """ @@ -833,7 +873,7 @@ class WebSearchInterceptionLogger(CustomLogger): Anthropic-native clients (Claude Desktop, the Anthropic SDK) can render citations / sources alongside the model's textual reply. """ - native_blocks: Final = plan.metadata.get(WEBSEARCH_NATIVE_BLOCKS_METADATA_KEY) + native_blocks: Final = _parse_native_blocks(plan.metadata.get(WEBSEARCH_NATIVE_BLOCKS_METADATA_KEY)) if not native_blocks: return response return self._inject_native_blocks(response, native_blocks) @@ -883,17 +923,20 @@ class WebSearchInterceptionLogger(CustomLogger): ) @staticmethod - def _inject_native_blocks(response: Any, native_blocks: Sequence[Mapping[str, object]]) -> Any: + def _inject_native_blocks( + response: _ResponseT, + native_blocks: Sequence[Mapping[str, object]], + ) -> _ResponseT: """Prepend native blocks to response content, dict or object form.""" if not native_blocks: return response - if isinstance(response, dict): - existing = response.get("content") or [] - response["content"] = list(native_blocks) + list(existing) + if _is_json_object(response): + existing_items: Final = _parse_object_items(response.get(_CONTENT_ATTR)) + response[_CONTENT_ATTR] = [*native_blocks, *existing_items] return response - existing = getattr(response, "content", None) or [] + existing_attr_items: Final = _parse_object_items(getattr(response, _CONTENT_ATTR, None)) try: - response.content = list(native_blocks) + list(existing) + setattr(response, _CONTENT_ATTR, [*native_blocks, *existing_attr_items]) except (AttributeError, TypeError): # Object refused write — fall through and leave the response # untouched rather than crash the request. @@ -1675,8 +1718,8 @@ class WebSearchInterceptionLogger(CustomLogger): @staticmethod def initialize_from_proxy_config( - litellm_settings: dict[str, Any], - callback_specific_params: dict[str, Any], + litellm_settings: Mapping[str, object], + callback_specific_params: Mapping[str, object], ) -> "WebSearchInterceptionLogger": """ Static method to initialize WebSearchInterceptionLogger from proxy config. @@ -1698,16 +1741,11 @@ class WebSearchInterceptionLogger(CustomLogger): ) """ # Get websearch_interception_params from litellm_settings or callback_specific_params - websearch_params: WebSearchInterceptionConfig = {} - if "websearch_interception_params" in litellm_settings: - websearch_params = litellm_settings["websearch_interception_params"] - elif "websearch_interception" in callback_specific_params and isinstance( - callback_specific_params["websearch_interception"], dict - ): - websearch_params = cast( - WebSearchInterceptionConfig, - callback_specific_params["websearch_interception"], - ) + raw_params: Final = ( + litellm_settings["websearch_interception_params"] + if "websearch_interception_params" in litellm_settings + else callback_specific_params.get("websearch_interception") + ) # Use classmethod to initialize from config - return WebSearchInterceptionLogger.from_config_yaml(websearch_params) + return WebSearchInterceptionLogger.from_config_yaml(_parse_websearch_config(raw_params)) diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 2e1e1a44594..031e42eebd1 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -9,10 +9,11 @@ from collections.abc import Awaitable, Callable, Mapping from datetime import datetime from functools import lru_cache from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Literal +from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, runtime_checkable import httpx from openai._streaming import SSEDecoder +from pydantic import TypeAdapter, ValidationError from typing_extensions import TypeIs import litellm @@ -33,6 +34,7 @@ from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfi from litellm.responses.utils import ResponseAPILoggingUtils, ResponsesAPIRequestUtils from litellm.types.llms.openai import ( PART_UNION_TYPES, + OpenAIChatCompletionLogprobsContent, ResponseAPIUsage, ResponsesAPIResponse, ResponsesAPIStreamEvents, @@ -57,6 +59,21 @@ def _get_openai_response_types(): return openai_types +@runtime_checkable +class _SupportsModelDump(Protocol): + def model_dump(self) -> Mapping[str, object]: ... + + +@runtime_checkable +class _SupportsModelDumpJson(Protocol): + def model_dump_json(self, *, exclude_none: bool = ...) -> str: ... + + +@runtime_checkable +class _SupportsModelDumpExcludeNone(Protocol): + def model_dump(self, *, exclude_none: bool = ...) -> Mapping[str, object]: ... + + def _is_json_object(value: object) -> TypeIs[dict[str, object]]: # guard-ok: trivial isinstance; JSON keys are str return isinstance(value, dict) @@ -69,7 +86,30 @@ def _is_str_mapping(value: object) -> TypeIs[dict[str, str]]: # guard-ok: verif return _is_json_object(value) and all(isinstance(item, str) for item in value.values()) -def _model_id_from_metadata(litellm_metadata: dict[str, object] | None) -> str | None: +def _json_array_items(value: object) -> tuple[object, ...]: + return tuple(value) if _is_json_array(value) else () + + +def _json_object_items(value: object) -> tuple[dict[str, object], ...]: + return tuple(item for item in _json_array_items(value) if _is_json_object(item)) + + +_LOGPROBS_ADAPTER: Final[TypeAdapter[list[OpenAIChatCompletionLogprobsContent] | None]] = TypeAdapter( + list[OpenAIChatCompletionLogprobsContent] | None +) + +_OBJECT_ADAPTER: Final[TypeAdapter[object]] = TypeAdapter(object) + +_JSON_OBJECT_ADAPTER: Final[TypeAdapter[dict[str, object]]] = TypeAdapter(dict[str, object]) + +_UNMASK_PII_TEXT_ATTR: Final = "_unmask_pii_text" + +_UnmaskPiiText = Callable[[str, dict[str, str]], str] + +_UNMASK_PII_TEXT_ADAPTER: Final[TypeAdapter[_UnmaskPiiText]] = TypeAdapter(_UnmaskPiiText) + + +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 @@ -228,10 +268,10 @@ class BaseResponsesAPIStreamingIterator: try: # Parse the JSON chunk - parsed_chunk: Final = json.loads(chunk) + parsed_chunk: Final = _OBJECT_ADAPTER.validate_python(json.loads(chunk)) # Format as ResponsesAPIStreamingResponse - if isinstance(parsed_chunk, dict): + if _is_json_object(parsed_chunk): if self.responses_api_provider_config is None: raise ValueError("responses_api_provider_config is required to process live streaming chunks") openai_responses_api_chunk: Final = self.responses_api_provider_config.transform_streaming_response( @@ -1029,8 +1069,8 @@ class CachedResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): return evt -def _dump_response_object(obj: Any) -> dict[str, Any]: - if hasattr(obj, "model_dump"): +def _dump_response_object(obj: object) -> Mapping[str, object]: + if isinstance(obj, _SupportsModelDump): return obj.model_dump() if _is_json_object(obj): return obj @@ -1059,7 +1099,7 @@ def _build_content_part_done_event( item_id: str, output_index: int, content_index: int, - part_payload: dict[str, Any], + part_payload: Mapping[str, object], ) -> ResponsesAPIStreamingResponse | None: openai_types: Final = _get_openai_response_types() part_type: Final = part_payload.get("type") @@ -1067,13 +1107,13 @@ def _build_content_part_done_event( if part_type == "output_text": annotations: Final = [ openai_types.BaseLiteLLMOpenAIResponseObject(**annotation) - for annotation in part_payload.get("annotations", []) or [] + for annotation in _json_object_items(part_payload.get("annotations")) ] part = openai_types.ContentPartDonePartOutputText( type="output_text", text=str(part_payload.get("text") or ""), annotations=annotations, - logprobs=part_payload.get("logprobs"), + logprobs=_LOGPROBS_ADAPTER.validate_python(part_payload.get("logprobs")), ) elif part_type == "refusal": part = openai_types.ContentPartDonePartRefusal( @@ -1103,7 +1143,7 @@ def _add_text_like_part_events( item_id: str, output_index: int, content_index: int, - part_payload: dict[str, Any], + part_payload: Mapping[str, object], chunk_size: int, ) -> None: openai_types: Final = _get_openai_response_types() @@ -1120,7 +1160,7 @@ def _add_text_like_part_events( delta=text[i : i + chunk_size], ) ) - for annotation_index, annotation in enumerate(part_payload.get("annotations", []) or []): + for annotation_index, annotation in enumerate(_json_object_items(part_payload.get("annotations"))): events.append( openai_types.OutputTextAnnotationAddedEvent( type=openai_types.ResponsesAPIStreamEvents.OUTPUT_TEXT_ANNOTATION_ADDED, @@ -1186,7 +1226,7 @@ def _build_synthetic_response_events( ] sequence_number = 0 - for output_index, output_item in enumerate(getattr(transformed, "output", []) or []): + for output_index, output_item in enumerate(_json_array_items(getattr(transformed, "output", None))): 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") @@ -1200,7 +1240,7 @@ def _build_synthetic_response_events( ) if item_type == "message": - for content_index, part in enumerate(output_item_payload.get("content", []) or []): + for content_index, part in enumerate(_json_array_items(output_item_payload.get("content"))): part_payload = _dump_response_object(part) events.append( openai_types.ContentPartAddedEvent( @@ -1247,7 +1287,7 @@ def _build_synthetic_response_events( ) ) elif item_type == "reasoning": - for summary_index, summary in enumerate(output_item_payload.get("summary", []) or []): + for summary_index, summary in enumerate(_json_array_items(output_item_payload.get("summary"))): summary_payload = _dump_response_object(summary) summary_text = str(summary_payload.get("text") or "") for i in range(0, len(summary_text), chunk_size): @@ -1340,7 +1380,7 @@ class ResponsesWebSocketStreaming: user_api_key_dict: UserAPIKeyAuth | None = None, request_data: dict[str, object] | None = None, first_message: str | None = None, - guardrail_callbacks: list[Any] | None = None, + guardrail_callbacks: list[PresidioGuardrailCallback] | None = None, output_guardrail_callbacks: list[PresidioGuardrailCallback] | None = None, authorized_model: str | None = None, ): @@ -1352,7 +1392,7 @@ class ResponsesWebSocketStreaming: self.messages: list[dict[str, object]] = [] self.input_messages: list[dict[str, object]] = [] self.first_message = first_message - self.guardrail_callbacks: list[Any] = guardrail_callbacks or [] + self.guardrail_callbacks: list[PresidioGuardrailCallback] = guardrail_callbacks or [] self.output_guardrail_callbacks: list[PresidioGuardrailCallback] = output_guardrail_callbacks or [] # Model name authorized at connection time; enforced on every # response.create frame to prevent deployment-substitution attacks. @@ -1362,24 +1402,22 @@ class ResponsesWebSocketStreaming: return event_obj.get("type") in RESPONSES_WS_LOGGED_EVENT_TYPES def _store_event(self, event: str | bytes | dict[str, object]) -> None: - if isinstance(event, bytes): - event = event.decode("utf-8") - if isinstance(event, str): + if isinstance(event, str | bytes): try: - event_obj = json.loads(event) - except (json.JSONDecodeError, TypeError): + parsed_event: dict[str, object] = _JSON_OBJECT_ADAPTER.validate_json(event) + except ValidationError: return else: - event_obj = event + parsed_event = event - if self._should_store_event(event_obj): - self.messages.append(event_obj) + if self._should_store_event(parsed_event): + self.messages.append(parsed_event) def _collect_input_from_client_event(self, message: object) -> None: """Extract user input content from response.create for logging.""" try: if isinstance(message, str): - msg_obj = json.loads(message) + msg_obj: dict[str, object] = _JSON_OBJECT_ADAPTER.validate_json(message) elif _is_json_object(message): msg_obj = message else: @@ -1407,7 +1445,7 @@ class ResponsesWebSocketStreaming: text = c.get("text", "") if text: self.input_messages.append({"role": "user", "content": text}) - except (json.JSONDecodeError, AttributeError, TypeError): + except (ValidationError, AttributeError, TypeError): pass def _store_input(self, message: object) -> None: @@ -1449,8 +1487,8 @@ class ResponsesWebSocketStreaming: # masked response.completed. if self.output_guardrail_callbacks: try: - _evt_type = json.loads(response_str).get("type") - except (json.JSONDecodeError, TypeError): + _evt_type = _JSON_OBJECT_ADAPTER.validate_json(response_str).get("type") + except ValidationError: _evt_type = None if _evt_type in self._DELTA_EVENT_TYPES or _evt_type in self._OUTPUT_DONE_EVENT_TYPES: continue @@ -1513,8 +1551,8 @@ class ResponsesWebSocketStreaming: Non-``response.create`` messages are returned unchanged. """ try: - msg_obj: Final = json.loads(message) - except (json.JSONDecodeError, TypeError): + msg_obj: Final = _JSON_OBJECT_ADAPTER.validate_json(message) + except ValidationError: return message if msg_obj.get("type") != "response.create": @@ -1641,11 +1679,12 @@ class ResponsesWebSocketStreaming: return response_str try: - evt_obj: Final = json.loads(response_str) - except (json.JSONDecodeError, TypeError): + evt_obj: Final = _JSON_OBJECT_ADAPTER.validate_json(response_str) + except ValidationError: return response_str cb: Final = self.guardrail_callbacks[0] + unmask_pii_text: Final = _UNMASK_PII_TEXT_ADAPTER.validate_python(getattr(cb, _UNMASK_PII_TEXT_ATTR)) event_type: Final = evt_obj.get("type") if event_type == "response.completed": @@ -1665,7 +1704,7 @@ class ResponsesWebSocketStreaming: continue text = content_block.get("text") if isinstance(text, str): - unmasked = cb._unmask_pii_text(text, pii_tokens) + unmasked = unmask_pii_text(text, pii_tokens) if unmasked != text: content_block["text"] = unmasked modified = True @@ -1674,7 +1713,7 @@ class ResponsesWebSocketStreaming: if event_type in self._DELTA_EVENT_TYPES: delta: Final = evt_obj.get("delta") if isinstance(delta, str): - unmasked = cb._unmask_pii_text(delta, pii_tokens) + unmasked = unmask_pii_text(delta, pii_tokens) if unmasked != delta: evt_obj["delta"] = unmasked return json.dumps(evt_obj) @@ -1697,8 +1736,8 @@ class ResponsesWebSocketStreaming: return response_str try: - evt_obj: Final[Mapping[str, object]] = json.loads(response_str) - except (json.JSONDecodeError, TypeError): + evt_obj: Final[Mapping[str, object]] = _JSON_OBJECT_ADAPTER.validate_json(response_str) + except ValidationError: return response_str if evt_obj.get("type") != "response.completed": @@ -1845,7 +1884,7 @@ class ManagedResponsesWebSocketHandler: model: str, logging_obj: LiteLLMLoggingObj, user_api_key_dict: UserAPIKeyAuth | None = None, - litellm_metadata: dict[str, Any] | None = None, + litellm_metadata: dict[str, object] | None = None, api_key: str | None = None, api_base: str | None = None, timeout: float | None = None, @@ -1857,10 +1896,11 @@ class ManagedResponsesWebSocketHandler: self.model = model self.logging_obj = logging_obj self.user_api_key_dict = user_api_key_dict - self.litellm_metadata: dict[str, Any] = litellm_metadata or {} - self.model_group: str | None = self.litellm_metadata.get("model_group") or self.litellm_metadata.get( + self.litellm_metadata: dict[str, object] = litellm_metadata or {} + raw_model_group: Final = self.litellm_metadata.get("model_group") or self.litellm_metadata.get( "deployment_model_name" ) + self.model_group: str | None = raw_model_group if isinstance(raw_model_group, str) else None self.api_key = api_key self.api_base = api_base self.timeout = timeout @@ -1880,14 +1920,14 @@ class ManagedResponsesWebSocketHandler: # ------------------------------------------------------------------ @staticmethod - def _serialize_chunk(chunk: Any) -> str | None: + def _serialize_chunk(chunk: object) -> str | None: """Serialize a streaming chunk to a JSON string for WebSocket transmission.""" try: - if hasattr(chunk, "model_dump_json"): + if isinstance(chunk, _SupportsModelDumpJson): return chunk.model_dump_json(exclude_none=True) - if hasattr(chunk, "model_dump"): + if isinstance(chunk, _SupportsModelDumpExcludeNone): return json.dumps(chunk.model_dump(exclude_none=True), default=str) - if isinstance(chunk, dict): + if _is_json_object(chunk): return json.dumps(chunk, default=str) return json.dumps(str(chunk)) except Exception as exc: @@ -1998,8 +2038,8 @@ class ManagedResponsesWebSocketHandler: async def _parse_message(self, raw_message: str) -> dict[str, object] | None: """Parse raw WS text; return the message dict or None (JSON error / ignored type).""" try: - msg_obj: Final = json.loads(raw_message) - except json.JSONDecodeError: + msg_obj: Final = _JSON_OBJECT_ADAPTER.validate_json(raw_message) + except ValidationError: await self._send_error("Invalid JSON in response.create event", "invalid_request_error") return None if msg_obj.get("type") != "response.create": @@ -2196,14 +2236,17 @@ class ManagedResponsesWebSocketHandler: if chunk is None: continue # Read type from the object before serializing to avoid double JSON parse - chunk_type = getattr(chunk, "type", None) or (chunk.get("type") if isinstance(chunk, dict) else None) - serialized = self._serialize_chunk(chunk) + chunk_obj = _OBJECT_ADAPTER.validate_python(chunk) + chunk_type = getattr(chunk_obj, "type", None) or ( + chunk_obj.get("type") if _is_json_object(chunk_obj) else None + ) + serialized = self._serialize_chunk(chunk_obj) if serialized is None: continue if chunk_type == "response.completed" and completed_event is None: try: - completed_event = json.loads(serialized) - except Exception: + completed_event = _JSON_OBJECT_ADAPTER.validate_json(serialized) + except ValidationError: pass try: await self.websocket.send_text(serialized)