mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
Merge 2bd661a6db into eb7eeb5419
This commit is contained in:
commit
5924f9d1e4
2 changed files with 92 additions and 16 deletions
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue