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
@@ -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,
|
||||
|
||||
@@ -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