This commit is contained in:
Mavik 2026-04-23 06:40:18 +00:00 • committed by GitHub
commit 6499e6680a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 81 additions and 0 deletions

View file

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

View file

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