Skip to content

Commit fd0d50f

Browse files
committed
fix(runner): shield toolset cleanup from cancellation
1 parent dec2182 commit fd0d50f

2 files changed

Lines changed: 78 additions & 3 deletions

File tree

src/google/adk/runners.py

Lines changed: 31 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -2095,10 +2095,12 @@ async def _cleanup_toolsets(self, toolsets_to_close: set[BaseToolset]):
20952095

20962096
# This maintains the same task context throughout cleanup
20972097
for toolset in toolsets_to_close:
2098+
cleanup_task = asyncio.create_task(
2099+
asyncio.wait_for(toolset.close(), timeout=10.0)
2100+
)
20982101
try:
20992102
logger.info('Closing toolset: %s', type(toolset).__name__)
2100-
# Use asyncio.wait_for to add timeout protection
2101-
await asyncio.wait_for(toolset.close(), timeout=10.0)
2103+
await asyncio.shield(cleanup_task)
21022104
logger.info('Successfully closed toolset: %s', type(toolset).__name__)
21032105
except asyncio.TimeoutError:
21042106
logger.warning('Toolset %s cleanup timed out', type(toolset).__name__)
@@ -2113,8 +2115,34 @@ async def _cleanup_toolsets(self, toolsets_to_close: set[BaseToolset]):
21132115
# improved context propagation across task boundaries, and better cancellation
21142116
# handling prevent the cross-task cancel scope violation.
21152117
logger.warning(
2116-
'Toolset %s cleanup cancelled: %s', type(toolset).__name__, e
2118+
'Toolset %s cleanup cancellation requested: %s',
2119+
type(toolset).__name__,
2120+
e,
21172121
)
2122+
try:
2123+
await cleanup_task
2124+
logger.info(
2125+
'Successfully closed toolset after cancellation request: %s',
2126+
type(toolset).__name__,
2127+
)
2128+
except asyncio.TimeoutError:
2129+
cleanup_task.cancel()
2130+
logger.warning(
2131+
'Toolset %s cleanup timed out after cancellation request',
2132+
type(toolset).__name__,
2133+
)
2134+
except asyncio.CancelledError as close_cancelled:
2135+
logger.warning(
2136+
'Toolset %s cleanup cancelled: %s',
2137+
type(toolset).__name__,
2138+
close_cancelled,
2139+
)
2140+
except Exception as close_error:
2141+
logger.error(
2142+
'Error closing toolset %s after cancellation request: %s',
2143+
type(toolset).__name__,
2144+
close_error,
2145+
)
21182146
except Exception as e:
21192147
logger.error('Error closing toolset %s: %s', type(toolset).__name__, e)
21202148

tests/unittests/test_runners.py

Lines changed: 47 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -36,6 +36,7 @@
3636
from google.adk.runners import Runner
3737
from google.adk.sessions.in_memory_session_service import InMemorySessionService
3838
from google.adk.sessions.session import Session
39+
from google.adk.tools.base_toolset import BaseToolset
3940
from google.genai import types
4041
import pytest
4142

@@ -1120,6 +1121,52 @@ async def test_runner_close_calls_plugin_close(self):
11201121

11211122
self.runner.plugin_manager.close.assert_awaited_once()
11221123

1124+
@pytest.mark.asyncio
1125+
async def test_runner_close_does_not_cancel_toolset_cleanup(self):
1126+
"""Caller cancellation should not cancel an in-flight toolset close."""
1127+
import asyncio
1128+
1129+
class SlowCloseToolset(BaseToolset):
1130+
1131+
def __init__(self):
1132+
super().__init__()
1133+
self.close_started = asyncio.Event()
1134+
self.close_finished = asyncio.Event()
1135+
self.close_cancelled = False
1136+
1137+
async def get_tools(self, readonly_context=None):
1138+
del readonly_context
1139+
return []
1140+
1141+
async def close(self) -> None:
1142+
self.close_started.set()
1143+
try:
1144+
await asyncio.sleep(0.05)
1145+
self.close_finished.set()
1146+
except asyncio.CancelledError:
1147+
self.close_cancelled = True
1148+
raise
1149+
1150+
toolset = SlowCloseToolset()
1151+
runner = Runner(
1152+
app_name="test_app",
1153+
agent=LlmAgent(
1154+
name="test_agent", model="gemini-1.5-pro", tools=[toolset]
1155+
),
1156+
session_service=self.session_service,
1157+
artifact_service=self.artifact_service,
1158+
)
1159+
1160+
close_task = asyncio.create_task(runner.close())
1161+
await toolset.close_started.wait()
1162+
close_task.cancel()
1163+
1164+
await close_task
1165+
1166+
assert close_task.cancelled() is False
1167+
assert toolset.close_cancelled is False
1168+
assert toolset.close_finished.is_set()
1169+
11231170
@pytest.mark.asyncio
11241171
async def test_runner_passes_plugin_close_timeout(self):
11251172
"""Test that runner passes plugin_close_timeout to PluginManager."""

0 commit comments

Comments
 (0)