mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
feat(anthropic.py): support 'cache_control' param for content when it is a string
This commit is contained in:
parent
3fac0349c2
commit
45891fed4f
5 changed files with 224 additions and 78 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue