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:
committed by
Copybara-Service
parent
068b7f0a83
commit
a69ba4fa74
@@ -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)
|
||||
|
||||
|
||||
@@ -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."""
|
||||
|
||||
Reference in New Issue
Block a user