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
9 changes: 7 additions & 2 deletions src/google/adk/tools/function_tool.py
Original file line number Diff line number Diff line change
Expand Up @@ -456,8 +456,13 @@ async def _invoke_callable(
return await target(**args_to_call)
runner = _SYNC_CALLABLE_RUNNER.get()
if runner is not None:
return await runner(target, args_to_call)
return target(**args_to_call)
result = await runner(target, args_to_call)
else:
result = target(**args_to_call)
# A sync decorator around an async function returns a coroutine.
if inspect.isawaitable(result):
result = await result
return result

def _get_mandatory_args(
self,
Expand Down
44 changes: 44 additions & 0 deletions tests/unittests/tools/test_function_tool.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,9 @@
# See the License for the specific language governing permissions and
# limitations under the License.

import asyncio
from enum import Enum
import functools
import inspect
from typing import Any
from typing import Optional
Expand All @@ -26,6 +28,7 @@
from google.adk.features._feature_registry import temporary_feature_override
from google.adk.sessions.session import Session
from google.adk.tools.function_tool import _build_declaration_cached
from google.adk.tools.function_tool import _use_sync_callable_runner
from google.adk.tools.function_tool import FunctionTool
from google.adk.tools.tool_confirmation import ToolConfirmation
from google.adk.tools.tool_context import ToolContext
Expand Down Expand Up @@ -205,6 +208,47 @@ async def test_run_async_without_tool_context_sync_func():
assert result == "test_value_1"


def _sync_wrapper(func):
"""A plain sync decorator, as commonly used for logging or retries."""

@functools.wraps(func)
def wrapper(*args, **kwargs):
return func(*args, **kwargs)

return wrapper


@_sync_wrapper
async def async_function_behind_sync_wrapper(item: str) -> dict:
"""Async function hidden behind a sync decorator."""
return {"item": item, "price": 3}


@pytest.mark.asyncio
async def test_run_async_awaits_async_function_behind_sync_wrapper():
"""Test that run_async awaits the coroutine returned by a sync wrapper."""
tool = FunctionTool(async_function_behind_sync_wrapper)
result = await tool.run_async(
args={"item": "apple"}, tool_context=MagicMock()
)
assert result == {"item": "apple", "price": 3}


@pytest.mark.asyncio
async def test_run_async_awaits_async_function_behind_sync_wrapper_with_runner():
"""Test that the coroutine is awaited when a sync callable runner is bound."""

async def run_in_thread(target, args):
return await asyncio.to_thread(target, **args)

tool = FunctionTool(async_function_behind_sync_wrapper)
with _use_sync_callable_runner(run_in_thread):
result = await tool.run_async(
args={"item": "apple"}, tool_context=MagicMock()
)
assert result == {"item": "apple", "price": 3}


@pytest.mark.asyncio
async def test_run_async_1_missing_arg_sync_func():
"""Test that run_async calls the function with 1 missing arg in signature (synchronous function)."""
Expand Down
Loading