mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(cohere): strip index from tool_calls and name from tool messages
Cohere's v2 chat API rejects the `index` field on tool_calls inside assistant messages and the `name` field on tool-role messages, even though these are standard OpenAI fields. When users append the raw assistant response back to the messages list for multi-turn tool calling, these fields cause 400 errors from Cohere. Strip both fields in transform_request before the request is sent. Fixes #24031 Made-with: Cursor
This commit is contained in:
parent
244bdffd1b
commit
74a1bdc3bf
2 changed files with 81 additions and 0 deletions
|
|
@ -161,6 +161,32 @@ class CohereV2ChatConfig(OpenAIGPTConfig):
|
|||
optional_params["seed"] = value
|
||||
return optional_params
|
||||
|
||||
@staticmethod
|
||||
def _strip_cohere_unsupported_fields(
|
||||
messages: List[AllMessageValues],
|
||||
) -> List[AllMessageValues]:
|
||||
"""Remove fields that Cohere's v2 API rejects.
|
||||
|
||||
* ``index`` on tool_calls inside assistant messages
|
||||
* ``name`` on tool-role messages
|
||||
"""
|
||||
cleaned: List[AllMessageValues] = []
|
||||
for msg in messages:
|
||||
if not isinstance(msg, dict):
|
||||
cleaned.append(msg)
|
||||
continue
|
||||
role = msg.get("role")
|
||||
if role == "assistant" and "tool_calls" in msg:
|
||||
msg = {**msg}
|
||||
msg["tool_calls"] = [
|
||||
{k: v for k, v in tc.items() if k != "index"}
|
||||
for tc in msg["tool_calls"]
|
||||
]
|
||||
elif role == "tool" and "name" in msg:
|
||||
msg = {k: v for k, v in msg.items() if k != "name"}
|
||||
cleaned.append(msg)
|
||||
return cleaned
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -172,6 +198,7 @@ class CohereV2ChatConfig(OpenAIGPTConfig):
|
|||
"""
|
||||
Cohere v2 chat api is in openai format, so we can use the openai transform request function to transform the request.
|
||||
"""
|
||||
messages = self._strip_cohere_unsupported_fields(messages)
|
||||
data = super().transform_request(
|
||||
model, messages, optional_params, litellm_params, headers
|
||||
)
|
||||
|
|
|
|||
|
|
@ -49,3 +49,57 @@ class TestCohereTransform:
|
|||
|
||||
# The function should properly map max_tokens if max_completion_tokens is not provided
|
||||
assert result == {"temperature": 0.7, "max_tokens": 200}
|
||||
|
||||
|
||||
class TestCohereV2StripFields:
|
||||
"""Regression tests for #24031: Cohere v2 API rejects ``index`` on
|
||||
tool_calls and ``name`` on tool-role messages."""
|
||||
|
||||
def test_strip_index_from_tool_calls(self):
|
||||
from litellm.llms.cohere.chat.v2_transformation import CohereV2ChatConfig
|
||||
|
||||
messages = [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_1",
|
||||
"index": 0,
|
||||
"type": "function",
|
||||
"function": {"name": "get_time", "arguments": "{}"},
|
||||
}
|
||||
],
|
||||
}
|
||||
]
|
||||
cleaned = CohereV2ChatConfig._strip_cohere_unsupported_fields(messages)
|
||||
tc = cleaned[0]["tool_calls"][0]
|
||||
assert "index" not in tc
|
||||
assert tc["id"] == "call_1"
|
||||
assert tc["function"]["name"] == "get_time"
|
||||
|
||||
def test_strip_name_from_tool_message(self):
|
||||
from litellm.llms.cohere.chat.v2_transformation import CohereV2ChatConfig
|
||||
|
||||
messages = [
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "call_1",
|
||||
"name": "get_time",
|
||||
"content": "12:00",
|
||||
}
|
||||
]
|
||||
cleaned = CohereV2ChatConfig._strip_cohere_unsupported_fields(messages)
|
||||
assert "name" not in cleaned[0]
|
||||
assert cleaned[0]["content"] == "12:00"
|
||||
assert cleaned[0]["tool_call_id"] == "call_1"
|
||||
|
||||
def test_non_tool_messages_unchanged(self):
|
||||
from litellm.llms.cohere.chat.v2_transformation import CohereV2ChatConfig
|
||||
|
||||
messages = [
|
||||
{"role": "user", "content": "hello"},
|
||||
{"role": "assistant", "content": "hi"},
|
||||
]
|
||||
cleaned = CohereV2ChatConfig._strip_cohere_unsupported_fields(messages)
|
||||
assert cleaned == messages
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue