From 8e5f36126498f751171bb2639c7f5a9e7dca2558 Mon Sep 17 00:00:00 2001 From: Google Team Member Date: Mon, 29 Sep 2025 11:42:06 -0700 Subject: [PATCH] fix: Update remote_a2a_agent to better handle streaming events and avoid duplicate responses Currently, the A2A Task -> ADK event conversion is producing the same events on the last two update events (the last is a status update marking the task complete) The change here based on A2AClientEvent(task, update): - if the update == None: handle the non-streaming task case and also streaming case for the initial task creation event - if the update = TaskStatusUpdateEvent AND a message is set: emit an event with that message - if a task status update AND no message is set: don't emit event (for example, the final status update) - if the update is ArtifactUpdateEvent and it's final artifact: emit the event PiperOrigin-RevId: 812878869 --- src/google/adk/agents/remote_a2a_agent.py | 41 ++- .../unittests/agents/test_remote_a2a_agent.py | 236 +++++++++++++++++- 2 files changed, 266 insertions(+), 11 deletions(-) diff --git a/src/google/adk/agents/remote_a2a_agent.py b/src/google/adk/agents/remote_a2a_agent.py index c073223f..8efeb103 100644 --- a/src/google/adk/agents/remote_a2a_agent.py +++ b/src/google/adk/agents/remote_a2a_agent.py @@ -36,6 +36,8 @@ try: from a2a.types import Message as A2AMessage from a2a.types import Part as A2APart from a2a.types import Role + from a2a.types import TaskArtifactUpdateEvent as A2ATaskArtifactUpdateEvent + from a2a.types import TaskStatusUpdateEvent as A2ATaskStatusUpdateEvent from a2a.types import TransportProtocol as A2ATransport except ImportError as e: import sys @@ -393,7 +395,7 @@ class RemoteA2aAgent(BaseAgent): async def _handle_a2a_response( self, a2a_response: A2AClientEvent | A2AMessage, ctx: InvocationContext - ) -> Event: + ) -> Optional[Event]: """Handle A2A response and convert to Event. Args: @@ -401,14 +403,37 @@ class RemoteA2aAgent(BaseAgent): ctx: The invocation context Returns: - Event object representing the response + Event object representing the response, or None if no event should be + emitted. """ try: if isinstance(a2a_response, tuple): - # ClientEvent is a tuple of the absolute Task state and the last update. - # We only need the Task state. - task = a2a_response[0] - event = convert_a2a_task_to_event(task, self.name, ctx) + task, update = a2a_response + if update is None: + # This is the initial response for a streaming task or the complete + # response for a non-streaming task, which is the full task state. + # We process this to get the initial message. + event = convert_a2a_task_to_event(task, self.name, ctx) + elif isinstance(update, A2ATaskStatusUpdateEvent) and update.message: + # This is a streaming task status update with a message. + event = convert_a2a_message_to_event(update.message, self.name, ctx) + elif isinstance(update, A2ATaskArtifactUpdateEvent) and ( + not update.append or update.last_chunk + ): + # This is a streaming task artifact update. + # We only handle full artifact updates and ignore partial updates. + # Note: Depends on the server implementation, there is no clear + # definition of what a partial update is currently. We use the two + # signals: + # 1. append: True for partial updates, False for full updates. + # 2. last_chunk: True for full updates, False for partial updates. + event = convert_a2a_task_to_event(task, self.name, ctx) + else: + # This is a streaming update without a message (e.g. status change) + # or an partial artifact update. We don't emit an event for these + # for now. + return None + event.custom_metadata = event.custom_metadata or {} event.custom_metadata[A2A_METADATA_PREFIX + "task_id"] = task.id if task.context_id: @@ -416,7 +441,7 @@ class RemoteA2aAgent(BaseAgent): task.context_id ) - # Otherwise, it's a regular A2AMessage. + # Otherwise, it's a regular A2AMessage for non-streaming responses. elif isinstance(a2a_response, A2AMessage): event = convert_a2a_message_to_event(a2a_response, self.name, ctx) event.custom_metadata = event.custom_metadata or {} @@ -492,6 +517,8 @@ class RemoteA2aAgent(BaseAgent): logger.debug(build_a2a_response_log(a2a_response)) event = await self._handle_a2a_response(a2a_response, ctx) + if not event: + continue # Add metadata about the request and response event.custom_metadata = event.custom_metadata or {} diff --git a/tests/unittests/agents/test_remote_a2a_agent.py b/tests/unittests/agents/test_remote_a2a_agent.py index 2afc65f5..65ccdd6b 100644 --- a/tests/unittests/agents/test_remote_a2a_agent.py +++ b/tests/unittests/agents/test_remote_a2a_agent.py @@ -39,9 +39,12 @@ try: from a2a.types import AgentCapabilities from a2a.types import AgentCard from a2a.types import AgentSkill + from a2a.types import Artifact from a2a.types import Message as A2AMessage from a2a.types import SendMessageSuccessResponse from a2a.types import Task as A2ATask + from a2a.types import TaskArtifactUpdateEvent + from a2a.types import TaskStatusUpdateEvent from google.adk.agents.invocation_context import InvocationContext from google.adk.agents.remote_a2a_agent import A2A_METADATA_PREFIX from google.adk.agents.remote_a2a_agent import AgentCardResolutionError @@ -60,6 +63,9 @@ except ImportError as e: A2AMessage = DummyTypes() SendMessageSuccessResponse = DummyTypes() A2ATask = DummyTypes() + TaskStatusUpdateEvent = DummyTypes() + Artifact = DummyTypes() + TaskArtifactUpdateEvent = DummyTypes() InvocationContext = DummyTypes() RemoteA2aAgent = DummyTypes() AgentCardResolutionError = Exception @@ -685,8 +691,8 @@ class TestRemoteA2aAgentMessageHandling: assert A2A_METADATA_PREFIX + "context_id" in result.custom_metadata @pytest.mark.asyncio - async def test_handle_a2a_response_success_with_task(self): - """Test successful A2A response handling with task.""" + async def test_handle_a2a_response_with_task_and_no_update(self): + """Test successful A2A response handling with task and no update.""" mock_a2a_task = Mock(spec=A2ATask) mock_a2a_task.id = "task-123" mock_a2a_task.context_id = "context-123" @@ -718,6 +724,116 @@ class TestRemoteA2aAgentMessageHandling: assert A2A_METADATA_PREFIX + "task_id" in result.custom_metadata assert A2A_METADATA_PREFIX + "context_id" in result.custom_metadata + @pytest.mark.asyncio + async def test_handle_a2a_response_with_task_status_update_with_message(self): + """Test handling of a task status update with a message.""" + mock_a2a_task = Mock(spec=A2ATask) + mock_a2a_task.id = "task-123" + mock_a2a_task.context_id = "context-123" + + mock_a2a_message = Mock(spec=A2AMessage) + mock_update = Mock(spec=TaskStatusUpdateEvent) + mock_update.message = mock_a2a_message + mock_update.status = "COMPLETED" + + # Create a proper Event mock that can handle custom_metadata + mock_event = Event( + author=self.agent.name, + invocation_id=self.mock_context.invocation_id, + branch=self.mock_context.branch, + ) + + with patch( + "google.adk.agents.remote_a2a_agent.convert_a2a_message_to_event" + ) as mock_convert: + mock_convert.return_value = mock_event + + result = await self.agent._handle_a2a_response( + (mock_a2a_task, mock_update), self.mock_context + ) + + assert result == mock_event + mock_convert.assert_called_once_with( + mock_a2a_message, + self.agent.name, + self.mock_context, + ) + # Check that metadata was added + assert result.custom_metadata is not None + assert A2A_METADATA_PREFIX + "task_id" in result.custom_metadata + assert A2A_METADATA_PREFIX + "context_id" in result.custom_metadata + + @pytest.mark.asyncio + async def test_handle_a2a_response_with_task_status_update_no_message(self): + """Test handling of a task status update with no message.""" + mock_a2a_task = Mock(spec=A2ATask) + mock_a2a_task.id = "task-123" + + mock_update = Mock(spec=TaskStatusUpdateEvent) + mock_update.message = None + mock_update.status = "COMPLETED" + + result = await self.agent._handle_a2a_response( + (mock_a2a_task, mock_update), self.mock_context + ) + + assert result is None + + @pytest.mark.asyncio + async def test_handle_a2a_response_with_artifact_update(self): + """Test successful A2A response handling with artifact update.""" + mock_a2a_task = Mock(spec=A2ATask) + mock_a2a_task.id = "task-123" + mock_a2a_task.context_id = "context-123" + + mock_artifact = Mock(spec=Artifact) + mock_update = Mock(spec=TaskArtifactUpdateEvent) + mock_update.artifact = mock_artifact + mock_update.append = False + mock_update.last_chunk = True + + # Create a proper Event mock that can handle custom_metadata + mock_event = Event( + author=self.agent.name, + invocation_id=self.mock_context.invocation_id, + branch=self.mock_context.branch, + ) + + with patch( + "google.adk.agents.remote_a2a_agent.convert_a2a_task_to_event" + ) as mock_convert: + mock_convert.return_value = mock_event + + result = await self.agent._handle_a2a_response( + (mock_a2a_task, mock_update), self.mock_context + ) + + assert result == mock_event + mock_convert.assert_called_once_with( + mock_a2a_task, self.agent.name, self.mock_context + ) + # Check that metadata was added + assert result.custom_metadata is not None + assert A2A_METADATA_PREFIX + "task_id" in result.custom_metadata + assert A2A_METADATA_PREFIX + "context_id" in result.custom_metadata + + @pytest.mark.asyncio + async def test_handle_a2a_response_with_partial_artifact_update(self): + """Test that partial artifact updates are ignored.""" + mock_a2a_task = Mock(spec=A2ATask) + mock_a2a_task.id = "task-123" + + mock_update = Mock(spec=TaskArtifactUpdateEvent) + mock_update.artifact = Mock(spec=Artifact) + mock_update.append = True + mock_update.last_chunk = False + + result = await self.agent._handle_a2a_response( + (mock_a2a_task, mock_update), self.mock_context + ) + + assert result is None + class TestRemoteA2aAgentMessageHandlingFromFactory: """Test message handling functionality.""" @@ -865,8 +981,8 @@ class TestRemoteA2aAgentMessageHandlingFromFactory: assert A2A_METADATA_PREFIX + "context_id" in result.custom_metadata @pytest.mark.asyncio - async def test_handle_a2a_response_success_with_task(self): - """Test successful A2A response handling with task.""" + async def test_handle_a2a_response_with_task_and_no_update(self): + """Test successful A2A response handling with task and no update.""" mock_a2a_task = Mock(spec=A2ATask) mock_a2a_task.id = "task-123" mock_a2a_task.context_id = "context-123" @@ -896,6 +1012,116 @@ class TestRemoteA2aAgentMessageHandlingFromFactory: assert A2A_METADATA_PREFIX + "task_id" in result.custom_metadata assert A2A_METADATA_PREFIX + "context_id" in result.custom_metadata + @pytest.mark.asyncio + async def test_handle_a2a_response_with_task_status_update_with_message(self): + """Test handling of a task status update with a message.""" + mock_a2a_task = Mock(spec=A2ATask) + mock_a2a_task.id = "task-123" + mock_a2a_task.context_id = "context-123" + + mock_a2a_message = Mock(spec=A2AMessage) + mock_update = Mock(spec=TaskStatusUpdateEvent) + mock_update.message = mock_a2a_message + mock_update.status = "COMPLETED" + + # Create a proper Event mock that can handle custom_metadata + mock_event = Event( + author=self.agent.name, + invocation_id=self.mock_context.invocation_id, + branch=self.mock_context.branch, + ) + + with patch( + "google.adk.agents.remote_a2a_agent.convert_a2a_message_to_event" + ) as mock_convert: + mock_convert.return_value = mock_event + + result = await self.agent._handle_a2a_response( + (mock_a2a_task, mock_update), self.mock_context + ) + + assert result == mock_event + mock_convert.assert_called_once_with( + mock_a2a_message, + self.agent.name, + self.mock_context, + ) + # Check that metadata was added + assert result.custom_metadata is not None + assert A2A_METADATA_PREFIX + "task_id" in result.custom_metadata + assert A2A_METADATA_PREFIX + "context_id" in result.custom_metadata + + @pytest.mark.asyncio + async def test_handle_a2a_response_with_task_status_update_no_message(self): + """Test handling of a task status update with no message.""" + mock_a2a_task = Mock(spec=A2ATask) + mock_a2a_task.id = "task-123" + + mock_update = Mock(spec=TaskStatusUpdateEvent) + mock_update.message = None + mock_update.status = "COMPLETED" + + result = await self.agent._handle_a2a_response( + (mock_a2a_task, mock_update), self.mock_context + ) + + assert result is None + + @pytest.mark.asyncio + async def test_handle_a2a_response_with_artifact_update(self): + """Test successful A2A response handling with artifact update.""" + mock_a2a_task = Mock(spec=A2ATask) + mock_a2a_task.id = "task-123" + mock_a2a_task.context_id = "context-123" + + mock_artifact = Mock(spec=Artifact) + mock_update = Mock(spec=TaskArtifactUpdateEvent) + mock_update.artifact = mock_artifact + mock_update.append = False + mock_update.last_chunk = True + + # Create a proper Event mock that can handle custom_metadata + mock_event = Event( + author=self.agent.name, + invocation_id=self.mock_context.invocation_id, + branch=self.mock_context.branch, + ) + + with patch( + "google.adk.agents.remote_a2a_agent.convert_a2a_task_to_event" + ) as mock_convert: + mock_convert.return_value = mock_event + + result = await self.agent._handle_a2a_response( + (mock_a2a_task, mock_update), self.mock_context + ) + + assert result == mock_event + mock_convert.assert_called_once_with( + mock_a2a_task, self.agent.name, self.mock_context + ) + # Check that metadata was added + assert result.custom_metadata is not None + assert A2A_METADATA_PREFIX + "task_id" in result.custom_metadata + assert A2A_METADATA_PREFIX + "context_id" in result.custom_metadata + + @pytest.mark.asyncio + async def test_handle_a2a_response_with_partial_artifact_update(self): + """Test that partial artifact updates are ignored.""" + mock_a2a_task = Mock(spec=A2ATask) + mock_a2a_task.id = "task-123" + + mock_update = Mock(spec=TaskArtifactUpdateEvent) + mock_update.artifact = Mock(spec=Artifact) + mock_update.append = True + mock_update.last_chunk = False + + result = await self.agent._handle_a2a_response( + (mock_a2a_task, mock_update), self.mock_context + ) + + assert result is None + class TestRemoteA2aAgentExecution: """Test agent execution functionality.""" @@ -1019,6 +1245,7 @@ class TestRemoteA2aAgentExecution: # Add model_dump to mock_response for metadata mock_response.model_dump.return_value = {"test": "response"} + # Execute events = [] async for event in self.agent._run_async_impl( self.mock_context @@ -1211,6 +1438,7 @@ class TestRemoteA2aAgentExecutionFromFactory: "test": "response" } + # Execute events = [] async for event in self.agent._run_async_impl( self.mock_context