diff --git a/litellm/llms/mistral/chat/transformation.py b/litellm/llms/mistral/chat/transformation.py index 33b567e9710..2c18a0aaaff 100644 --- a/litellm/llms/mistral/chat/transformation.py +++ b/litellm/llms/mistral/chat/transformation.py @@ -6,7 +6,7 @@ Why separate file? Make it easy to see how transformation works Docs - https://docs.mistral.ai/api/ """ -from collections.abc import AsyncIterator, Coroutine, Iterator +from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping from typing import TYPE_CHECKING, Any, Final, Literal, cast, get_type_hints, overload import httpx @@ -27,7 +27,11 @@ from litellm.router_utils.reasoning_effort_capability import ( ) from litellm.secret_managers.main import get_secret_str from litellm.types.llms.mistral import MistralThinkingBlock, MistralToolCallMessage -from litellm.types.llms.openai import AllMessageValues +from litellm.types.llms.openai import ( + AllMessageValues, + ChatCompletionToolChoiceFunctionParam, + ChatCompletionToolChoiceObjectParam, +) from litellm.types.utils import ModelResponse, ModelResponseStream from litellm.utils import convert_to_model_response_object, supports_reasoning @@ -66,7 +70,7 @@ class MistralConfig(OpenAIGPTConfig): - `tools` (list or null): A list of available tools for the model. Use this to specify functions for which the model can generate JSON inputs. - - `tool_choice` (string - 'auto'/'any'/'none' or null): Specifies if/how functions are called. If set to none the model won't call a function and will generate a message instead. If set to auto the model can choose to either generate a message or call a function. If set to any the model is forced to call a function. Default - 'auto'. + - `tool_choice` (string - 'auto'/'any'/'none', an object naming one function such as {"type": "function", "function": {"name": "my_function"}}, or null): Specifies if/how functions are called. If set to none the model won't call a function and will generate a message instead. If set to auto the model can choose to either generate a message or call a function. If set to any the model is forced to call a function. Default - 'auto'. - `stop` (string or array of strings): Stop generation if this token is detected. Or if one of these tokens is detected when providing an array @@ -81,7 +85,7 @@ class MistralConfig(OpenAIGPTConfig): top_p: int | None = None max_tokens: int | None = None tools: list | None = None - tool_choice: Literal["auto", "any", "none"] | None = None + tool_choice: Literal["auto", "any", "none"] | ChatCompletionToolChoiceObjectParam | None = None random_seed: int | None = None safe_prompt: bool | None = None response_format: dict | None = None @@ -93,7 +97,7 @@ class MistralConfig(OpenAIGPTConfig): top_p: int | None = None, max_tokens: int | None = None, tools: list | None = None, - tool_choice: Literal["auto", "any", "none"] | None = None, + tool_choice: Literal["auto", "any", "none"] | ChatCompletionToolChoiceObjectParam | None = None, random_seed: int | None = None, safe_prompt: bool | None = None, response_format: dict | None = None, @@ -133,13 +137,19 @@ class MistralConfig(OpenAIGPTConfig): *(("reasoning_effort",) if accepts_reasoning_effort else ()), ] - def _map_tool_choice(self, tool_choice: str) -> str: - if tool_choice == "auto" or tool_choice == "none": - return tool_choice - elif tool_choice == "required": - return "any" - else: # openai 'tool_choice' object param not supported by Mistral API - return "any" + @staticmethod + def _map_tool_choice(tool_choice: str | Mapping[str, object]) -> str | ChatCompletionToolChoiceObjectParam | None: + match tool_choice: + case "auto" | "none": + return tool_choice + case {"type": "function", "function": {"name": str(name)}} if name: + return ChatCompletionToolChoiceObjectParam( + type="function", function=ChatCompletionToolChoiceFunctionParam(name=name) + ) + case str(): + return "any" + case _: + return None @staticmethod def _get_mistral_reasoning_system_prompt() -> str: @@ -191,8 +201,12 @@ class MistralConfig(OpenAIGPTConfig): optional_params["top_p"] = value if param == "stop": optional_params["stop"] = value - if param == "tool_choice" and isinstance(value, str): - optional_params["tool_choice"] = self._map_tool_choice(tool_choice=value) + if ( + param == "tool_choice" + and isinstance(value, (str, dict)) + and (mapped_tool_choice := self._map_tool_choice(tool_choice=value)) is not None + ): + optional_params["tool_choice"] = mapped_tool_choice if param == "seed": optional_params["extra_body"] = {"random_seed": value} if param == "response_format": diff --git a/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py b/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py index 8fb3b3c43df..fb4faf46c94 100644 --- a/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py +++ b/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py @@ -1,10 +1,11 @@ -from typing import List, cast +from collections.abc import Mapping, Sequence +from typing import Final, List, cast from unittest.mock import MagicMock, patch import pytest from litellm.litellm_core_utils.prompt_templates.common_utils import TOOL_RESULT_IMAGE_BOUNDARY -from litellm.types.llms.openai import AllMessageValues +from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam from litellm.llms.mistral.chat.transformation import ( @@ -926,3 +927,64 @@ def test_mistral_transform_request_hoists_tool_message_image(): {"type": "text", "text": TOOL_RESULT_IMAGE_BOUNDARY}, {"type": "image_url", "image_url": {"url": data_uri}}, ] + + +class TestMistralToolChoice: + """Mistral's tool_choice is either an enum or a named function object: https://docs.mistral.ai/api/""" + + TOOLS: Final[Sequence[ChatCompletionToolParam]] = [ + { + "type": "function", + "function": {"name": "get_weather", "parameters": {"type": "object", "properties": {}}}, + } + ] + + def _transform(self, tool_choice: str | Mapping[str, object]) -> Mapping[str, object]: + config: Final = MistralConfig() + optional_params: Final = config.map_openai_params( + non_default_params={"tool_choice": tool_choice, "tools": list(self.TOOLS)}, + optional_params={}, + model="mistral-large-latest", + drop_params=False, + ) + return config.transform_request( + model="mistral-large-latest", + messages=[{"role": "user", "content": "what is the weather in Madrid"}], + optional_params=optional_params, + litellm_params={}, + headers={}, + ) + + def test_named_function_is_forwarded_to_mistral(self) -> None: + request: Final = self._transform({"type": "function", "function": {"name": "get_weather"}}) + + assert request["tool_choice"] == {"type": "function", "function": {"name": "get_weather"}} + + def test_keys_mistral_forbids_are_not_replayed(self) -> None: + request: Final = self._transform( + { + "type": "function", + "function": {"name": "get_weather", "strict": True}, + "cache_control": {"type": "ephemeral"}, + } + ) + + assert request["tool_choice"] == {"type": "function", "function": {"name": "get_weather"}} + + @pytest.mark.parametrize( + "tool_choice, expected", + [("auto", "auto"), ("none", "none"), ("required", "any")], + ) + def test_enum_tool_choice_is_mapped(self, tool_choice: str, expected: str) -> None: + assert self._transform(tool_choice)["tool_choice"] == expected + + @pytest.mark.parametrize( + "tool_choice", + [ + {"type": "function", "function": {}}, + {"type": "function", "function": {"name": ""}}, + {"type": "function", "function": {"name": 123}}, + ], + ) + def test_function_objects_without_a_usable_name_are_left_out(self, tool_choice: Mapping[str, object]) -> None: + assert "tool_choice" not in self._transform(tool_choice)