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:
Sameer Kankute 2026-03-18 17:13:42 +05:30
parent 22fc08d602
commit f29b4981a0
2 changed files with 65 additions and 10 deletions

View file

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

View file

@ -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?"},