mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
feat(tools): Add an option to disallow propagating runner plugins to AgentTool runner
Merge https://github.com/google/adk-python/pull/2779 Fixes #2780 ### testing plan not available as is doesn't introduce new functionality Co-authored-by: Wei Sun (Jack) <weisun@google.com> COPYBARA_INTEGRATE_REVIEW=https://github.com/google/adk-python/pull/2779 from davidkl97:feature/agent-tool-plugins a602c808789f3daeed6244e352a6fb8fb6972de3 PiperOrigin-RevId: 835366974
This commit is contained in:
committed by
Copybara-Service
parent
2247a45922
commit
777dba3033
@@ -570,3 +570,112 @@ def test_agent_tool_response_schema_with_input_schema_no_output_vertex_ai(
|
||||
# Should have string response schema for VERTEX_AI when no output_schema
|
||||
assert declaration.response is not None
|
||||
assert declaration.response.type == types.Type.STRING
|
||||
|
||||
|
||||
def test_include_plugins_default_true():
|
||||
"""Test that plugins are propagated by default (include_plugins=True)."""
|
||||
|
||||
# Create a test plugin that tracks callbacks
|
||||
class TrackingPlugin(BasePlugin):
|
||||
|
||||
def __init__(self, name: str):
|
||||
super().__init__(name)
|
||||
self.before_agent_calls = 0
|
||||
|
||||
async def before_agent_callback(self, **kwargs):
|
||||
self.before_agent_calls += 1
|
||||
|
||||
tracking_plugin = TrackingPlugin(name='tracking')
|
||||
|
||||
mock_model = testing_utils.MockModel.create(
|
||||
responses=[function_call_no_schema, 'response1', 'response2']
|
||||
)
|
||||
|
||||
tool_agent = Agent(
|
||||
name='tool_agent',
|
||||
model=mock_model,
|
||||
)
|
||||
|
||||
root_agent = Agent(
|
||||
name='root_agent',
|
||||
model=mock_model,
|
||||
tools=[AgentTool(agent=tool_agent)], # Default include_plugins=True
|
||||
)
|
||||
|
||||
runner = testing_utils.InMemoryRunner(root_agent, plugins=[tracking_plugin])
|
||||
runner.run('test1')
|
||||
|
||||
# Plugin should be called for both root_agent and tool_agent
|
||||
assert tracking_plugin.before_agent_calls == 2
|
||||
|
||||
|
||||
def test_include_plugins_explicit_true():
|
||||
"""Test that plugins are propagated when include_plugins=True."""
|
||||
|
||||
class TrackingPlugin(BasePlugin):
|
||||
|
||||
def __init__(self, name: str):
|
||||
super().__init__(name)
|
||||
self.before_agent_calls = 0
|
||||
|
||||
async def before_agent_callback(self, **kwargs):
|
||||
self.before_agent_calls += 1
|
||||
|
||||
tracking_plugin = TrackingPlugin(name='tracking')
|
||||
|
||||
mock_model = testing_utils.MockModel.create(
|
||||
responses=[function_call_no_schema, 'response1', 'response2']
|
||||
)
|
||||
|
||||
tool_agent = Agent(
|
||||
name='tool_agent',
|
||||
model=mock_model,
|
||||
)
|
||||
|
||||
root_agent = Agent(
|
||||
name='root_agent',
|
||||
model=mock_model,
|
||||
tools=[AgentTool(agent=tool_agent, include_plugins=True)],
|
||||
)
|
||||
|
||||
runner = testing_utils.InMemoryRunner(root_agent, plugins=[tracking_plugin])
|
||||
runner.run('test1')
|
||||
|
||||
# Plugin should be called for both root_agent and tool_agent
|
||||
assert tracking_plugin.before_agent_calls == 2
|
||||
|
||||
|
||||
def test_include_plugins_false():
|
||||
"""Test that plugins are NOT propagated when include_plugins=False."""
|
||||
|
||||
class TrackingPlugin(BasePlugin):
|
||||
|
||||
def __init__(self, name: str):
|
||||
super().__init__(name)
|
||||
self.before_agent_calls = 0
|
||||
|
||||
async def before_agent_callback(self, **kwargs):
|
||||
self.before_agent_calls += 1
|
||||
|
||||
tracking_plugin = TrackingPlugin(name='tracking')
|
||||
|
||||
mock_model = testing_utils.MockModel.create(
|
||||
responses=[function_call_no_schema, 'response1', 'response2']
|
||||
)
|
||||
|
||||
tool_agent = Agent(
|
||||
name='tool_agent',
|
||||
model=mock_model,
|
||||
)
|
||||
|
||||
root_agent = Agent(
|
||||
name='root_agent',
|
||||
model=mock_model,
|
||||
tools=[AgentTool(agent=tool_agent, include_plugins=False)],
|
||||
)
|
||||
|
||||
runner = testing_utils.InMemoryRunner(root_agent, plugins=[tracking_plugin])
|
||||
runner.run('test1')
|
||||
|
||||
# Plugin should only be called for root_agent, not tool_agent
|
||||
assert tracking_plugin.before_agent_calls == 1
|
||||
|
||||
Reference in New Issue
Block a user