feat: support Tool Choice for Gemini in Python (#3551)

Co-authored-by: strandly-the-agent <strandly-the-agent@users.noreply.github.com>
This commit is contained in:
opieter-aws
2026-08-04 11:40:24 -04:00
committed by GitHub
co-authored by strandly-the-agent
parent c93666c64a
commit 0ab31d7c12
3 changed files with 274 additions and 6 deletions
+57 -6
View File
@@ -19,7 +19,7 @@ from ..types.content import ContentBlock, ContentBlockStartToolUse, Messages, Sy
from ..types.event_loop import Usage
from ..types.exceptions import ContextWindowOverflowException, ModelThrottledException, ProviderTokenCountError
from ..types.streaming import StreamEvent
from ..types.tools import ToolChoice, ToolSpec
from ..types.tools import ToolChoice, ToolChoiceToolDict, ToolSpec
from ._defaults import resolve_config_metadata
from ._validation import _has_location_source, validate_config_keys
from .model import BaseModelConfig, Model
@@ -288,11 +288,46 @@ class GeminiModel(Model):
tools.extend(self.config["gemini_tools"])
return tools
@staticmethod
def _format_tool_choice(tool_choice: ToolChoice | None) -> genai.types.ToolConfig | None:
"""Format a tool choice into a Gemini tool config.
- Docs: https://googleapis.github.io/python-genai/genai.html#genai.types.ToolConfig
Args:
tool_choice: Selection strategy for tool invocation.
Returns:
Gemini tool config, or None when no recognized strategy is requested.
"""
if tool_choice is None:
return None
allowed_function_names: list[str] | None = None
if "auto" in tool_choice:
mode = genai.types.FunctionCallingConfigMode.AUTO
elif "any" in tool_choice:
mode = genai.types.FunctionCallingConfigMode.ANY
elif "tool" in tool_choice:
# Gemini has no single-tool mode, so a specific tool is ANY narrowed to that one function.
mode = genai.types.FunctionCallingConfigMode.ANY
allowed_function_names = [cast(ToolChoiceToolDict, tool_choice)["tool"]["name"]]
else:
return None
return genai.types.ToolConfig(
function_calling_config=genai.types.FunctionCallingConfig(
mode=mode,
allowed_function_names=allowed_function_names,
),
)
def _format_request_config(
self,
tool_specs: list[ToolSpec] | None,
system_prompt: str | None,
params: dict[str, Any] | None,
tool_choice: ToolChoice | None = None,
) -> genai.types.GenerateContentConfig:
"""Format Gemini request config.
@@ -302,14 +337,24 @@ class GeminiModel(Model):
tool_specs: List of tool specifications to make available to the model.
system_prompt: System prompt to provide context to the model.
params: Additional model parameters (e.g., temperature).
tool_choice: Selection strategy for tool invocation.
Returns:
Gemini request config.
"""
config_params = dict(params or {})
tool_config = self._format_tool_choice(tool_choice) if tool_specs else None
if tool_config is not None:
# A tool config set in params wins, matching the other providers and the TypeScript SDK. A tool
# config expresses more than a ToolChoice can, so the explicit one is left whole rather than
# merged with, or replaced by, the narrower per-request choice.
config_params.setdefault("tool_config", tool_config)
return genai.types.GenerateContentConfig(
system_instruction=system_prompt,
tools=self._format_request_tools(tool_specs),
**(params or {}),
**config_params,
)
def _format_request(
@@ -318,6 +363,7 @@ class GeminiModel(Model):
tool_specs: list[ToolSpec] | None,
system_prompt: str | None,
params: dict[str, Any] | None,
tool_choice: ToolChoice | None = None,
) -> dict[str, Any]:
"""Format a Gemini streaming request.
@@ -328,12 +374,13 @@ class GeminiModel(Model):
tool_specs: List of tool specifications to make available to the model.
system_prompt: System prompt to provide context to the model.
params: Additional model parameters (e.g., temperature).
tool_choice: Selection strategy for tool invocation.
Returns:
A Gemini streaming request.
"""
return {
"config": self._format_request_config(tool_specs, system_prompt, params).to_json_dict(),
"config": self._format_request_config(tool_specs, system_prompt, params, tool_choice).to_json_dict(),
"contents": [content.to_json_dict() for content in self._format_request_content(messages)],
"model": self.config["model_id"],
}
@@ -518,6 +565,7 @@ class GeminiModel(Model):
messages: Messages,
tool_specs: list[ToolSpec] | None = None,
system_prompt: str | None = None,
*,
tool_choice: ToolChoice | None = None,
**kwargs: Any,
) -> AsyncGenerator[StreamEvent, None]:
@@ -527,8 +575,9 @@ class GeminiModel(Model):
messages: List of message objects to be processed by the model.
tool_specs: List of tool specifications to make available to the model.
system_prompt: System prompt to provide context to the model.
tool_choice: Selection strategy for tool invocation.
Note: Currently unused.
tool_choice: Selection strategy for tool invocation. Applied only when tool specs are provided,
since there is nothing to choose from without them, and only when params sets no tool config of
its own - an explicit tool config takes precedence.
**kwargs: Additional keyword arguments for future extensibility.
Yields:
@@ -537,7 +586,9 @@ class GeminiModel(Model):
Raises:
ModelThrottledException: If the request is throttled by Gemini.
"""
request = self._format_request(messages, tool_specs, system_prompt, self.config.get("params"))
request = self._format_request(
messages, tool_specs, system_prompt, self.config.get("params"), tool_choice=tool_choice
)
client = self._get_client().aio
@@ -1059,6 +1059,178 @@ async def test_stream_request_with_gemini_tools_and_function_tools(gemini_client
gemini_client.aio.models.generate_content_stream.assert_called_with(**exp_request)
@pytest.mark.parametrize(
("tool_choice", "exp_function_calling_config"),
[
({"auto": {}}, {"mode": "AUTO"}),
({"any": {}}, {"mode": "ANY"}),
({"tool": {"name": "name"}}, {"allowed_function_names": ["name"], "mode": "ANY"}),
],
)
@pytest.mark.asyncio
async def test_stream_request_with_tool_choice(
gemini_client, model, messages, tool_spec, model_id, tool_choice, exp_function_calling_config
):
await anext(model.stream(messages, tool_specs=[tool_spec], tool_choice=tool_choice))
exp_request = {
"config": {
"tools": [
{
"function_declarations": [
{
"description": tool_spec["description"],
"name": tool_spec["name"],
"parameters_json_schema": tool_spec["inputSchema"]["json"],
}
]
}
],
"tool_config": {"function_calling_config": exp_function_calling_config},
},
"contents": [{"parts": [{"text": "test"}], "role": "user"}],
"model": model_id,
}
gemini_client.aio.models.generate_content_stream.assert_called_with(**exp_request)
@pytest.mark.asyncio
async def test_stream_request_with_tool_choice_and_no_tool_specs(gemini_client, model, messages, model_id):
await anext(model.stream(messages, tool_choice={"any": {}}))
exp_request = {
"config": {},
"contents": [{"parts": [{"text": "test"}], "role": "user"}],
"model": model_id,
}
gemini_client.aio.models.generate_content_stream.assert_called_with(**exp_request)
@pytest.fixture
def tool_config_param_model(gemini_client, model_id):
_ = gemini_client
tool_config = genai.types.ToolConfig(
function_calling_config=genai.types.FunctionCallingConfig(
mode=genai.types.FunctionCallingConfigMode.NONE,
allowed_function_names=["safe_tool"],
),
# A full tool config, so the assertions below pin that a tool choice leaves all of it in place.
retrieval_config=genai.types.RetrievalConfig(language_code="en-GB"),
)
return GeminiModel(model_id=model_id, params={"tool_config": tool_config})
@pytest.mark.parametrize(
"tool_choice",
[None, {"auto": {}}, {"any": {}}, {"tool": {"name": "name"}}],
ids=["no-choice", "auto", "any", "tool"],
)
@pytest.mark.asyncio
async def test_stream_request_tool_config_param_takes_precedence(
gemini_client, tool_config_param_model, messages, tool_spec, model_id, tool_choice
):
"""An explicit tool config wins over any per-request choice, matching the other providers."""
await anext(tool_config_param_model.stream(messages, tool_specs=[tool_spec], tool_choice=tool_choice))
exp_request = {
"config": {
"tools": [
{
"function_declarations": [
{
"description": tool_spec["description"],
"name": tool_spec["name"],
"parameters_json_schema": tool_spec["inputSchema"]["json"],
}
]
}
],
"tool_config": {
"function_calling_config": {"mode": "NONE", "allowed_function_names": ["safe_tool"]},
"retrieval_config": {"language_code": "en-GB"},
},
},
"contents": [{"parts": [{"text": "test"}], "role": "user"}],
"model": model_id,
}
gemini_client.aio.models.generate_content_stream.assert_called_with(**exp_request)
@pytest.mark.asyncio
async def test_stream_tool_config_param_set_to_none_still_takes_precedence(
gemini_client, messages, tool_spec, model_id
):
"""Params owns the key, so an explicit None keeps a tool choice out of the request.
Matches the sibling providers, which spread params last and therefore let an explicit None win.
"""
model = GeminiModel(model_id=model_id, params={"tool_config": None})
await anext(model.stream(messages, tool_specs=[tool_spec], tool_choice={"any": {}}))
exp_request = {
"config": {
"tools": [
{
"function_declarations": [
{
"description": tool_spec["description"],
"name": tool_spec["name"],
"parameters_json_schema": tool_spec["inputSchema"]["json"],
}
]
}
],
},
"contents": [{"parts": [{"text": "test"}], "role": "user"}],
"model": model_id,
}
gemini_client.aio.models.generate_content_stream.assert_called_with(**exp_request)
@pytest.mark.asyncio
async def test_stream_tool_choice_does_not_persist_into_the_next_request(gemini_client, messages, tool_spec, model_id):
"""A tool choice configures its own request only, so it never lands in the model's own params."""
model = GeminiModel(model_id=model_id, params={"temperature": 0.5})
await anext(model.stream(messages, tool_specs=[tool_spec], tool_choice={"any": {}}))
await anext(model.stream(messages, tool_specs=[tool_spec]))
exp_request = {
"config": {
"temperature": 0.5,
"tools": [
{
"function_declarations": [
{
"description": tool_spec["description"],
"name": tool_spec["name"],
"parameters_json_schema": tool_spec["inputSchema"]["json"],
}
]
}
],
},
"contents": [{"parts": [{"text": "test"}], "role": "user"}],
"model": model_id,
}
gemini_client.aio.models.generate_content_stream.assert_called_with(**exp_request)
def test_format_tool_choice_unrecognized_strategy(model):
tru_tool_config = model._format_tool_choice({"unrecognized": {}})
exp_tool_config = None
assert tru_tool_config == exp_tool_config
@pytest.mark.asyncio
async def test_stream_tool_choice_no_warning(model, messages, tool_spec, captured_warnings):
await anext(model.stream(messages, tool_specs=[tool_spec], tool_choice={"auto": {}}))
assert len(captured_warnings) == 0
@pytest.mark.asyncio
async def test_stream_handles_non_json_error(gemini_client, model, messages, alist):
error_message = "Invalid API key"
@@ -204,6 +204,51 @@ def test_agent_with_gemini_code_execution_tool(gemini_tool_model):
assert "5117" in str(result_turn2)
@pytest.fixture
def tool_specs():
return [
{
"name": "tool_time",
"description": "Get the current time for a city",
"inputSchema": {"json": {"type": "object", "properties": {"city": {"type": "string"}}}},
},
{
"name": "tool_weather",
"description": "Get the current weather for a city",
"inputSchema": {"json": {"type": "object", "properties": {"city": {"type": "string"}}}},
},
]
@pytest.mark.asyncio
async def test_model_stream_tool_choice_any_forces_a_tool_use(model, tool_specs, alist):
messages = [{"role": "user", "content": [{"text": "Hello there!"}]}]
events = await alist(model.stream(messages, tool_specs=tool_specs, tool_choice={"any": {}}))
tru_stop_reason = next(event["messageStop"]["stopReason"] for event in events if "messageStop" in event)
exp_stop_reason = "tool_use"
assert tru_stop_reason == exp_stop_reason
@pytest.mark.asyncio
async def test_model_stream_tool_choice_tool_forces_the_named_tool(model, tool_specs, alist):
messages = [{"role": "user", "content": [{"text": "What is the weather in New York?"}]}]
events = await alist(
model.stream(messages, tool_specs=tool_specs, tool_choice={"tool": {"name": "tool_time"}}),
)
# ANY mode may emit more than one call, so only the narrowing to tool_time is guaranteed.
tru_tool_names = {
event["contentBlockStart"]["start"]["toolUse"]["name"]
for event in events
if "contentBlockStart" in event and "toolUse" in event["contentBlockStart"]["start"]
}
exp_tool_names = {"tool_time"}
assert tru_tool_names == exp_tool_names
def test_agent_with_reasoning_content(model, assistant_agent):
"""Test that reasoning content is captured in message history."""