Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions src/any_llm/providers/gemini/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -165,8 +165,8 @@ def _convert_completion_params(params: CompletionParams, **kwargs: Any) -> dict[
kwargs["temperature"] = params.temperature
if params.tools is not None:
kwargs["tools"] = _convert_tool_spec(params.tools, provider_name)
if isinstance(params.tool_choice, str):
kwargs["tool_config"] = _convert_tool_choice(params.tool_choice)
if params.tool_choice is not None:
kwargs["tool_config"] = _convert_tool_choice(params.tool_choice, provider_name)
if params.top_p is not None:
kwargs["top_p"] = params.top_p
if params.stop is not None:
Expand Down
41 changes: 38 additions & 3 deletions src/any_llm/providers/gemini/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@
from google.genai.pagers import Pager
from pydantic import ValidationError

from any_llm.exceptions import InvalidRequestError
from any_llm.exceptions import InvalidRequestError, UnsupportedParameterError
from any_llm.logging import logger
from any_llm.types.batch import Batch, BatchRequestCounts, BatchResult, BatchResultError, BatchResultItem
from any_llm.types.completion import (
Expand Down Expand Up @@ -128,13 +128,48 @@ def _convert_tool_spec(tools: list[dict[str, Any] | Any], provider_name: str) ->
return converted_tools


def _convert_tool_choice(tool_choice: str) -> types.ToolConfig:
def _convert_tool_choice(tool_choice: str | dict[str, Any], provider_name: str) -> types.ToolConfig:
error_message = "tool_choice"
additional_message = f"Unsupported tool_choice: {tool_choice}"

if isinstance(tool_choice, dict):
if tool_choice.get("type") == "allowed_tools":
allowed = tool_choice.get("allowed_tools")
# Gemini only honors allowed_function_names in ANY mode, so mode="auto" has no equivalent.
if not isinstance(allowed, dict) or allowed.get("mode") != "required":
raise UnsupportedParameterError(error_message, provider_name, additional_message)
allowed_tools = allowed.get("tools")
if not isinstance(allowed_tools, list):
raise UnsupportedParameterError(error_message, provider_name, additional_message)
# Every entry is kept so that an unusable one fails the name check below rather than
# being dropped, which would silently narrow the set of tools the caller asked for.
functions = [
tool.get("function") if isinstance(tool, dict) and tool.get("type") == "function" else None
for tool in allowed_tools
]
else:
functions = [tool_choice.get("function")] if tool_choice.get("type") == "function" else []
raw_names = [function.get("name") if isinstance(function, dict) else None for function in functions]
names = [name for name in raw_names if isinstance(name, str) and name]
if not names or len(names) != len(raw_names):
raise UnsupportedParameterError(error_message, provider_name, additional_message)
return types.ToolConfig(
function_calling_config=types.FunctionCallingConfig(
mode=types.FunctionCallingConfigMode.ANY,
allowed_function_names=names,
)
)

tool_choice_to_mode = {
"required": types.FunctionCallingConfigMode.ANY,
"auto": types.FunctionCallingConfigMode.AUTO,
"none": types.FunctionCallingConfigMode.NONE,
}
mode = tool_choice_to_mode.get(tool_choice)
if mode is None:
raise UnsupportedParameterError(error_message, provider_name, additional_message)

return types.ToolConfig(function_calling_config=types.FunctionCallingConfig(mode=tool_choice_to_mode[tool_choice]))
return types.ToolConfig(function_calling_config=types.FunctionCallingConfig(mode=mode))


def _parse_data_uri(data_uri: str, field_name: str, provider_name: str) -> tuple[str, bytes]:
Expand Down
126 changes: 126 additions & 0 deletions tests/unit/providers/test_gemini_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -301,6 +301,132 @@ async def test_completion_with_tool_choice_auto(tool_choice: str, expected_mode:
assert generation_config.tool_config.function_calling_config.mode.value == expected_mode


@pytest.mark.asyncio
async def test_completion_with_tool_choice_none_disables_function_calling() -> None:
"""tool_choice='none' must map to Gemini's NONE mode instead of raising."""
messages = [{"role": "user", "content": "Hello"}]

with mock_gemini_provider() as mock_genai:
provider = GeminiProvider(api_key="test-api-key")
await provider._acompletion(
CompletionParams(model_id="gemini-pro", messages=messages, tool_choice="none"),
)

_, call_kwargs = mock_genai.return_value.aio.models.generate_content.call_args
generation_config = call_kwargs["config"]

assert generation_config.tool_config.function_calling_config.mode.value == "NONE"


@pytest.mark.asyncio
async def test_completion_with_named_function_tool_choice() -> None:
"""The OpenAI named-function tool_choice must force that function via allowed_function_names."""
messages = [{"role": "user", "content": "Hello"}]

with mock_gemini_provider() as mock_genai:
provider = GeminiProvider(api_key="test-api-key")
await provider._acompletion(
CompletionParams(
model_id="gemini-pro",
messages=messages,
tool_choice={"type": "function", "function": {"name": "get_weather"}},
),
)

_, call_kwargs = mock_genai.return_value.aio.models.generate_content.call_args
function_calling_config = call_kwargs["config"].tool_config.function_calling_config

assert function_calling_config.mode.value == "ANY"
assert function_calling_config.allowed_function_names == ["get_weather"]


@pytest.mark.asyncio
async def test_completion_with_allowed_tools_tool_choice() -> None:
"""A required allowed_tools choice must forward every listed name in ANY mode."""
messages = [{"role": "user", "content": "Hello"}]

with mock_gemini_provider() as mock_genai:
provider = GeminiProvider(api_key="test-api-key")
await provider._acompletion(
CompletionParams(
model_id="gemini-pro",
messages=messages,
tool_choice={
"type": "allowed_tools",
"allowed_tools": {
"mode": "required",
"tools": [
{"type": "function", "function": {"name": "get_weather"}},
{"type": "function", "function": {"name": "get_time"}},
],
},
},
),
)

_, call_kwargs = mock_genai.return_value.aio.models.generate_content.call_args
function_calling_config = call_kwargs["config"].tool_config.function_calling_config

assert function_calling_config.mode.value == "ANY"
assert function_calling_config.allowed_function_names == ["get_weather", "get_time"]


@pytest.mark.parametrize(
"tool_choice",
[
"sometimes",
{"type": "custom", "custom": {"name": "get_weather"}},
{"type": "function"},
{"type": "function", "function": {"name": 1}},
# Gemini honors allowed_function_names only in ANY mode, so "auto" cannot be expressed.
{
"type": "allowed_tools",
"allowed_tools": {"mode": "auto", "tools": [{"type": "function", "function": {"name": "get_weather"}}]},
},
{"type": "allowed_tools", "allowed_tools": {"mode": "required", "tools": []}},
{"type": "allowed_tools", "allowed_tools": {"mode": "required"}},
{"type": "allowed_tools", "allowed_tools": {"mode": "required", "tools": None}},
{"type": "allowed_tools", "allowed_tools": {"mode": "required", "tools": 1}},
{
"type": "allowed_tools",
"allowed_tools": {
"mode": "required",
"tools": [{"type": "function", "function": {"name": "get_weather"}}, "get_time"],
},
},
{
"type": "allowed_tools",
"allowed_tools": {
"mode": "required",
"tools": [
{"type": "function", "function": {"name": "get_weather"}},
{"type": "function", "function": {"name": ""}},
],
},
},
{
"type": "allowed_tools",
"allowed_tools": {"mode": "required", "tools": [{"type": "function", "function": {"name": 1}}]},
},
{
"type": "allowed_tools",
"allowed_tools": {"mode": "required", "tools": [{"type": "custom", "function": {"name": "get_weather"}}]},
},
],
)
@pytest.mark.asyncio
async def test_completion_with_unsupported_tool_choice_raises(tool_choice: str | dict[str, Any]) -> None:
"""Unrecognized tool_choice values report the offending value rather than leaking a KeyError."""
messages = [{"role": "user", "content": "Hello"}]

with mock_gemini_provider():
provider = GeminiProvider(api_key="test-api-key")
with pytest.raises(UnsupportedParameterError, match="tool_choice"):
await provider._acompletion(
CompletionParams(model_id="gemini-pro", messages=messages, tool_choice=tool_choice),
)


@pytest.mark.asyncio
async def test_completion_without_tool_choice() -> None:
"""Test that completion works correctly without tool_choice."""
Expand Down