mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(prompting): preserve separator for assistant(tc)->assistant edge case
When scanning backward over counted messages, preserve old behavior for adjacent assistant turns by inserting user_continue if the immediate previous raw message is assistant. This handles malformed assistant(tool_calls)->assistant(no-tool-calls) inputs without splitting valid assistant(tool_calls)->tool chains. Made-with: Cursor
This commit is contained in:
parent
22fc08d602
commit
f29b4981a0
2 changed files with 65 additions and 10 deletions
|
|
@ -298,16 +298,23 @@ def _insert_user_continue_message(
|
|||
curr_message = result_messages[i]
|
||||
inserted_continue_message = False
|
||||
if _counts_for_alternation(curr_message) and curr_message["role"] == "assistant":
|
||||
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
|
||||
# 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
|
||||
|
||||
|
|
|
|||
|
|
@ -903,6 +903,54 @@ def test_ensure_alternating_roles_does_not_split_tool_call_chain():
|
|||
]
|
||||
|
||||
|
||||
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?"},
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue