mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
feat: add new callbacks to handle tool and model errors
This CL add new callbacks in plugin system: - `on_tool_error_callback` - `on_model_error_callback` This allow the user to create plugins that can handle errors. PiperOrigin-RevId: 786469646
This commit is contained in:
committed by
Copybara-Service
parent
dfc25c17a9
commit
00afaaf2fc
@@ -24,19 +24,35 @@ from google.adk.tools.base_tool import BaseTool
|
||||
from google.adk.tools.function_tool import FunctionTool
|
||||
from google.adk.tools.tool_context import ToolContext
|
||||
from google.genai import types
|
||||
from google.genai.errors import ClientError
|
||||
import pytest
|
||||
|
||||
from ... import testing_utils
|
||||
|
||||
mock_error = ClientError(
|
||||
code=429,
|
||||
response_json={
|
||||
"error": {
|
||||
"code": 429,
|
||||
"message": "Quota exceeded.",
|
||||
"status": "RESOURCE_EXHAUSTED",
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
class MockPlugin(BasePlugin):
|
||||
before_tool_response = {"MockPlugin": "before_tool_response from MockPlugin"}
|
||||
after_tool_response = {"MockPlugin": "after_tool_response from MockPlugin"}
|
||||
on_tool_error_response = {
|
||||
"MockPlugin": "on_tool_error_response from MockPlugin"
|
||||
}
|
||||
|
||||
def __init__(self, name="mock_plugin"):
|
||||
self.name = name
|
||||
self.enable_before_tool_callback = False
|
||||
self.enable_after_tool_callback = False
|
||||
self.enable_on_tool_error_callback = False
|
||||
|
||||
async def before_tool_callback(
|
||||
self,
|
||||
@@ -61,6 +77,18 @@ class MockPlugin(BasePlugin):
|
||||
return None
|
||||
return self.after_tool_response
|
||||
|
||||
async def on_tool_error_callback(
|
||||
self,
|
||||
*,
|
||||
tool: BaseTool,
|
||||
tool_args: dict[str, Any],
|
||||
tool_context: ToolContext,
|
||||
error: Exception,
|
||||
) -> Optional[dict]:
|
||||
if not self.enable_on_tool_error_callback:
|
||||
return None
|
||||
return self.on_tool_error_response
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_tool():
|
||||
@@ -70,6 +98,14 @@ def mock_tool():
|
||||
return FunctionTool(simple_fn)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_error_tool():
|
||||
def raise_error_fn(**kwargs) -> Dict[str, Any]:
|
||||
raise mock_error
|
||||
|
||||
return FunctionTool(raise_error_fn)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_plugin():
|
||||
return MockPlugin()
|
||||
@@ -124,5 +160,30 @@ async def test_async_after_tool_callback(mock_tool, mock_plugin):
|
||||
assert part.function_response.response == mock_plugin.after_tool_response
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_on_tool_error_use_plugin_response(
|
||||
mock_error_tool, mock_plugin
|
||||
):
|
||||
mock_plugin.enable_on_tool_error_callback = True
|
||||
|
||||
result_event = await invoke_tool_with_plugin(mock_error_tool, mock_plugin)
|
||||
|
||||
assert result_event is not None
|
||||
part = result_event.content.parts[0]
|
||||
assert part.function_response.response == mock_plugin.on_tool_error_response
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_on_tool_error_fallback_to_runner(
|
||||
mock_error_tool, mock_plugin
|
||||
):
|
||||
mock_plugin.enable_on_tool_error_callback = False
|
||||
|
||||
try:
|
||||
await invoke_tool_with_plugin(mock_error_tool, mock_plugin)
|
||||
except Exception as e:
|
||||
assert e == mock_error
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__])
|
||||
|
||||
Reference in New Issue
Block a user