mirror of
https://github.com/strands-agents/harness-sdk.git
synced 2026-10-02 02:44:48 +08:00
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:
co-authored by
strandly-the-agent
parent
c93666c64a
commit
0ab31d7c12
@@ -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."""
|
||||
|
||||
|
||||
Reference in New Issue
Block a user