This commit is contained in:
Pasqual Troncone 2026-09-24 01:12:05 +02:00 • committed by GitHub
commit 5924f9d1e4
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 92 additions and 16 deletions

View file

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

View file

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