From 3018abc77c093fbbf5fe781f796a30838655178d Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 11 Aug 2026 13:45:53 +0000 Subject: [PATCH] chore(typing): clear basedpyright Any errors in websearch interception and anthropic guardrail translation Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../websearch_interception/handler.py | 99 +++- .../chat/guardrail_translation/handler.py | 436 +++++++++--------- 2 files changed, 304 insertions(+), 231 deletions(-) diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index f7f27459768..5077528af50 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -10,7 +10,7 @@ import asyncio import math import uuid from collections.abc import AsyncIterator, Mapping, Sequence -from typing import TYPE_CHECKING, Any, Final, TypedDict, cast +from typing import TYPE_CHECKING, Any, Final, Protocol, TypedDict, TypeVar, cast, runtime_checkable import litellm from litellm._logging import verbose_logger @@ -85,6 +85,26 @@ class _WebSearchSettingsView(TypedDict): websearch_interception_params: WebSearchInterceptionConfig +_ResponseT = TypeVar("_ResponseT") + + +@runtime_checkable +class _MappingLike(Protocol): + def get(self, key: str, /) -> object | None: ... + + +@runtime_checkable +class _MutableMappingLike(Protocol): + def get(self, key: str, /) -> object | None: ... + + def __setitem__(self, key: str, value: object, /) -> None: ... + + +@runtime_checkable +class _HasContent(Protocol): + content: object + + class WebSearchInterceptionLogger(CustomLogger): """ CustomLogger that intercepts WebSearch tool calls for models that don't @@ -260,7 +280,32 @@ class WebSearchInterceptionLogger(CustomLogger): ) return response - async def async_pre_call_deployment_hook(self, kwargs: dict[str, Any], call_type: CallTypes | None) -> dict | None: + @staticmethod + def _mapping_str(container: object, key: str) -> str: + """String value at ``key`` when ``container`` is mapping-like, else empty.""" + if not isinstance(container, _MappingLike): + return "" + value: Final = container.get(key) + return value if isinstance(value, str) else "" + + @staticmethod + def _tool_list(value: object) -> Sequence[dict[str, object]]: + """The request's tool list, empty when absent or not a list.""" + return value if isinstance(value, list) else () + + @staticmethod + def _provider_from_model(kwargs: Mapping[str, object]) -> str: + """Provider derived from the request's model name, empty when undeterminable.""" + model: Final = kwargs.get("model") + try: + _, custom_llm_provider, _, _ = litellm.get_llm_provider(model=model if isinstance(model, str) else "") + except Exception: + return "" + return custom_llm_provider + + async def async_pre_call_deployment_hook( + self, kwargs: dict[str, object], call_type: CallTypes | None + ) -> dict | None: """ Pre-call hook to convert native Anthropic web_search tools to regular tools. @@ -270,19 +315,16 @@ class WebSearchInterceptionLogger(CustomLogger): """ # Check if this is for an enabled provider # Try top-level kwargs first, then nested litellm_params, then derive from model name - custom_llm_provider = kwargs.get("custom_llm_provider", "") or kwargs.get("litellm_params", {}).get( - "custom_llm_provider", "" + custom_llm_provider: Final = ( + self._mapping_str(kwargs, "custom_llm_provider") + or self._mapping_str(kwargs.get("litellm_params"), "custom_llm_provider") + or self._provider_from_model(kwargs) ) - if not custom_llm_provider: - try: - _, custom_llm_provider, _, _ = litellm.get_llm_provider(model=kwargs.get("model", "")) - except Exception: - custom_llm_provider = "" if custom_llm_provider not in self.enabled_providers: return None # Check if request has tools with native web_search - tools: Final[Sequence[dict[str, object]] | None] = kwargs.get("tools") + tools: Final = self._tool_list(kwargs.get("tools")) if not tools: return None @@ -898,24 +940,31 @@ class WebSearchInterceptionLogger(CustomLogger): ) @staticmethod - def _inject_native_blocks(response: Any, native_blocks: Sequence[Mapping[str, object]]) -> Any: + def _prepended_content(existing: object, native_blocks: Sequence[Mapping[str, object]]) -> list[object]: + return [*native_blocks, *existing] if isinstance(existing, (list, tuple)) else list(native_blocks) + + @staticmethod + 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) - return response - existing = getattr(response, "content", None) or [] - try: - response.content = list(native_blocks) + list(existing) - except (AttributeError, TypeError): - # Object refused write — fall through and leave the response - # untouched rather than crash the request. - verbose_logger.debug( - "WebSearchInterception: could not inject native blocks into response of type %s", - type(response).__name__, + if isinstance(response, _MutableMappingLike) and isinstance(response, dict): + response["content"] = WebSearchInterceptionLogger._prepended_content( + response.get("content"), native_blocks ) + return response + if isinstance(response, _HasContent): + try: + response.content = WebSearchInterceptionLogger._prepended_content(response.content, native_blocks) + return response + except (AttributeError, TypeError): + pass + # Object refused write — fall through and leave the response + # untouched rather than crash the request. + verbose_logger.debug( + "WebSearchInterception: could not inject native blocks into response of type %s", + type(response).__name__, + ) return response async def async_run_chat_completion_agentic_loop( @@ -1417,7 +1466,7 @@ class WebSearchInterceptionLogger(CustomLogger): valid_token=user_api_key_auth, ) - team_id: Final = getattr(user_api_key_auth, "team_id", None) + team_id: Final = user_api_key_auth.team_id if team_id: from litellm.proxy.proxy_server import ( prisma_client, diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py index e4a4d23b438..0d65d4058cd 100644 --- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py @@ -13,11 +13,12 @@ Pattern Overview: """ import json -from collections.abc import Mapping +from collections.abc import Iterable, Mapping, MutableMapping, Sequence from copy import deepcopy from dataclasses import dataclass -from typing import TYPE_CHECKING, Any, Final, cast +from typing import TYPE_CHECKING, Final, Protocol, cast, runtime_checkable +from pydantic import TypeAdapter, ValidationError from typing_extensions import assert_never from litellm._logging import verbose_proxy_logger @@ -61,6 +62,7 @@ if TYPE_CHECKING: ModifyResponseException, ) from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.proxy._types import UserAPIKeyAuth from litellm.types.llms.anthropic_messages.anthropic_response import ( AnthropicMessagesResponse, ) @@ -109,6 +111,58 @@ class ExtractedInput: EMPTY_EXTRACTED_INPUT: Final = ExtractedInput(scanned=(), images=()) +_JSON_OBJECT_ADAPTER: Final = TypeAdapter(dict[str, object]) + + +def _parse_json_object(payload: str) -> tuple[dict[str, object], ...]: + """The JSON object carried by an SSE ``data:`` payload, empty when it is not one.""" + try: + return (_JSON_OBJECT_ADAPTER.validate_json(payload),) + except ValidationError: + return () + + +@runtime_checkable +class _MappingLike(Protocol): + def get(self, key: str, /) -> object | None: ... + + +@runtime_checkable +class _MutableMappingLike(Protocol): + def get(self, key: str, /) -> object | None: ... + + def __setitem__(self, key: str, value: object, /) -> None: ... + + +@runtime_checkable +class _ItemsLike(Protocol): + def items(self) -> Iterable[tuple[str, object]]: ... + + +@runtime_checkable +class _HasContent(Protocol): + content: object + + +@runtime_checkable +class _HasModel(Protocol): + model: object + + +@runtime_checkable +class _HasType(Protocol): + type: object + + +@runtime_checkable +class _HasText(Protocol): + text: object + + +@runtime_checkable +class _HasModelDump(Protocol): + def model_dump(self) -> dict[str, object]: ... + class AnthropicMessagesHandler(BaseTranslation): """Process Anthropic messages with guardrails. @@ -123,7 +177,7 @@ class AnthropicMessagesHandler(BaseTranslation): @staticmethod def _build_streaming_usage_response( - responses_so_far: list[Any], + responses_so_far: Sequence[object], request_data: dict | None, ) -> ModelResponse | None: chunks: Final = tuple(response for response in responses_so_far if isinstance(response, (str, bytes))) @@ -141,7 +195,7 @@ class AnthropicMessagesHandler(BaseTranslation): self, exc: "ModifyResponseException", stream_started: bool = False, - responses_so_far: list[Any] | None = None, + responses_so_far: Sequence[object] | None = None, ) -> list[bytes]: """ Build an Anthropic SSE sequence delivering the guardrail block message @@ -159,7 +213,7 @@ class AnthropicMessagesHandler(BaseTranslation): would make Anthropic clients reject the stream. """ if stream_started: - return self._block_continuation_chunks(exc, responses_so_far or []) + return self._block_continuation_chunks(exc, responses_so_far or ()) return self._standalone_block_chunks(exc) def _standalone_block_chunks(self, exc: "ModifyResponseException") -> list[bytes]: @@ -184,7 +238,9 @@ class AnthropicMessagesHandler(BaseTranslation): ) return list(FakeAnthropicMessagesStreamIterator(response=block_response)) - def _block_continuation_chunks(self, exc: "ModifyResponseException", responses_so_far: list[Any]) -> list[bytes]: + def _block_continuation_chunks( + self, exc: "ModifyResponseException", responses_so_far: Sequence[object] + ) -> list[bytes]: """Continue an already-started message: close the open content block, append the block message as a new text block, then end the message -- without a second message_start.""" @@ -193,7 +249,7 @@ class AnthropicMessagesHandler(BaseTranslation): blocked_response_usage, ) - def _sse(event_type: str, payload: dict) -> bytes: + def _sse(event_type: str, payload: Mapping[str, object]) -> bytes: return f"event: {event_type}\ndata: {json.dumps(payload)}\n\n".encode() output_tokens: Final = blocked_response_usage(getattr(exc, "original_response", None))["output_tokens"] @@ -234,7 +290,7 @@ class AnthropicMessagesHandler(BaseTranslation): @staticmethod def _content_block_state( - responses_so_far: list[Any], + responses_so_far: Sequence[object], ) -> tuple[int | None, int | None]: """From the SSE chunks already sent to the client, return (open content-block index or None, highest content-block index seen or None). @@ -260,30 +316,24 @@ class AnthropicMessagesHandler(BaseTranslation): return open_index, max_index @staticmethod - def _iter_sse_events(item: Any) -> list[dict]: + def _iter_sse_events(item: object) -> tuple[Mapping[str, object], ...]: """Yield the event-data dicts in one stream chunk. Handles both formats this stream can carry (see ``get_streaming_string_so_far``): raw SSE ``bytes`` -- which may bundle several events separated by a blank line -- and an already-parsed event ``dict``.""" - if isinstance(item, dict): - return [item] + if isinstance(item, _ItemsLike): + return (dict(item.items()),) if not isinstance(item, (bytes, bytearray)): - return [] - events: Final[list[dict]] = [] - for block in item.decode("utf-8", errors="replace").split("\n\n"): - for line in block.split("\n"): - line = line.strip() - if not line.startswith("data:"): - continue - try: - parsed = json.loads(line[len("data:") :].strip()) - except json.JSONDecodeError: - continue - if isinstance(parsed, dict): - events.append(parsed) - return events + return () + return tuple( + parsed + for block in item.decode("utf-8", errors="replace").split("\n\n") + for line in block.split("\n") + if (stripped := line.strip()).startswith("data:") + for parsed in _parse_json_object(stripped[len("data:") :].strip()) + ) def _translate_to_openai(self, data: dict) -> ChatCompletionRequest: """Translate Anthropic request to OpenAI chat completion format.""" @@ -315,8 +365,8 @@ class AnthropicMessagesHandler(BaseTranslation): self, data: dict, guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Any | None = None, - ) -> Any: + litellm_logging_obj: "LiteLLMLoggingObj | None" = None, + ) -> dict: """ Process input messages by applying guardrails to text content. """ @@ -467,8 +517,8 @@ class AnthropicMessagesHandler(BaseTranslation): @staticmethod def _openai_system_message_to_anthropic( - message: dict[str, Any], - ) -> dict[str, Any] | None: # mutable-ok: API message payload + message: Mapping[str, object], + ) -> dict[str, object] | None: # mutable-ok: API message payload """Convert an OpenAI system message to the client's Anthropic-shaped entry.""" content: Final = message.get("content") if isinstance(content, str): @@ -477,14 +527,14 @@ class AnthropicMessagesHandler(BaseTranslation): ) # mutable-ok: API message payload if not isinstance(content, list): return None - blocks: Final[list[dict[str, Any]]] = [] # mutable-ok: API message payload + blocks: Final[list[dict[str, object]]] = [] # mutable-ok: API message payload for block in content: if not isinstance(block, dict) or block.get("type") != "text": continue text = block.get("text") if not isinstance(text, str) or not text: continue - anthropic_block: dict[str, Any] = { # mutable-ok: API message payload + anthropic_block: dict[str, object] = { # mutable-ok: API message payload "type": "text", "text": text, } # mutable-ok: API message payload @@ -602,7 +652,7 @@ class AnthropicMessagesHandler(BaseTranslation): @staticmethod def _extract_midturn_system_text( - message: dict[str, Any], # mutable-ok: API message payload + message: Mapping[str, object], msg_idx: int, ) -> ExtractedInput: """Match the adapter's filtering so positional guardrail write-back stays aligned.""" @@ -636,7 +686,7 @@ class AnthropicMessagesHandler(BaseTranslation): @classmethod def _extract_input_text_and_images( cls, - message: dict[str, Any], + message: Mapping[str, object], msg_idx: int, skip_system_message: bool = False, skip_tool_message: bool = False, @@ -682,7 +732,7 @@ class AnthropicMessagesHandler(BaseTranslation): @classmethod def _extract_content_block( cls, - content_item: Mapping[str, Any], + content_item: Mapping[str, object], msg_idx: int, content_idx: int, skip_tool_message: bool, @@ -699,7 +749,9 @@ class AnthropicMessagesHandler(BaseTranslation): text_str: Final = content_item.get("text", None) return ExtractedInput( scanned=( - () if text_str is None else (ScannedText(text_str, ContentBlockTextTarget(msg_idx, content_idx)),) + (ScannedText(text_str, ContentBlockTextTarget(msg_idx, content_idx)),) + if isinstance(text_str, str) + else () ), images=cls._image_sources(content_item) if content_item.get("type") == "image" else (), ) @@ -707,7 +759,7 @@ class AnthropicMessagesHandler(BaseTranslation): @classmethod def _extract_tool_result( cls, - content_item: Mapping[str, Any], + content_item: Mapping[str, object], msg_idx: int, content_idx: int, ) -> ExtractedInput: @@ -736,18 +788,18 @@ class AnthropicMessagesHandler(BaseTranslation): ) @staticmethod - def _image_sources(block: Mapping[str, Any]) -> tuple[str, ...]: + def _image_sources(block: Mapping[str, object]) -> tuple[str, ...]: source: Final = block.get("source") - if not isinstance(source, Mapping): + if not isinstance(source, _MappingLike): return () # Could be base64 or url data: Final = source.get("data") - return (data,) if data else () + return (data,) if isinstance(data, str) and data else () async def _apply_guardrail_responses_to_input( self, - messages: list[dict[str, Any]], - responses: list[str], + messages: Sequence[MutableMapping[str, object]], + responses: Sequence[str], scanned: tuple[ScannedText, ...], ) -> None: """ @@ -788,10 +840,10 @@ class AnthropicMessagesHandler(BaseTranslation): self, response: "AnthropicMessagesResponse", guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Any | None = None, - user_api_key_dict: Any | None = None, + litellm_logging_obj: "LiteLLMLoggingObj | None" = None, + user_api_key_dict: "UserAPIKeyAuth | None" = None, request_data: dict | None = None, - ) -> Any: + ) -> "AnthropicMessagesResponse": """ Process output response by applying guardrails to text content and tool calls. @@ -867,12 +919,12 @@ class AnthropicMessagesHandler(BaseTranslation): async def process_output_streaming_response( self, - responses_so_far: list[Any], + responses_so_far: Sequence[object], guardrail_to_apply: "CustomGuardrail", - litellm_logging_obj: Any | None = None, - user_api_key_dict: Any | None = None, + litellm_logging_obj: "LiteLLMLoggingObj | None" = None, + user_api_key_dict: "UserAPIKeyAuth | None" = None, request_data: dict | None = None, - ) -> list[Any]: + ) -> Sequence[object]: """ Process output streaming response by applying guardrails to text content. @@ -884,7 +936,7 @@ class AnthropicMessagesHandler(BaseTranslation): if has_ended: # build the model response from the responses_so_far built_response: Final = AnthropicPassthroughLoggingHandler._build_complete_streaming_response( - all_chunks=responses_so_far, + all_chunks=tuple(chunk for chunk in responses_so_far if isinstance(chunk, (str, bytes))), litellm_logging_obj=cast("LiteLLMLoggingObj", litellm_logging_obj), model="", ) @@ -950,8 +1002,8 @@ class AnthropicMessagesHandler(BaseTranslation): def _prepare_request_data( self, request_data: dict | None, - response: Any, - user_api_key_dict: Any | None, + response: object, + user_api_key_dict: "UserAPIKeyAuth | None", key: str, ) -> dict: """Ensure request_data has the response/responses_so_far key and metadata.""" @@ -968,17 +1020,35 @@ class AnthropicMessagesHandler(BaseTranslation): return request_data @staticmethod - def _get_response_content(response: Any) -> list[Any]: - """Extract content list from a dict or object response.""" - if isinstance(response, dict): - return response.get("content", []) or [] - elif hasattr(response, "content"): - return getattr(response, "content", None) or [] - return [] + def _get_response_content(response: object) -> Sequence[object]: + """Extract content blocks from a dict or object response.""" + raw: Final = ( + response.get("content") + if isinstance(response, _MappingLike) + else response.content + if isinstance(response, _HasContent) + else None + ) + content: Final[Sequence[object]] = raw if isinstance(raw, (list, tuple)) else () + return content + + @staticmethod + def _content_block_as_dict(content_block: object) -> dict[str, object] | None: + """The block as a plain dict, dict or Pydantic form, None when it is neither.""" + if isinstance(content_block, _ItemsLike): + return dict(content_block.items()) + if not isinstance(content_block, _HasType): + return None + if isinstance(content_block, _HasModelDump): + return content_block.model_dump() + return { + "type": content_block.type, + "text": content_block.text if isinstance(content_block, _HasText) else None, + } def _extract_from_content_blocks( self, - response_content: list[Any], + response_content: Sequence[object], texts_to_check: list[str], images_to_check: list[str], task_mappings: list[tuple[int, int | None]], @@ -986,38 +1056,24 @@ class AnthropicMessagesHandler(BaseTranslation): ) -> None: """Extract text, images, and tool calls from content blocks.""" for content_idx, content_block in enumerate(response_content): - block_dict: dict[str, Any] = {} - if isinstance(content_block, dict): - block_type = content_block.get("type") - block_dict = cast(dict[str, Any], content_block) - elif hasattr(content_block, "type"): - block_type = getattr(content_block, "type", None) - if hasattr(content_block, "model_dump"): - block_dict = content_block.model_dump() - else: - block_dict = { - "type": block_type, - "text": getattr(content_block, "text", None), - } - else: + block_dict = self._content_block_as_dict(content_block) + if block_dict is None or block_dict.get("type") not in ("text", "tool_use"): continue - - if block_type in ["text", "tool_use"]: - self._extract_output_text_and_images( - content_block=block_dict, - content_idx=content_idx, - texts_to_check=texts_to_check, - images_to_check=images_to_check, - task_mappings=task_mappings, - tool_calls_to_check=tool_calls_to_check, - ) + self._extract_output_text_and_images( + content_block=block_dict, + content_idx=content_idx, + texts_to_check=texts_to_check, + images_to_check=images_to_check, + task_mappings=task_mappings, + tool_calls_to_check=tool_calls_to_check, + ) @staticmethod def _build_guardrail_inputs( texts_to_check: list[str], images_to_check: list[str], tool_calls_to_check: list["ChatCompletionToolCallChunk"], - response: Any, + response: object, ) -> "GenericGuardrailAPIInputs": """Build GenericGuardrailAPIInputs with optional images, tool calls, model.""" inputs: Final = GenericGuardrailAPIInputs(texts=texts_to_check) @@ -1025,16 +1081,18 @@ class AnthropicMessagesHandler(BaseTranslation): inputs["images"] = images_to_check if tool_calls_to_check: inputs["tool_calls"] = tool_calls_to_check - response_model = None - if isinstance(response, dict): - response_model = response.get("model") - elif hasattr(response, "model"): - response_model = getattr(response, "model", None) - if response_model: + response_model: Final = ( + response.get("model") + if isinstance(response, _MappingLike) + else response.model + if isinstance(response, _HasModel) + else None + ) + if isinstance(response_model, str) and response_model: inputs["model"] = response_model return inputs - def get_streaming_string_so_far(self, responses_so_far: list[Any]) -> str: + def get_streaming_string_so_far(self, responses_so_far: Sequence[object]) -> str: """ Parse streaming responses and extract accumulated text content. @@ -1055,19 +1113,43 @@ class AnthropicMessagesHandler(BaseTranslation): } } """ - text_so_far = "" - for response in responses_so_far: - # Handle raw bytes in SSE format - if isinstance(response, bytes): - text_so_far += self._extract_text_from_sse(response) - # Handle already-parsed dict format - elif isinstance(response, dict): - delta = response.get("delta") if response.get("delta") else None - if delta and delta.get("type") == "text_delta": - text = delta.get("text", "") - if text: - text_so_far += text - return text_so_far + return "".join( + self._extract_text_from_sse(response) if isinstance(response, bytes) else self._parsed_event_text(response) + for response in responses_so_far + ) + + @staticmethod + def _parsed_event_text(response: object) -> str: + """Text delta carried by an already-parsed streaming event.""" + if not isinstance(response, _MappingLike): + return "" + delta: Final = response.get("delta") + if not isinstance(delta, _MappingLike) or delta.get("type") != "text_delta": + return "" + text: Final = delta.get("text") + return text if isinstance(text, str) else "" + + @staticmethod + def _sse_event_fields(event: str) -> tuple[str | None, str | None]: + """The ``event:`` type and ``data:`` payload of one SSE event, last line winning.""" + lines: Final = tuple(reversed(event.strip().split("\n"))) + return ( + next((line[len("event:") :].strip() for line in lines if line.startswith("event:")), None), + next((line[len("data:") :].strip() for line in lines if line.startswith("data:")), None), + ) + + @staticmethod + def _sse_event_delta(event: str, expected_event_type: str) -> _MappingLike | None: + """The ``delta`` object of an SSE event of the expected type.""" + event_type, data_line = AnthropicMessagesHandler._sse_event_fields(event) + if event_type != expected_event_type or not data_line: + return None + parsed: Final = _parse_json_object(data_line) + if not parsed: + verbose_proxy_logger.warning("Failed to parse JSON from SSE data: %s", data_line) + return None + delta: Final = parsed[0].get("delta") + return delta if isinstance(delta, _MappingLike) else None def _extract_text_from_sse(self, sse_bytes: bytes) -> str: """ @@ -1079,45 +1161,25 @@ class AnthropicMessagesHandler(BaseTranslation): Returns: Accumulated text from all content_block_delta events """ - text = "" try: - # Decode bytes to string sse_string: Final = sse_bytes.decode("utf-8") - - # Split by double newline to get individual events - events: Final = sse_string.split("\n\n") - - for event in events: - if not event.strip(): - continue - - # Parse event lines - lines = event.strip().split("\n") - event_type = None - data_line = None - - for line in lines: - if line.startswith("event:"): - event_type = line[6:].strip() - elif line.startswith("data:"): - data_line = line[5:].strip() - - # Only process content_block_delta events - if event_type == "content_block_delta" and data_line: - try: - data = json.loads(data_line) - delta = data.get("delta", {}) - if delta.get("type") == "text_delta": - text += delta.get("text", "") - except json.JSONDecodeError: - verbose_proxy_logger.warning("Failed to parse JSON from SSE data: %s", data_line) - + return "".join( + self._sse_text_delta(event) for event in sse_string.split("\n\n") if event.strip() + ) except Exception as e: verbose_proxy_logger.error("Error extracting text from SSE: %s", e) + return "" - return text + @staticmethod + def _sse_text_delta(event: str) -> str: + """Text of a ``content_block_delta`` event, empty for every other event.""" + delta: Final = AnthropicMessagesHandler._sse_event_delta(event, "content_block_delta") + if delta is None or delta.get("type") != "text_delta": + return "" + text: Final = delta.get("text") + return text if isinstance(text, str) else "" - def _check_streaming_has_ended(self, responses_so_far: list[Any]) -> bool: + def _check_streaming_has_ended(self, responses_so_far: Sequence[object]) -> bool: """ Check if streaming response has ended by looking for non-null stop_reason. @@ -1140,54 +1202,30 @@ class AnthropicMessagesHandler(BaseTranslation): Returns: True if stop_reason is set to a non-null value, indicating stream has ended """ - for response in responses_so_far: - # Handle raw bytes in SSE format - if isinstance(response, bytes): - try: - # Decode bytes to string - sse_string = response.decode("utf-8") + return any(self._chunk_signals_end(response) for response in responses_so_far) - # Split by double newline to get individual events - events = sse_string.split("\n\n") + @staticmethod + def _chunk_signals_end(response: object) -> bool: + """Whether one streamed chunk carries a non-null ``stop_reason``.""" + if isinstance(response, bytes): + try: + return any( + AnthropicMessagesHandler._sse_stop_reason_present(event) + for event in response.decode("utf-8").split("\n\n") + if event.strip() + ) + except Exception as e: + verbose_proxy_logger.error("Error checking streaming end in SSE: %s", e) + return False + if not isinstance(response, _MappingLike) or response.get("type") != "message_delta": + return False + delta: Final = response.get("delta") + return isinstance(delta, _MappingLike) and delta.get("stop_reason") is not None - for event in events: - if not event.strip(): - continue - - # Parse event lines - lines = event.strip().split("\n") - event_type = None - data_line = None - - for line in lines: - if line.startswith("event:"): - event_type = line[6:].strip() - elif line.startswith("data:"): - data_line = line[5:].strip() - - # Check for message_delta event with stop_reason - if event_type == "message_delta" and data_line: - try: - data = json.loads(data_line) - delta = data.get("delta", {}) - stop_reason = delta.get("stop_reason") - if stop_reason is not None: - return True - except json.JSONDecodeError: - verbose_proxy_logger.warning("Failed to parse JSON from SSE data: %s", data_line) - - except Exception as e: - verbose_proxy_logger.error("Error checking streaming end in SSE: %s", e) - - # Handle already-parsed dict format - elif isinstance(response, dict): - if response.get("type") == "message_delta": - delta = response.get("delta", {}) - stop_reason = delta.get("stop_reason") - if stop_reason is not None: - return True - - return False + @staticmethod + def _sse_stop_reason_present(event: str) -> bool: + delta: Final = AnthropicMessagesHandler._sse_event_delta(event, "message_delta") + return delta is not None and delta.get("stop_reason") is not None def _has_text_content(self, response: "AnthropicMessagesResponse") -> bool: """ @@ -1212,7 +1250,7 @@ class AnthropicMessagesHandler(BaseTranslation): def _extract_output_text_and_images( self, - content_block: dict[str, Any], + content_block: dict[str, object], content_idx: int, texts_to_check: list[str], images_to_check: list[str], @@ -1247,8 +1285,8 @@ class AnthropicMessagesHandler(BaseTranslation): async def _apply_guardrail_responses_to_output( self, response: "AnthropicMessagesResponse", - responses: list[str], - task_mappings: list[tuple[int, int | None]], + responses: Sequence[str], + task_mappings: Sequence[tuple[int, int | None]], ) -> None: """ Apply guardrail responses back to output response. @@ -1256,23 +1294,9 @@ class AnthropicMessagesHandler(BaseTranslation): Override this method to customize how responses are applied. """ for task_idx, guardrail_response in enumerate(responses): - mapping = task_mappings[task_idx] - content_idx = cast(int, mapping[0]) + content_idx = task_mappings[task_idx][0] + response_content = self._get_response_content(response) - # Handle both dict and object responses - response_content: list[Any] = [] - if isinstance(response, dict): - response_content = response.get("content", []) or [] - elif hasattr(response, "content"): - content = getattr(response, "content", None) - response_content = content or [] - else: - continue - - if not response_content: - continue - - # Get the content block at the index if content_idx >= len(response_content): continue @@ -1280,10 +1304,10 @@ class AnthropicMessagesHandler(BaseTranslation): # Verify it's a text block and update the text field # Handle both dict and Pydantic object content blocks - if isinstance(content_block, dict): + if isinstance(content_block, _MutableMappingLike): if content_block.get("type") == "text": - cast(dict[str, Any], content_block)["text"] = guardrail_response - elif hasattr(content_block, "type") and getattr(content_block, "type", None) == "text": + content_block["text"] = guardrail_response + elif isinstance(content_block, _HasType) and content_block.type == "text": # Update Pydantic object's text attribute - if hasattr(content_block, "text"): + if isinstance(content_block, _HasText): content_block.text = guardrail_response