feat(anthropic.py): support 'cache_control' param for content when it is a string

This commit is contained in:
Krrish Dholakia 2024-09-04 16:01:12 -07:00
parent 3fac0349c2
commit 45891fed4f
5 changed files with 224 additions and 78 deletions

View file

@ -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:

View file

@ -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

View file

@ -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

View file

@ -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):

View file

@ -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,