Skip to content
Merged
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
18 changes: 12 additions & 6 deletions python/packages/core/agent_framework/_tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -614,9 +614,16 @@ def __call__(self, *args: Any, **kwargs: Any) -> Any:
if func is None:
raise ToolException(f"Function '{self.name}' has no implementation.")
# If we have a bound instance, call the function with self
if self._instance is not None:
return func(self._instance, *args, **kwargs)
return func(*args, **kwargs)
result = func(self._instance, *args, **kwargs) if self._instance is not None else func(*args, **kwargs)
return self._await_invocation_result(result) if inspect.isawaitable(result) else result
except Exception:
self.invocation_exception_count += 1
raise

async def _await_invocation_result(self, result: Any) -> Any:
"""Await a function result and count exceptions raised by the awaitable."""
try:
return await result
except Exception:
self.invocation_exception_count += 1
raise
Expand All @@ -626,9 +633,8 @@ async def _invoke_function(self, call_kwargs: Mapping[str, Any]) -> Any:
func = self.func.func if isinstance(self.func, FunctionTool) else self.func
if inspect.iscoroutinefunction(func) or getattr(self, "_invoke_sync_on_event_loop", False):
res = self.__call__(**call_kwargs)
return await res if inspect.isawaitable(res) else res

res = await asyncio.to_thread(self.__call__, **call_kwargs)
else:
res = await asyncio.to_thread(self.__call__, **call_kwargs)
return await res if inspect.isawaitable(res) else res

@overload
Expand Down
66 changes: 66 additions & 0 deletions python/packages/core/tests/core/test_tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -343,6 +343,72 @@ async def async_test_tool(x: int, y: int) -> int:
assert (await async_test_tool(1, 2)) == 3


async def test_async_tool_exception_limit_counts_awaited_failures() -> None:
"""Async tool failures count toward the configured exception limit."""
from agent_framework.exceptions import ToolException

@tool(name="failing_async_tool", max_invocation_exceptions=1)
async def failing_async_tool() -> str:
raise RuntimeError("boom")

with pytest.raises(RuntimeError, match="boom"):
await failing_async_tool.invoke(skip_parsing=True)

assert failing_async_tool.invocation_count == 1
assert failing_async_tool.invocation_exception_count == 1

with pytest.raises(ToolException, match="maximum exception limit"):
await failing_async_tool.invoke(skip_parsing=True)

assert failing_async_tool.invocation_count == 1
assert failing_async_tool.invocation_exception_count == 1


async def test_direct_async_tool_exception_limit_counts_awaited_failures() -> None:
"""Direct async tool calls count failures toward the configured exception limit."""
from agent_framework.exceptions import ToolException

@tool(name="failing_direct_async_tool", max_invocation_exceptions=1)
async def failing_direct_async_tool() -> str:
raise RuntimeError("boom")

with pytest.raises(RuntimeError, match="boom"):
await failing_direct_async_tool()

assert failing_direct_async_tool.invocation_count == 1
assert failing_direct_async_tool.invocation_exception_count == 1

with pytest.raises(ToolException, match="maximum exception limit"):
await failing_direct_async_tool()

assert failing_direct_async_tool.invocation_count == 1
assert failing_direct_async_tool.invocation_exception_count == 1


async def test_sync_awaitable_tool_exception_limit_counts_awaited_failures() -> None:
"""Sync tools returning awaitables count failures during async invocation."""
from agent_framework.exceptions import ToolException

@tool(name="failing_sync_awaitable_tool", max_invocation_exceptions=1)
def failing_sync_awaitable_tool() -> Any:
async def fail_later() -> str:
raise RuntimeError("boom")

return fail_later()

with pytest.raises(RuntimeError, match="boom"):
await failing_sync_awaitable_tool.invoke(skip_parsing=True)

assert failing_sync_awaitable_tool.invocation_count == 1
assert failing_sync_awaitable_tool.invocation_exception_count == 1

with pytest.raises(ToolException, match="maximum exception limit"):
await failing_sync_awaitable_tool.invoke(skip_parsing=True)

assert failing_sync_awaitable_tool.invocation_count == 1
assert failing_sync_awaitable_tool.invocation_exception_count == 1


def test_tool_decorator_in_class():
"""Test the tool decorator."""

Expand Down
Loading