Skip to content
Closed
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
2 changes: 1 addition & 1 deletion strands-py/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -47,7 +47,7 @@ dependencies = [

[project.optional-dependencies]
anthropic = ["anthropic>=0.21.0,<1.0.0"]
gemini = ["google-genai>=1.67.0,<3.0.0"]
gemini = ["google-genai>=1.68.0,<3.0.0"]
# The latest litellm v1.92.0 do not support python 3.14, for the short term we will keep
# the version lower than 1.92.0 until litellm has the python 3.14 supported version.
litellm = ["litellm>=1.75.9,<=1.95.0", "openai>=1.68.0,<3.0.0"]
Expand Down
41 changes: 38 additions & 3 deletions strands-py/src/strands/models/gemini.py
Original file line number Diff line number Diff line change
Expand Up @@ -328,6 +328,8 @@ def _format_request_config(
system_prompt: str | None,
params: dict[str, Any] | None,
tool_choice: ToolChoice | None = None,
*,
is_vertex: bool = False,
) -> genai.types.GenerateContentConfig:
"""Format Gemini request config.

Expand All @@ -338,6 +340,8 @@ def _format_request_config(
system_prompt: System prompt to provide context to the model.
params: Additional model parameters (e.g., temperature).
tool_choice: Selection strategy for tool invocation.
is_vertex: Whether the request is sent through Vertex AI, which does not support
`include_server_side_tool_invocations`.

Returns:
Gemini request config.
Expand All @@ -351,6 +355,21 @@ def _format_request_config(
# merged with, or replaced by, the narrower per-request choice.
config_params.setdefault("tool_config", tool_config)

if tool_specs and self.config.get("gemini_tools") and not is_vertex:
Comment thread
Solaris-star marked this conversation as resolved.
# Gemini requires this flag when server-side built-in tools and function declarations are combined.
# An explicit params key, including None, owns the request. Normalize a supplied dict through
# the SDK so both snake_case and the REST-facing camelCase aliases work.
existing_tool_config = config_params.get("tool_config")
if "tool_config" not in config_params:
config_params["tool_config"] = genai.types.ToolConfig(include_server_side_tool_invocations=True)
elif existing_tool_config is not None:
merged_tool_config = genai.types.ToolConfig.model_validate(existing_tool_config)
if merged_tool_config.include_server_side_tool_invocations is None:
merged_tool_config = merged_tool_config.model_copy(
update={"include_server_side_tool_invocations": True},
)
config_params["tool_config"] = merged_tool_config

return genai.types.GenerateContentConfig(
system_instruction=system_prompt,
tools=self._format_request_tools(tool_specs),
Expand All @@ -364,6 +383,8 @@ def _format_request(
system_prompt: str | None,
params: dict[str, Any] | None,
tool_choice: ToolChoice | None = None,
*,
is_vertex: bool = False,
) -> dict[str, Any]:
"""Format a Gemini streaming request.

Expand All @@ -375,12 +396,20 @@ def _format_request(
system_prompt: System prompt to provide context to the model.
params: Additional model parameters (e.g., temperature).
tool_choice: Selection strategy for tool invocation.
is_vertex: Whether the request is sent through Vertex AI, which does not support
`include_server_side_tool_invocations`.

Returns:
A Gemini streaming request.
"""
return {
"config": self._format_request_config(tool_specs, system_prompt, params, tool_choice).to_json_dict(),
"config": self._format_request_config(
tool_specs,
system_prompt,
params,
tool_choice,
is_vertex=is_vertex,
).to_json_dict(),
"contents": [content.to_json_dict() for content in self._format_request_content(messages)],
"model": self.config["model_id"],
}
Expand Down Expand Up @@ -586,11 +615,17 @@ async def stream(
Raises:
ModelThrottledException: If the request is throttled by Gemini.
"""
gemini_client = self._get_client()
request = self._format_request(
messages, tool_specs, system_prompt, self.config.get("params"), tool_choice=tool_choice
messages,
tool_specs,
system_prompt,
self.config.get("params"),
tool_choice=tool_choice,
is_vertex=bool(getattr(gemini_client, "vertexai", False)),
)

client = self._get_client().aio
client = gemini_client.aio

try:
response = await client.models.generate_content_stream(**request)
Expand Down
85 changes: 82 additions & 3 deletions strands-py/tests/strands/models/test_gemini.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@ def gemini_client():
with unittest.mock.patch.object(strands.models.gemini.genai, "Client") as mock_client_cls:
mock_client = mock_client_cls.return_value
mock_client.aio = unittest.mock.AsyncMock()
mock_client.vertexai = False
yield mock_client


Expand Down Expand Up @@ -1033,6 +1034,7 @@ async def test_stream_request_with_gemini_tools(gemini_client, messages, model_i

@pytest.mark.asyncio
async def test_stream_request_with_gemini_tools_and_function_tools(gemini_client, messages, tool_spec, model_id):
"""Regression coverage for #3639: built-in and function tools coexist."""
code_execution_tool = genai.types.Tool(code_execution=genai.types.ToolCodeExecution())
model = GeminiModel(model_id=model_id, gemini_tools=[code_execution_tool])

Expand All @@ -1051,14 +1053,86 @@ async def test_stream_request_with_gemini_tools_and_function_tools(gemini_client
]
},
{"code_execution": {}},
]
],
"tool_config": {"include_server_side_tool_invocations": True},
},
"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_gemini_tools_and_function_tools_for_vertex(
gemini_client, messages, tool_spec, model_id
):
gemini_client.vertexai = True
code_execution_tool = genai.types.Tool(code_execution=genai.types.ToolCodeExecution())
model = GeminiModel(model_id=model_id, client=gemini_client, gemini_tools=[code_execution_tool])

await anext(model.stream(messages, tool_specs=[tool_spec]))

request = gemini_client.aio.models.generate_content_stream.call_args.kwargs
assert "include_server_side_tool_invocations" not in (request["config"].get("tool_config") or {})


@pytest.mark.asyncio
async def test_stream_request_with_empty_function_tools_does_not_enable_server_side_invocations(
gemini_client, messages, model_id
):
code_execution_tool = genai.types.Tool(code_execution=genai.types.ToolCodeExecution())
model = GeminiModel(model_id=model_id, client=gemini_client, gemini_tools=[code_execution_tool])

await anext(model.stream(messages, tool_specs=[]))

request = gemini_client.aio.models.generate_content_stream.call_args.kwargs
assert "tool_config" not in request["config"]


@pytest.mark.asyncio
async def test_stream_request_with_gemini_tools_merges_camel_case_tool_config(
gemini_client, messages, tool_spec, model_id
):
code_execution_tool = genai.types.Tool(code_execution=genai.types.ToolCodeExecution())
model = GeminiModel(
model_id=model_id,
client=gemini_client,
gemini_tools=[code_execution_tool],
params={"tool_config": {"includeServerSideToolInvocations": False}},
)

await anext(model.stream(messages, tool_specs=[tool_spec]))

request = gemini_client.aio.models.generate_content_stream.call_args.kwargs
assert request["config"]["tool_config"] == {"include_server_side_tool_invocations": False}


@pytest.mark.asyncio
async def test_stream_request_with_gemini_tools_merges_existing_tool_config(
gemini_client, messages, tool_spec, model_id
):
code_execution_tool = genai.types.Tool(code_execution=genai.types.ToolCodeExecution())
tool_config = genai.types.ToolConfig(
function_calling_config=genai.types.FunctionCallingConfig(
mode=genai.types.FunctionCallingConfigMode.AUTO,
),
)
model = GeminiModel(
model_id=model_id,
client=gemini_client,
gemini_tools=[code_execution_tool],
params={"tool_config": tool_config},
)

await anext(model.stream(messages, tool_specs=[tool_spec]))

request = gemini_client.aio.models.generate_content_stream.call_args.kwargs
assert request["config"]["tool_config"] == {
"function_calling_config": {"mode": "AUTO"},
"include_server_side_tool_invocations": True,
}


@pytest.mark.parametrize(
("tool_choice", "exp_function_calling_config"),
[
Expand Down Expand Up @@ -1165,7 +1239,11 @@ async def test_stream_tool_config_param_set_to_none_still_takes_precedence(

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})
model = GeminiModel(
model_id=model_id,
gemini_tools=[genai.types.Tool(code_execution=genai.types.ToolCodeExecution())],
params={"tool_config": None},
)

await anext(model.stream(messages, tool_specs=[tool_spec], tool_choice={"any": {}}))

Expand All @@ -1180,7 +1258,8 @@ async def test_stream_tool_config_param_set_to_none_still_takes_precedence(
"parameters_json_schema": tool_spec["inputSchema"]["json"],
}
]
}
},
{"code_execution": {}},
],
},
"contents": [{"parts": [{"text": "test"}], "role": "user"}],
Expand Down
Loading