Skip to content
Open
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
31 changes: 28 additions & 3 deletions packages/nvidia_nat_adk/src/nat/plugins/adk/tool_wrapper.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down Expand Up @@ -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
Expand Down
108 changes: 108 additions & 0 deletions packages/nvidia_nat_adk/tests/test_adk_tool_wrapper.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Comment thread
coderabbitai[bot] marked this conversation as resolved.


class DummyOutput(BaseModel):
result: int

Expand Down Expand Up @@ -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):
Expand Down