From 249216e8904ddd849918b5be8894a5f48b0f1225 Mon Sep 17 00:00:00 2001 From: Google Team Member Date: Tue, 11 Nov 2025 13:04:37 -0800 Subject: [PATCH] feat: Add Graceful Plugin Shutdown to Runner This change introduces a shutdown lifecycle hook for plugins. The `PluginManager` now has an `async def shutdown()` method that will call `await plugin.shutdown()` on any registered plugins that implement the method. This is called from `Runner.close()`, allowing plugins to perform cleanup tasks like flushing logs or closing connections when the runner instance is being closed. This improves the reliability of plugins that perform background operations. PiperOrigin-RevId: 831037737 --- src/google/adk/plugins/base_plugin.py | 8 +++ .../bigquery_agent_analytics_plugin.py | 2 +- src/google/adk/plugins/plugin_manager.py | 50 ++++++++++++++++++- src/google/adk/runners.py | 19 ++++++- .../test_bigquery_agent_analytics_plugin.py | 6 +-- .../unittests/plugins/test_plugin_manager.py | 50 +++++++++++++++++++ tests/unittests/test_runners.py | 24 +++++++++ 7 files changed, 152 insertions(+), 7 deletions(-) diff --git a/src/google/adk/plugins/base_plugin.py b/src/google/adk/plugins/base_plugin.py index fb3e3c00..f75c33ec 100644 --- a/src/google/adk/plugins/base_plugin.py +++ b/src/google/adk/plugins/base_plugin.py @@ -187,6 +187,14 @@ class BasePlugin(ABC): """ pass + async def close(self) -> None: + """Method executed when the runner is closed. + + This method is used for cleanup tasks such as closing network connections + or releasing resources. + """ + pass + async def before_agent_callback( self, *, agent: BaseAgent, callback_context: CallbackContext ) -> Optional[types.Content]: diff --git a/src/google/adk/plugins/bigquery_agent_analytics_plugin.py b/src/google/adk/plugins/bigquery_agent_analytics_plugin.py index cc1bf6a7..65fd5398 100644 --- a/src/google/adk/plugins/bigquery_agent_analytics_plugin.py +++ b/src/google/adk/plugins/bigquery_agent_analytics_plugin.py @@ -455,7 +455,7 @@ class BigQueryAgentAnalyticsPlugin(BasePlugin): self._background_tasks.add(task) task.add_done_callback(self._background_tasks.discard) - async def shutdown(self): + async def close(self): """Flushes pending logs and closes client.""" # 1. Wait for pending background logs (best effort, 2s timeout) if self._background_tasks: diff --git a/src/google/adk/plugins/plugin_manager.py b/src/google/adk/plugins/plugin_manager.py index 634f9043..650583c2 100644 --- a/src/google/adk/plugins/plugin_manager.py +++ b/src/google/adk/plugins/plugin_manager.py @@ -14,7 +14,9 @@ from __future__ import annotations +import asyncio import logging +import sys from typing import Any from typing import List from typing import Literal @@ -70,13 +72,19 @@ class PluginManager: tool calls, or model requests. """ - def __init__(self, plugins: Optional[List[BasePlugin]] = None): + def __init__( + self, + plugins: Optional[List[BasePlugin]] = None, + close_timeout: float = 5.0, + ): """Initializes the plugin service. Args: plugins: An optional list of plugins to register upon initialization. + close_timeout: The timeout in seconds for each plugin's close method. """ self.plugins: List[BasePlugin] = [] + self._close_timeout = close_timeout if plugins: for plugin in plugins: self.register_plugin(plugin) @@ -297,3 +305,43 @@ class PluginManager: raise RuntimeError(error_message) from e return None + + async def close(self) -> None: + """Calls the close method on all registered plugins concurrently. + + Raises: + RuntimeError: If one or more plugins failed to close, containing + details of all failures. + """ + exceptions = {} + # We iterate sequentially to avoid creating new tasks which can cause issues + # with some libraries (like anyio/mcp) that rely on task-local context. + for plugin in self.plugins: + try: + if sys.version_info >= (3, 11): + async with asyncio.timeout(self._close_timeout): + await plugin.close() + else: + # For Python < 3.11, we use wait_for which creates a new task. + # This might still cause issues with task-local contexts, but + # asyncio.timeout is not available. + await asyncio.wait_for(plugin.close(), timeout=self._close_timeout) + except Exception as e: + exceptions[plugin.name] = e + if isinstance(e, (asyncio.TimeoutError, asyncio.CancelledError)): + logger.warning( + "Timeout/Cancelled while closing plugin: %s", plugin.name + ) + else: + logger.error( + "Error during close of plugin %s: %s", + plugin.name, + e, + exc_info=e, + ) + + if exceptions: + error_summary = ", ".join( + f"'{name}': {type(exc).__name__}" for name, exc in exceptions.items() + ) + raise RuntimeError(f"Failed to close plugins: {error_summary}") diff --git a/src/google/adk/runners.py b/src/google/adk/runners.py index 828d9c2c..492ef2ea 100644 --- a/src/google/adk/runners.py +++ b/src/google/adk/runners.py @@ -115,6 +115,7 @@ class Runner: session_service: BaseSessionService, memory_service: Optional[BaseMemoryService] = None, credential_service: Optional[BaseCredentialService] = None, + plugin_close_timeout: float = 5.0, ): """Initializes the Runner. @@ -134,6 +135,7 @@ class Runner: session_service: The session service for the runner. memory_service: The memory service for the runner. credential_service: The credential service for the runner. + plugin_close_timeout: The timeout in seconds for plugin close methods. Raises: ValueError: If `app` is provided along with `app_name` or `plugins`, or @@ -151,7 +153,9 @@ class Runner: self.session_service = session_service self.memory_service = memory_service self.credential_service = credential_service - self.plugin_manager = PluginManager(plugins=plugins) + self.plugin_manager = PluginManager( + plugins=plugins, close_timeout=plugin_close_timeout + ) ( self._agent_origin_app_name, self._agent_origin_dir, @@ -1297,8 +1301,16 @@ class Runner: async def close(self): """Closes the runner.""" + logger.info('Closing runner...') + # Close Toolsets await self._cleanup_toolsets(self._collect_toolset(self.agent)) + # Close Plugins + if self.plugin_manager: + await self.plugin_manager.close() + + logger.info('Runner closed.') + async def __aenter__(self): """Async context manager entry.""" return self @@ -1329,6 +1341,7 @@ class InMemoryRunner(Runner): app_name: Optional[str] = None, plugins: Optional[list[BasePlugin]] = None, app: Optional[App] = None, + plugin_close_timeout: float = 5.0, ): """Initializes the InMemoryRunner. @@ -1336,6 +1349,9 @@ class InMemoryRunner(Runner): agent: The root agent to run. app_name: The application name of the runner. Defaults to 'InMemoryRunner'. + plugins: Optional list of plugins for the runner. + app: Optional App instance. + plugin_close_timeout: The timeout in seconds for plugin close methods. """ if app is None and app_name is None: app_name = 'InMemoryRunner' @@ -1347,4 +1363,5 @@ class InMemoryRunner(Runner): app=app, session_service=InMemorySessionService(), memory_service=InMemoryMemoryService(), + plugin_close_timeout=plugin_close_timeout, ) diff --git a/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py b/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py index f3251119..8696d744 100644 --- a/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py +++ b/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py @@ -513,10 +513,8 @@ class TestBigQueryAgentAnalyticsPlugin: mock_write_client.append_rows.assert_called_once() @pytest.mark.asyncio - async def test_shutdown( - self, bq_plugin_inst, mock_bq_client, mock_write_client - ): - await bq_plugin_inst.shutdown() + async def test_close(self, bq_plugin_inst, mock_bq_client, mock_write_client): + await bq_plugin_inst.close() mock_write_client.transport.close.assert_called_once() mock_bq_client.close.assert_called_once() diff --git a/tests/unittests/plugins/test_plugin_manager.py b/tests/unittests/plugins/test_plugin_manager.py index e3edfa83..87e0b8cb 100644 --- a/tests/unittests/plugins/test_plugin_manager.py +++ b/tests/unittests/plugins/test_plugin_manager.py @@ -16,6 +16,8 @@ from __future__ import annotations +import asyncio +from unittest.mock import AsyncMock from unittest.mock import Mock from google.adk.models.llm_response import LlmResponse @@ -267,3 +269,51 @@ async def test_all_callbacks_are_supported( "on_model_error_callback", ] assert set(plugin1.call_log) == set(expected_callbacks) + + +@pytest.mark.asyncio +async def test_close_calls_plugin_close( + service: PluginManager, plugin1: TestPlugin +): + """Tests that close calls the close method on registered plugins.""" + plugin1.close = AsyncMock() + service.register_plugin(plugin1) + + await service.close() + + plugin1.close.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_close_raises_runtime_error_on_plugin_exception( + service: PluginManager, plugin1: TestPlugin +): + """Tests that close raises a RuntimeError if a plugin's close fails.""" + plugin1.close = AsyncMock(side_effect=ValueError("Shutdown error")) + service.register_plugin(plugin1) + + with pytest.raises( + RuntimeError, match="Failed to close plugins: 'plugin1': ValueError" + ): + await service.close() + + plugin1.close.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_close_with_timeout(plugin1: TestPlugin): + """Tests that close respects the timeout and raises on failure.""" + service = PluginManager(close_timeout=0.1) + + async def slow_close(): + await asyncio.sleep(0.2) + + plugin1.close = slow_close + service.register_plugin(plugin1) + + with pytest.raises(RuntimeError) as excinfo: + await service.close() + + assert "Failed to close plugins: 'plugin1': TimeoutError" in str( + excinfo.value + ) diff --git a/tests/unittests/test_runners.py b/tests/unittests/test_runners.py index d4b7c02c..7c4de3ec 100644 --- a/tests/unittests/test_runners.py +++ b/tests/unittests/test_runners.py @@ -15,6 +15,7 @@ from pathlib import Path import textwrap from typing import Optional +from unittest.mock import AsyncMock from google.adk.agents.base_agent import BaseAgent from google.adk.agents.context_cache_config import ContextCacheConfig @@ -562,6 +563,29 @@ class TestRunnerWithPlugins: assert modified_event_message == MockPlugin.ON_EVENT_CALLBACK_MSG + @pytest.mark.asyncio + async def test_runner_close_calls_plugin_close(self): + """Test that runner.close() calls plugin manager close.""" + # Mock the plugin manager's close method + self.runner.plugin_manager.close = AsyncMock() + + await self.runner.close() + + self.runner.plugin_manager.close.assert_awaited_once() + + @pytest.mark.asyncio + async def test_runner_passes_plugin_close_timeout(self): + """Test that runner passes plugin_close_timeout to PluginManager.""" + runner = Runner( + app_name="test_app", + agent=MockLlmAgent("test_agent"), + session_service=self.session_service, + artifact_service=self.artifact_service, + plugins=[self.plugin], + plugin_close_timeout=10.0, + ) + assert runner.plugin_manager._close_timeout == 10.0 + def test_runner_init_raises_error_with_app_and_app_name_and_agent(self): """Test that ValueError is raised when app, app_name and agent are provided.""" with pytest.raises(