From 9e2b47d2b5c5d9489b10836ccf0693756c96ad87 Mon Sep 17 00:00:00 2001 From: Nilay <118994073+NilayYadav@users.noreply.github.com> Date: Tue, 28 Jul 2026 01:05:34 +0530 Subject: [PATCH] Studio: split parallel tool calls for Llama 3.x chat templates (#7426) * split parallel tool calls for single-call-only chat templates * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> Co-authored-by: Lee Jackson <130007945+Imagineer99@users.noreply.github.com> --- .../core/inference/chat_template_helpers.py | 72 +++++++- .../test_chat_template_tool_arguments.py | 154 ++++++++++++++++++ 2 files changed, 220 insertions(+), 6 deletions(-) diff --git a/studio/backend/core/inference/chat_template_helpers.py b/studio/backend/core/inference/chat_template_helpers.py index 528c059fbc..3a8463855b 100644 --- a/studio/backend/core/inference/chat_template_helpers.py +++ b/studio/backend/core/inference/chat_template_helpers.py @@ -326,6 +326,58 @@ def _normalize_tool_call_arguments(messages: list) -> list: return out if mutated else messages +def _take_tool_result(pending: list, call_id) -> Optional[dict]: + if call_id: + for i, result in enumerate(pending): + if result.get("tool_call_id") == call_id: + return pending.pop(i) + for i, result in enumerate(pending): + if not result.get("tool_call_id"): + return pending.pop(i) + return None + + +def _split_parallel_tool_calls(messages: list) -> list: + """Llama 3.x templates render one call per message, so split parallel calls + into consecutive single-call messages, each followed by its own result.""" + if not any(isinstance(m, dict) and len(m.get("tool_calls") or ()) > 1 for m in messages): + return messages + + out: list = [] + i = 0 + total = len(messages) + while i < total: + msg = messages[i] + calls = msg.get("tool_calls") if isinstance(msg, dict) else None + if not calls or len(calls) <= 1: + out.append(msg) + i += 1 + continue + + # Tool results right after this message answer its calls. + j = i + 1 + pending: list = [] + while ( + j < total + and isinstance(messages[j], dict) + and messages[j].get("role") in ("tool", "ipython") + ): + pending.append(messages[j]) + j += 1 + + for idx, call in enumerate(calls): + piece = {**msg, "tool_calls": [call]} + if idx: + piece["content"] = "" + out.append(piece) + result = _take_tool_result(pending, call.get("id") if isinstance(call, dict) else None) + if result is not None: + out.append(result) + out.extend(pending) + i = j + return out + + def apply_chat_template_for_generation( tokenizer, messages: list, @@ -378,13 +430,21 @@ def apply_chat_template_for_generation( try: return _render(messages) except Exception: - # Strict tool templates reject the JSON-string ``arguments`` form via - # TypeError or a broad Jinja raise_exception, so retry with dicts coerced. - # Original messages render first, so working templates stay byte-identical. + # Retry with repairs applied cumulatively. Originals render first, so + # working templates stay byte-identical. + candidates: list = [] normalized = _normalize_tool_call_arguments(messages) - if normalized is messages: - raise - return _render(normalized) + if normalized is not messages: + candidates.append(normalized) + split = _split_parallel_tool_calls(normalized) + if split is not normalized: + candidates.append(split) + for candidate in candidates: + try: + return _render(candidate) + except Exception: + continue + raise def render_native_template( diff --git a/studio/backend/tests/test_chat_template_tool_arguments.py b/studio/backend/tests/test_chat_template_tool_arguments.py index 13d1ecabaa..8a927ea93c 100644 --- a/studio/backend/tests/test_chat_template_tool_arguments.py +++ b/studio/backend/tests/test_chat_template_tool_arguments.py @@ -6,10 +6,14 @@ from the OpenAI JSON-string form to a dict before rendering. Strict tool templates (e.g. mlx-community Qwen3.5 checkpoints) iterate arguments.items() and raise "Can only get item pairs from a mapping." on the string form when a prior tool call is re-rendered on the next turn (MLX + transformers paths). + +It must likewise split parallel tool calls for templates that render only one +call per message (Llama 3.x). """ from __future__ import annotations +import json import sys from pathlib import Path @@ -21,6 +25,7 @@ if str(_BACKEND) not in sys.path: from core.inference.chat_template_helpers import ( # noqa: E402 _normalize_tool_call_arguments, + _split_parallel_tool_calls, apply_chat_template_for_generation, ) @@ -155,3 +160,152 @@ def test_unrelated_template_error_still_propagates_with_dict_args(): with pytest.raises(ValueError, match = "broken"): apply_chat_template_for_generation(_AlwaysRaises(), _conv({"query": "x"})) + + +def _parallel_conv( + *, + ids = ("c1", "c2"), + results_have_ids = True, + content = "sure", +): + a, b = ids + return [ + {"role": "user", "content": "search then render"}, + { + "role": "assistant", + "content": content, + "tool_calls": [ + { + "type": "function", + "id": a, + "function": {"name": "web_search", "arguments": {"query": "x"}}, + }, + { + "type": "function", + "id": b, + "function": {"name": "render_html", "arguments": {"html": ""}}, + }, + ], + }, + { + "role": "tool", + "name": "web_search", + **({"tool_call_id": a} if results_have_ids else {}), + "content": "no text", + }, + { + "role": "tool", + "name": "render_html", + **({"tool_call_id": b} if results_have_ids else {}), + "content": "ok", + }, + ] + + +class _SingleToolCallTokenizer: + """Mimics the Llama 3.x template: rejects >1 call per message.""" + + def apply_chat_template( + self, + messages, + *, + tokenize = False, + add_generation_prompt = True, + **kw, + ): + for msg in messages: + if len(msg.get("tool_calls") or ()) > 1: + raise ValueError("This model only supports single tool-calls at once!") + return "RENDERED" + + +def test_parallel_calls_split_into_sequential_single_call_turns(): + out = _split_parallel_tool_calls(_parallel_conv()) + assert [(m["role"], m.get("name")) for m in out] == [ + ("user", None), + ("assistant", None), + ("tool", "web_search"), + ("assistant", None), + ("tool", "render_html"), + ] + assert [len(m["tool_calls"]) for m in out if m.get("tool_calls")] == [1, 1] + assert out[1]["tool_calls"][0]["function"]["name"] == "web_search" + assert out[3]["tool_calls"][0]["function"]["name"] == "render_html" + + +def test_split_pairs_results_by_tool_call_id_not_position(): + conv = _parallel_conv() + conv[2], conv[3] = conv[3], conv[2] # results arrive out of order + out = _split_parallel_tool_calls(conv) + assert out[1]["tool_calls"][0]["id"] == "c1" and out[2]["tool_call_id"] == "c1" + assert out[3]["tool_calls"][0]["id"] == "c2" and out[4]["tool_call_id"] == "c2" + + +def test_split_falls_back_to_order_when_results_have_no_ids(): + out = _split_parallel_tool_calls(_parallel_conv(results_have_ids = False)) + assert [m["role"] for m in out] == ["user", "assistant", "tool", "assistant", "tool"] + assert out[2]["name"] == "web_search" and out[4]["name"] == "render_html" + + +def test_split_keeps_content_on_first_piece_only(): + out = _split_parallel_tool_calls(_parallel_conv(content = "sure")) + assert out[1]["content"] == "sure" + assert out[3]["content"] == "" + + +def test_split_keeps_unmatched_results_after_the_split(): + conv = _parallel_conv() + del conv[3] # second call never returned a result + out = _split_parallel_tool_calls(conv) + assert [m["role"] for m in out] == ["user", "assistant", "tool", "assistant"] + + +def test_split_leaves_later_turns_intact(): + conv = _parallel_conv() + [ + {"role": "assistant", "content": "done"}, + {"role": "user", "content": "thanks"}, + ] + out = _split_parallel_tool_calls(conv) + assert [m["role"] for m in out[-2:]] == ["assistant", "user"] + assert out[-2]["content"] == "done" + + +def test_single_call_and_plain_conversations_pass_through_unchanged(): + conv = _conv({"query": "x"}) + assert _split_parallel_tool_calls(conv) is conv + plain = [{"role": "user", "content": "hi"}] + assert _split_parallel_tool_calls(plain) is plain + + +def test_render_succeeds_on_single_call_template_with_parallel_calls(): + # Regression: two calls in one turn used to break every later render. + result = apply_chat_template_for_generation(_SingleToolCallTokenizer(), _parallel_conv()) + assert result == "RENDERED" + + +def test_string_arguments_and_parallel_calls_are_repaired_together(): + conv = _parallel_conv() + for call in conv[1]["tool_calls"]: + call["function"]["arguments"] = json.dumps(call["function"]["arguments"]) + + class _StrictAndSingleCall(_SingleToolCallTokenizer): + def apply_chat_template(self, messages, **kw): + for msg in messages: + for call in msg.get("tool_calls", []) or []: + if isinstance(call.get("function", {}).get("arguments"), str): + raise TypeError("Can only get item pairs from a mapping.") + return super().apply_chat_template(messages, **kw) + + assert apply_chat_template_for_generation(_StrictAndSingleCall(), conv) == "RENDERED" + + +def test_lenient_template_never_sees_a_split_conversation(): + seen = {} + + class _Lenient: + def apply_chat_template(self, messages, **kw): + seen["n"] = len(messages) + return "RENDERED" + + apply_chat_template_for_generation(_Lenient(), _parallel_conv()) + assert seen["n"] == 4 # unsplit