diff --git a/litellm/llms/prompt_templates/factory.py b/litellm/llms/prompt_templates/factory.py index 1a06123bd55..dc4c417cd37 100644 --- a/litellm/llms/prompt_templates/factory.py +++ b/litellm/llms/prompt_templates/factory.py @@ -16,7 +16,11 @@ from litellm.types.completion import ( ChatCompletionUserMessageParam, ChatCompletionSystemMessageParam, ChatCompletionMessageParam, + ChatCompletionFunctionMessageParam, + ChatCompletionMessageToolCallParam, ) +from litellm.types.llms.anthropic import * +import uuid def default_pt(messages): @@ -947,8 +951,10 @@ def anthropic_messages_pt(messages: list): # reformat messages to ensure user/assistant are alternating, if there's either 2 consecutive 'user' messages or 2 consecutive 'assistant' message, merge them. new_messages = [] msg_i = 0 + tool_use_param = False while msg_i < len(messages): user_content = [] + init_msg_i = msg_i ## 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): @@ -995,9 +1001,48 @@ def anthropic_messages_pt(messages: list): msg_i += 1 + ## MERGE CONSECUTIVE FUNCTION CONTENT ## + while msg_i < len(messages) and messages[msg_i]["role"] == "function": + """ + Anthropic function message: "role", "name", "input", "id" + OpenAI function message: "content", "name", "role" + + - Check if received message is a tool call input or model text response + """ + tool_use_param = True + _message = ChatCompletionFunctionMessageParam(**messages[msg_i]) # type: ignore + anthropic_function_message: Optional[ + AnthropicMessagesAssistantMessageValues + ] = None + try: + anthropic_function_message = ( + AnthopicMessagesAssistantMessageToolCallParam(type="tool_use") + ) + anthropic_function_message["input"] = json.loads(_message["content"]) + anthropic_function_message["id"] = str(uuid.uuid4()) + anthropic_function_message["name"] = _message["name"] + except Exception as e: + litellm.print_verbose( + "Invalid dictionary content. Treating as text instead." + ) + anthropic_function_message = ( + AnthopicMessagesAssistantMessageTextContentParam(type="text") + ) + anthropic_function_message["text"] = _message["content"] + + assistant_content.append(anthropic_function_message) # type: ignore + + msg_i += 1 + if assistant_content: new_messages.append({"role": "assistant", "content": assistant_content}) + if msg_i == init_msg_i: # prevent infinite loops + raise Exception( + "Invalid Message passed in - {}. File an issue https://github.com/BerriAI/litellm/issues".format( + messages[msg_i] + ) + ) if not new_messages or new_messages[0]["role"] != "user": if litellm.modify_params: new_messages.insert( @@ -1009,12 +1054,26 @@ def anthropic_messages_pt(messages: list): ) if new_messages[-1]["role"] == "assistant": - for content in new_messages[-1]["content"]: - if isinstance(content, dict) and content["type"] == "text": - content["text"] = content[ - "text" - ].rstrip() # no trailing whitespace for final assistant message - + if tool_use_param == True: + """ + Final assistant message cannot be a tool use param. + """ + if litellm.modify_params: + new_messages.append( + {"role": "user", "content": [{"type": "text", "text": "."}]} + ) + else: + raise Exception( + "AnthropicError: Invalid last message. Your API request included an `assistant` message in the final position, which would pre-fill the `assistant` response. When using tools, pre-filling the `assistant` response is not supported. set 'litellm.modify_params = True' or 'litellm_settings:modify_params = True' on proxy, to insert a placeholder user message - '.' as the last message, " + ) + if isinstance(new_messages[-1]["content"], str): + new_messages[-1]["content"] = new_messages[-1]["content"].rstrip() + elif isinstance(new_messages[-1]["content"], list): + for content in new_messages[-1]["content"]: + if isinstance(content, dict) and content["type"] == "text": + content["text"] = content[ + "text" + ].rstrip() # no trailing whitespace for final assistant message return new_messages diff --git a/litellm/tests/test_completion.py b/litellm/tests/test_completion.py index 2b6c05558a3..1108db860d4 100644 --- a/litellm/tests/test_completion.py +++ b/litellm/tests/test_completion.py @@ -2355,6 +2355,59 @@ def test_completion_with_fallbacks(): # test_completion_with_fallbacks() + + +@pytest.mark.parametrize( + "function_call", + [ + [{"role": "function", "name": "get_capital", "content": "Kokoko"}], + [ + {"role": "function", "name": "get_capital", "content": "Kokoko"}, + {"role": "function", "name": "get_capital", "content": "Kokoko"}, + ], + ], +) +def test_completion_anthropic_hanging(function_call): + litellm.modify_params = True + messages = [ + { + "role": "user", + "content": "What's the capital of fictional country Ubabababababaaba? Use your tools.", + }, + { + "role": "assistant", + "function_call": { + "name": "get_capital", + "arguments": '{"country": "Ubabababababaaba"}', + }, + }, + ] + messages = messages + function_call + litellm.completion( + model="claude-3-haiku-20240307", + messages=messages, + tools=[ + { + "function": { + "name": "get_capital", + "description": "Get the capital of a country", + "parameters": { + "title": "GetCapitalToolArgs", + "type": "object", + "properties": { + "country": {"title": "Country", "type": "string"} + }, + "required": ["country"], + }, + }, + "type": "function", + } + ], + tool_choice="auto", + temperature=0.0, + ) + + def test_completion_anyscale_api(): try: # litellm.set_verbose=True diff --git a/litellm/types/completion.py b/litellm/types/completion.py index 9df860f5853..c5148b16896 100644 --- a/litellm/types/completion.py +++ b/litellm/types/completion.py @@ -93,6 +93,28 @@ class Function(TypedDict, total=False): """The name of the function to call.""" +class ChatCompletionToolMessageParam(TypedDict, total=False): + content: Required[str] + """The contents of the tool message.""" + + role: Required[Literal["tool"]] + """The role of the messages author, in this case `tool`.""" + + tool_call_id: Required[str] + """Tool call that this message is responding to.""" + + +class ChatCompletionFunctionMessageParam(TypedDict, total=False): + content: Required[Optional[str]] + """The contents of the function message.""" + + name: Required[str] + """The name of the function to call.""" + + role: Required[Literal["function"]] + """The role of the messages author, in this case `function`.""" + + class ChatCompletionMessageToolCallParam(TypedDict, total=False): id: Required[str] """The ID of the tool call.""" @@ -136,6 +158,8 @@ ChatCompletionMessageParam = Union[ ChatCompletionSystemMessageParam, ChatCompletionUserMessageParam, ChatCompletionAssistantMessageParam, + ChatCompletionFunctionMessageParam, + ChatCompletionToolMessageParam, ] diff --git a/litellm/types/llms/anthropic.py b/litellm/types/llms/anthropic.py new file mode 100644 index 00000000000..faf4aa3564a --- /dev/null +++ b/litellm/types/llms/anthropic.py @@ -0,0 +1,42 @@ +from typing import List, Optional, Union, Iterable + +from pydantic import BaseModel, validator + +from typing_extensions import Literal, Required, TypedDict + + +class AnthopicMessagesAssistantMessageTextContentParam(TypedDict, total=False): + type: Required[Literal["text"]] + + text: str + + +class AnthopicMessagesAssistantMessageToolCallParam(TypedDict, total=False): + type: Required[Literal["tool_use"]] + + id: str + + name: str + + input: dict + + +AnthropicMessagesAssistantMessageValues = Union[ + AnthopicMessagesAssistantMessageTextContentParam, + AnthopicMessagesAssistantMessageToolCallParam, +] + + +class AnthopicMessagesAssistantMessageParam(TypedDict, total=False): + content: Required[Union[str, Iterable[AnthropicMessagesAssistantMessageValues]]] + """The contents of the system message.""" + + role: Required[Literal["assistant"]] + """The role of the messages author, in this case `author`.""" + + name: str + """An optional name for the participant. + + Provides the model information to differentiate between participants of the same + role. + """