diff --git a/litellm/llms/anthropic/chat.py b/litellm/llms/anthropic/chat.py index c87fd5f9736..dd7ab58c1d5 100644 --- a/litellm/llms/anthropic/chat.py +++ b/litellm/llms/anthropic/chat.py @@ -228,6 +228,54 @@ class AnthropicConfig: return False + def translate_system_message( + self, messages: List[AllMessageValues] + ) -> List[AnthropicSystemMessageContent]: + system_prompt_indices = [] + anthropic_system_message_list: List[AnthropicSystemMessageContent] = [] + for idx, message in enumerate(messages): + if message["role"] == "system": + valid_content: bool = False + system_message_block = ChatCompletionSystemMessage(**message) + if isinstance(system_message_block["content"], str): + anthropic_system_message_content = AnthropicSystemMessageContent( + type="text", + text=system_message_block["content"], + ) + if "cache_control" in system_message_block: + anthropic_system_message_content["cache_control"] = ( + system_message_block["cache_control"] + ) + anthropic_system_message_list.append( + anthropic_system_message_content + ) + valid_content = True + elif isinstance(message["content"], list): + for _content in message["content"]: + anthropic_system_message_content = ( + AnthropicSystemMessageContent( + type=_content.get("type"), + text=_content.get("text"), + ) + ) + if "cache_control" in _content: + anthropic_system_message_content["cache_control"] = ( + _content["cache_control"] + ) + + anthropic_system_message_list.append( + anthropic_system_message_content + ) + valid_content = True + + if valid_content: + system_prompt_indices.append(idx) + if len(system_prompt_indices) > 0: + for idx in reversed(system_prompt_indices): + messages.pop(idx) + + return anthropic_system_message_list + ### FOR [BETA] `/v1/messages` endpoint support def translatable_anthropic_params(self) -> List: @@ -314,7 +362,7 @@ class AnthropicConfig: new_messages.append(user_message) if len(new_user_content_list) > 0: - new_messages.append({"role": "user", "content": new_user_content_list}) + new_messages.append({"role": "user", "content": new_user_content_list}) # type: ignore if len(tool_message_list) > 0: new_messages.extend(tool_message_list) @@ -940,45 +988,11 @@ class AnthropicChatCompletion(BaseLLM): ) else: # Separate system prompt from rest of message - system_prompt_indices = [] - system_prompt = "" - anthropic_system_message_list = None - for idx, message in enumerate(messages): - if message["role"] == "system": - valid_content: bool = False - if isinstance(message["content"], str): - system_prompt += message["content"] - valid_content = True - elif isinstance(message["content"], list): - for _content in message["content"]: - anthropic_system_message_content = ( - AnthropicSystemMessageContent( - type=_content.get("type"), - text=_content.get("text"), - ) - ) - if "cache_control" in _content: - anthropic_system_message_content["cache_control"] = ( - _content["cache_control"] - ) - - if anthropic_system_message_list is None: - anthropic_system_message_list = [] - anthropic_system_message_list.append( - anthropic_system_message_content - ) - valid_content = True - - if valid_content: - system_prompt_indices.append(idx) - if len(system_prompt_indices) > 0: - for idx in reversed(system_prompt_indices): - messages.pop(idx) - if len(system_prompt) > 0: - optional_params["system"] = system_prompt - + anthropic_system_message_list = AnthropicConfig().translate_system_message( + messages=messages + ) # Handling anthropic API Prompt Caching - if anthropic_system_message_list is not None: + if len(anthropic_system_message_list) > 0: optional_params["system"] = anthropic_system_message_list # Format rest of message according to anthropic guidelines try: diff --git a/litellm/llms/prompt_templates/factory.py b/litellm/llms/prompt_templates/factory.py index 228a0fdec33..d2b9db03713 100644 --- a/litellm/llms/prompt_templates/factory.py +++ b/litellm/llms/prompt_templates/factory.py @@ -27,10 +27,13 @@ from litellm.types.completion import ( from litellm.types.llms.anthropic import * from litellm.types.llms.bedrock import MessageBlock as BedrockMessageBlock from litellm.types.llms.openai import ( + AllMessageValues, ChatCompletionAssistantMessage, + ChatCompletionAssistantToolCall, ChatCompletionFunctionMessage, ChatCompletionToolCallFunctionChunk, ChatCompletionToolMessage, + ChatCompletionUserMessage, ) from litellm.types.utils import GenericImageParsingChunk @@ -1170,7 +1173,9 @@ def convert_to_gemini_tool_call_result( return _part -def convert_to_anthropic_tool_result(message: dict) -> AnthropicMessagesToolResultParam: +def convert_to_anthropic_tool_result( + message: Union[dict, ChatCompletionToolMessage, ChatCompletionFunctionMessage] +) -> AnthropicMessagesToolResultParam: """ OpenAI message with a tool result looks like: { @@ -1214,7 +1219,7 @@ def convert_to_anthropic_tool_result(message: dict) -> AnthropicMessagesToolResu return anthropic_tool_result if message["role"] == "function": content = message.get("content") # type: ignore - tool_call_id = message.get("tool_call_id") or str(uuid.uuid4()) + tool_call_id = message.get("tool_call_id") or str(uuid.uuid4()) # type: ignore anthropic_tool_result = AnthropicMessagesToolResultParam( type="tool_result", tool_use_id=tool_call_id, content=content ) @@ -1229,7 +1234,7 @@ def convert_to_anthropic_tool_result(message: dict) -> AnthropicMessagesToolResu def convert_function_to_anthropic_tool_invoke( - function_call, + function_call: Union[dict, ChatCompletionToolCallFunctionChunk], ) -> List[AnthropicMessagesToolUseParam]: try: anthropic_tool_invoke = [ @@ -1246,7 +1251,7 @@ def convert_function_to_anthropic_tool_invoke( def convert_to_anthropic_tool_invoke( - tool_calls: list, + tool_calls: List[ChatCompletionAssistantToolCall], ) -> List[AnthropicMessagesToolUseParam]: """ OpenAI tool invokes: @@ -1306,17 +1311,19 @@ def add_cache_control_to_content( anthropic_content_element: Union[ dict, AnthropicMessagesImageParam, AnthropicMessagesTextParam ], - orignal_content_element: dict, + orignal_content_element: Union[dict, AllMessageValues], ): - if "cache_control" in orignal_content_element: - anthropic_content_element["cache_control"] = orignal_content_element[ - "cache_control" - ] + cache_control_param = orignal_content_element.get("cache_control") + if cache_control_param is not None and isinstance(cache_control_param, dict): + transformed_param = ChatCompletionCachedContent(**cache_control_param) # type: ignore + + anthropic_content_element["cache_control"] = transformed_param + return anthropic_content_element def anthropic_messages_pt( - messages: list, + messages: List[AllMessageValues], model: str, llm_provider: str, ) -> List[ @@ -1347,10 +1354,21 @@ def anthropic_messages_pt( while msg_i < len(messages): user_content: List[AnthropicMessagesUserMessageValues] = [] init_msg_i = msg_i + if isinstance(messages[msg_i], BaseModel): + messages[msg_i] = dict(messages[msg_i]) # type: ignore ## MERGE CONSECUTIVE USER CONTENT ## while msg_i < len(messages) and messages[msg_i]["role"] in user_message_types: - if isinstance(messages[msg_i]["content"], list): - for m in messages[msg_i]["content"]: + user_message_types_block: Union[ + ChatCompletionToolMessage, + ChatCompletionUserMessage, + ChatCompletionFunctionMessage, + ] = messages[ + msg_i + ] # type: ignore + if user_message_types_block["content"] and isinstance( + user_message_types_block["content"], list + ): + for m in user_message_types_block["content"]: if m.get("type", "") == "image_url": image_chunk = convert_to_anthropic_image_obj( m["image_url"]["url"] @@ -1381,15 +1399,24 @@ def anthropic_messages_pt( ) user_content.append(anthropic_content_element) elif ( - messages[msg_i]["role"] == "tool" - or messages[msg_i]["role"] == "function" + user_message_types_block["role"] == "tool" + or user_message_types_block["role"] == "function" ): # OpenAI's tool message content will always be a string - user_content.append(convert_to_anthropic_tool_result(messages[msg_i])) - else: user_content.append( - {"type": "text", "text": messages[msg_i]["content"]} + convert_to_anthropic_tool_result(user_message_types_block) ) + elif isinstance(user_message_types_block["content"], str): + _anthropic_content_text_element: AnthropicMessagesTextParam = { + "type": "text", + "text": user_message_types_block["content"], + } + anthropic_content_element = add_cache_control_to_content( + anthropic_content_element=_anthropic_content_text_element, + orignal_content_element=user_message_types_block, + ) + + user_content.append(anthropic_content_element) msg_i += 1 @@ -1399,10 +1426,11 @@ def anthropic_messages_pt( assistant_content: List[AnthropicMessagesAssistantMessageValues] = [] ## MERGE CONSECUTIVE ASSISTANT CONTENT ## while msg_i < len(messages) and messages[msg_i]["role"] == "assistant": - if "content" in messages[msg_i] and isinstance( - messages[msg_i]["content"], list + assistant_content_block: ChatCompletionAssistantMessage = messages[msg_i] # type: ignore + if "content" in assistant_content_block and isinstance( + assistant_content_block["content"], list ): - for m in messages[msg_i]["content"]: + for m in assistant_content_block["content"]: # handle text if ( m.get("type", "") == "text" and len(m.get("text", "")) > 0 @@ -1416,35 +1444,37 @@ def anthropic_messages_pt( ) assistant_content.append(anthropic_message) elif ( - "content" in messages[msg_i] - and isinstance(messages[msg_i]["content"], str) - and len(messages[msg_i]["content"]) - > 0 # don't pass empty text blocks. anthropic api raises errors. + "content" in assistant_content_block + and isinstance(assistant_content_block["content"], str) + and assistant_content_block[ + "content" + ] # don't pass empty text blocks. anthropic api raises errors. ): _anthropic_text_content_element = { "type": "text", - "text": messages[msg_i]["content"], + "text": assistant_content_block["content"], } anthropic_content_element = add_cache_control_to_content( anthropic_content_element=_anthropic_text_content_element, - orignal_content_element=messages[msg_i], + orignal_content_element=assistant_content_block, ) assistant_content.append(anthropic_content_element) - if messages[msg_i].get( - "tool_calls", [] + assistant_tool_calls = assistant_content_block.get("tool_calls") + if ( + assistant_tool_calls is not None ): # support assistant tool invoke conversion assistant_content.extend( - convert_to_anthropic_tool_invoke(messages[msg_i]["tool_calls"]) + convert_to_anthropic_tool_invoke(assistant_tool_calls) ) - if messages[msg_i].get("function_call"): + assistant_function_call = assistant_content_block.get("function_call") + + if assistant_function_call is not None: assistant_content.extend( - convert_function_to_anthropic_tool_invoke( - messages[msg_i]["function_call"] - ) + convert_function_to_anthropic_tool_invoke(assistant_function_call) ) msg_i += 1 diff --git a/litellm/tests/test_anthropic_prompt_caching.py b/litellm/tests/test_anthropic_prompt_caching.py index b9c70f0c3b8..06f6916ed97 100644 --- a/litellm/tests/test_anthropic_prompt_caching.py +++ b/litellm/tests/test_anthropic_prompt_caching.py @@ -222,6 +222,94 @@ async def test_anthropic_api_prompt_caching_basic(): ) +@pytest.mark.asyncio() +async def test_anthropic_api_prompt_caching_with_content_str(): + from litellm.llms.prompt_templates.factory import anthropic_messages_pt + + system_message = [ + { + "role": "system", + "content": "Here is the full text of a complex legal agreement", + "cache_control": {"type": "ephemeral"}, + }, + ] + translated_system_message = litellm.AnthropicConfig().translate_system_message( + messages=system_message + ) + + assert translated_system_message == [ + # System Message + { + "type": "text", + "text": "Here is the full text of a complex legal agreement", + "cache_control": {"type": "ephemeral"}, + } + ] + user_messages = [ + # marked for caching with the cache_control parameter, so that this checkpoint can read from the previous cache. + { + "role": "user", + "content": "What are the key terms and conditions in this agreement?", + "cache_control": {"type": "ephemeral"}, + }, + { + "role": "assistant", + "content": "Certainly! the key terms and conditions are the following: the contract is 1 year long for $10/mo", + }, + # The final turn is marked with cache-control, for continuing in followups. + { + "role": "user", + "content": "What are the key terms and conditions in this agreement?", + "cache_control": {"type": "ephemeral"}, + }, + ] + + translated_messages = anthropic_messages_pt( + messages=user_messages, + model="claude-3-5-sonnet-20240620", + llm_provider="anthropic", + ) + + expected_messages = [ + { + "role": "user", + "content": [ + { + "type": "text", + "text": "What are the key terms and conditions in this agreement?", + "cache_control": {"type": "ephemeral"}, + } + ], + }, + { + "role": "assistant", + "content": [ + { + "type": "text", + "text": "Certainly! the key terms and conditions are the following: the contract is 1 year long for $10/mo", + } + ], + }, + # The final turn is marked with cache-control, for continuing in followups. + { + "role": "user", + "content": [ + { + "type": "text", + "text": "What are the key terms and conditions in this agreement?", + "cache_control": {"type": "ephemeral"}, + } + ], + }, + ] + + assert len(translated_messages) == len(expected_messages) + for idx, i in enumerate(translated_messages): + assert ( + i == expected_messages[idx] + ), "Error on idx={}. Got={}, Expected={}".format(idx, i, expected_messages[idx]) + + @pytest.mark.asyncio() async def test_anthropic_api_prompt_caching_no_headers(): litellm.set_verbose = True diff --git a/litellm/types/llms/anthropic.py b/litellm/types/llms/anthropic.py index 7b856a284f1..720abf8dde4 100644 --- a/litellm/types/llms/anthropic.py +++ b/litellm/types/llms/anthropic.py @@ -3,6 +3,8 @@ from typing import Any, Dict, Iterable, List, Optional, Union from pydantic import BaseModel, validator from typing_extensions import Literal, Required, TypedDict +from .openai import ChatCompletionCachedContent + class AnthropicMessagesToolChoice(TypedDict, total=False): type: Required[Literal["auto", "any", "tool"]] @@ -18,7 +20,7 @@ class AnthropicMessagesTool(TypedDict, total=False): class AnthropicMessagesTextParam(TypedDict, total=False): type: Literal["text"] text: str - cache_control: Optional[dict] + cache_control: Optional[Union[dict, ChatCompletionCachedContent]] class AnthropicMessagesToolUseParam(TypedDict): @@ -58,7 +60,7 @@ class AnthropicImageParamSource(TypedDict): class AnthropicMessagesImageParam(TypedDict, total=False): type: Literal["image"] source: AnthropicImageParamSource - cache_control: Optional[dict] + cache_control: Optional[Union[dict, ChatCompletionCachedContent]] class AnthropicMessagesToolResultContent(TypedDict): @@ -97,7 +99,7 @@ class AnthropicMetadata(TypedDict, total=False): class AnthropicSystemMessageContent(TypedDict, total=False): type: str text: str - cache_control: Optional[dict] + cache_control: Optional[Union[dict, ChatCompletionCachedContent]] class AnthropicMessagesRequest(TypedDict, total=False): diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 0219145c642..788199c00d5 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -354,14 +354,18 @@ class ChatCompletionImageObject(TypedDict): image_url: ChatCompletionImageUrlObject -class ChatCompletionUserMessage(TypedDict): +class OpenAIChatCompletionUserMessage(TypedDict): role: Literal["user"] content: Union[ str, Iterable[Union[ChatCompletionTextObject, ChatCompletionImageObject]] ] -class ChatCompletionAssistantMessage(TypedDict, total=False): +class ChatCompletionUserMessage(OpenAIChatCompletionUserMessage, total=False): + cache_control: ChatCompletionCachedContent + + +class OpenAIChatCompletionAssistantMessage(TypedDict, total=False): role: Required[Literal["assistant"]] content: Optional[Union[str, Iterable[ChatCompletionTextObject]]] name: Optional[str] @@ -369,6 +373,10 @@ class ChatCompletionAssistantMessage(TypedDict, total=False): function_call: Optional[ChatCompletionToolCallFunctionChunk] +class ChatCompletionAssistantMessage(OpenAIChatCompletionAssistantMessage, total=False): + cache_control: ChatCompletionCachedContent + + class ChatCompletionToolMessage(TypedDict): role: Literal["tool"] content: str @@ -381,12 +389,16 @@ class ChatCompletionFunctionMessage(TypedDict): name: str -class ChatCompletionSystemMessage(TypedDict, total=False): +class OpenAIChatCompletionSystemMessage(TypedDict, total=False): role: Required[Literal["system"]] content: Required[Union[str, List]] name: str +class ChatCompletionSystemMessage(OpenAIChatCompletionSystemMessage, total=False): + cache_control: ChatCompletionCachedContent + + AllMessageValues = Union[ ChatCompletionUserMessage, ChatCompletionAssistantMessage,