fix(databricks): split parallel tool calls so each tool message follows tool_calls (#31633)

* fix(databricks): split parallel tool calls so each tool message follows tool_calls

Databricks OpenAI-compatible serving (e.g. GPT models) 400s with "messages with
role 'tool' must be a response to a preceeding message with 'tool_calls'" when an
assistant turn makes parallel tool calls. LiteLLM faithfully sends one assistant
message holding all tool_calls followed by one 'tool' message per result, so every
result after the first is preceded by another 'tool' message rather than the
assistant tool_calls message, which Databricks rejects.

Re-emit each result immediately after an assistant message that carries only its
matching tool_call, turning assistant(tool_calls=[A, B]), tool(A), tool(B) into
assistant(tool_calls=[A]), tool(A), assistant(tool_calls=[B]), tool(B). The
rewrite is a no-op when the turn is already valid (single call), the group is
incomplete, or ids don't line up, so no tool call is ever dropped. Scoped to
non-Claude models, matching the existing OpenAI-shaped transformation path.

* style(databricks): use builtin list generics in parallel tool-call split

Switch the List[...] annotations introduced by _split_parallel_tool_calls
to lowercase list[...] so the UP006 strict-rule budget stays within its
ceiling.
This commit is contained in:
Mateo Wang 2026-06-29 13:46:53 -07:00 • committed by GitHub
parent 5e5b09709c
commit 20dabb781a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 220 additions and 6 deletions

View file

@ -40,10 +40,13 @@ from litellm.types.llms.databricks import (
)
from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionAssistantMessage,
ChatCompletionAssistantToolCall,
ChatCompletionRedactedThinkingBlock,
ChatCompletionThinkingBlock,
ChatCompletionToolChoiceFunctionParam,
ChatCompletionToolChoiceObjectParam,
ChatCompletionToolMessage,
ChatCompletionToolParam,
)
from litellm.types.utils import (
@ -92,6 +95,58 @@ def _sanitize_empty_content(message_dict: dict[str, Any]) -> None:
message_dict["content"] = filtered
def _split_parallel_tool_calls(messages: list[AllMessageValues]) -> list[AllMessageValues]:
"""
Databricks (OpenAI-compatible serving) rejects a ``tool`` message unless the
message immediately before it carries ``tool_calls``. A single assistant turn
with parallel tool calls is followed by one ``tool`` message per call, so every
result after the first is preceded by another ``tool`` message and 400s. Re-emit
each result right after an assistant message holding only its matching call:
``assistant(tool_calls=[A, B]), tool(A), tool(B)`` becomes
``assistant(tool_calls=[A]), tool(A), assistant(tool_calls=[B]), tool(B)``.
Left untouched (no-op) when the turn is already valid or the history is
malformed, so no tool call is ever dropped.
"""
def _expand(
assistant: ChatCompletionAssistantMessage,
calls_by_id: dict[Optional[str], ChatCompletionAssistantToolCall],
tool_messages: list[ChatCompletionToolMessage],
) -> Iterator[AllMessageValues]:
for position, tool_message in enumerate(tool_messages):
matched_call = calls_by_id[tool_message["tool_call_id"]]
if position == 0:
yield cast(AllMessageValues, {**assistant, "tool_calls": [matched_call]})
else:
yield ChatCompletionAssistantMessage(role="assistant", tool_calls=[matched_call])
yield tool_message
def _generate() -> Iterator[AllMessageValues]:
index = 0
while index < len(messages):
message = messages[index]
tool_calls = message.get("tool_calls") if message["role"] == "assistant" else None
if not tool_calls or len(tool_calls) < 2:
yield message
index += 1
continue
end = index + 1
while end < len(messages) and messages[end]["role"] == "tool":
end += 1
tool_messages = cast(list[ChatCompletionToolMessage], messages[index + 1 : end])
calls_by_id = {call["id"]: call for call in tool_calls}
result_ids = {tool_message["tool_call_id"] for tool_message in tool_messages}
if len(tool_messages) == len(tool_calls) and set(calls_by_id) == result_ids:
yield from _expand(cast(ChatCompletionAssistantMessage, message), calls_by_id, tool_messages)
index = end
else:
yield message
index += 1
return list(_generate())
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
@ -385,6 +440,9 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig):
_sanitize_empty_content(cast(dict[str, Any], _message))
new_messages.append(_message)
if "claude" not in model:
new_messages = _split_parallel_tool_calls(cast(list[AllMessageValues], new_messages))
if is_async:
return super()._transform_messages(messages=new_messages, model=model, is_async=cast(Literal[True], True))
else:

View file

@ -5,9 +5,7 @@ import sys
import pytest
from fastapi.testclient import TestClient
sys.path.insert(
0, os.path.abspath("../../../../..")
) # Adds the parent directory to the system path
sys.path.insert(0, os.path.abspath("../../../../..")) # Adds the parent directory to the system path
from unittest.mock import MagicMock, patch
from litellm.llms.databricks.chat.transformation import (
@ -255,8 +253,166 @@ def test_transform_messages_sanitizes_empty_content():
{"role": "user", "content": [{"type": "text", "text": ""}]},
{"role": "user", "content": "Hi"},
]
result = config._transform_messages(
messages=messages, model="databricks-claude", is_async=False
)
result = config._transform_messages(messages=messages, model="databricks-claude", is_async=False)
assert "content" not in result[0]
assert result[1]["content"] == "Hi"
def _parallel_tool_calls():
return [
{
"id": "call_A",
"type": "function",
"function": {"name": "get_weather", "arguments": '{"city": "SF"}'},
},
{
"id": "call_B",
"type": "function",
"function": {"name": "get_weather", "arguments": '{"city": "NYC"}'},
},
]
def _assert_every_tool_message_follows_tool_calls(messages):
for index, message in enumerate(messages):
if message.get("role") == "tool":
previous = messages[index - 1] if index > 0 else {}
assert previous.get("role") == "assistant" and previous.get("tool_calls"), (
f"tool message at index {index} is not preceded by an assistant message with tool_calls: {messages}"
)
def _declared_tool_call_ids(messages):
return sorted(
call["id"]
for message in messages
if message.get("role") == "assistant" and message.get("tool_calls")
for call in message["tool_calls"]
)
def test_transform_request_splits_parallel_tool_calls_for_gpt():
"""Regression for LIT-3984: Databricks 400s with 'messages with role tool must
be a response to a preceeding message with tool_calls' because parallel tool
calls send consecutive tool messages. Each result must be re-paired with an
assistant tool_calls message holding only its matching call."""
config = DatabricksConfig()
messages = [
{"role": "user", "content": "weather in SF and NYC?"},
{"role": "assistant", "content": "checking", "tool_calls": _parallel_tool_calls()},
{"role": "tool", "tool_call_id": "call_A", "content": "sunny"},
{"role": "tool", "tool_call_id": "call_B", "content": "rainy"},
]
result = config.transform_request(
model="gpt-5.4-mini",
messages=messages,
optional_params={},
litellm_params={},
headers={},
)["messages"]
_assert_every_tool_message_follows_tool_calls(result)
assert _declared_tool_call_ids(result) == ["call_A", "call_B"]
assistant_tool_call_messages = [m for m in result if m.get("role") == "assistant" and m.get("tool_calls")]
assert all(len(m["tool_calls"]) == 1 for m in assistant_tool_call_messages), (
"each split assistant message must declare exactly one tool call"
)
tool_messages = [m for m in result if m.get("role") == "tool"]
assert [m["tool_call_id"] for m in tool_messages] == ["call_A", "call_B"]
for tool_message, assistant_message in zip(tool_messages, assistant_tool_call_messages):
assert assistant_message["tool_calls"][0]["id"] == tool_message["tool_call_id"]
def test_transform_request_pairs_out_of_order_parallel_results():
config = DatabricksConfig()
messages = [
{"role": "user", "content": "weather?"},
{"role": "assistant", "content": "checking", "tool_calls": _parallel_tool_calls()},
{"role": "tool", "tool_call_id": "call_B", "content": "rainy"},
{"role": "tool", "tool_call_id": "call_A", "content": "sunny"},
]
result = config.transform_request(
model="gpt-5.4-mini",
messages=messages,
optional_params={},
litellm_params={},
headers={},
)["messages"]
_assert_every_tool_message_follows_tool_calls(result)
for index, message in enumerate(result):
if message.get("role") == "tool":
assert result[index - 1]["tool_calls"][0]["id"] == message["tool_call_id"]
def test_transform_request_leaves_single_tool_call_untouched():
config = DatabricksConfig()
messages = [
{"role": "user", "content": "weather?"},
{
"role": "assistant",
"content": "",
"tool_calls": [
{
"id": "call_A",
"type": "function",
"function": {"name": "get_weather", "arguments": "{}"},
}
],
},
{"role": "tool", "tool_call_id": "call_A", "content": "sunny"},
]
result = config.transform_request(
model="gpt-5.4-mini",
messages=messages,
optional_params={},
litellm_params={},
headers={},
)["messages"]
assert len(result) == 3
_assert_every_tool_message_follows_tool_calls(result)
assert _declared_tool_call_ids(result) == ["call_A"]
def test_transform_request_does_not_drop_tool_calls_on_incomplete_results():
config = DatabricksConfig()
messages = [
{"role": "user", "content": "weather?"},
{"role": "assistant", "content": "checking", "tool_calls": _parallel_tool_calls()},
{"role": "tool", "tool_call_id": "call_A", "content": "sunny"},
{"role": "user", "content": "thanks"},
]
result = config.transform_request(
model="gpt-5.4-mini",
messages=messages,
optional_params={},
litellm_params={},
headers={},
)["messages"]
assert _declared_tool_call_ids(result) == ["call_A", "call_B"]
def test_transform_request_keeps_parallel_tool_calls_for_claude():
config = DatabricksConfig()
messages = [
{"role": "user", "content": "weather?"},
{"role": "assistant", "content": "checking", "tool_calls": _parallel_tool_calls()},
{"role": "tool", "tool_call_id": "call_A", "content": "sunny"},
{"role": "tool", "tool_call_id": "call_B", "content": "rainy"},
]
result = config.transform_request(
model="databricks-claude-3-7-sonnet",
messages=messages,
optional_params={},
litellm_params={},
headers={},
)["messages"]
assert len([m for m in result if m.get("role") == "assistant"]) == 1