diff --git a/packages/nvidia_nat_adk/src/nat/plugins/adk/tool_wrapper.py b/packages/nvidia_nat_adk/src/nat/plugins/adk/tool_wrapper.py index 44ddb3742c..be67f59e12 100644 --- a/packages/nvidia_nat_adk/src/nat/plugins/adk/tool_wrapper.py +++ b/packages/nvidia_nat_adk/src/nat/plugins/adk/tool_wrapper.py @@ -68,6 +68,23 @@ def google_adk_tool_wrapper( """ import inspect + def _field_default(field: Any) -> Any: + """Return a Pydantic field's default for an inspect signature.""" + is_required = getattr(field, "is_required", None) + if is_required is not None and is_required(): + return inspect.Parameter.empty + if getattr(field, "required", False): + return inspect.Parameter.empty + default_factory = getattr(field, "default_factory", None) + if default_factory is not None: + if getattr(field, "default_factory_takes_validated_data", False): + return None + try: + return default_factory() + except TypeError: + return None + return getattr(field, "default", inspect.Parameter.empty) + async def callable_ainvoke(*args: Any, **kwargs: Any) -> Any: """Async function to invoke the NAT function. @@ -137,16 +154,24 @@ def decorator(func_to_wrap: Callable[..., Any]) -> Callable[..., Any]: if input_schema is not None: model_fields = getattr(input_schema, "model_fields", None) if model_fields is not None: - field_items = ((n, f.annotation) for n, f in model_fields.items()) + field_items = ((n, f.annotation, _field_default(f)) for n, f in model_fields.items()) else: - field_items = getattr(input_schema, "__annotations__", {}).items() - for param_name, param_annotation in field_items: + legacy_fields = getattr(input_schema, "__fields__", None) + if legacy_fields is not None: + field_items = ((n, getattr(f, "outer_type_", f.annotation), _field_default(f)) + for n, f in legacy_fields.items()) + else: + field_items = ((n, a, getattr(input_schema, n, inspect.Parameter.empty)) + for n, a in getattr(input_schema, "__annotations__", {}).items()) + for param_name, param_annotation, default in field_items: params.append( inspect.Parameter( param_name, inspect.Parameter.POSITIONAL_OR_KEYWORD, annotation=resolve_type(param_annotation), + default=default, )) + params.sort(key=lambda param: param.default is not inspect.Parameter.empty) setattr(func_to_wrap, "__signature__", inspect.Signature(parameters=params)) return func_to_wrap diff --git a/packages/nvidia_nat_adk/tests/test_adk_tool_wrapper.py b/packages/nvidia_nat_adk/tests/test_adk_tool_wrapper.py index 921b1211d9..8d2574053f 100644 --- a/packages/nvidia_nat_adk/tests/test_adk_tool_wrapper.py +++ b/packages/nvidia_nat_adk/tests/test_adk_tool_wrapper.py @@ -14,11 +14,13 @@ # limitations under the License. from typing import Any +from typing import ClassVar from unittest.mock import MagicMock from unittest.mock import patch import pytest from pydantic import BaseModel +from pydantic import Field from nat.plugins.adk.tool_wrapper import google_adk_tool_wrapper from nat.plugins.adk.tool_wrapper import resolve_type @@ -32,6 +34,13 @@ class DummyInput(BaseModel): value: int +class OptionalInput(BaseModel): + """Input model with one required and one optional field.""" + + optional_value: int = 42 + required_value: str + + class DummyOutput(BaseModel): result: int @@ -194,6 +203,105 @@ async def test_google_adk_tool_wrapper_nested_function(mock_function_tool): assert call_args.__doc__ == "Nested ADK function" +@patch('google.adk.tools.function_tool.FunctionTool') +def test_google_adk_tool_wrapper_preserves_field_defaults(mock_function_tool): + """Optional input fields must remain optional in the ADK signature.""" + import inspect + + class OptionalFunction: + description = "Optional ADK function" + has_single_output = True + has_streaming_output = False + input_schema = OptionalInput + + async def acall_invoke(self, *_args, **_kwargs): + return None + + google_adk_tool_wrapper("optional_adk_func", OptionalFunction(), MagicMock()) + + callable_tool = mock_function_tool.call_args[0][0] + signature = inspect.signature(callable_tool) + + assert list(signature.parameters) == ["required_value", "optional_value"] + assert signature.parameters["required_value"].default is inspect.Parameter.empty + assert signature.parameters["optional_value"].default == 42 + + +@patch('google.adk.tools.function_tool.FunctionTool') +def test_google_adk_tool_wrapper_handles_data_aware_default_factory(mock_function_tool): + """A default factory requiring validated data must not break signature creation.""" + import inspect + + class FactoryInput(BaseModel): + required_value: str + generated: list[str] = Field(default_factory=list) + derived: str = Field(default_factory=lambda data: data["required_value"]) + + class FactoryFunction: + description = "Factory ADK function" + has_single_output = True + has_streaming_output = False + input_schema = FactoryInput + + async def acall_invoke(self, *_args, **_kwargs): + return None + + google_adk_tool_wrapper("factory_adk_func", FactoryFunction(), MagicMock()) + + signature = inspect.signature(mock_function_tool.call_args[0][0]) + + assert list(signature.parameters) == ["required_value", "generated", "derived"] + assert signature.parameters["generated"].default == [] + assert signature.parameters["derived"].default is None + + +@patch('google.adk.tools.function_tool.FunctionTool') +def test_google_adk_tool_wrapper_supports_legacy_and_annotation_fields(mock_function_tool): + """Legacy Pydantic and annotation-only schemas must preserve defaults.""" + import inspect + + class LegacyField: + annotation = int + outer_type_ = int + required = False + default = 7 + default_factory = None + + class LegacyInput: + __fields__: ClassVar[dict[str, object]] = { + "required_value": type("RequiredField", (), {"annotation": str, "outer_type_": str, "required": True})(), + "optional_value": LegacyField(), + } + + class AnnotationInput: + __annotations__ = {"required_value": str, "optional_value": int} + optional_value = 7 + + class LegacyFunction: + description = "Legacy ADK function" + has_single_output = True + has_streaming_output = False + input_schema = LegacyInput + + async def acall_invoke(self, *_args, **_kwargs): + return None + + class AnnotationFunction(LegacyFunction): + input_schema = AnnotationInput + + google_adk_tool_wrapper("legacy_adk_func", LegacyFunction(), MagicMock()) + legacy_signature = inspect.signature(mock_function_tool.call_args[0][0]) + assert list(legacy_signature.parameters) == ["required_value", "optional_value"] + assert legacy_signature.parameters["required_value"].default is inspect.Parameter.empty + assert legacy_signature.parameters["optional_value"].default == 7 + + google_adk_tool_wrapper("annotation_adk_func", AnnotationFunction(), MagicMock()) + annotation_signature = inspect.signature(mock_function_tool.call_args[0][0]) + assert list(annotation_signature.parameters) == ["required_value", "optional_value"] + assert annotation_signature.parameters["required_value"].default is inspect.Parameter.empty + assert annotation_signature.parameters["optional_value"].default == 7 + + @patch('google.adk.tools.function_tool.FunctionTool') @pytest.mark.asyncio async def test_google_adk_tool_wrapper_streaming_function(mock_function_tool):