mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge pull request #24015 from Sameerlite/litellm_fix_ensure_alternating_roles
Litellm fix ensure alternating roles
This commit is contained in:
commit
9272483f77
2 changed files with 268 additions and 33 deletions
|
|
@ -257,6 +257,15 @@ def detect_first_expected_role(
|
|||
return None
|
||||
|
||||
|
||||
def _counts_for_alternation(message: AllMessageValues) -> bool:
|
||||
role = message.get("role")
|
||||
if role == "user":
|
||||
return True
|
||||
if role == "assistant":
|
||||
return not bool(message.get("tool_calls"))
|
||||
return False
|
||||
|
||||
|
||||
def _insert_user_continue_message(
|
||||
messages: List[AllMessageValues],
|
||||
user_continue_message: Optional[ChatCompletionUserMessage],
|
||||
|
|
@ -269,8 +278,8 @@ def _insert_user_continue_message(
|
|||
2. Final assistant message
|
||||
3. Consecutive assistant messages
|
||||
|
||||
Only inserts messages between consecutive assistant messages,
|
||||
ignoring all other role types.
|
||||
Skips tool messages and assistant messages with tool calls in the
|
||||
alternation check, matching strict templates like llama.cpp.
|
||||
"""
|
||||
if not messages:
|
||||
return messages
|
||||
|
|
@ -278,25 +287,39 @@ def _insert_user_continue_message(
|
|||
result_messages = messages.copy() # Don't modify the input list
|
||||
continue_message = user_continue_message or DEFAULT_USER_CONTINUE_MESSAGE
|
||||
|
||||
# Handle first message if it's an assistant message
|
||||
# Handle first message if it's an assistant message — always prepend
|
||||
# user_continue regardless of tool_calls, to preserve backward compatibility.
|
||||
if result_messages[0]["role"] == "assistant":
|
||||
result_messages.insert(0, continue_message)
|
||||
|
||||
# Handle consecutive assistant messages and final message
|
||||
i = 1 # Start from second message since we handled first message
|
||||
# Handle consecutive assistant messages in the counted sequence
|
||||
i = 1
|
||||
while i < len(result_messages):
|
||||
curr_message = result_messages[i]
|
||||
prev_message = result_messages[i - 1]
|
||||
|
||||
# Only check for consecutive assistant messages
|
||||
# Ignore all other role types
|
||||
if curr_message["role"] == "assistant" and prev_message["role"] == "assistant":
|
||||
result_messages.insert(i, continue_message)
|
||||
i += 2 # Skip over the message we just inserted
|
||||
else:
|
||||
inserted_continue_message = False
|
||||
if _counts_for_alternation(curr_message) and curr_message["role"] == "assistant":
|
||||
# Preserve old behavior for malformed adjacent assistant sequences like
|
||||
# assistant(tool_calls) -> assistant(no-tool-calls) with no tool message.
|
||||
if i > 0 and result_messages[i - 1].get("role") == "assistant":
|
||||
result_messages.insert(i, continue_message)
|
||||
i += 2
|
||||
inserted_continue_message = True
|
||||
else:
|
||||
j = i - 1
|
||||
while j >= 0:
|
||||
previous_message = result_messages[j]
|
||||
if _counts_for_alternation(previous_message):
|
||||
if previous_message["role"] == "assistant":
|
||||
result_messages.insert(i, continue_message)
|
||||
i += 2
|
||||
inserted_continue_message = True
|
||||
break
|
||||
j -= 1
|
||||
if not inserted_continue_message:
|
||||
i += 1
|
||||
|
||||
# Handle final message
|
||||
# Handle final message — append user_continue after any trailing assistant,
|
||||
# including ones with tool_calls, to preserve backward compatibility.
|
||||
if result_messages[-1]["role"] == "assistant" and ensure_alternating_roles:
|
||||
result_messages.append(continue_message)
|
||||
|
||||
|
|
@ -311,34 +334,24 @@ def _insert_assistant_continue_message(
|
|||
"""
|
||||
Add assistant continuation messages between consecutive user messages.
|
||||
|
||||
Args:
|
||||
messages: List of message dictionaries
|
||||
assistant_continue_message: Optional custom assistant message
|
||||
ensure_alternating_roles: Whether to enforce alternating roles
|
||||
|
||||
Returns:
|
||||
Modified list of messages with inserted assistant messages
|
||||
Only checks directly adjacent messages to preserve backward compatibility.
|
||||
"""
|
||||
if not ensure_alternating_roles or len(messages) <= 1:
|
||||
return messages
|
||||
|
||||
# Create a new list to store modified messages
|
||||
continue_message = assistant_continue_message or DEFAULT_ASSISTANT_CONTINUE_MESSAGE
|
||||
|
||||
modified_messages: List[AllMessageValues] = []
|
||||
|
||||
for i, message in enumerate(messages):
|
||||
modified_messages.append(message)
|
||||
|
||||
# Check if we need to insert an assistant message
|
||||
if (
|
||||
i < len(messages) - 1 # Not the last message
|
||||
and message.get("role") == "user" # Current is user
|
||||
i < len(messages) - 1
|
||||
and message.get("role") == "user"
|
||||
and messages[i + 1].get("role") == "user"
|
||||
): # Next is user
|
||||
# Insert assistant message
|
||||
continue_message = (
|
||||
assistant_continue_message or DEFAULT_ASSISTANT_CONTINUE_MESSAGE
|
||||
)
|
||||
):
|
||||
modified_messages.append(message)
|
||||
modified_messages.append(continue_message)
|
||||
else:
|
||||
modified_messages.append(message)
|
||||
|
||||
return modified_messages
|
||||
|
||||
|
|
|
|||
|
|
@ -775,6 +775,228 @@ def test_ensure_alternating_roles(
|
|||
assert messages == expected_messages
|
||||
|
||||
|
||||
def test_ensure_alternating_roles_with_tool_calls():
|
||||
"""Fixes Regression in #18685 """
|
||||
messages = [
|
||||
{"role": "user", "content": "What's the weather?"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_123",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"arguments": '{"location": "NYC"}',
|
||||
},
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "tool", "tool_call_id": "call_123", "content": "72F, sunny"},
|
||||
{"role": "assistant", "content": "It's 72F and sunny in NYC."},
|
||||
{"role": "user", "content": "What about tomorrow?"},
|
||||
{"role": "user", "content": "And the day after?"},
|
||||
{"role": "user", "content": "What about next week?"},
|
||||
]
|
||||
|
||||
transformed_messages = get_completion_messages(
|
||||
messages=messages,
|
||||
assistant_continue_message=None,
|
||||
user_continue_message=None,
|
||||
ensure_alternating_roles=True,
|
||||
)
|
||||
|
||||
assert transformed_messages == [
|
||||
{"role": "user", "content": "What's the weather?"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_123",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"arguments": '{"location": "NYC"}',
|
||||
},
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "tool", "tool_call_id": "call_123", "content": "72F, sunny"},
|
||||
{"role": "assistant", "content": "It's 72F and sunny in NYC."},
|
||||
{"role": "user", "content": "What about tomorrow?"},
|
||||
{"role": "assistant", "content": "Please continue."},
|
||||
{"role": "user", "content": "And the day after?"},
|
||||
{"role": "assistant", "content": "Please continue."},
|
||||
{"role": "user", "content": "What about next week?"},
|
||||
]
|
||||
|
||||
|
||||
def test_ensure_alternating_roles_three_consecutive_assistants():
|
||||
messages = [
|
||||
{"role": "user", "content": "Hello"},
|
||||
{"role": "assistant", "content": "A1"},
|
||||
{"role": "assistant", "content": "A2"},
|
||||
{"role": "assistant", "content": "A3"},
|
||||
]
|
||||
|
||||
transformed_messages = get_completion_messages(
|
||||
messages=messages,
|
||||
assistant_continue_message=None,
|
||||
user_continue_message=None,
|
||||
ensure_alternating_roles=True,
|
||||
)
|
||||
|
||||
assert transformed_messages == [
|
||||
{"role": "user", "content": "Hello"},
|
||||
{"role": "assistant", "content": "A1"},
|
||||
{"role": "user", "content": "Please continue."},
|
||||
{"role": "assistant", "content": "A2"},
|
||||
{"role": "user", "content": "Please continue."},
|
||||
{"role": "assistant", "content": "A3"},
|
||||
{"role": "user", "content": "Please continue."},
|
||||
]
|
||||
|
||||
|
||||
def test_ensure_alternating_roles_does_not_split_tool_call_chain():
|
||||
"""Tool-call chains [user, assistant(tc), tool, user] are preserved as-is."""
|
||||
messages = [
|
||||
{"role": "user", "content": "Search for X"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "c1",
|
||||
"type": "function",
|
||||
"function": {"name": "search", "arguments": "{}"},
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "tool", "tool_call_id": "c1", "content": "results"},
|
||||
{"role": "user", "content": "Thanks, now do Y"},
|
||||
]
|
||||
|
||||
transformed_messages = get_completion_messages(
|
||||
messages=messages,
|
||||
assistant_continue_message=None,
|
||||
user_continue_message=None,
|
||||
ensure_alternating_roles=True,
|
||||
)
|
||||
|
||||
assert transformed_messages == [
|
||||
{"role": "user", "content": "Search for X"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "c1",
|
||||
"type": "function",
|
||||
"function": {"name": "search", "arguments": "{}"},
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "tool", "tool_call_id": "c1", "content": "results"},
|
||||
{"role": "user", "content": "Thanks, now do Y"},
|
||||
]
|
||||
|
||||
|
||||
def test_ensure_alternating_roles_assistant_tool_call_then_assistant():
|
||||
"""
|
||||
Preserve old behavior for malformed adjacent assistant turns:
|
||||
[assistant(tool_calls), assistant(no-tool-calls), user] should insert
|
||||
user_continue between assistant messages.
|
||||
"""
|
||||
messages = [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "c1",
|
||||
"type": "function",
|
||||
"function": {"name": "search", "arguments": "{}"},
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "assistant", "content": "Here's what I found."},
|
||||
{"role": "user", "content": "Thanks"},
|
||||
]
|
||||
|
||||
transformed_messages = get_completion_messages(
|
||||
messages=messages,
|
||||
assistant_continue_message=None,
|
||||
user_continue_message=None,
|
||||
ensure_alternating_roles=True,
|
||||
)
|
||||
|
||||
assert transformed_messages == [
|
||||
{"role": "user", "content": "Please continue."},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "c1",
|
||||
"type": "function",
|
||||
"function": {"name": "search", "arguments": "{}"},
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "user", "content": "Please continue."},
|
||||
{"role": "assistant", "content": "Here's what I found."},
|
||||
{"role": "user", "content": "Thanks"},
|
||||
]
|
||||
|
||||
|
||||
def test_ensure_alternating_roles_trailing_tool_call_assistant():
|
||||
messages = [
|
||||
{"role": "user", "content": "What's the weather?"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_abc",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"arguments": '{"location": "NYC"}',
|
||||
},
|
||||
}
|
||||
],
|
||||
},
|
||||
]
|
||||
|
||||
transformed_messages = get_completion_messages(
|
||||
messages=messages,
|
||||
assistant_continue_message=None,
|
||||
user_continue_message=None,
|
||||
ensure_alternating_roles=True,
|
||||
)
|
||||
|
||||
assert transformed_messages == [
|
||||
{"role": "user", "content": "What's the weather?"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_abc",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"arguments": '{"location": "NYC"}',
|
||||
},
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "user", "content": "Please continue."},
|
||||
]
|
||||
|
||||
|
||||
def test_alternating_roles_e2e():
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
import json
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue