Skip to content
Open
Show file tree
Hide file tree
Changes from 4 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
24 changes: 21 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,21 @@ 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:
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 +152,19 @@ 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:
field_items = ((n, a, inspect.Parameter.empty)
for n, a in getattr(input_schema, "__annotations__", {}).items())
Comment thread
coderabbitai[bot] marked this conversation as resolved.
Outdated
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
58 changes: 58 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 @@ -19,6 +19,7 @@

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 +33,11 @@ class DummyInput(BaseModel):
value: int


class OptionalInput(BaseModel):
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 +200,58 @@ 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')
@pytest.mark.asyncio
async def test_google_adk_tool_wrapper_streaming_function(mock_function_tool):
Expand Down