fix: add the missing required tool parameters for Anthropic models

Fixes #1692

PiperOrigin-RevId: 795125801
This commit is contained in:
Xuan Yang
2025-08-14 11:31:00 -07:00
committed by Copybara-Service
parent 9ba8eec220
commit e63e2a7106
2 changed files with 213 additions and 12 deletions
+18 -12
View File
@@ -216,25 +216,31 @@ def _update_type_string(value_dict: dict[str, Any]):
def function_declaration_to_tool_param(
function_declaration: types.FunctionDeclaration,
) -> anthropic_types.ToolParam:
"""Converts a function declaration to an Anthropic tool param."""
assert function_declaration.name
properties = {}
if (
function_declaration.parameters
and function_declaration.parameters.properties
):
for key, value in function_declaration.parameters.properties.items():
value_dict = value.model_dump(exclude_none=True)
_update_type_string(value_dict)
properties[key] = value_dict
required_params = []
if function_declaration.parameters:
if function_declaration.parameters.properties:
for key, value in function_declaration.parameters.properties.items():
value_dict = value.model_dump(exclude_none=True)
_update_type_string(value_dict)
properties[key] = value_dict
if function_declaration.parameters.required:
required_params = function_declaration.parameters.required
input_schema = {
"type": "object",
"properties": properties,
}
if required_params:
input_schema["required"] = required_params
return anthropic_types.ToolParam(
name=function_declaration.name,
description=function_declaration.description or "",
input_schema={
"type": "object",
"properties": properties,
},
input_schema=input_schema,
)
@@ -20,6 +20,7 @@ from anthropic import types as anthropic_types
from google.adk import version as adk_version
from google.adk.models import anthropic_llm
from google.adk.models.anthropic_llm import Claude
from google.adk.models.anthropic_llm import function_declaration_to_tool_param
from google.adk.models.llm_request import LlmRequest
from google.adk.models.llm_response import LlmResponse
from google.genai import types
@@ -96,6 +97,200 @@ def test_supported_models():
assert models[1] == r"claude-.*-4.*"
function_declaration_test_cases = [
(
"function_with_no_parameters",
types.FunctionDeclaration(
name="get_current_time",
description="Gets the current time.",
),
anthropic_types.ToolParam(
name="get_current_time",
description="Gets the current time.",
input_schema={"type": "object", "properties": {}},
),
),
(
"function_with_one_optional_parameter",
types.FunctionDeclaration(
name="get_weather",
description="Gets weather information for a given location.",
parameters=types.Schema(
type=types.Type.OBJECT,
properties={
"location": types.Schema(
type=types.Type.STRING,
description="City and state, e.g., San Francisco, CA",
)
},
),
),
anthropic_types.ToolParam(
name="get_weather",
description="Gets weather information for a given location.",
input_schema={
"type": "object",
"properties": {
"location": {
"type": "string",
"description": (
"City and state, e.g., San Francisco, CA"
),
}
},
},
),
),
(
"function_with_one_required_parameter",
types.FunctionDeclaration(
name="get_stock_price",
description="Gets the current price for a stock ticker.",
parameters=types.Schema(
type=types.Type.OBJECT,
properties={
"ticker": types.Schema(
type=types.Type.STRING,
description="The stock ticker, e.g., AAPL",
)
},
required=["ticker"],
),
),
anthropic_types.ToolParam(
name="get_stock_price",
description="Gets the current price for a stock ticker.",
input_schema={
"type": "object",
"properties": {
"ticker": {
"type": "string",
"description": "The stock ticker, e.g., AAPL",
}
},
"required": ["ticker"],
},
),
),
(
"function_with_multiple_mixed_parameters",
types.FunctionDeclaration(
name="submit_order",
description="Submits a product order.",
parameters=types.Schema(
type=types.Type.OBJECT,
properties={
"product_id": types.Schema(
type=types.Type.STRING, description="The product ID"
),
"quantity": types.Schema(
type=types.Type.INTEGER,
description="The order quantity",
),
"notes": types.Schema(
type=types.Type.STRING,
description="Optional order notes",
),
},
required=["product_id", "quantity"],
),
),
anthropic_types.ToolParam(
name="submit_order",
description="Submits a product order.",
input_schema={
"type": "object",
"properties": {
"product_id": {
"type": "string",
"description": "The product ID",
},
"quantity": {
"type": "integer",
"description": "The order quantity",
},
"notes": {
"type": "string",
"description": "Optional order notes",
},
},
"required": ["product_id", "quantity"],
},
),
),
(
"function_with_complex_nested_parameter",
types.FunctionDeclaration(
name="create_playlist",
description="Creates a playlist from a list of songs.",
parameters=types.Schema(
type=types.Type.OBJECT,
properties={
"playlist_name": types.Schema(
type=types.Type.STRING,
description="The name for the new playlist",
),
"songs": types.Schema(
type=types.Type.ARRAY,
description="A list of songs to add to the playlist",
items=types.Schema(
type=types.Type.OBJECT,
properties={
"title": types.Schema(type=types.Type.STRING),
"artist": types.Schema(type=types.Type.STRING),
},
required=["title", "artist"],
),
),
},
required=["playlist_name", "songs"],
),
),
anthropic_types.ToolParam(
name="create_playlist",
description="Creates a playlist from a list of songs.",
input_schema={
"type": "object",
"properties": {
"playlist_name": {
"type": "string",
"description": "The name for the new playlist",
},
"songs": {
"type": "array",
"description": "A list of songs to add to the playlist",
"items": {
"type": "object",
"properties": {
"title": {"type": "string"},
"artist": {"type": "string"},
},
"required": ["title", "artist"],
},
},
},
"required": ["playlist_name", "songs"],
},
),
),
]
@pytest.mark.parametrize(
"_, function_declaration, expected_tool_param",
function_declaration_test_cases,
ids=[case[0] for case in function_declaration_test_cases],
)
async def test_function_declaration_to_tool_param(
_, function_declaration, expected_tool_param
):
"""Test function_declaration_to_tool_param."""
assert (
function_declaration_to_tool_param(function_declaration)
== expected_tool_param
)
@pytest.mark.asyncio
async def test_generate_content_async(
claude_llm, llm_request, generate_content_response, generate_llm_response