mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge remote-tracking branch 'origin/litellm_daily_any_cleanup_08_11_2026_3Mn0k_c' into litellm_daily_any_cleanup_08_11_2026_3Mn0k
This commit is contained in:
commit
b534354941
2 changed files with 300 additions and 231 deletions
|
|
@ -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,29 @@ 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)
|
||||
if isinstance(response, _MutableMappingLike) and isinstance(response, dict):
|
||||
response["content"] = WebSearchInterceptionLogger._prepended_content(response.get("content"), native_blocks)
|
||||
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, _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 +1464,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,
|
||||
|
|
|
|||
|
|
@ -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,23 @@ 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 +1200,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 +1248,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 +1283,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 +1292,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 +1302,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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue