add test for large context in system message for anthropic

This commit is contained in:
Ishaan Jaff 2024-08-14 17:03:10 -07:00
parent 76a5f5d433
commit 78a2013e51
4 changed files with 128 additions and 17 deletions

View file

@ -35,6 +35,7 @@ from litellm.types.llms.anthropic import (
AnthropicResponseContentBlockText,
AnthropicResponseContentBlockToolUse,
AnthropicResponseUsageBlock,
AnthropicSystemMessageContent,
ContentBlockDelta,
ContentBlockStart,
ContentBlockStop,
@ -907,7 +908,7 @@ class AnthropicChatCompletion(BaseLLM):
# Separate system prompt from rest of message
system_prompt_indices = []
system_prompt = ""
system_prompt_dict = None
anthropic_system_message_list = None
for idx, message in enumerate(messages):
if message["role"] == "system":
valid_content: bool = False
@ -915,19 +916,24 @@ class AnthropicChatCompletion(BaseLLM):
system_prompt += message["content"]
valid_content = True
elif isinstance(message["content"], list):
for content in message["content"]:
system_prompt += content.get("text", "")
valid_content = True
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"]
)
# Handle Anthropic API context caching
if "cache_control" in message:
system_prompt_dict = [
{
"cache_control": message["cache_control"],
"text": system_prompt,
"type": "text",
}
]
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)
@ -938,8 +944,8 @@ class AnthropicChatCompletion(BaseLLM):
optional_params["system"] = system_prompt
# Handling anthropic API Prompt Caching
if system_prompt_dict is not None:
optional_params["system"] = system_prompt_dict
if anthropic_system_message_list is not None:
optional_params["system"] = anthropic_system_message_list
# Format rest of message according to anthropic guidelines
try:
messages = prompt_factory(

View file

@ -220,3 +220,102 @@ async def test_anthropic_api_prompt_caching_basic():
assert (response.usage.cache_read_input_tokens > 0) or (
response.usage.cache_creation_input_tokens > 0
)
@pytest.mark.asyncio
async def test_litellm_anthropic_prompt_caching_system():
# https://docs.anthropic.com/en/docs/build-with-claude/prompt-caching#prompt-caching-examples
# LArge Context Caching Example
mock_response = AsyncMock()
def return_val():
return {
"id": "msg_01XFDUDYJgAACzvnptvVoYEL",
"type": "message",
"role": "assistant",
"content": [{"type": "text", "text": "Hello!"}],
"model": "claude-3-5-sonnet-20240620",
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 12, "output_tokens": 6},
}
mock_response.json = return_val
litellm.set_verbose = True
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
return_value=mock_response,
) as mock_post:
# Act: Call the litellm.acompletion function
response = await litellm.acompletion(
api_key="mock_api_key",
model="anthropic/claude-3-5-sonnet-20240620",
messages=[
{
"role": "system",
"content": [
{
"type": "text",
"text": "You are an AI assistant tasked with analyzing legal documents.",
},
{
"type": "text",
"text": "Here is the full text of a complex legal agreement",
"cache_control": {"type": "ephemeral"},
},
],
},
{
"role": "user",
"content": "what are the key terms and conditions in this agreement?",
},
],
extra_headers={
"anthropic-version": "2023-06-01",
"anthropic-beta": "prompt-caching-2024-07-31",
},
)
# Print what was called on the mock
print("call args=", mock_post.call_args)
expected_url = "https://api.anthropic.com/v1/messages"
expected_headers = {
"accept": "application/json",
"content-type": "application/json",
"anthropic-version": "2023-06-01",
"anthropic-beta": "prompt-caching-2024-07-31",
"x-api-key": "mock_api_key",
}
expected_json = {
"system": [
{
"type": "text",
"text": "You are an AI assistant tasked with analyzing legal documents.",
},
{
"type": "text",
"text": "Here is the full text of a complex legal agreement",
"cache_control": {"type": "ephemeral"},
},
],
"messages": [
{
"role": "user",
"content": [
{
"type": "text",
"text": "what are the key terms and conditions in this agreement?",
}
],
}
],
"max_tokens": 4096,
"model": "claude-3-5-sonnet-20240620",
}
mock_post.assert_called_once_with(
expected_url, json=expected_json, headers=expected_headers, timeout=600.0
)

View file

@ -94,6 +94,12 @@ class AnthropicMetadata(TypedDict, total=False):
user_id: str
class AnthropicSystemMessageContent(TypedDict, total=False):
type: str
text: str
cache_control: Optional[dict]
class AnthropicMessagesRequest(TypedDict, total=False):
model: Required[str]
messages: Required[
@ -108,7 +114,7 @@ class AnthropicMessagesRequest(TypedDict, total=False):
metadata: AnthropicMetadata
stop_sequences: List[str]
stream: bool
system: str
system: Union[str, List]
temperature: float
tool_choice: AnthropicMessagesToolChoice
tools: List[AnthropicMessagesTool]

View file

@ -361,7 +361,7 @@ class ChatCompletionToolMessage(TypedDict):
class ChatCompletionSystemMessage(TypedDict, total=False):
role: Required[Literal["system"]]
content: Required[str]
content: Required[Union[str, List]]
name: str