diff --git a/litellm/llms/mistral/chat/transformation.py b/litellm/llms/mistral/chat/transformation.py index a76a8a3e98c..f9a8b5d986b 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 @@ -107,13 +107,14 @@ class MistralConfig(OpenAIGPTConfig): return supported_params - 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" + def _map_tool_choice(self, tool_choice: str | Mapping[str, object]) -> str | Mapping[str, object]: + match tool_choice: + case "auto" | "none": + return tool_choice + case {"type": "function", "function": {"name": str()}}: + return tool_choice + case _: + return "any" @staticmethod def _get_mistral_reasoning_system_prompt() -> str: @@ -165,7 +166,7 @@ class MistralConfig(OpenAIGPTConfig): optional_params["top_p"] = value if param == "stop": optional_params["stop"] = value - if param == "tool_choice" and isinstance(value, str): + if param == "tool_choice" and isinstance(value, (str, Mapping)): optional_params["tool_choice"] = self._map_tool_choice(tool_choice=value) if param == "seed": optional_params["extra_body"] = {"random_seed": value} 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 15694d9f218..a3d00322b03 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,15 @@ -from typing import List, cast +from collections.abc import Mapping +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, + ChatCompletionToolChoiceObjectParam, + ChatCompletionToolParam, +) from litellm.llms.mistral.chat.transformation import ( @@ -844,3 +849,47 @@ 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/""" + + NAMED_TOOL_CHOICE: Final[ChatCompletionToolChoiceObjectParam] = { + "type": "function", + "function": {"name": "get_weather"}, + } + TOOLS: Final[list[ChatCompletionToolParam]] = [ + { + "type": "function", + "function": {"name": "get_weather", "parameters": {"type": "object", "properties": {}}}, + } + ] + + def _transform(self, tool_choice: str | Mapping[str, object]) -> dict[str, object]: + config: Final = MistralConfig() + optional_params: Final = config.map_openai_params( + non_default_params={"tool_choice": tool_choice, "tools": 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: + assert self._transform(self.NAMED_TOOL_CHOICE)["tool_choice"] == self.NAMED_TOOL_CHOICE + + @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 + + def test_object_shapes_mistral_does_not_accept_fall_back_to_any(self) -> None: + assert self._transform({"type": "allowed_tools", "allowed_tools": {"mode": "auto"}})["tool_choice"] == "any"