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
53 changes: 50 additions & 3 deletions src/google/adk/runners.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@
import logging
from pathlib import Path
import queue
import threading
from types import TracebackType
from typing import Any
from typing import AsyncGenerator
Expand All @@ -31,6 +32,7 @@
from typing import Optional
from typing import TYPE_CHECKING
import warnings
import weakref

from google.genai import types
from opentelemetry import context
Expand Down Expand Up @@ -77,6 +79,37 @@

_EventQueueItem = tuple[object, asyncio.Event | None]

# Toolsets are often shared across Runners (parallel AgentTool, two Runners
# on the same agent). Holders are weak, so a Runner dropped without close()
# does not pin the toolset open for the process lifetime.
_toolset_holders: weakref.WeakKeyDictionary[
BaseToolset, weakref.WeakSet[Runner]
] = weakref.WeakKeyDictionary()
_toolset_holders_lock = threading.Lock()


def _register_toolset_runner(toolset: BaseToolset, runner: Runner) -> None:
with _toolset_holders_lock:
holders = _toolset_holders.get(toolset)
if holders is None:
holders = weakref.WeakSet()
_toolset_holders[toolset] = holders
holders.add(runner)


def _unregister_toolset_runner(toolset: BaseToolset, runner: Runner) -> bool:
"""Returns True if no live runner still holds the toolset."""
with _toolset_holders_lock:
holders = _toolset_holders.get(toolset)
if holders is None:
return True
holders.discard(runner)
if holders:
return False
_toolset_holders.pop(toolset, None)
return True


# Silence unused warning.
# tracer is imported for backwards compatibility, to avoid breaking change in the API.
_ = tracer
Expand Down Expand Up @@ -292,6 +325,12 @@ def __init__(
self._app_name_alignment_hint: Optional[str] = None
self._enforce_app_name_alignment()
self._warn_uncached_agent_transfer()
self._held_toolsets: set[BaseToolset] = set()
self._toolsets_released = False
if isinstance(self.agent, BaseAgent):
self._held_toolsets = self._collect_toolset(self.agent)
for toolset in self._held_toolsets:
_register_toolset_runner(toolset, self)

def _require_root_agent(self) -> BaseAgent:
"""Returns the root as an agent for agent-only execution paths."""
Expand Down Expand Up @@ -2140,9 +2179,17 @@ async def _cleanup_toolsets(
async def close(self) -> None:
"""Closes the runner."""
logger.info('Closing runner...')
# Close Toolsets
if isinstance(self.agent, BaseAgent):
await self._cleanup_toolsets(self._collect_toolset(self.agent))
# Close Toolsets only when this runner was the last live holder.
if isinstance(self.agent, BaseAgent) and not self._toolsets_released:
self._toolsets_released = True
collected = self._collect_toolset(self.agent)
to_close = {
toolset
for toolset in collected | self._held_toolsets
if _unregister_toolset_runner(toolset, self)
}
self._held_toolsets = set()
await self._cleanup_toolsets(to_close)

# Close Plugins
if self.plugin_manager:
Expand Down
45 changes: 24 additions & 21 deletions src/google/adk/tools/agent_tool.py
Original file line number Diff line number Diff line change
Expand Up @@ -317,27 +317,30 @@ async def run_async(
last_content = None
last_error_message = None
last_grounding_metadata = None
async with Aclosing(
runner.run_async(
user_id=session.user_id,
session_id=session.id,
new_message=content,
run_config=nested_run_config,
)
) as agen:
async for event in agen:
# Forward state delta to parent session.
if event.actions.state_delta:
tool_context.state.update(event.actions.state_delta)
if event.error_message:
last_error_message = event.error_message
if event.content:
last_content = event.content
last_grounding_metadata = event.grounding_metadata

# Clean up runner resources (especially MCP sessions)
# to avoid "Attempted to exit cancel scope in a different task" errors
await runner.close()
try:
async with Aclosing(
runner.run_async(
user_id=session.user_id,
session_id=session.id,
new_message=content,
run_config=nested_run_config,
)
) as agen:
async for event in agen:
# Forward state delta to parent session.
if event.actions.state_delta:
tool_context.state.update(event.actions.state_delta)
if event.error_message:
last_error_message = event.error_message
if event.content:
last_content = event.content
last_grounding_metadata = event.grounding_metadata
finally:
# Same-task close so MCP cancel scopes unwind on this task even
# when the nested run raises. Runner.close only closes a toolset
# once no other live Runner holds it, so this is safe when the
# toolset is shared with the parent.
await runner.close()

if last_content is None or last_content.parts is None:
return last_error_message or ''
Expand Down
98 changes: 98 additions & 0 deletions tests/unittests/test_runners.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@

import asyncio
from contextlib import aclosing
import gc
import importlib
import logging
from pathlib import Path
Expand Down Expand Up @@ -60,6 +61,21 @@
TEST_SESSION_ID = "test_session"


class CountingToolset(BaseToolset):
"""Toolset that records how many times close() ran."""

def __init__(self):
super().__init__()
self.close_count = 0

async def get_tools(self, readonly_context=None):
del readonly_context
return []

async def close(self) -> None:
self.close_count += 1


class MockAgent(BaseAgent):
"""Mock agent for unit testing."""

Expand Down Expand Up @@ -1274,6 +1290,88 @@ async def close(self) -> None:
assert toolset.close_cancelled is False
assert toolset.close_finished.is_set()

@pytest.mark.asyncio
async def test_shared_toolset_stays_open_until_last_runner_closes(self):
"""A toolset shared by two Runners must not close with the first one."""

toolset = CountingToolset()
agent = LlmAgent(
name="shared_agent", model="gemini-1.5-pro", tools=[toolset]
)
runner_one = Runner(
app_name="test_app",
agent=agent,
session_service=self.session_service,
artifact_service=self.artifact_service,
)
runner_two = Runner(
app_name="test_app",
agent=agent,
session_service=self.session_service,
artifact_service=self.artifact_service,
)

await runner_one.close()
assert toolset.close_count == 0
await runner_two.close()
assert toolset.close_count == 1

@pytest.mark.asyncio
async def test_single_runner_still_closes_its_toolset(self):
toolset = CountingToolset()
runner = Runner(
app_name="test_app",
agent=LlmAgent(
name="solo_agent", model="gemini-1.5-pro", tools=[toolset]
),
session_service=self.session_service,
artifact_service=self.artifact_service,
)

await runner.close()
assert toolset.close_count == 1
await runner.close()
assert toolset.close_count == 1

@pytest.mark.asyncio
async def test_dropped_runner_does_not_pin_shared_toolset(self):
toolset = CountingToolset()
agent = LlmAgent(
name="shared_agent", model="gemini-1.5-pro", tools=[toolset]
)
dropped = Runner(
app_name="test_app",
agent=agent,
session_service=self.session_service,
artifact_service=self.artifact_service,
)
surviving = Runner(
app_name="test_app",
agent=agent,
session_service=self.session_service,
artifact_service=self.artifact_service,
)
del dropped
gc.collect()

await surviving.close()
assert toolset.close_count == 1

@pytest.mark.asyncio
async def test_late_added_toolset_is_closed(self):
toolset = CountingToolset()
agent = LlmAgent(name="late_agent", model="gemini-1.5-pro", tools=[])
runner = Runner(
app_name="test_app",
agent=agent,
session_service=self.session_service,
artifact_service=self.artifact_service,
)
agent.tools.append(toolset)

await runner.close()
assert toolset.close_count == 1

@pytest.mark.asyncio
async def test_runner_passes_plugin_close_timeout(self):
"""Test that runner passes plugin_close_timeout to PluginManager."""
Expand Down
101 changes: 101 additions & 0 deletions tests/unittests/tools/test_agent_tool.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,7 @@
import google.adk.runners as _runners_module
from google.adk.sessions.in_memory_session_service import InMemorySessionService
from google.adk.tools.agent_tool import AgentTool
from google.adk.tools.base_toolset import BaseToolset
from google.adk.tools.tool_context import ToolContext
from google.adk.utils._schema_utils import validate_node_data
from google.adk.utils.variant_utils import GoogleLLMVariant
Expand All @@ -50,6 +51,37 @@

from .. import testing_utils


class CountingToolset(BaseToolset):
"""Toolset that records how many times close() ran."""

def __init__(self):
super().__init__()
self.close_count = 0

async def get_tools(self, readonly_context=None):
del readonly_context
return []

async def close(self):
self.close_count += 1


async def _tool_context_for(root_agent: Agent) -> ToolContext:
session_service = InMemorySessionService()
session = await session_service.create_session(
app_name='test_app', user_id='test_user'
)
invocation_context = InvocationContext(
invocation_id='invocation_id',
agent=root_agent,
session=session,
session_service=session_service,
plugin_manager=PluginManager(),
)
return ToolContext(invocation_context=invocation_context)


function_call_custom = Part.from_function_call(
name='tool_agent', args={'custom_input': 'test1'}
)
Expand Down Expand Up @@ -921,6 +953,75 @@ async def close(self):
assert parent_plugin.close_calls == 0


@mark.asyncio
async def test_parallel_agent_tools_do_not_close_shared_toolset_early():
toolset = CountingToolset()
started = asyncio.Event()
release = asyncio.Event()

async def _hold(_callback_context):
started.set()
await release.wait()
return None

mock_model = testing_utils.MockModel.create(responses=['ok-a', 'ok-b'])
agent_a = Agent(name='tool_a', model=mock_model, tools=[toolset])
agent_b = Agent(
name='tool_b',
model=mock_model,
tools=[toolset],
before_agent_callback=_hold,
)
root_agent = Agent(name='root_agent', model='test-model')
tool_context = await _tool_context_for(root_agent)

slow = asyncio.create_task(
AgentTool(agent=agent_b).run_async(
args={'request': 'b'}, tool_context=tool_context
)
)
await started.wait()
await AgentTool(agent=agent_a).run_async(
args={'request': 'a'}, tool_context=tool_context
)
assert toolset.close_count == 0
release.set()
await slow
assert toolset.close_count == 1


@mark.asyncio
async def test_failed_agent_tool_still_closes_shared_toolset():
async def _boom(_callback_context):
raise RuntimeError('boom')

toolset = CountingToolset()
boom_agent = Agent(
name='boom_agent',
model='test-model',
tools=[toolset],
before_agent_callback=_boom,
)
ok_agent = Agent(
name='ok_agent',
model=testing_utils.MockModel.create(responses=['ok']),
tools=[toolset],
)
root_agent = Agent(name='root_agent', model='test-model')
tool_context = await _tool_context_for(root_agent)

with pytest.raises(RuntimeError, match='boom'):
await AgentTool(agent=boom_agent).run_async(
args={'request': 'x'}, tool_context=tool_context
)
assert toolset.close_count == 1

await AgentTool(agent=ok_agent).run_async(
args={'request': 'y'}, tool_context=tool_context
)
assert toolset.close_count == 2


def test_agent_tool_description_with_input_schema():
"""Test that agent description is propagated when using input_schema."""

Expand Down