Merge pull request #24015 from Sameerlite/litellm_fix_ensure_alternating_roles

Litellm fix ensure alternating roles
This commit is contained in:
Sameer Kankute 2026-03-20 18:37:23 +05:30 • committed by GitHub
commit 9272483f77
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 268 additions and 33 deletions

View file

@ -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

View file

@ -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