mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
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:
committed by
Copybara-Service
parent
5d5708b2ab
commit
d12468ee5a
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user