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
@@ -20,19 +20,33 @@ from google.adk.models.llm_request import LlmRequest
|
||||
from google.adk.models.llm_response import LlmResponse
|
||||
from google.adk.plugins.base_plugin import BasePlugin
|
||||
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_model_text = 'before_model_text from MockPlugin'
|
||||
after_model_text = 'after_model_text from MockPlugin'
|
||||
on_model_error_text = 'on_model_error_text from MockPlugin'
|
||||
|
||||
def __init__(self, name='mock_plugin'):
|
||||
self.name = name
|
||||
self.enable_before_model_callback = False
|
||||
self.enable_after_model_callback = False
|
||||
self.enable_on_model_error_callback = False
|
||||
self.before_model_response = LlmResponse(
|
||||
content=testing_utils.ModelContent(
|
||||
[types.Part.from_text(text=self.before_model_text)]
|
||||
@@ -43,6 +57,11 @@ class MockPlugin(BasePlugin):
|
||||
[types.Part.from_text(text=self.after_model_text)]
|
||||
)
|
||||
)
|
||||
self.on_model_error_response = LlmResponse(
|
||||
content=testing_utils.ModelContent(
|
||||
[types.Part.from_text(text=self.on_model_error_text)]
|
||||
)
|
||||
)
|
||||
|
||||
async def before_model_callback(
|
||||
self, *, callback_context: CallbackContext, llm_request: LlmRequest
|
||||
@@ -58,6 +77,17 @@ class MockPlugin(BasePlugin):
|
||||
return None
|
||||
return self.after_model_response
|
||||
|
||||
async def on_model_error_callback(
|
||||
self,
|
||||
*,
|
||||
callback_context: CallbackContext,
|
||||
llm_request: LlmRequest,
|
||||
error: Exception,
|
||||
) -> Optional[LlmResponse]:
|
||||
if not self.enable_on_model_error_callback:
|
||||
return None
|
||||
return self.on_model_error_response
|
||||
|
||||
|
||||
CANONICAL_MODEL_CALLBACK_CONTENT = 'canonical_model_callback_content'
|
||||
|
||||
@@ -124,5 +154,36 @@ def test_before_model_callback_fallback_model(mock_plugin):
|
||||
]
|
||||
|
||||
|
||||
def test_on_model_error_callback_with_plugin(mock_plugin):
|
||||
"""Tests that the model error is handled by the plugin."""
|
||||
mock_model = testing_utils.MockModel.create(error=mock_error, responses=[])
|
||||
mock_plugin.enable_on_model_error_callback = True
|
||||
agent = Agent(
|
||||
name='root_agent',
|
||||
model=mock_model,
|
||||
)
|
||||
|
||||
runner = testing_utils.InMemoryRunner(agent, plugins=[mock_plugin])
|
||||
|
||||
assert testing_utils.simplify_events(runner.run('test')) == [
|
||||
('root_agent', mock_plugin.on_model_error_text),
|
||||
]
|
||||
|
||||
|
||||
def test_on_model_error_callback_fallback_to_runner(mock_plugin):
|
||||
"""Tests that the model error is not handled and falls back to raise from runner."""
|
||||
mock_model = testing_utils.MockModel.create(error=mock_error, responses=[])
|
||||
mock_plugin.enable_on_model_error_callback = False
|
||||
agent = Agent(
|
||||
name='root_agent',
|
||||
model=mock_model,
|
||||
)
|
||||
|
||||
try:
|
||||
testing_utils.InMemoryRunner(agent, plugins=[mock_plugin])
|
||||
except Exception as e:
|
||||
assert e == mock_error
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
pytest.main([__file__])
|
||||
|
||||
Reference in New Issue
Block a user