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
+22 -15
View File
@@ -20,6 +20,7 @@ import logging
from pathlib import Path
from typing import Any
from typing import AsyncGenerator
from typing import Callable
from typing import Optional
from typing import Union
from urllib.parse import urlparse
@@ -125,12 +126,16 @@ class RemoteA2aAgent(BaseAgent):
self,
name: str,
agent_card: Union[AgentCard, str],
*,
description: str = "",
httpx_client: Optional[httpx.AsyncClient] = None,
timeout: float = DEFAULT_TIMEOUT,
genai_part_converter: GenAIPartToA2APartConverter = convert_genai_part_to_a2a_part,
a2a_part_converter: A2APartToGenAIPartConverter = convert_a2a_part_to_genai_part,
a2a_client_factory: Optional[A2AClientFactory] = None,
a2a_request_meta_provider: Optional[
Callable[[InvocationContext, A2AMessage], dict[str, Any]]
] = None,
**kwargs: Any,
) -> None:
"""Initialize RemoteA2aAgent.
@@ -144,6 +149,9 @@ class RemoteA2aAgent(BaseAgent):
timeout: HTTP timeout in seconds
a2a_client_factory: Optional A2AClientFactory object (will create own if
not provided)
a2a_request_meta_provider: Optional callable that takes InvocationContext
and A2AMessage and returns a metadata object to attach to the A2A
request.
**kwargs: Additional arguments passed to BaseAgent
Raises:
@@ -169,6 +177,7 @@ class RemoteA2aAgent(BaseAgent):
self._genai_part_converter = genai_part_converter
self._a2a_part_converter = a2a_part_converter
self._a2a_client_factory: Optional[A2AClientFactory] = a2a_client_factory
self._a2a_request_meta_provider = a2a_request_meta_provider
# Validate and store agent card reference
if isinstance(agent_card, AgentCard):
@@ -318,7 +327,7 @@ class RemoteA2aAgent(BaseAgent):
def _create_a2a_request_for_user_function_response(
self, ctx: InvocationContext
) -> tuple[Optional[A2AMessage], Optional[dict[str, Any]]]:
) -> Optional[A2AMessage]:
"""Create A2A request for user function response if applicable.
Args:
@@ -328,26 +337,24 @@ class RemoteA2aAgent(BaseAgent):
SendMessageRequest if function response found, None otherwise
"""
if not ctx.session.events or ctx.session.events[-1].author != "user":
return None, None
return None
function_call_event = find_matching_function_call(ctx.session.events)
if not function_call_event:
return None, None
return 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, message_metadata
return a2a_message
def _construct_message_parts_from_session(
self, ctx: InvocationContext
) -> tuple[list[A2APart], Optional[str], Optional[dict[str, Any]]]:
) -> tuple[list[A2APart], Optional[str]]:
"""Construct A2A message parts from session events.
Args:
@@ -359,7 +366,6 @@ class RemoteA2aAgent(BaseAgent):
"""
message_parts: list[A2APart] = []
context_id = None
request_metadata = None
events_to_process = []
for event in reversed(ctx.session.events):
@@ -369,7 +375,6 @@ 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)
@@ -390,7 +395,7 @@ class RemoteA2aAgent(BaseAgent):
else:
logger.warning("Failed to convert part to A2A format: %s", part)
return message_parts, context_id, request_metadata
return message_parts, context_id
async def _handle_a2a_response(
self, a2a_response: A2AClientEvent | A2AMessage, ctx: InvocationContext
@@ -498,12 +503,10 @@ class RemoteA2aAgent(BaseAgent):
return
# Create A2A request for function response or regular message
a2a_request, request_metadata = (
self._create_a2a_request_for_user_function_response(ctx)
)
a2a_request = self._create_a2a_request_for_user_function_response(ctx)
if not a2a_request:
message_parts, context_id, request_metadata = (
self._construct_message_parts_from_session(ctx)
message_parts, context_id = self._construct_message_parts_from_session(
ctx
)
if not message_parts:
@@ -528,6 +531,10 @@ class RemoteA2aAgent(BaseAgent):
logger.debug(build_a2a_request_log(a2a_request))
try:
request_metadata = None
if self._a2a_request_meta_provider:
request_metadata = self._a2a_request_meta_provider(ctx, a2a_request)
async for a2a_response in self._a2a_client.send_message(
request=a2a_request,
request_metadata=request_metadata,
+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