diff --git a/src/any_llm/providers/gemini/base.py b/src/any_llm/providers/gemini/base.py index 2c57e1c77..0d106f5dc 100644 --- a/src/any_llm/providers/gemini/base.py +++ b/src/any_llm/providers/gemini/base.py @@ -175,8 +175,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) - 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: diff --git a/src/any_llm/providers/gemini/utils.py b/src/any_llm/providers/gemini/utils.py index 6403cc264..59682354b 100644 --- a/src/any_llm/providers/gemini/utils.py +++ b/src/any_llm/providers/gemini/utils.py @@ -8,7 +8,7 @@ from google.genai import types from google.genai.pagers import Pager -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 ( @@ -118,13 +118,45 @@ def _convert_tool_spec(tools: list[dict[str, Any] | Any]) -> list[types.Tool]: 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) 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]: diff --git a/tests/unit/providers/test_gemini_provider.py b/tests/unit/providers/test_gemini_provider.py index 19fc54c9f..898f05b35 100644 --- a/tests/unit/providers/test_gemini_provider.py +++ b/tests/unit/providers/test_gemini_provider.py @@ -284,6 +284,128 @@ 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}}]}, + }, + ], +) +@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."""