Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
232 changes: 170 additions & 62 deletions litellm/llms/anthropic/chat/guardrail_translation/handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,8 +13,12 @@
"""

import json
from collections.abc import Mapping
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Final, cast

from typing_extensions import assert_never

from litellm._logging import verbose_proxy_logger
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (
Expand Down Expand Up @@ -58,6 +62,50 @@
)


@dataclass(frozen=True, slots=True)
class MessageContentTarget:
msg_idx: int


@dataclass(frozen=True, slots=True)
class ContentBlockTextTarget:
msg_idx: int
content_idx: int


@dataclass(frozen=True, slots=True)
class ToolResultStringTarget:
msg_idx: int
content_idx: int


@dataclass(frozen=True, slots=True)
class ToolResultBlockTextTarget:
msg_idx: int
content_idx: int
block_idx: int


InputWriteBackTarget = (
MessageContentTarget | ContentBlockTextTarget | ToolResultStringTarget | ToolResultBlockTextTarget
)


@dataclass(frozen=True, slots=True)
class ScannedText:
text: str
target: InputWriteBackTarget


@dataclass(frozen=True, slots=True)
class ExtractedInput:
scanned: tuple[ScannedText, ...]
images: tuple[str, ...]


EMPTY_EXTRACTED_INPUT: Final = ExtractedInput(scanned=(), images=())


class AnthropicMessagesHandler(BaseTranslation):
"""
Handler for processing Anthropic messages with guardrails.
Expand Down Expand Up @@ -290,22 +338,23 @@ async def process_input_messages(
if skip_tool:
structured_messages = openai_messages_without_tool(structured_messages)

texts_to_check: Final[list[str]] = []
images_to_check: Final[list[str]] = []
tools_to_check: Final[list[ChatCompletionToolParam]] = chat_completion_compatible_request.get("tools", [])
task_mappings: Final[list[tuple[int, int | None]]] = []

# Step 1: Extract all text content and images
for msg_idx, message in enumerate(messages):
extracted: Final = tuple(
self._extract_input_text_and_images(
message=message,
msg_idx=msg_idx,
texts_to_check=texts_to_check,
images_to_check=images_to_check,
task_mappings=task_mappings,
skip_system_message=skip_system,
skip_tool_message=skip_tool,
)
for msg_idx, message in enumerate(messages)
)
scanned: Final = tuple(item for one_message in extracted for item in one_message.scanned)
texts_to_check: Final = [item.text for item in scanned] # mutable-ok: GenericGuardrailAPIInputs takes list[str]
Comment thread
greptile-apps[bot] marked this conversation as resolved.
images_to_check: Final = [
image for one_message in extracted for image in one_message.images
] # mutable-ok: GenericGuardrailAPIInputs takes list[str]

# Step 2: Apply guardrail to all texts in batch
if texts_to_check:
Expand Down Expand Up @@ -352,7 +401,7 @@ async def process_input_messages(
await self._apply_guardrail_responses_to_input(
messages=messages,
responses=guardrailed_texts,
task_mappings=task_mappings,
scanned=scanned,
)

verbose_proxy_logger.debug("Anthropic Messages: Processed input messages: %s", messages)
Expand Down Expand Up @@ -405,55 +454,102 @@ def extract_request_tool_names(self, data: dict) -> list[str]:
names.append(str(tool["name"]))
return names

@classmethod
def _extract_input_text_and_images(
self,
cls,
message: dict[str, Any],
msg_idx: int,
texts_to_check: list[str],
images_to_check: list[str],
task_mappings: list[tuple[int, int | None]],
skip_system_message: bool = False,
skip_tool_message: bool = False,
) -> None:
) -> ExtractedInput:
"""
Extract text content and images from a message.

Override this method to customize text/image extraction logic.
"""
role: Final = str(message.get("role") or "").lower()
if skip_system_message and role == "system":
return
if skip_tool_message and role == "tool":
return
if (skip_system_message and role == "system") or (skip_tool_message and role == "tool"):
return EMPTY_EXTRACTED_INPUT

content: Final = message.get("content", None)
tools: Final = message.get("tools", None)
if content is None and tools is None:
return

## CHECK FOR TEXT + IMAGES
if content is not None and isinstance(content, str):
# Simple string content
texts_to_check.append(content)
task_mappings.append((msg_idx, None))

elif content is not None and isinstance(content, list):
# List content (e.g., multimodal with text and images)
for content_idx, content_item in enumerate(content):
# Extract text
text_str = content_item.get("text", None)
if text_str is not None:
texts_to_check.append(text_str)
task_mappings.append((msg_idx, int(content_idx)))

# Extract images
if content_item.get("type") == "image":
source = content_item.get("source", {})
if isinstance(source, dict):
# Could be base64 or url
data = source.get("data")
if data:
images_to_check.append(data)
if isinstance(content, str):
return ExtractedInput(scanned=(ScannedText(content, MessageContentTarget(msg_idx)),), images=())
if not isinstance(content, list):
return EMPTY_EXTRACTED_INPUT

blocks: Final = tuple(
cls._extract_content_block(
content_item=content_item,
msg_idx=msg_idx,
content_idx=content_idx,
skip_tool_message=skip_tool_message,
)
for content_idx, content_item in enumerate(content)
if isinstance(content_item, dict)
)
return ExtractedInput(
scanned=tuple(item for block in blocks for item in block.scanned),
images=tuple(image for block in blocks for image in block.images),
)

@classmethod
def _extract_content_block(
cls,
content_item: Mapping[str, Any],
msg_idx: int,
content_idx: int,
skip_tool_message: bool,
) -> ExtractedInput:
if content_item.get("type") == "tool_result":
if skip_tool_message:
return EMPTY_EXTRACTED_INPUT
return cls._extract_tool_result(content_item=content_item, msg_idx=msg_idx, content_idx=content_idx)

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)),)
),
images=cls._image_sources(content_item) if content_item.get("type") == "image" else (),
)

@classmethod
def _extract_tool_result(
cls,
content_item: Mapping[str, Any],
msg_idx: int,
content_idx: int,
) -> ExtractedInput:
tool_result_content: Final = content_item.get("content")

if isinstance(tool_result_content, str):
return ExtractedInput(
scanned=(ScannedText(tool_result_content, ToolResultStringTarget(msg_idx, content_idx)),),
images=(),
)
if not isinstance(tool_result_content, list):
return EMPTY_EXTRACTED_INPUT

blocks: Final = tuple(
(block_idx, block) for block_idx, block in enumerate(tool_result_content) if isinstance(block, dict)
)
return ExtractedInput(
scanned=tuple(
ScannedText(block["text"], ToolResultBlockTextTarget(msg_idx, content_idx, block_idx))
for block_idx, block in blocks
if isinstance(block.get("text"), str)
),
images=tuple(
image for _, block in blocks if block.get("type") == "image" for image in cls._image_sources(block)
),
)

@staticmethod
def _image_sources(block: Mapping[str, Any]) -> tuple[str, ...]:
source: Final = block.get("source")
if not isinstance(source, Mapping):
return ()
# Could be base64 or url
data: Final = source.get("data")
return (data,) if data else ()

def _extract_input_tools(
self,
Expand All @@ -475,29 +571,41 @@ async def _apply_guardrail_responses_to_input(
self,
messages: list[dict[str, Any]],
responses: list[str],
task_mappings: list[tuple[int, int | None]],
scanned: tuple[ScannedText, ...],
) -> None:
"""
Apply guardrail responses back to input messages.

Override this method to customize how responses are applied.
"""
for task_idx, guardrail_response in enumerate(responses):
mapping = task_mappings[task_idx]
msg_idx = cast(int, mapping[0])
content_idx_optional = cast(int | None, mapping[1])

content = messages[msg_idx].get("content", None)
for item, guardrail_response in zip(scanned, responses):
target = item.target
message = messages[target.msg_idx]
content = message.get("content", None)
if content is None:
continue

if isinstance(content, str) and content_idx_optional is None:
# Replace string content with guardrail response
messages[msg_idx]["content"] = guardrail_response

elif isinstance(content, list) and content_idx_optional is not None:
# Replace specific text item in list content
messages[msg_idx]["content"][content_idx_optional]["text"] = guardrail_response
match target:
case MessageContentTarget():
if isinstance(content, str):
message["content"] = (
guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place
)
case ContentBlockTextTarget(content_idx=content_idx):
if isinstance(content, list):
content[content_idx]["text"] = (
guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place
)
case ToolResultStringTarget(content_idx=content_idx):
if isinstance(content, list):
content[content_idx]["content"] = (
guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place
)
case ToolResultBlockTextTarget(content_idx=content_idx, block_idx=block_idx):
if isinstance(content, list):
content[content_idx]["content"][block_idx]["text"] = (
guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place
)
case _:
assert_never(target)

async def process_output_response(
self,
Expand Down
Loading
Loading