mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
fix: Add push notification config store to agent_to_a2a
This change allows users to provide a custom PushNotificationConfigStore when converting an ADK agent to an A2A Starlette application. If no custom store is provided, an InMemoryPushNotificationConfigStore is used by default, thios now lets A2A push notification configuration RPCs Close #4126 Co-authored-by: George Weale <gweale@google.com> PiperOrigin-RevId: 874118109
This commit is contained in:
committed by
Copybara-Service
parent
ffbcc0a626
commit
4ca904f111
@@ -12,12 +12,14 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from unittest.mock import ANY
|
||||
from unittest.mock import AsyncMock
|
||||
from unittest.mock import Mock
|
||||
from unittest.mock import patch
|
||||
|
||||
from a2a.server.apps import A2AStarletteApplication
|
||||
from a2a.server.request_handlers import DefaultRequestHandler
|
||||
from a2a.server.tasks import InMemoryPushNotificationConfigStore
|
||||
from a2a.server.tasks import InMemoryTaskStore
|
||||
from a2a.types import AgentCard
|
||||
from google.adk.a2a.executor.a2a_agent_executor import A2aAgentExecutor
|
||||
@@ -77,7 +79,9 @@ class TestToA2A:
|
||||
mock_task_store_class.assert_called_once()
|
||||
mock_agent_executor_class.assert_called_once()
|
||||
mock_request_handler_class.assert_called_once_with(
|
||||
agent_executor=mock_agent_executor, task_store=mock_task_store
|
||||
agent_executor=mock_agent_executor,
|
||||
push_config_store=ANY,
|
||||
task_store=mock_task_store,
|
||||
)
|
||||
mock_card_builder_class.assert_called_once_with(
|
||||
agent=self.mock_agent, rpc_url="http://localhost:8000/"
|
||||
@@ -122,7 +126,9 @@ class TestToA2A:
|
||||
mock_task_store_class.assert_called_once()
|
||||
mock_agent_executor_class.assert_called_once_with(runner=custom_runner)
|
||||
mock_request_handler_class.assert_called_once_with(
|
||||
agent_executor=mock_agent_executor, task_store=mock_task_store
|
||||
agent_executor=mock_agent_executor,
|
||||
push_config_store=ANY,
|
||||
task_store=mock_task_store,
|
||||
)
|
||||
mock_card_builder_class.assert_called_once_with(
|
||||
agent=self.mock_agent, rpc_url="http://localhost:8000/"
|
||||
@@ -131,6 +137,42 @@ class TestToA2A:
|
||||
"startup", mock_app.add_event_handler.call_args[0][1]
|
||||
)
|
||||
|
||||
@patch("google.adk.a2a.utils.agent_to_a2a.A2aAgentExecutor")
|
||||
@patch("google.adk.a2a.utils.agent_to_a2a.DefaultRequestHandler")
|
||||
@patch("google.adk.a2a.utils.agent_to_a2a.InMemoryTaskStore")
|
||||
@patch("google.adk.a2a.utils.agent_to_a2a.AgentCardBuilder")
|
||||
@patch("google.adk.a2a.utils.agent_to_a2a.Starlette")
|
||||
def test_to_a2a_passes_custom_push_config_store(
|
||||
self,
|
||||
mock_starlette_class,
|
||||
mock_card_builder_class,
|
||||
mock_task_store_class,
|
||||
mock_request_handler_class,
|
||||
mock_agent_executor_class,
|
||||
):
|
||||
"""Test to_a2a forwards a custom push config store."""
|
||||
mock_app = Mock(spec=Starlette)
|
||||
mock_starlette_class.return_value = mock_app
|
||||
mock_task_store = Mock(spec=InMemoryTaskStore)
|
||||
mock_task_store_class.return_value = mock_task_store
|
||||
mock_agent_executor = Mock(spec=A2aAgentExecutor)
|
||||
mock_agent_executor_class.return_value = mock_agent_executor
|
||||
mock_request_handler = Mock(spec=DefaultRequestHandler)
|
||||
mock_request_handler_class.return_value = mock_request_handler
|
||||
mock_card_builder = Mock(spec=AgentCardBuilder)
|
||||
mock_card_builder_class.return_value = mock_card_builder
|
||||
|
||||
custom_push_store = InMemoryPushNotificationConfigStore()
|
||||
|
||||
result = to_a2a(self.mock_agent, push_config_store=custom_push_store)
|
||||
|
||||
assert result == mock_app
|
||||
mock_request_handler_class.assert_called_once_with(
|
||||
agent_executor=mock_agent_executor,
|
||||
push_config_store=custom_push_store,
|
||||
task_store=mock_task_store,
|
||||
)
|
||||
|
||||
@patch("google.adk.a2a.utils.agent_to_a2a.A2aAgentExecutor")
|
||||
@patch("google.adk.a2a.utils.agent_to_a2a.DefaultRequestHandler")
|
||||
@patch("google.adk.a2a.utils.agent_to_a2a.InMemoryTaskStore")
|
||||
|
||||
@@ -15,7 +15,6 @@
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from pathlib import Path
|
||||
import signal
|
||||
import tempfile
|
||||
@@ -677,6 +676,7 @@ def test_app_with_a2a(
|
||||
mock_eval_sets_manager,
|
||||
mock_eval_set_results_manager,
|
||||
temp_agents_dir_with_a2a,
|
||||
monkeypatch,
|
||||
):
|
||||
"""Create a TestClient for the FastAPI app with A2A enabled."""
|
||||
# Mock A2A related classes
|
||||
@@ -728,26 +728,22 @@ def test_app_with_a2a(
|
||||
mock_a2a_app.return_value = mock_app_instance
|
||||
|
||||
# Change to temp directory
|
||||
original_cwd = os.getcwd()
|
||||
os.chdir(temp_agents_dir_with_a2a)
|
||||
monkeypatch.chdir(temp_agents_dir_with_a2a)
|
||||
|
||||
try:
|
||||
app = get_fast_api_app(
|
||||
agents_dir=".",
|
||||
web=True,
|
||||
session_service_uri="",
|
||||
artifact_service_uri="",
|
||||
memory_service_uri="",
|
||||
allow_origins=["*"],
|
||||
a2a=True,
|
||||
host="127.0.0.1",
|
||||
port=8000,
|
||||
)
|
||||
app = get_fast_api_app(
|
||||
agents_dir=".",
|
||||
web=True,
|
||||
session_service_uri="",
|
||||
artifact_service_uri="",
|
||||
memory_service_uri="",
|
||||
allow_origins=["*"],
|
||||
a2a=True,
|
||||
host="127.0.0.1",
|
||||
port=8000,
|
||||
)
|
||||
|
||||
client = TestClient(app)
|
||||
yield client
|
||||
finally:
|
||||
os.chdir(original_cwd)
|
||||
client = TestClient(app)
|
||||
yield client
|
||||
|
||||
|
||||
#################################################
|
||||
@@ -1406,6 +1402,86 @@ def test_a2a_agent_discovery(test_app_with_a2a):
|
||||
logger.info("A2A agent discovery test passed")
|
||||
|
||||
|
||||
def test_a2a_request_handler_uses_push_config_store(
|
||||
mock_session_service,
|
||||
mock_artifact_service,
|
||||
mock_memory_service,
|
||||
mock_agent_loader,
|
||||
mock_eval_sets_manager,
|
||||
mock_eval_set_results_manager,
|
||||
temp_agents_dir_with_a2a,
|
||||
monkeypatch,
|
||||
):
|
||||
"""Test A2A request handler gets push config store when supported."""
|
||||
with (
|
||||
patch("signal.signal", return_value=None),
|
||||
patch(
|
||||
"google.adk.cli.fast_api.create_session_service_from_options",
|
||||
return_value=mock_session_service,
|
||||
),
|
||||
patch(
|
||||
"google.adk.cli.fast_api.create_artifact_service_from_options",
|
||||
return_value=mock_artifact_service,
|
||||
),
|
||||
patch(
|
||||
"google.adk.cli.fast_api.create_memory_service_from_options",
|
||||
return_value=mock_memory_service,
|
||||
),
|
||||
patch(
|
||||
"google.adk.cli.fast_api.AgentLoader",
|
||||
return_value=mock_agent_loader,
|
||||
),
|
||||
patch(
|
||||
"google.adk.cli.fast_api.LocalEvalSetsManager",
|
||||
return_value=mock_eval_sets_manager,
|
||||
),
|
||||
patch(
|
||||
"google.adk.cli.fast_api.LocalEvalSetResultsManager",
|
||||
return_value=mock_eval_set_results_manager,
|
||||
),
|
||||
patch("a2a.server.tasks.InMemoryTaskStore") as mock_task_store,
|
||||
patch(
|
||||
"a2a.server.tasks.InMemoryPushNotificationConfigStore"
|
||||
) as mock_push_config_store_class,
|
||||
patch(
|
||||
"google.adk.a2a.executor.a2a_agent_executor.A2aAgentExecutor"
|
||||
) as mock_executor,
|
||||
patch(
|
||||
"a2a.server.request_handlers.DefaultRequestHandler"
|
||||
) as mock_handler,
|
||||
patch("a2a.server.apps.A2AStarletteApplication") as mock_a2a_app,
|
||||
):
|
||||
mock_task_store_instance = MagicMock()
|
||||
mock_task_store.return_value = mock_task_store_instance
|
||||
mock_push_config_store = MagicMock()
|
||||
mock_push_config_store_class.return_value = mock_push_config_store
|
||||
mock_executor_instance = MagicMock()
|
||||
mock_executor.return_value = mock_executor_instance
|
||||
mock_handler.return_value = MagicMock()
|
||||
mock_a2a_app_instance = MagicMock()
|
||||
mock_a2a_app_instance.routes.return_value = []
|
||||
mock_a2a_app.return_value = mock_a2a_app_instance
|
||||
|
||||
monkeypatch.chdir(temp_agents_dir_with_a2a)
|
||||
_ = get_fast_api_app(
|
||||
agents_dir=".",
|
||||
web=True,
|
||||
session_service_uri="",
|
||||
artifact_service_uri="",
|
||||
memory_service_uri="",
|
||||
allow_origins=["*"],
|
||||
a2a=True,
|
||||
host="127.0.0.1",
|
||||
port=8000,
|
||||
)
|
||||
|
||||
mock_handler.assert_called_once_with(
|
||||
agent_executor=mock_executor_instance,
|
||||
push_config_store=mock_push_config_store,
|
||||
task_store=mock_task_store_instance,
|
||||
)
|
||||
|
||||
|
||||
def test_a2a_disabled_by_default(test_app):
|
||||
"""Test that A2A functionality is disabled by default."""
|
||||
# The regular test_app fixture has a2a=False
|
||||
|
||||
Reference in New Issue
Block a user