mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(oci): report finish_reason="tool_calls" when the response carries tool calls
OCI non-streaming responses derived finish_reason solely from the provider's raw stop reason via _normalize_oci_finish_reason. The OCI Cohere protocol has no TOOL_CALLS finishReason — a tool-calling response comes back as "COMPLETE", which normalizes to "stop" — so handle_cohere_response returned finish_reason="stop" while message.tool_calls was populated, violating OpenAI semantics (finish_reason must be "tool_calls" when the assistant message contains tool calls). Override finish_reason to "tool_calls" whenever tool calls are present, in both the Cohere and GENERIC response handlers.
This commit is contained in:
parent
69b0dd2da0
commit
b3d3444bac
3 changed files with 107 additions and 0 deletions
|
|
@ -247,6 +247,13 @@ def handle_cohere_response(
|
|||
for i, tc in enumerate(cohere_response.chatResponse.toolCalls)
|
||||
]
|
||||
|
||||
# Per OpenAI semantics a response carrying tool calls must report
|
||||
# finish_reason="tool_calls". OCI's Cohere protocol has no TOOL_CALLS
|
||||
# finishReason value (it returns "COMPLETE"), so _normalize_oci_finish_reason
|
||||
# yields "stop"; override it when tool calls are present.
|
||||
if tool_calls:
|
||||
finish_reason = "tool_calls"
|
||||
|
||||
content: Optional[str] = response_text if response_text else None
|
||||
|
||||
# Only include ``tool_calls`` in the message dict when actually present.
|
||||
|
|
|
|||
|
|
@ -377,6 +377,10 @@ def handle_generic_response(
|
|||
model_response.choices[0].finish_reason = _normalize_oci_finish_reason( # type: ignore[union-attr,assignment]
|
||||
response_choice.finishReason
|
||||
)
|
||||
# OpenAI semantics: a response carrying tool calls reports
|
||||
# finish_reason="tool_calls" regardless of the provider's raw stop reason.
|
||||
if message.tool_calls:
|
||||
model_response.choices[0].finish_reason = "tool_calls" # type: ignore[union-attr]
|
||||
|
||||
oci_usage = completion_response.chatResponse.usage
|
||||
reasoning_tokens: Optional[int] = None
|
||||
|
|
|
|||
96
tests/test_litellm/llms/oci/chat/test_oci_finish_reason.py
Normal file
96
tests/test_litellm/llms/oci/chat/test_oci_finish_reason.py
Normal file
|
|
@ -0,0 +1,96 @@
|
|||
import datetime
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm import ModelResponse
|
||||
from litellm.llms.oci.chat.transformation import OCIChatConfig
|
||||
|
||||
|
||||
def _transform(model: str, body: dict) -> ModelResponse:
|
||||
response = httpx.Response(
|
||||
status_code=200, json=body, headers={"Content-Type": "application/json"}
|
||||
)
|
||||
return OCIChatConfig().transform_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
model_response=ModelResponse(),
|
||||
logging_obj={}, # type: ignore
|
||||
request_data={},
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
encoding={},
|
||||
)
|
||||
|
||||
|
||||
def _cohere_body(with_tool_calls: bool) -> dict:
|
||||
chat_response: dict = {
|
||||
"apiFormat": "COHERE",
|
||||
"text": "I will look up the weather in Tokyo.",
|
||||
"finishReason": "COMPLETE",
|
||||
"usage": {"promptTokens": 26, "completionTokens": 22, "totalTokens": 48},
|
||||
}
|
||||
if with_tool_calls:
|
||||
chat_response["toolCalls"] = [
|
||||
{"name": "get_weather", "parameters": {"location": "Tokyo"}}
|
||||
]
|
||||
return {
|
||||
"modelId": "cohere.command-latest",
|
||||
"modelVersion": "1.0",
|
||||
"chatResponse": chat_response,
|
||||
}
|
||||
|
||||
|
||||
def _generic_body_with_tool_calls() -> dict:
|
||||
created = (
|
||||
datetime.datetime.now(datetime.timezone.utc).isoformat().replace("+00:00", "Z")
|
||||
)
|
||||
return {
|
||||
"modelId": "meta.llama-3.3-70b-instruct",
|
||||
"modelVersion": "1.0",
|
||||
"chatResponse": {
|
||||
"apiFormat": "GENERIC",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {
|
||||
"role": "ASSISTANT",
|
||||
"content": [{"type": "TEXT", "text": "Calling a tool."}],
|
||||
"toolCalls": [
|
||||
{
|
||||
"id": "call_0",
|
||||
"type": "FUNCTION",
|
||||
"name": "get_weather",
|
||||
"arguments": '{"location": "Tokyo"}',
|
||||
}
|
||||
],
|
||||
},
|
||||
# OCI GENERIC may report STOP even when tool calls are present.
|
||||
"finishReason": "STOP",
|
||||
}
|
||||
],
|
||||
"timeCreated": created,
|
||||
"usage": {"promptTokens": 5, "completionTokens": 10, "totalTokens": 15},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def test_cohere_tool_calls_report_finish_reason_tool_calls():
|
||||
# OCI Cohere returns finishReason="COMPLETE" for tool calls; finish_reason
|
||||
# must still be "tool_calls" when the message carries them.
|
||||
result = _transform("cohere.command-latest", _cohere_body(with_tool_calls=True))
|
||||
assert result.choices[0].message.tool_calls is not None
|
||||
assert result.choices[0].finish_reason == "tool_calls"
|
||||
|
||||
|
||||
def test_cohere_without_tool_calls_finish_reason_stop():
|
||||
# Sanity: a plain COMPLETE response still maps to "stop".
|
||||
result = _transform("cohere.command-latest", _cohere_body(with_tool_calls=False))
|
||||
assert result.choices[0].message.tool_calls is None
|
||||
assert result.choices[0].finish_reason == "stop"
|
||||
|
||||
|
||||
def test_generic_tool_calls_report_finish_reason_tool_calls():
|
||||
result = _transform("meta.llama-3.3-70b-instruct", _generic_body_with_tool_calls())
|
||||
assert result.choices[0].message.tool_calls is not None
|
||||
assert result.choices[0].finish_reason == "tool_calls"
|
||||
Loading…
Add table
Reference in a new issue