From d85a301039003210b61a572ca3d98183d717768a Mon Sep 17 00:00:00 2001 From: Google Team Member Date: Tue, 11 Nov 2025 16:49:11 -0800 Subject: [PATCH] feat: Add support for passing request metadata in RemoteA2AAgent This change updates `RemoteA2AAgent` to extract and forward custom metadata from session events to the `a2a-sdk`'s `send_message` method. The metadata is looked for under the key `A2A_METADATA_PREFIX + "metadata"` within the `custom_metadata` of the relevant session events. The `a2a-sdk` dependency is also updated to a version that supports this feature. This feature was added in v0.3.11 of the a2a-sdk library: https://github.com/a2aproject/a2a-python/releases/tag/v0.3.11 PiperOrigin-RevId: 831120978 --- pyproject.toml | 2 +- src/google/adk/agents/remote_a2a_agent.py | 30 +++-- .../unittests/agents/test_remote_a2a_agent.py | 109 ++++++++++-------- 3 files changed, 84 insertions(+), 57 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 2414cbd1..af287984 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -92,7 +92,7 @@ dev = [ a2a = [ # go/keep-sorted start - "a2a-sdk>=0.3.4,<0.4.0;python_version>='3.10'", + "a2a-sdk>=0.3.11,<0.4.0;python_version>='3.10'", # go/keep-sorted end ] diff --git a/src/google/adk/agents/remote_a2a_agent.py b/src/google/adk/agents/remote_a2a_agent.py index 06713574..d15fbb5c 100644 --- a/src/google/adk/agents/remote_a2a_agent.py +++ b/src/google/adk/agents/remote_a2a_agent.py @@ -318,7 +318,7 @@ class RemoteA2aAgent(BaseAgent): def _create_a2a_request_for_user_function_response( self, ctx: InvocationContext - ) -> Optional[A2AMessage]: + ) -> tuple[Optional[A2AMessage], Optional[dict[str, Any]]]: """Create A2A request for user function response if applicable. Args: @@ -328,34 +328,38 @@ class RemoteA2aAgent(BaseAgent): SendMessageRequest if function response found, None otherwise """ if not ctx.session.events or ctx.session.events[-1].author != "user": - return None + return None, None function_call_event = find_matching_function_call(ctx.session.events) if not function_call_event: - return None + return None, None a2a_message = convert_event_to_a2a_message( ctx.session.events[-1], ctx, Role.user, self._genai_part_converter ) + message_metadata = None if function_call_event.custom_metadata: metadata = function_call_event.custom_metadata a2a_message.task_id = metadata.get(A2A_METADATA_PREFIX + "task_id") a2a_message.context_id = metadata.get(A2A_METADATA_PREFIX + "context_id") + message_metadata = metadata.get(A2A_METADATA_PREFIX + "metadata") - return a2a_message + return a2a_message, message_metadata def _construct_message_parts_from_session( self, ctx: InvocationContext - ) -> tuple[list[A2APart], Optional[str]]: + ) -> tuple[list[A2APart], Optional[str], Optional[dict[str, Any]]]: """Construct A2A message parts from session events. Args: ctx: The invocation context Returns: - List of A2A parts extracted from session events, context ID + List of A2A parts extracted from session events, context ID, + request metadata """ message_parts: list[A2APart] = [] context_id = None + request_metadata = None events_to_process = [] for event in reversed(ctx.session.events): @@ -365,6 +369,7 @@ class RemoteA2aAgent(BaseAgent): if event.custom_metadata: metadata = event.custom_metadata context_id = metadata.get(A2A_METADATA_PREFIX + "context_id") + request_metadata = metadata.get(A2A_METADATA_PREFIX + "metadata") break events_to_process.append(event) @@ -385,7 +390,7 @@ class RemoteA2aAgent(BaseAgent): else: logger.warning("Failed to convert part to A2A format: %s", part) - return message_parts, context_id + return message_parts, context_id, request_metadata async def _handle_a2a_response( self, a2a_response: A2AClientEvent | A2AMessage, ctx: InvocationContext @@ -493,10 +498,12 @@ class RemoteA2aAgent(BaseAgent): return # Create A2A request for function response or regular message - a2a_request = self._create_a2a_request_for_user_function_response(ctx) + a2a_request, request_metadata = ( + self._create_a2a_request_for_user_function_response(ctx) + ) if not a2a_request: - message_parts, context_id = self._construct_message_parts_from_session( - ctx + message_parts, context_id, request_metadata = ( + self._construct_message_parts_from_session(ctx) ) if not message_parts: @@ -522,7 +529,8 @@ class RemoteA2aAgent(BaseAgent): try: async for a2a_response in self._a2a_client.send_message( - request=a2a_request + request=a2a_request, + request_metadata=request_metadata, ): logger.debug(build_a2a_response_log(a2a_response)) diff --git a/tests/unittests/agents/test_remote_a2a_agent.py b/tests/unittests/agents/test_remote_a2a_agent.py index c9fece49..6526b63e 100644 --- a/tests/unittests/agents/test_remote_a2a_agent.py +++ b/tests/unittests/agents/test_remote_a2a_agent.py @@ -582,7 +582,7 @@ class TestRemoteA2aAgentMessageHandling: ) as mock_find: mock_find.return_value = None - result = self.agent._create_a2a_request_for_user_function_response( + result, _ = self.agent._create_a2a_request_for_user_function_response( self.mock_context ) @@ -593,7 +593,8 @@ class TestRemoteA2aAgentMessageHandling: # Mock function call event mock_function_event = Mock() mock_function_event.custom_metadata = { - A2A_METADATA_PREFIX + "task_id": "task-123" + A2A_METADATA_PREFIX + "task_id": "task-123", + A2A_METADATA_PREFIX + "metadata": {"foo": "bar"}, } # Mock latest event with function response - set proper author @@ -614,13 +615,17 @@ class TestRemoteA2aAgentMessageHandling: mock_a2a_message.task_id = None # Will be set by the method mock_convert.return_value = mock_a2a_message - result = self.agent._create_a2a_request_for_user_function_response( - self.mock_context + result, metadata = ( + self.agent._create_a2a_request_for_user_function_response( + self.mock_context + ) ) assert result is not None assert result == mock_a2a_message assert mock_a2a_message.task_id == "task-123" + assert metadata is not None + assert metadata == {"foo": "bar"} def test_construct_message_parts_from_session_success(self): """Test successful message parts construction from session.""" @@ -644,14 +649,14 @@ class TestRemoteA2aAgentMessageHandling: mock_a2a_part = Mock() self.mock_genai_part_converter.return_value = mock_a2a_part - result = self.agent._construct_message_parts_from_session( - self.mock_context + parts, context_id, metadata = ( + self.agent._construct_message_parts_from_session(self.mock_context) ) - assert len(result) == 2 # Returns tuple of (parts, context_id) - assert len(result[0]) == 1 # parts list - assert result[0][0] == mock_a2a_part - assert result[1] is None # context_id + assert len(parts) == 1 + assert parts[0] == mock_a2a_part + assert context_id is None + assert metadata is None def test_construct_message_parts_from_session_success_multiple_parts(self): """Test successful message parts construction from session.""" @@ -679,24 +684,25 @@ class TestRemoteA2aAgentMessageHandling: mock_a2a_part2, ] - result = self.agent._construct_message_parts_from_session( - self.mock_context + parts, context_id, metadata = ( + self.agent._construct_message_parts_from_session(self.mock_context) ) - assert len(result) == 2 # Returns tuple of (parts, context_id) - assert len(result[0]) == 2 # parts list - assert result[0] == [mock_a2a_part1, mock_a2a_part2] - assert result[1] is None # context_id + assert parts == [mock_a2a_part1, mock_a2a_part2] + assert context_id is None + assert metadata is None def test_construct_message_parts_from_session_empty_events(self): """Test message parts construction with empty events.""" self.mock_session.events = [] - result = self.agent._construct_message_parts_from_session(self.mock_context) + parts, context_id, metadata = ( + self.agent._construct_message_parts_from_session(self.mock_context) + ) - assert len(result) == 2 # Returns tuple of (parts, context_id) - assert result[0] == [] # empty parts list - assert result[1] is None # context_id + assert parts == [] + assert context_id is None + assert metadata is None @pytest.mark.asyncio async def test_handle_a2a_response_success_with_message(self): @@ -818,14 +824,14 @@ class TestRemoteA2aAgentMessageHandling: self.mock_genai_part_converter.side_effect = mock_converter - result = self.agent._construct_message_parts_from_session( - self.mock_context + parts, context_id, metadata = ( + self.agent._construct_message_parts_from_session(self.mock_context) ) # Verify the parts are in correct order - assert len(result) == 2 # Returns tuple of (parts, context_id) - assert len(result[0]) == 3 # 1 user part + 2 other agent parts - assert result[1] is None # context_id + assert len(parts) == 3 # 1 user part + 2 other agent parts + assert context_id is None + assert metadata is None # Verify order: user part, then "For context:", then agent message assert converted_parts[0].original_text == "User question" @@ -1068,7 +1074,7 @@ class TestRemoteA2aAgentMessageHandlingFromFactory: ) as mock_find: mock_find.return_value = None - result = self.agent._create_a2a_request_for_user_function_response( + result, _ = self.agent._create_a2a_request_for_user_function_response( self.mock_context ) @@ -1079,7 +1085,8 @@ class TestRemoteA2aAgentMessageHandlingFromFactory: # Mock function call event mock_function_event = Mock() mock_function_event.custom_metadata = { - A2A_METADATA_PREFIX + "task_id": "task-123" + A2A_METADATA_PREFIX + "task_id": "task-123", + A2A_METADATA_PREFIX + "metadata": {"foo": "bar"}, } # Mock latest event with function response - set proper author @@ -1100,13 +1107,17 @@ class TestRemoteA2aAgentMessageHandlingFromFactory: mock_a2a_message.task_id = None # Will be set by the method mock_convert.return_value = mock_a2a_message - result = self.agent._create_a2a_request_for_user_function_response( - self.mock_context + result, metadata = ( + self.agent._create_a2a_request_for_user_function_response( + self.mock_context + ) ) assert result is not None assert result == mock_a2a_message assert mock_a2a_message.task_id == "task-123" + assert metadata is not None + assert metadata == {"foo": "bar"} def test_construct_message_parts_from_session_success(self): """Test successful message parts construction from session.""" @@ -1133,24 +1144,26 @@ class TestRemoteA2aAgentMessageHandlingFromFactory: mock_a2a_part = Mock() mock_convert_part.return_value = mock_a2a_part - result = self.agent._construct_message_parts_from_session( - self.mock_context + parts, context_id, metadata = ( + self.agent._construct_message_parts_from_session(self.mock_context) ) - assert len(result) == 2 # Returns tuple of (parts, context_id) - assert len(result[0]) == 1 # parts list - assert result[0][0] == mock_a2a_part - assert result[1] is None # context_id + assert len(parts) == 1 + assert parts[0] == mock_a2a_part + assert context_id is None + assert metadata is None def test_construct_message_parts_from_session_empty_events(self): """Test message parts construction with empty events.""" self.mock_session.events = [] - result = self.agent._construct_message_parts_from_session(self.mock_context) + parts, context_id, metadata = ( + self.agent._construct_message_parts_from_session(self.mock_context) + ) - assert len(result) == 2 # Returns tuple of (parts, context_id) - assert result[0] == [] # empty parts list - assert result[1] is None # context_id + assert parts == [] + assert context_id is None + assert metadata is None @pytest.mark.asyncio async def test_handle_a2a_response_success_with_message(self): @@ -1469,7 +1482,7 @@ class TestRemoteA2aAgentExecution: with patch.object( self.agent, "_create_a2a_request_for_user_function_response" ) as mock_create_func: - mock_create_func.return_value = None + mock_create_func.return_value = (None, None) with patch.object( self.agent, "_construct_message_parts_from_session" @@ -1477,6 +1490,7 @@ class TestRemoteA2aAgentExecution: mock_construct.return_value = ( [], None, + None, ) # Tuple with empty parts and no context_id events = [] @@ -1494,7 +1508,7 @@ class TestRemoteA2aAgentExecution: with patch.object( self.agent, "_create_a2a_request_for_user_function_response" ) as mock_create_func: - mock_create_func.return_value = None + mock_create_func.return_value = (None, None) with patch.object( self.agent, "_construct_message_parts_from_session" @@ -1507,6 +1521,7 @@ class TestRemoteA2aAgentExecution: mock_construct.return_value = ( [mock_a2a_part], "context-123", + {"foo": "bar"}, ) # Tuple with parts and context_id # Mock A2A client @@ -1567,7 +1582,7 @@ class TestRemoteA2aAgentExecution: with patch.object( self.agent, "_create_a2a_request_for_user_function_response" ) as mock_create_func: - mock_create_func.return_value = None + mock_create_func.return_value = None, None with patch.object( self.agent, "_construct_message_parts_from_session" @@ -1579,6 +1594,7 @@ class TestRemoteA2aAgentExecution: mock_construct.return_value = ( [mock_a2a_part], "context-123", + {"foo": "bar"}, ) # Tuple with parts and context_id # Mock A2A client that throws an exception @@ -1660,7 +1676,7 @@ class TestRemoteA2aAgentExecutionFromFactory: with patch.object( self.agent, "_create_a2a_request_for_user_function_response" ) as mock_create_func: - mock_create_func.return_value = None + mock_create_func.return_value = None, None with patch.object( self.agent, "_construct_message_parts_from_session" @@ -1668,6 +1684,7 @@ class TestRemoteA2aAgentExecutionFromFactory: mock_construct.return_value = ( [], None, + None, ) # Tuple with empty parts and no context_id events = [] @@ -1685,7 +1702,7 @@ class TestRemoteA2aAgentExecutionFromFactory: with patch.object( self.agent, "_create_a2a_request_for_user_function_response" ) as mock_create_func: - mock_create_func.return_value = None + mock_create_func.return_value = None, None with patch.object( self.agent, "_construct_message_parts_from_session" @@ -1698,6 +1715,7 @@ class TestRemoteA2aAgentExecutionFromFactory: mock_construct.return_value = ( [mock_a2a_part], "context-123", + None, ) # Tuple with parts and context_id # Mock A2A client @@ -1760,7 +1778,7 @@ class TestRemoteA2aAgentExecutionFromFactory: with patch.object( self.agent, "_create_a2a_request_for_user_function_response" ) as mock_create_func: - mock_create_func.return_value = None + mock_create_func.return_value = None, None with patch.object( self.agent, "_construct_message_parts_from_session" @@ -1772,6 +1790,7 @@ class TestRemoteA2aAgentExecutionFromFactory: mock_construct.return_value = ( [mock_a2a_part], "context-123", + None, ) # Tuple with parts and context_id # Mock A2A client that throws an exception