diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index aa2e07234d..739c3119cc 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -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], @@ -275,14 +284,6 @@ def _insert_user_continue_message( if not messages: return messages - 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 - result_messages = messages.copy() # Don't modify the input list continue_message = user_continue_message or DEFAULT_USER_CONTINUE_MESSAGE @@ -346,37 +347,33 @@ def _insert_assistant_continue_message( """ if not ensure_alternating_roles or len(messages) <= 1: return messages - - 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 - - # Create a new list to store modified messages - modified_messages: List[AllMessageValues] = [] + continue_message = assistant_continue_message or DEFAULT_ASSISTANT_CONTINUE_MESSAGE + insert_before_indexes = set() for i, message in enumerate(messages): + if message.get("role") != "user": + continue + + next_counted_index = i + 1 + while next_counted_index < len(messages) and not _counts_for_alternation( + messages[next_counted_index] + ): + next_counted_index += 1 + + if ( + next_counted_index < len(messages) + and messages[next_counted_index].get("role") == "user" + ): + # Insert before the next counted user turn. + # This avoids splitting assistant tool-call -> tool chains. + insert_before_indexes.add(next_counted_index) + + modified_messages: List[AllMessageValues] = [] + for idx, message in enumerate(messages): + if idx in insert_before_indexes: + modified_messages.append(continue_message) modified_messages.append(message) - if message.get("role") == "user" and _counts_for_alternation(message): - next_counted_index = i + 1 - while next_counted_index < len(messages) and not _counts_for_alternation( - messages[next_counted_index] - ): - next_counted_index += 1 - - if ( - next_counted_index < len(messages) - and messages[next_counted_index].get("role") == "user" - ): - continue_message = ( - assistant_continue_message or DEFAULT_ASSISTANT_CONTINUE_MESSAGE - ) - modified_messages.append(continue_message) - return modified_messages diff --git a/tests/llm_translation/test_prompt_factory.py b/tests/llm_translation/test_prompt_factory.py index 12efb47e06..355c23ae17 100644 --- a/tests/llm_translation/test_prompt_factory.py +++ b/tests/llm_translation/test_prompt_factory.py @@ -858,6 +858,50 @@ def test_ensure_alternating_roles_three_consecutive_assistants(): ] +def test_ensure_alternating_roles_does_not_split_tool_call_chain(): + 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": "assistant", "content": "Please continue."}, + {"role": "user", "content": "Thanks, now do Y"}, + ] + + def test_alternating_roles_e2e(): from litellm.llms.custom_httpx.http_handler import HTTPHandler import json