diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py index 0db559e9794..c74203b7544 100644 --- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py @@ -44,8 +44,10 @@ from litellm.types.llms.anthropic import ( ) from litellm.types.llms.openai import ( AllMessageValues, + ChatCompletionAssistantToolCall, ChatCompletionRequest, ChatCompletionToolCallChunk, + ChatCompletionToolCallFunctionChunk, ChatCompletionToolParam, ) from litellm.types.utils import ( @@ -316,9 +318,9 @@ class AnthropicMessagesHandler(BaseTranslation): def scoped_request_conversation( self, - request_data: dict, + request_data: dict, # mutable-ok: API request payload guardrail_to_apply: "CustomGuardrail", - ) -> list[AllMessageValues] | None: + ) -> tuple[AllMessageValues, ...] | None: """ Mirror the request scan's scoping: translate without the trusted top-level prompt, hoist it back unless skip_system, and keep @@ -330,27 +332,26 @@ class AnthropicMessagesHandler(BaseTranslation): translation_source: Final = { # mutable-ok: API message payload key: value for key, value in request_data.items() if key != "system" } - full_structured_messages: Final = cast( - list[AllMessageValues], - self._translate_to_openai(translation_source).get("messages", []), - ) + translated_messages: Final = self._translate_to_openai(translation_source).get("messages") hoisted_system_message: Final = None if skip_system else self._hoisted_top_level_system_message(request_data) - if hoisted_system_message is not None: - full_structured_messages.insert(0, hoisted_system_message) + full_structured_messages: Final = ( + (hoisted_system_message, *(translated_messages or ())) + if hoisted_system_message is not None + else tuple(translated_messages or ()) + ) scoped_indices: Final = scoped_structured_message_indices( full_structured_messages, scan_only_tool_results=effective_scan_only_tool_results_for_guardrail(guardrail_to_apply), skip_system=False, skip_tool=effective_skip_tool_message_for_guardrail(guardrail_to_apply), ) - scoped: Final = [full_structured_messages[index] for index in scoped_indices] - return scoped or None + return tuple(full_structured_messages[index] for index in scoped_indices) or None def request_tools_for_guardrail( self, - request_data: dict, + request_data: dict, # mutable-ok: API request payload guardrail_to_apply: "CustomGuardrail", - ) -> list[ChatCompletionToolParam] | None: + ) -> tuple[ChatCompletionToolParam, ...] | None: if effective_scan_only_tool_results_for_guardrail(guardrail_to_apply): return None if not request_data.get("tools"): @@ -358,8 +359,8 @@ class AnthropicMessagesHandler(BaseTranslation): translation_source: Final = { # mutable-ok: API message payload key: value for key, value in request_data.items() if key != "system" } - tools: Final = self._translate_to_openai(translation_source).get("tools", []) - return list(tools) if tools else None + tools: Final = self._translate_to_openai(translation_source).get("tools") + return tuple(tools) if tools else None async def process_input_messages( self, @@ -949,10 +950,12 @@ class AnthropicMessagesHandler(BaseTranslation): self.assistant_turn_from_extraction(texts_to_check, tool_calls_to_check), ) if structured_conversation: - inputs["structured_messages"] = structured_conversation + inputs["structured_messages"] = list( + structured_conversation + ) # mutable-ok: GenericGuardrailAPIInputs takes list response_scan_tools: Final = self.request_tools_for_guardrail(request_data, guardrail_to_apply) if response_scan_tools: - inputs["tools"] = response_scan_tools + inputs["tools"] = list(response_scan_tools) # mutable-ok: GenericGuardrailAPIInputs takes list guardrailed_inputs: Final = await guardrail_to_apply.apply_guardrail( inputs=inputs, @@ -1020,24 +1023,36 @@ class AnthropicMessagesHandler(BaseTranslation): user_api_key_dict, key="response", ) - stream_tool_call_dicts: Final = tuple( - tool_call.model_dump() for tool_call in tool_calls_list or () + stream_tool_calls: Final = tuple( + ChatCompletionAssistantToolCall( + id=tool_call.id, + type="function", + function=ChatCompletionToolCallFunctionChunk( + name=tool_call.function.name, + arguments=tool_call.function.arguments, + ), + ) + for tool_call in tool_calls_list or () ) structured_conversation: Final = self.response_scan_conversation( prepared_request_data, guardrail_to_apply, self.assistant_turn_from_extraction( - [string_so_far] if isinstance(string_so_far, str) and string_so_far else [], - stream_tool_call_dicts, + (string_so_far,) if isinstance(string_so_far, str) and string_so_far else (), + stream_tool_calls, ), ) if structured_conversation: - guardrail_inputs["structured_messages"] = structured_conversation + guardrail_inputs["structured_messages"] = list( + structured_conversation + ) # mutable-ok: GenericGuardrailAPIInputs takes list response_scan_tools: Final = self.request_tools_for_guardrail( prepared_request_data, guardrail_to_apply ) if response_scan_tools: - guardrail_inputs["tools"] = response_scan_tools + guardrail_inputs["tools"] = list( + response_scan_tools + ) # mutable-ok: GenericGuardrailAPIInputs takes list _guardrailed_inputs = await guardrail_to_apply.apply_guardrail( inputs=guardrail_inputs, request_data=prepared_request_data, diff --git a/litellm/llms/base_llm/guardrail_translation/base_translation.py b/litellm/llms/base_llm/guardrail_translation/base_translation.py index da37cbf1a93..c59d8d0f841 100644 --- a/litellm/llms/base_llm/guardrail_translation/base_translation.py +++ b/litellm/llms/base_llm/guardrail_translation/base_translation.py @@ -1,7 +1,7 @@ from abc import ABC, abstractmethod -from collections.abc import Mapping, Sequence +from collections.abc import Sequence from dataclasses import dataclass, field -from typing import TYPE_CHECKING, Any, Final, Optional, cast +from typing import TYPE_CHECKING, Any, Final, Optional from litellm.llms.base_llm.guardrail_translation.utils import ( effective_scan_only_tool_results_for_guardrail, @@ -9,6 +9,10 @@ from litellm.llms.base_llm.guardrail_translation.utils import ( effective_skip_tool_message_for_guardrail, scoped_structured_message_indices, ) +from litellm.types.llms.openai import ( + ChatCompletionAssistantMessage, + ChatCompletionAssistantToolCall, +) if TYPE_CHECKING: from litellm.integrations.custom_guardrail import ( @@ -167,9 +171,9 @@ class BaseTranslation(ABC): def scoped_request_conversation( self, - request_data: dict, + request_data: dict, # mutable-ok: API request payload guardrail_to_apply: "CustomGuardrail", - ) -> list["AllMessageValues"] | None: + ) -> tuple["AllMessageValues", ...] | None: """ The request conversation as the guardrail's request scan saw it: the handler's structured messages with the operator scoping flags applied. @@ -186,15 +190,14 @@ class BaseTranslation(ABC): skip_system=effective_skip_system_message_for_guardrail(guardrail_to_apply), skip_tool=effective_skip_tool_message_for_guardrail(guardrail_to_apply), ) - scoped: Final = [structured_messages[index] for index in scoped_indices] - return scoped or None + return tuple(structured_messages[index] for index in scoped_indices) or None def response_scan_conversation( self, - request_data: dict | None, + request_data: dict | None, # mutable-ok: API request payload guardrail_to_apply: "CustomGuardrail", response_turns: Sequence["AllMessageValues"], - ) -> list["AllMessageValues"] | None: + ) -> tuple["AllMessageValues", ...] | None: """ Full conversation for a response scan: the scoped request conversation with the model's response turns appended. @@ -208,13 +211,13 @@ class BaseTranslation(ABC): request_conversation: Final = self.scoped_request_conversation(request_data, guardrail_to_apply) if request_conversation is None: return None - return [*request_conversation, *response_turns] + return (*request_conversation, *response_turns) def request_tools_for_guardrail( self, - request_data: dict, + request_data: dict, # mutable-ok: API request payload guardrail_to_apply: "CustomGuardrail", - ) -> list["ChatCompletionToolParam"] | None: + ) -> tuple["ChatCompletionToolParam", ...] | None: """ The request's tool definitions in the shape the request scan sends them. @@ -225,21 +228,24 @@ class BaseTranslation(ABC): @staticmethod def assistant_turn_from_extraction( texts: Sequence[str], - tool_calls: Sequence[Mapping[str, object]] | None = None, - ) -> list["AllMessageValues"]: + tool_calls: Sequence["ChatCompletionAssistantToolCall"] | None = None, + ) -> tuple["ChatCompletionAssistantMessage", ...]: """ One OpenAI-shape assistant turn built from the texts and tool calls a handler's response extraction collected; empty when there is nothing. """ tool_call_items: Final = tuple(tool_calls or ()) if not texts and not tool_call_items: - return [] - turn: Final = { - "role": "assistant", - "content": "\n".join(texts), - **({"tool_calls": list(tool_call_items)} if tool_call_items else {}), - } - return [cast("AllMessageValues", turn)] + return () + if tool_call_items: + turn_with_tools: Final[ChatCompletionAssistantMessage] = { + "role": "assistant", + "content": "\n".join(texts), + "tool_calls": list(tool_call_items), # mutable-ok: the message field type is a list + } + return (turn_with_tools,) + turn: Final[ChatCompletionAssistantMessage] = {"role": "assistant", "content": "\n".join(texts)} + return (turn,) def extract_request_tool_names(self, data: dict) -> list[str]: """ diff --git a/litellm/llms/openai/chat/guardrail_translation/handler.py b/litellm/llms/openai/chat/guardrail_translation/handler.py index 5352be2a0ff..1699597d120 100644 --- a/litellm/llms/openai/chat/guardrail_translation/handler.py +++ b/litellm/llms/openai/chat/guardrail_translation/handler.py @@ -33,11 +33,18 @@ from litellm.llms.base_llm.guardrail_translation.utils import ( scoped_structured_message_indices, ) from litellm.main import stream_chunk_builder -from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam +from litellm.types.llms.openai import ( + AllMessageValues, + ChatCompletionAssistantMessage, + ChatCompletionAssistantToolCall, + ChatCompletionToolCallFunctionChunk, + ChatCompletionToolParam, +) from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import ( coerce_stream_holdback_value, ) from litellm.types.utils import ( + ChatCompletionMessageToolCall, Choices, GenericGuardrailAPIInputs, ModelResponse, @@ -400,10 +407,12 @@ class OpenAIChatCompletionsHandler(BaseTranslation): request_data, guardrail_to_apply, self._build_response_turns(response) ) if structured_conversation: - inputs["structured_messages"] = structured_conversation + inputs["structured_messages"] = list( + structured_conversation + ) # mutable-ok: GenericGuardrailAPIInputs takes list response_scan_tools: Final = self.request_tools_for_guardrail(request_data, guardrail_to_apply) if response_scan_tools: - inputs["tools"] = response_scan_tools + inputs["tools"] = list(response_scan_tools) # mutable-ok: GenericGuardrailAPIInputs takes list guardrailed_inputs: Final = await guardrail_to_apply.apply_guardrail( inputs=inputs, @@ -843,39 +852,43 @@ class OpenAIChatCompletionsHandler(BaseTranslation): def request_tools_for_guardrail( self, - request_data: dict, + request_data: dict, # mutable-ok: API request payload guardrail_to_apply: "CustomGuardrail", - ) -> list[ChatCompletionToolParam] | None: + ) -> tuple[ChatCompletionToolParam, ...] | None: if effective_scan_only_tool_results_for_guardrail(guardrail_to_apply): return None tools: Final = request_data.get("tools") - return cast(list[ChatCompletionToolParam], tools) if tools else None + return tuple(tools) if tools else None - def _build_response_turns(self, response: "ModelResponse") -> list[AllMessageValues]: + def _build_response_turns(self, response: "ModelResponse") -> tuple[ChatCompletionAssistantMessage, ...]: """Assistant turns for the response-scan conversation, one per choice.""" - return [ - turn - for choice in response.choices - if isinstance(choice, litellm.Choices) and (turn := self._choice_assistant_turn(choice)) is not None - ] + return tuple(turn for choice in response.choices if (turn := self._choice_assistant_turn(choice)) is not None) - def _choice_assistant_turn(self, choice: Choices) -> AllMessageValues | None: - tool_call_dicts: Final = tuple( - converted - for tool_call in (choice.message.tool_calls or []) - if (converted := self._convert_tool_call_to_dict(tool_call)) is not None + def _choice_assistant_turn(self, choice: Choices) -> ChatCompletionAssistantMessage | None: + tool_calls: Final = tuple( + ChatCompletionAssistantToolCall( + id=tool_call.id, + type="function", + function=ChatCompletionToolCallFunctionChunk( + name=tool_call.function.name, + arguments=tool_call.function.arguments, + ), + ) + for tool_call in choice.message.tool_calls or () + if isinstance(tool_call, ChatCompletionMessageToolCall) ) content: Final = choice.message.content - if content is None and not tool_call_dicts: + if content is None and not tool_calls: return None - return cast( - AllMessageValues, - { + if tool_calls: + turn_with_tools: Final[ChatCompletionAssistantMessage] = { "role": "assistant", "content": content, - **({"tool_calls": list(tool_call_dicts)} if tool_call_dicts else {}), - }, - ) + "tool_calls": list(tool_calls), # mutable-ok: the message field type is a list + } + return turn_with_tools + turn: Final[ChatCompletionAssistantMessage] = {"role": "assistant", "content": content} + return turn def _convert_tool_call_to_dict(self, tool_call: dict[str, Any] | Any) -> dict[str, Any] | None: """ diff --git a/litellm/llms/openai/responses/guardrail_translation/handler.py b/litellm/llms/openai/responses/guardrail_translation/handler.py index 2d747d43968..8776893e85f 100644 --- a/litellm/llms/openai/responses/guardrail_translation/handler.py +++ b/litellm/llms/openai/responses/guardrail_translation/handler.py @@ -250,15 +250,15 @@ class OpenAIResponsesHandler(BaseTranslation): def request_tools_for_guardrail( self, - request_data: dict, + request_data: dict, # mutable-ok: API request payload guardrail_to_apply: "CustomGuardrail", - ) -> list[ChatCompletionToolParam] | None: + ) -> tuple[ChatCompletionToolParam, ...] | None: raw_tools: Final = request_data.get("tools") if not raw_tools: return None - tools_to_check: Final[list[ChatCompletionToolParam]] = [] + tools_to_check: Final[list[ChatCompletionToolParam]] = [] # mutable-ok: _extract_and_transform_tools appends self._extract_and_transform_tools(raw_tools, tools_to_check) - return tools_to_check or None + return tuple(tools_to_check) or None def _remap_tools_to_responses_api_format(self, guardrailed_tools: list[Any]) -> list[dict[str, object]]: """ @@ -473,10 +473,12 @@ class OpenAIResponsesHandler(BaseTranslation): self.assistant_turn_from_extraction(texts_to_check, tool_calls_to_check), ) if structured_conversation: - inputs["structured_messages"] = structured_conversation + inputs["structured_messages"] = list( + structured_conversation + ) # mutable-ok: GenericGuardrailAPIInputs takes list response_scan_tools: Final = self.request_tools_for_guardrail(request_data, guardrail_to_apply) if response_scan_tools: - inputs["tools"] = response_scan_tools + inputs["tools"] = list(response_scan_tools) # mutable-ok: GenericGuardrailAPIInputs takes list guardrailed_inputs: Final = await guardrail_to_apply.apply_guardrail( inputs=inputs, @@ -574,10 +576,12 @@ class OpenAIResponsesHandler(BaseTranslation): self.assistant_turn_from_extraction(texts_to_check, tool_calls_to_check), ) if structured_conversation: - inputs["structured_messages"] = structured_conversation + inputs["structured_messages"] = list( + structured_conversation + ) # mutable-ok: GenericGuardrailAPIInputs takes list response_scan_tools: Final = self.request_tools_for_guardrail(request_data, guardrail_to_apply) if response_scan_tools: - inputs["tools"] = response_scan_tools + inputs["tools"] = list(response_scan_tools) # mutable-ok: GenericGuardrailAPIInputs takes list guardrailed_inputs: Final = await guardrail_to_apply.apply_guardrail( inputs=inputs, diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_moderations.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_moderations.py index 3187f1d322a..52575dc640c 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_moderations.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_moderations.py @@ -994,8 +994,10 @@ async def test_openai_moderation_initialize_guardrail_forwards_streaming_flags() @pytest.mark.asyncio async def test_openai_moderation_response_scan_moderates_output_not_user_prompt(): """On response scans the conversation is available in structured_messages, - but moderation must still target the model output carried in texts.""" - from unittest.mock import AsyncMock + but moderation must still target the model output carried in texts. The fake + moderation endpoint flags only the harmful model answer: moderating the + (benign) user prompt instead would let the flagged output through.""" + from fastapi import HTTPException from litellm.types.utils import GenericGuardrailAPIInputs @@ -1004,34 +1006,33 @@ async def test_openai_moderation_response_scan_moderates_output_not_user_prompt( guardrail_name="test-openai-moderation", ) - mock_response = OpenAIModerationResponse( - id="modr-123", - model="omni-moderation-latest", - results=[ - OpenAIModerationResult( - flagged=False, - categories={}, - category_scores={}, - category_applied_input_types={}, - ) - ], - ) - - with patch.object( - guardrail, "async_make_request", new_callable=AsyncMock, return_value=mock_response - ) as mock_request: - inputs = GenericGuardrailAPIInputs( - texts=["the model's answer"], - structured_messages=[ - {"role": "user", "content": "the user's question"}, - {"role": "assistant", "content": "the model's answer"}, + async def moderate(input_text: str) -> OpenAIModerationResponse: + flagged = "harmful answer" in input_text + return OpenAIModerationResponse( + id="modr-123", + model="omni-moderation-latest", + results=[ + OpenAIModerationResult( + flagged=flagged, + categories={"violence": flagged}, + category_scores={"violence": 0.99 if flagged else 0.0}, + category_applied_input_types={"violence": []}, + ) ], ) - await guardrail.apply_guardrail( - inputs=inputs, - request_data={}, - input_type="response", + with patch.object(guardrail, "async_make_request", side_effect=moderate): + inputs = GenericGuardrailAPIInputs( + texts=["a harmful answer"], + structured_messages=[ + {"role": "user", "content": "a benign question"}, + {"role": "assistant", "content": "a harmful answer"}, + ], ) - assert mock_request.call_args.kwargs["input_text"] == "the model's answer" + with pytest.raises(HTTPException): + await guardrail.apply_guardrail( + inputs=inputs, + request_data={}, + input_type="response", + )