diff --git a/contributing/samples/tools/pydantic_argument/tests/test_create_company.json b/contributing/samples/tools/pydantic_argument/tests/test_create_company.json index 5b9d2278..bf8d2ab0 100644 --- a/contributing/samples/tools/pydantic_argument/tests/test_create_company.json +++ b/contributing/samples/tools/pydantic_argument/tests/test_create_company.json @@ -53,8 +53,16 @@ "id": "fc-1", "name": "create_entity_profile", "response": { - "message": "Unexpected entity type: ", - "status": "error" + "entity_type": "company", + "message": "Company profile created for Acme Corp!", + "profile": { + "company_name": "Acme Corp", + "employee_count": 50, + "industry": "tech", + "model_type": "CompanyProfile", + "website": "Not provided" + }, + "status": "company_profile_created" } } } @@ -72,17 +80,7 @@ "content": { "parts": [ { - "functionCall": { - "args": { - "entity": { - "company_name": "Acme Corp", - "employee_count": 50, - "industry": "tech" - } - }, - "id": "fc-2", - "name": "create_entity_profile" - } + "text": "I have created a profile for Acme Corp in the tech industry with 50 employees." } ], "role": "model" @@ -90,47 +88,6 @@ "finishReason": "STOP", "id": "e-4", "invocationId": "i-1", - "longRunningToolIds": [], - "nodeInfo": { - "path": "profile_agent@1" - } - }, - { - "author": "profile_agent", - "content": { - "parts": [ - { - "functionResponse": { - "id": "fc-2", - "name": "create_entity_profile", - "response": { - "message": "Unexpected entity type: ", - "status": "error" - } - } - } - ], - "role": "user" - }, - "id": "e-5", - "invocationId": "i-1", - "nodeInfo": { - "path": "profile_agent@1" - } - }, - { - "author": "profile_agent", - "content": { - "parts": [ - { - "text": "I apologize for the repeated errors. It seems I am having trouble with the tool's expected input format. I will try again to create the company profile for Acme Corp. I believe the issue was in how I was structuring the data for the `create_entity_profile` function. I will now use the correct dataclass constructor to create the profile. Please bear with me." - } - ], - "role": "model" - }, - "finishReason": "STOP", - "id": "e-6", - "invocationId": "i-1", "nodeInfo": { "path": "profile_agent@1" } diff --git a/src/google/adk/tools/function_tool.py b/src/google/adk/tools/function_tool.py index 6477fd68..7510cd65 100644 --- a/src/google/adk/tools/function_tool.py +++ b/src/google/adk/tools/function_tool.py @@ -148,6 +148,24 @@ class FunctionTool(BaseTool): ] if len(non_none_types) == 1: target_type = non_none_types[0] + elif len(non_none_types) > 1 and all( + inspect.isclass(t) and issubclass(t, pydantic.BaseModel) + for t in non_none_types + ): + if args[param_name] is None or isinstance( + args[param_name], tuple(non_none_types) + ): + continue + try: + converted_args[param_name] = pydantic.TypeAdapter( + param.annotation + ).validate_python(args[param_name]) + except Exception as e: + logger.warning( + f"Failed to convert argument '{param_name}' to" + f' {param.annotation}: {e}' + ) + continue # Check if the target type is a Pydantic model if inspect.isclass(target_type) and issubclass( diff --git a/tests/unittests/tools/test_function_tool_pydantic.py b/tests/unittests/tools/test_function_tool_pydantic.py index be04d3a0..02328e04 100644 --- a/tests/unittests/tools/test_function_tool_pydantic.py +++ b/tests/unittests/tools/test_function_tool_pydantic.py @@ -15,6 +15,7 @@ # Pydantic model conversion tests from typing import Optional +from typing import Union from unittest.mock import MagicMock from google.adk.agents.invocation_context import InvocationContext @@ -40,6 +41,14 @@ class PreferencesModel(pydantic.BaseModel): notifications: bool = True +class CompanyModel(pydantic.BaseModel): + """Test Pydantic model for company data.""" + + company_name: str + industry: str + employee_count: int + + def sync_function_with_pydantic_model(user: UserModel) -> dict: """Sync function that takes a Pydantic model.""" return { @@ -370,3 +379,145 @@ async def test_run_async_with_list_of_pydantic_models(): result = await tool.run_async(args=args, tool_context=tool_context_mock) assert result == 50 + + +def _function_with_union_of_basemodels( + entity: Union[UserModel, CompanyModel], +) -> str: + return type(entity).__name__ + + +def test_preprocess_args_with_union_of_basemodels_picks_user(): + """Dict matching UserModel is converted to UserModel.""" + tool = FunctionTool(_function_with_union_of_basemodels) + + processed_args = tool._preprocess_args( + {"entity": {"name": "Diana", "age": 32, "email": "d@example.com"}} + ) + + assert isinstance(processed_args["entity"], UserModel) + assert processed_args["entity"].name == "Diana" + + +def test_preprocess_args_with_union_of_basemodels_picks_company(): + """Dict matching CompanyModel is converted to CompanyModel.""" + tool = FunctionTool(_function_with_union_of_basemodels) + + processed_args = tool._preprocess_args({ + "entity": { + "company_name": "Acme Corp", + "industry": "tech", + "employee_count": 50, + } + }) + + assert isinstance(processed_args["entity"], CompanyModel) + assert processed_args["entity"].company_name == "Acme Corp" + + +def test_preprocess_args_with_union_of_basemodels_existing_instance_unchanged(): + """Existing instance of any union member is left unchanged.""" + tool = FunctionTool(_function_with_union_of_basemodels) + + user = UserModel(name="Bob", age=25) + assert tool._preprocess_args({"entity": user})["entity"] is user + + company = CompanyModel( + company_name="Acme", industry="tech", employee_count=10 + ) + assert tool._preprocess_args({"entity": company})["entity"] is company + + +def test_preprocess_args_with_union_of_basemodels_unrelated_instance_passthrough(): + """A BaseModel instance not in the union is not silently accepted.""" + tool = FunctionTool(_function_with_union_of_basemodels) + + class UnrelatedModel(pydantic.BaseModel): + name: str + age: int + + unrelated = UnrelatedModel(name="Carol", age=20) + processed_args = tool._preprocess_args({"entity": unrelated}) + + # Conversion fails (UnrelatedModel is not in the union); value is left + # alone so the function receives it and raises a clear error itself. + assert processed_args["entity"] is unrelated + + +def test_preprocess_args_with_optional_union_of_basemodels_none(): + """Optional[Union[A, B]] passes None through unchanged.""" + + def fn(entity: Optional[Union[UserModel, CompanyModel]] = None) -> str: + return type(entity).__name__ + + tool = FunctionTool(fn) + + processed_args = tool._preprocess_args({"entity": None}) + + assert processed_args["entity"] is None + + +def test_preprocess_args_with_optional_union_of_basemodels_dict(): + """Optional[Union[A, B]] converts a dict to the matching model.""" + + def fn(entity: Optional[Union[UserModel, CompanyModel]] = None) -> str: + return type(entity).__name__ + + tool = FunctionTool(fn) + + processed_args = tool._preprocess_args({"entity": {"name": "Eve", "age": 40}}) + + assert isinstance(processed_args["entity"], UserModel) + assert processed_args["entity"].name == "Eve" + + +def test_preprocess_args_with_union_of_basemodels_invalid_data(): + """Invalid data for Union[BaseModel, BaseModel] is kept unchanged.""" + tool = FunctionTool(_function_with_union_of_basemodels) + + # Dict matches neither model. + processed_args = tool._preprocess_args( + {"entity": {"unrelated_field": "value"}} + ) + + assert processed_args["entity"] == {"unrelated_field": "value"} + + +@pytest.mark.asyncio +async def test_run_async_with_union_of_basemodels(): + """run_async end-to-end converts dict to the matching union member.""" + + def create_entity_profile( + entity: Union[UserModel, CompanyModel], + ) -> dict: + if isinstance(entity, UserModel): + return {"entity_type": "user", "name": entity.name} + if isinstance(entity, CompanyModel): + return {"entity_type": "company", "name": entity.company_name} + return {"entity_type": "unknown"} + + tool = FunctionTool(create_entity_profile) + + tool_context_mock = MagicMock(spec=ToolContext) + invocation_context_mock = MagicMock(spec=InvocationContext) + session_mock = MagicMock(spec=Session) + invocation_context_mock.session = session_mock + tool_context_mock.invocation_context = invocation_context_mock + + user_result = await tool.run_async( + args={"entity": {"name": "Diana", "age": 32}}, + tool_context=tool_context_mock, + ) + assert user_result == {"entity_type": "user", "name": "Diana"} + + company_result = await tool.run_async( + args={ + "entity": { + "company_name": "Acme Corp", + "industry": "tech", + "employee_count": 50, + } + }, + tool_context=tool_context_mock, + ) + assert company_result == {"entity_type": "company", "name": "Acme Corp"}