mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-06 08:16:43 +00:00
fix(factory.py): support 'function' openai message role for anthropic
Fixes https://github.com/BerriAI/litellm/issues/3446
This commit is contained in:
parent
b7ca9a53c9
commit
33472bfd2b
4 changed files with 184 additions and 6 deletions
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
42
litellm/types/llms/anthropic.py
Normal file
42
litellm/types/llms/anthropic.py
Normal file
|
|
@ -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.
|
||||
"""
|
||||
Loading…
Add table
Reference in a new issue