feat: add a2a_request_meta_provider to RemoteAgent init

The change adds an extension point for controlling which request metadata gets attached to A2A requests made by a RemoteAgent.
Instead of taking metadata from custom_metadata of session events users can construct payloads using a2a_request_meta_provider.

request_meta 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: 831506364
This commit is contained in:
Google Team Member
2025-11-12 13:04:06 -08:00
committed by Copybara-Service
parent 5d5708b2ab
commit d12468ee5a
2 changed files with 131 additions and 63 deletions
+109 -48
View File
@@ -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,8 +593,7 @@ 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 + "metadata": {"foo": "bar"},
A2A_METADATA_PREFIX + "task_id": "task-123"
}
# Mock latest event with function response - set proper author
@@ -615,17 +614,13 @@ class TestRemoteA2aAgentMessageHandling:
mock_a2a_message.task_id = None # Will be set by the method
mock_convert.return_value = mock_a2a_message
result, metadata = (
self.agent._create_a2a_request_for_user_function_response(
self.mock_context
)
result = 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."""
@@ -649,14 +644,13 @@ class TestRemoteA2aAgentMessageHandling:
mock_a2a_part = Mock()
self.mock_genai_part_converter.return_value = mock_a2a_part
parts, context_id, metadata = (
self.agent._construct_message_parts_from_session(self.mock_context)
parts, context_id = self.agent._construct_message_parts_from_session(
self.mock_context
)
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."""
@@ -684,25 +678,23 @@ class TestRemoteA2aAgentMessageHandling:
mock_a2a_part2,
]
parts, context_id, metadata = (
self.agent._construct_message_parts_from_session(self.mock_context)
parts, context_id = self.agent._construct_message_parts_from_session(
self.mock_context
)
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 = []
parts, context_id, metadata = (
self.agent._construct_message_parts_from_session(self.mock_context)
parts, context_id = self.agent._construct_message_parts_from_session(
self.mock_context
)
assert parts == []
assert context_id is None
assert metadata is None
@pytest.mark.asyncio
async def test_handle_a2a_response_success_with_message(self):
@@ -824,14 +816,13 @@ class TestRemoteA2aAgentMessageHandling:
self.mock_genai_part_converter.side_effect = mock_converter
parts, context_id, metadata = (
self.agent._construct_message_parts_from_session(self.mock_context)
parts, context_id = self.agent._construct_message_parts_from_session(
self.mock_context
)
# Verify the parts are in correct order
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"
@@ -1074,7 +1065,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
)
@@ -1085,8 +1076,7 @@ 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 + "metadata": {"foo": "bar"},
A2A_METADATA_PREFIX + "task_id": "task-123"
}
# Mock latest event with function response - set proper author
@@ -1107,17 +1097,13 @@ class TestRemoteA2aAgentMessageHandlingFromFactory:
mock_a2a_message.task_id = None # Will be set by the method
mock_convert.return_value = mock_a2a_message
result, metadata = (
self.agent._create_a2a_request_for_user_function_response(
self.mock_context
)
result = 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."""
@@ -1144,26 +1130,24 @@ class TestRemoteA2aAgentMessageHandlingFromFactory:
mock_a2a_part = Mock()
mock_convert_part.return_value = mock_a2a_part
parts, context_id, metadata = (
self.agent._construct_message_parts_from_session(self.mock_context)
parts, context_id = self.agent._construct_message_parts_from_session(
self.mock_context
)
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 = []
parts, context_id, metadata = (
self.agent._construct_message_parts_from_session(self.mock_context)
parts, context_id = self.agent._construct_message_parts_from_session(
self.mock_context
)
assert parts == []
assert context_id is None
assert metadata is None
@pytest.mark.asyncio
async def test_handle_a2a_response_success_with_message(self):
@@ -1482,7 +1466,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, None)
mock_create_func.return_value = None
with patch.object(
self.agent, "_construct_message_parts_from_session"
@@ -1490,7 +1474,6 @@ class TestRemoteA2aAgentExecution:
mock_construct.return_value = (
[],
None,
None,
) # Tuple with empty parts and no context_id
events = []
@@ -1508,7 +1491,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, None)
mock_create_func.return_value = None
with patch.object(
self.agent, "_construct_message_parts_from_session"
@@ -1521,7 +1504,6 @@ class TestRemoteA2aAgentExecution:
mock_construct.return_value = (
[mock_a2a_part],
"context-123",
{"foo": "bar"},
) # Tuple with parts and context_id
# Mock A2A client
@@ -1582,7 +1564,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, None
mock_create_func.return_value = None
with patch.object(
self.agent, "_construct_message_parts_from_session"
@@ -1594,7 +1576,6 @@ 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
@@ -1631,6 +1612,89 @@ class TestRemoteA2aAgentExecution:
async for _ in self.agent._run_live_impl(self.mock_context):
pass
@pytest.mark.asyncio
async def test_run_async_impl_with_meta_provider(self):
"""Test _run_async_impl with a2a_request_meta_provider."""
mock_meta_provider = Mock()
request_metadata = {"custom_meta": "value"}
mock_meta_provider.return_value = request_metadata
agent = RemoteA2aAgent(
name="test_agent",
agent_card=self.agent_card,
genai_part_converter=self.mock_genai_part_converter,
a2a_part_converter=self.mock_a2a_part_converter,
a2a_request_meta_provider=mock_meta_provider,
)
with patch.object(agent, "_ensure_resolved"):
with patch.object(
agent, "_create_a2a_request_for_user_function_response"
) as mock_create_func:
mock_create_func.return_value = None
with patch.object(
agent, "_construct_message_parts_from_session"
) as mock_construct:
# Create proper A2A part mocks
from a2a.client import Client as A2AClient
from a2a.types import TextPart
mock_a2a_part = Mock(spec=TextPart)
mock_construct.return_value = (
[mock_a2a_part],
"context-123",
) # Tuple with parts and context_id
# Mock A2A client
mock_a2a_client = create_autospec(spec=A2AClient, instance=True)
mock_response = Mock()
mock_send_message = AsyncMock()
mock_send_message.__aiter__.return_value = [mock_response]
mock_a2a_client.send_message.return_value = mock_send_message
agent._a2a_client = mock_a2a_client
mock_event = Event(
author=agent.name,
invocation_id=self.mock_context.invocation_id,
branch=self.mock_context.branch,
)
with patch.object(agent, "_handle_a2a_response") as mock_handle:
mock_handle.return_value = mock_event
# Mock the logging functions to avoid iteration issues
with patch(
"google.adk.agents.remote_a2a_agent.build_a2a_request_log"
) as mock_req_log:
with patch(
"google.adk.agents.remote_a2a_agent.build_a2a_response_log"
) as mock_resp_log:
mock_req_log.return_value = "Mock request log"
mock_resp_log.return_value = "Mock response log"
# Mock the A2AMessage constructor
with patch(
"google.adk.agents.remote_a2a_agent.A2AMessage"
) as mock_message_class:
mock_message = Mock(spec=A2AMessage)
mock_message_class.return_value = mock_message
# Add model_dump to mock_response for metadata
mock_response.model_dump.return_value = {"test": "response"}
# Execute
events = []
async for event in agent._run_async_impl(self.mock_context):
events.append(event)
assert len(events) == 1
mock_meta_provider.assert_called_once_with(
self.mock_context, mock_message
)
mock_a2a_client.send_message.assert_called_once_with(
request=mock_message,
request_metadata=request_metadata,
)
class TestRemoteA2aAgentExecutionFromFactory:
"""Test agent execution functionality."""
@@ -1676,7 +1740,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, None
mock_create_func.return_value = None
with patch.object(
self.agent, "_construct_message_parts_from_session"
@@ -1684,7 +1748,6 @@ class TestRemoteA2aAgentExecutionFromFactory:
mock_construct.return_value = (
[],
None,
None,
) # Tuple with empty parts and no context_id
events = []
@@ -1702,7 +1765,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, None
mock_create_func.return_value = None
with patch.object(
self.agent, "_construct_message_parts_from_session"
@@ -1715,7 +1778,6 @@ class TestRemoteA2aAgentExecutionFromFactory:
mock_construct.return_value = (
[mock_a2a_part],
"context-123",
None,
) # Tuple with parts and context_id
# Mock A2A client
@@ -1778,7 +1840,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, None
mock_create_func.return_value = None
with patch.object(
self.agent, "_construct_message_parts_from_session"
@@ -1790,7 +1852,6 @@ 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