fix: shield toolset cleanup from cancellation

Shields each toolset cleanup task during Runner.close() so caller cancellation does not cancel the underlying toolset.close() coroutine.

PiperOrigin-RevId: 945791744
This commit is contained in:
Sehlani042
2026-07-10 10:52:05 -07:00
committed by Copybara-Service
parent 068b7f0a83
commit a69ba4fa74
2 changed files with 80 additions and 3 deletions
+32 -3
View File
@@ -2185,10 +2185,12 @@ class Runner:
# This maintains the same task context throughout cleanup
for toolset in toolsets_to_close:
cleanup_task = asyncio.create_task(
asyncio.wait_for(toolset.close(), timeout=10.0)
)
try:
logger.info('Closing toolset: %s', type(toolset).__name__)
# Use asyncio.wait_for to add timeout protection
await asyncio.wait_for(toolset.close(), timeout=10.0)
await asyncio.shield(cleanup_task)
logger.info('Successfully closed toolset: %s', type(toolset).__name__)
except asyncio.TimeoutError:
logger.warning('Toolset %s cleanup timed out', type(toolset).__name__)
@@ -2203,8 +2205,35 @@ class Runner:
# improved context propagation across task boundaries, and better cancellation
# handling prevent the cross-task cancel scope violation.
logger.warning(
'Toolset %s cleanup cancelled: %s', type(toolset).__name__, e
'Toolset %s cleanup cancellation requested: %s',
type(toolset).__name__,
e,
)
try:
await cleanup_task
logger.info(
'Successfully closed toolset after cancellation request: %s',
type(toolset).__name__,
)
except asyncio.TimeoutError:
cleanup_task.cancel()
logger.warning(
'Toolset %s cleanup timed out after cancellation request',
type(toolset).__name__,
)
except asyncio.CancelledError as close_cancelled:
logger.warning(
'Toolset %s cleanup cancelled: %s',
type(toolset).__name__,
close_cancelled,
)
except Exception as close_error:
logger.error(
'Error closing toolset %s after cancellation request: %s',
type(toolset).__name__,
close_error,
)
raise
except Exception as e:
logger.error('Error closing toolset %s: %s', type(toolset).__name__, e)
+48
View File
@@ -12,6 +12,7 @@
# See the License for the specific language governing permissions and
# limitations under the License.
import asyncio
from contextlib import aclosing
import importlib
from pathlib import Path
@@ -36,6 +37,7 @@ from google.adk.plugins.base_plugin import BasePlugin
from google.adk.runners import Runner
from google.adk.sessions.in_memory_session_service import InMemorySessionService
from google.adk.sessions.session import Session
from google.adk.tools.base_toolset import BaseToolset
from google.genai import types
import pytest
@@ -1120,6 +1122,52 @@ class TestRunnerWithPlugins:
self.runner.plugin_manager.close.assert_awaited_once()
@pytest.mark.asyncio
async def test_runner_close_does_not_cancel_toolset_cleanup(self):
"""Caller cancellation should not cancel an in-flight toolset close."""
class SlowCloseToolset(BaseToolset):
def __init__(self):
super().__init__()
self.close_started = asyncio.Event()
self.close_finished = asyncio.Event()
self.close_cancelled = False
async def get_tools(self, readonly_context=None):
del readonly_context
return []
async def close(self) -> None:
self.close_started.set()
try:
await asyncio.sleep(0.05)
self.close_finished.set()
except asyncio.CancelledError:
self.close_cancelled = True
raise
toolset = SlowCloseToolset()
runner = Runner(
app_name="test_app",
agent=LlmAgent(
name="test_agent", model="gemini-1.5-pro", tools=[toolset]
),
session_service=self.session_service,
artifact_service=self.artifact_service,
)
close_task = asyncio.create_task(runner.close())
await toolset.close_started.wait()
close_task.cancel()
with pytest.raises(asyncio.CancelledError):
await close_task
assert close_task.cancelled() is True
assert toolset.close_cancelled is False
assert toolset.close_finished.is_set()
@pytest.mark.asyncio
async def test_runner_passes_plugin_close_timeout(self):
"""Test that runner passes plugin_close_timeout to PluginManager."""