chore: Update a2a-sdk to 0.2.16

Convert a2a types to use snake_case fields

https://github.com/a2aproject/a2a-python/releases/tag/v0.2.16

PiperOrigin-RevId: 786279179
This commit is contained in:
Holt Skinner
2025-07-23 07:54:03 -07:00
committed by Copybara-Service
parent ce7253f63f
commit a911469616
14 changed files with 196 additions and 194 deletions
+1 -1
View File
@@ -80,7 +80,7 @@ dev = [
a2a = [
# go/keep-sorted start
"a2a-sdk>=0.2.11;python_version>='3.10'"
"a2a-sdk>=0.2.16;python_version>='3.10'"
# go/keep-sorted end
]
@@ -193,7 +193,7 @@ def convert_a2a_task_to_event(
message = None
if a2a_task.artifacts:
message = Message(
messageId="", role=Role.agent, parts=a2a_task.artifacts[-1].parts
message_id="", role=Role.agent, parts=a2a_task.artifacts[-1].parts
)
elif a2a_task.status and a2a_task.status.message:
message = a2a_task.status.message
@@ -353,7 +353,7 @@ def convert_event_to_a2a_message(
_process_long_running_tool(a2a_part, event)
if a2a_parts:
return Message(messageId=str(uuid.uuid4()), role=role, parts=a2a_parts)
return Message(message_id=str(uuid.uuid4()), role=role, parts=a2a_parts)
except Exception as e:
logger.error("Failed to convert event to status message: %s", e)
@@ -387,13 +387,13 @@ def _create_error_status_event(
event_metadata[_get_adk_metadata_key("error_code")] = str(event.error_code)
return TaskStatusUpdateEvent(
taskId=task_id,
contextId=context_id,
task_id=task_id,
context_id=context_id,
metadata=event_metadata,
status=TaskStatus(
state=TaskState.failed,
message=Message(
messageId=str(uuid.uuid4()),
message_id=str(uuid.uuid4()),
role=Role.agent,
parts=[TextPart(text=error_message)],
metadata={
@@ -463,8 +463,8 @@ def _create_status_update_event(
status.state = TaskState.input_required
return TaskStatusUpdateEvent(
taskId=task_id,
contextId=context_id,
task_id=task_id,
context_id=context_id,
status=status,
metadata=_get_context_metadata(event, invocation_context),
final=False,
@@ -64,7 +64,7 @@ def convert_a2a_part_to_genai_part(
if isinstance(part.file, a2a_types.FileWithUri):
return genai_types.Part(
file_data=genai_types.FileData(
file_uri=part.file.uri, mime_type=part.file.mimeType
file_uri=part.file.uri, mime_type=part.file.mime_type
)
)
@@ -72,7 +72,7 @@ def convert_a2a_part_to_genai_part(
return genai_types.Part(
inline_data=genai_types.Blob(
data=base64.b64decode(part.file.bytes),
mime_type=part.file.mimeType,
mime_type=part.file.mime_type,
)
)
else:
@@ -157,7 +157,7 @@ def convert_genai_part_to_a2a_part(
root=a2a_types.FilePart(
file=a2a_types.FileWithUri(
uri=part.file_data.file_uri,
mimeType=part.file_data.mime_type,
mime_type=part.file_data.mime_type,
)
)
)
@@ -166,7 +166,7 @@ def convert_genai_part_to_a2a_part(
a2a_part = a2a_types.FilePart(
file=a2a_types.FileWithBytes(
bytes=base64.b64encode(part.inline_data.data).decode('utf-8'),
mimeType=part.inline_data.mime_type,
mime_type=part.inline_data.mime_type,
)
)
@@ -133,13 +133,13 @@ class A2aAgentExecutor(AgentExecutor):
if not context.current_task:
await event_queue.enqueue_event(
TaskStatusUpdateEvent(
taskId=context.task_id,
task_id=context.task_id,
status=TaskStatus(
state=TaskState.submitted,
message=context.message,
timestamp=datetime.now(timezone.utc).isoformat(),
),
contextId=context.context_id,
context_id=context.context_id,
final=False,
)
)
@@ -153,17 +153,17 @@ class A2aAgentExecutor(AgentExecutor):
try:
await event_queue.enqueue_event(
TaskStatusUpdateEvent(
taskId=context.task_id,
task_id=context.task_id,
status=TaskStatus(
state=TaskState.failed,
timestamp=datetime.now(timezone.utc).isoformat(),
message=Message(
messageId=str(uuid.uuid4()),
message_id=str(uuid.uuid4()),
role=Role.agent,
parts=[TextPart(text=str(e))],
),
),
contextId=context.context_id,
context_id=context.context_id,
final=True,
)
)
@@ -196,12 +196,12 @@ class A2aAgentExecutor(AgentExecutor):
# publish the task working event
await event_queue.enqueue_event(
TaskStatusUpdateEvent(
taskId=context.task_id,
task_id=context.task_id,
status=TaskStatus(
state=TaskState.working,
timestamp=datetime.now(timezone.utc).isoformat(),
),
contextId=context.context_id,
context_id=context.context_id,
final=False,
metadata={
_get_adk_metadata_key('app_name'): runner.app_name,
@@ -229,11 +229,11 @@ class A2aAgentExecutor(AgentExecutor):
# the final result according to a2a protocol.
await event_queue.enqueue_event(
TaskArtifactUpdateEvent(
taskId=context.task_id,
lastChunk=True,
contextId=context.context_id,
task_id=context.task_id,
last_chunk=True,
context_id=context.context_id,
artifact=Artifact(
artifactId=str(uuid.uuid4()),
artifact_id=str(uuid.uuid4()),
parts=task_result_aggregator.task_status_message.parts,
),
)
@@ -241,25 +241,25 @@ class A2aAgentExecutor(AgentExecutor):
# public the final status update event
await event_queue.enqueue_event(
TaskStatusUpdateEvent(
taskId=context.task_id,
task_id=context.task_id,
status=TaskStatus(
state=TaskState.completed,
timestamp=datetime.now(timezone.utc).isoformat(),
),
contextId=context.context_id,
context_id=context.context_id,
final=True,
)
)
else:
await event_queue.enqueue_event(
TaskStatusUpdateEvent(
taskId=context.task_id,
task_id=context.task_id,
status=TaskStatus(
state=task_result_aggregator.task_state,
timestamp=datetime.now(timezone.utc).isoformat(),
message=task_result_aggregator.task_status_message,
),
contextId=context.context_id,
context_id=context.context_id,
final=True,
)
)
+13 -13
View File
@@ -172,10 +172,10 @@ Method: {req.method}
JSON-RPC: {req.jsonrpc}
-----------------------------------------------------------
Message:
ID: {req.params.message.messageId}
ID: {req.params.message.message_id}
Role: {req.params.message.role}
Task ID: {req.params.message.taskId}
Context ID: {req.params.message.contextId}{message_metadata_section}
Task ID: {req.params.message.task_id}
Context ID: {req.params.message.context_id}{message_metadata_section}
-----------------------------------------------------------
Message Parts:
{_NEW_LINE.join(message_parts_logs) if message_parts_logs else "No parts"}
@@ -221,7 +221,7 @@ JSON-RPC: {resp.root.jsonrpc}
if _is_a2a_task(result):
result_details.extend([
f"Task ID: {result.id}",
f"Context ID: {result.contextId}",
f"Context ID: {result.context_id}",
f"Status State: {result.status.state}",
f"Status Timestamp: {result.status.timestamp}",
f"History Length: {len(result.history) if result.history else 0}",
@@ -238,10 +238,10 @@ JSON-RPC: {resp.root.jsonrpc}
elif _is_a2a_message(result):
result_details.extend([
f"Message ID: {result.messageId}",
f"Message ID: {result.message_id}",
f"Role: {result.role}",
f"Task ID: {result.taskId}",
f"Context ID: {result.contextId}",
f"Task ID: {result.task_id}",
f"Context ID: {result.context_id}",
])
# Add message parts
@@ -288,10 +288,10 @@ JSON-RPC: {resp.root.jsonrpc}
Metadata:
{json.dumps(result.status.message.metadata, indent=2)}"""
status_message_section = f"""ID: {result.status.message.messageId}
status_message_section = f"""ID: {result.status.message.message_id}
Role: {result.status.message.role}
Task ID: {result.status.message.taskId}
Context ID: {result.status.message.contextId}
Task ID: {result.status.message.task_id}
Context ID: {result.status.message.context_id}
Message Parts:
{_NEW_LINE.join(status_parts_logs) if status_parts_logs else "No parts"}{status_metadata_section}"""
@@ -317,10 +317,10 @@ Message Parts:
history_logs.append(
f"""Message {i + 1}:
ID: {message.messageId}
ID: {message.message_id}
Role: {message.role}
Task ID: {message.taskId}
Context ID: {message.contextId}
Task ID: {message.task_id}
Context ID: {message.context_id}
Message Parts:
{_NEW_LINE.join(message_parts_logs) if message_parts_logs else " No parts"}{message_metadata_section}"""
)
+19 -19
View File
@@ -90,11 +90,11 @@ class AgentCardBuilder:
version=self._agent_version,
capabilities=self._capabilities,
skills=all_skills,
defaultInputModes=['text/plain'],
defaultOutputModes=['text/plain'],
supportsAuthenticatedExtendedCard=False,
default_input_modes=['text/plain'],
default_output_modes=['text/plain'],
supports_authenticated_extended_card=False,
provider=self._provider,
securitySchemes=self._security_schemes,
security_schemes=self._security_schemes,
)
except Exception as e:
raise RuntimeError(
@@ -125,8 +125,8 @@ async def _build_llm_agent_skills(agent: LlmAgent) -> List[AgentSkill]:
name='model',
description=agent_description,
examples=agent_examples,
inputModes=_get_input_modes(agent),
outputModes=_get_output_modes(agent),
input_modes=_get_input_modes(agent),
output_modes=_get_output_modes(agent),
tags=['llm'],
)
)
@@ -160,8 +160,8 @@ async def _build_sub_agent_skills(agent: BaseAgent) -> List[AgentSkill]:
name=f'{sub_agent.name}: {skill.name}',
description=skill.description,
examples=skill.examples,
inputModes=skill.inputModes,
outputModes=skill.outputModes,
input_modes=skill.input_modes,
output_modes=skill.output_modes,
tags=[f'sub_agent:{sub_agent.name}'] + (skill.tags or []),
)
sub_agent_skills.append(aggregated_skill)
@@ -197,8 +197,8 @@ async def _build_tool_skills(agent: LlmAgent) -> List[AgentSkill]:
name=tool_name,
description=getattr(tool, 'description', f'Tool: {tool_name}'),
examples=None,
inputModes=None,
outputModes=None,
input_modes=None,
output_modes=None,
tags=['llm', 'tools'],
)
)
@@ -213,8 +213,8 @@ def _build_planner_skill(agent: LlmAgent) -> AgentSkill:
name='planning',
description='Can think about the tasks to do and make plans',
examples=None,
inputModes=None,
outputModes=None,
input_modes=None,
output_modes=None,
tags=['llm', 'planning'],
)
@@ -226,8 +226,8 @@ def _build_code_executor_skill(agent: LlmAgent) -> AgentSkill:
name='code-execution',
description='Can execute codes',
examples=None,
inputModes=None,
outputModes=None,
input_modes=None,
output_modes=None,
tags=['llm', 'code_execution'],
)
@@ -250,8 +250,8 @@ async def _build_non_llm_agent_skills(agent: BaseAgent) -> List[AgentSkill]:
name=agent_name,
description=agent_description,
examples=agent_examples,
inputModes=_get_input_modes(agent),
outputModes=_get_output_modes(agent),
input_modes=_get_input_modes(agent),
output_modes=_get_output_modes(agent),
tags=[agent_type],
)
)
@@ -282,8 +282,8 @@ def _build_orchestration_skill(
name='sub-agents',
description='Orchestrates: ' + '; '.join(sub_agent_descriptions),
examples=None,
inputModes=None,
outputModes=None,
input_modes=None,
output_modes=None,
tags=[agent_type, 'orchestration'],
)
@@ -525,7 +525,7 @@ def _get_input_modes(agent: BaseAgent) -> Optional[List[str]]:
return None
# This could be enhanced to check model capabilities
# For now, return None to use defaultInputModes
# For now, return None to use default_input_modes
return None
+8 -8
View File
@@ -301,14 +301,14 @@ class RemoteA2aAgent(BaseAgent):
ctx.session.events[-1], ctx, Role.user
)
if function_call_event.custom_metadata:
a2a_message.taskId = (
a2a_message.task_id = (
function_call_event.custom_metadata.get(
A2A_METADATA_PREFIX + "task_id"
)
if function_call_event.custom_metadata
else None
)
a2a_message.contextId = (
a2a_message.context_id = (
function_call_event.custom_metadata.get(
A2A_METADATA_PREFIX + "context_id"
)
@@ -392,14 +392,14 @@ class RemoteA2aAgent(BaseAgent):
a2a_response.root.result, self.name, ctx
)
event.custom_metadata = event.custom_metadata or {}
if a2a_response.root.result.taskId:
if a2a_response.root.result.task_id:
event.custom_metadata[A2A_METADATA_PREFIX + "task_id"] = (
a2a_response.root.result.taskId
a2a_response.root.result.task_id
)
if a2a_response.root.result.contextId:
if a2a_response.root.result.context_id:
event.custom_metadata[A2A_METADATA_PREFIX + "context_id"] = (
a2a_response.root.result.contextId
a2a_response.root.result.context_id
)
else:
@@ -473,10 +473,10 @@ class RemoteA2aAgent(BaseAgent):
id=str(uuid.uuid4()),
params=A2AMessageSendParams(
message=A2AMessage(
messageId=str(uuid.uuid4()),
message_id=str(uuid.uuid4()),
parts=message_parts,
role="user",
contextId=context_id,
context_id=context_id,
)
),
)
@@ -532,8 +532,8 @@ class TestEventConverter:
)
assert isinstance(result, TaskStatusUpdateEvent)
assert result.taskId == task_id
assert result.contextId == context_id
assert result.task_id == task_id
assert result.context_id == context_id
assert result.status.state == TaskState.auth_required
def test_create_status_update_event_with_input_required_state(self):
@@ -596,8 +596,8 @@ class TestEventConverter:
)
assert isinstance(result, TaskStatusUpdateEvent)
assert result.taskId == task_id
assert result.contextId == context_id
assert result.task_id == task_id
assert result.context_id == context_id
assert result.status.state == TaskState.input_required
@@ -79,7 +79,7 @@ class TestConvertA2aPartToGenaiPart:
a2a_part = a2a_types.Part(
root=a2a_types.FilePart(
file=a2a_types.FileWithUri(
uri="gs://bucket/file.txt", mimeType="text/plain"
uri="gs://bucket/file.txt", mime_type="text/plain"
)
)
)
@@ -105,7 +105,7 @@ class TestConvertA2aPartToGenaiPart:
a2a_part = a2a_types.Part(
root=a2a_types.FilePart(
file=a2a_types.FileWithBytes(
bytes=base64_encoded, mimeType="text/plain"
bytes=base64_encoded, mime_type="text/plain"
)
)
)
@@ -307,7 +307,7 @@ class TestConvertGenaiPartToA2aPart:
assert isinstance(result.root, a2a_types.FilePart)
assert isinstance(result.root.file, a2a_types.FileWithUri)
assert result.root.file.uri == "gs://bucket/file.txt"
assert result.root.file.mimeType == "text/plain"
assert result.root.file.mime_type == "text/plain"
def test_convert_inline_data_part(self):
"""Test conversion of GenAI inline_data Part to A2A Part."""
@@ -330,7 +330,7 @@ class TestConvertGenaiPartToA2aPart:
expected_base64 = base64.b64encode(test_bytes).decode("utf-8")
assert result.root.file.bytes == expected_base64
assert result.root.file.mimeType == "text/plain"
assert result.root.file.mime_type == "text/plain"
def test_convert_inline_data_part_with_video_metadata(self):
"""Test conversion of GenAI inline_data Part with video metadata to A2A Part."""
@@ -496,7 +496,7 @@ class TestRoundTripConversions:
a2a_part = a2a_types.Part(
root=a2a_types.FilePart(
file=a2a_types.FileWithUri(
uri=original_uri, mimeType=original_mime_type
uri=original_uri, mime_type=original_mime_type
)
)
)
@@ -511,7 +511,7 @@ class TestRoundTripConversions:
assert isinstance(result_a2a_part.root, a2a_types.FilePart)
assert isinstance(result_a2a_part.root.file, a2a_types.FileWithUri)
assert result_a2a_part.root.file.uri == original_uri
assert result_a2a_part.root.file.mimeType == original_mime_type
assert result_a2a_part.root.file.mime_type == original_mime_type
def test_file_bytes_round_trip(self):
"""Test round-trip conversion for file parts with bytes."""
@@ -683,7 +683,7 @@ class TestA2aAgentExecutor:
from a2a.types import TextPart
test_message = Mock(spec=Message)
test_message.messageId = "test-message-id"
test_message.message_id = "test-message-id"
test_message.role = Role.agent
test_message.parts = [Mock(spec=TextPart)]
@@ -764,7 +764,7 @@ class TestA2aAgentExecutor:
from a2a.types import TextPart
test_message = Mock(spec=Message)
test_message.messageId = "test-message-id"
test_message.message_id = "test-message-id"
test_message.role = Role.agent
test_message.parts = [Mock(spec=TextPart)]
@@ -849,7 +849,7 @@ class TestA2aAgentExecutor:
from a2a.types import TextPart
test_message = Mock(spec=Message)
test_message.messageId = "test-message-id"
test_message.message_id = "test-message-id"
test_message.role = Role.agent
test_message.parts = [Part(root=TextPart(text="test content"))]
@@ -911,12 +911,12 @@ class TestA2aAgentExecutor:
call[0][0]
for call in self.mock_event_queue.enqueue_event.call_args_list
if hasattr(call[0][0], "artifact")
and call[0][0].lastChunk == True
and call[0][0].last_chunk == True
]
assert len(artifact_events) == 1
artifact_event = artifact_events[0]
assert artifact_event.taskId == "test-task-id"
assert artifact_event.contextId == "test-context-id"
assert artifact_event.task_id == "test-task-id"
assert artifact_event.context_id == "test-context-id"
# Check that artifact parts correspond to message parts
assert len(artifact_event.artifact.parts) == len(test_message.parts)
assert artifact_event.artifact.parts == test_message.parts
@@ -930,8 +930,8 @@ class TestA2aAgentExecutor:
assert len(final_events) >= 1
final_event = final_events[-1] # Get the last final event
assert final_event.status.state == TaskState.completed
assert final_event.taskId == "test-task-id"
assert final_event.contextId == "test-context-id"
assert final_event.task_id == "test-task-id"
assert final_event.context_id == "test-context-id"
@pytest.mark.asyncio
async def test_handle_request_with_non_working_state_publishes_status_only(
@@ -949,7 +949,7 @@ class TestA2aAgentExecutor:
from a2a.types import TextPart
test_message = Mock(spec=Message)
test_message.messageId = "test-message-id"
test_message.message_id = "test-message-id"
test_message.role = Role.agent
test_message.parts = [Part(root=TextPart(text="test content"))]
@@ -1011,7 +1011,7 @@ class TestA2aAgentExecutor:
call[0][0]
for call in self.mock_event_queue.enqueue_event.call_args_list
if hasattr(call[0][0], "artifact")
and call[0][0].lastChunk == True
and call[0][0].last_chunk == True
]
assert len(artifact_events) == 0
@@ -1025,5 +1025,5 @@ class TestA2aAgentExecutor:
final_event = final_events[-1] # Get the last final event
assert final_event.status.state == TaskState.auth_required
assert final_event.status.message == test_message
assert final_event.taskId == "test-task-id"
assert final_event.contextId == "test-context-id"
assert final_event.task_id == "test-task-id"
assert final_event.context_id == "test-context-id"
@@ -50,7 +50,7 @@ except ImportError as e:
def create_test_message(text: str) -> Message:
"""Helper function to create a test Message object."""
return Message(
messageId="test-msg",
message_id="test-msg",
role=Role.agent,
parts=[Part(root=TextPart(text=text))],
)
@@ -72,8 +72,8 @@ class TestTaskResultAggregator:
"""Test processing a failed event."""
status_message = create_test_message("Failed to process")
event = TaskStatusUpdateEvent(
taskId="test-task",
contextId="test-context",
task_id="test-task",
context_id="test-context",
status=TaskStatus(state=TaskState.failed, message=status_message),
final=True,
)
@@ -88,8 +88,8 @@ class TestTaskResultAggregator:
"""Test processing an auth_required event."""
status_message = create_test_message("Authentication needed")
event = TaskStatusUpdateEvent(
taskId="test-task",
contextId="test-context",
task_id="test-task",
context_id="test-context",
status=TaskStatus(
state=TaskState.auth_required, message=status_message
),
@@ -106,8 +106,8 @@ class TestTaskResultAggregator:
"""Test processing an input_required event."""
status_message = create_test_message("Input required")
event = TaskStatusUpdateEvent(
taskId="test-task",
contextId="test-context",
task_id="test-task",
context_id="test-context",
status=TaskStatus(
state=TaskState.input_required, message=status_message
),
@@ -123,8 +123,8 @@ class TestTaskResultAggregator:
def test_status_message_with_none_message(self):
"""Test that status message handles None message properly."""
event = TaskStatusUpdateEvent(
taskId="test-task",
contextId="test-context",
task_id="test-task",
context_id="test-context",
status=TaskStatus(state=TaskState.failed, message=None),
final=True,
)
@@ -138,8 +138,8 @@ class TestTaskResultAggregator:
# First set auth_required
auth_message = create_test_message("Auth required")
auth_event = TaskStatusUpdateEvent(
taskId="test-task",
contextId="test-context",
task_id="test-task",
context_id="test-context",
status=TaskStatus(state=TaskState.auth_required, message=auth_message),
final=False,
)
@@ -150,8 +150,8 @@ class TestTaskResultAggregator:
# Then process failed - should override
failed_message = create_test_message("Failed")
failed_event = TaskStatusUpdateEvent(
taskId="test-task",
contextId="test-context",
task_id="test-task",
context_id="test-context",
status=TaskStatus(state=TaskState.failed, message=failed_message),
final=True,
)
@@ -164,8 +164,8 @@ class TestTaskResultAggregator:
# First set input_required
input_message = create_test_message("Input needed")
input_event = TaskStatusUpdateEvent(
taskId="test-task",
contextId="test-context",
task_id="test-task",
context_id="test-context",
status=TaskStatus(
state=TaskState.input_required, message=input_message
),
@@ -178,8 +178,8 @@ class TestTaskResultAggregator:
# Then process auth_required - should override
auth_message = create_test_message("Auth needed")
auth_event = TaskStatusUpdateEvent(
taskId="test-task",
contextId="test-context",
task_id="test-task",
context_id="test-context",
status=TaskStatus(state=TaskState.auth_required, message=auth_message),
final=False,
)
@@ -204,8 +204,8 @@ class TestTaskResultAggregator:
# First set failed state
failed_message = create_test_message("Failure message")
failed_event = TaskStatusUpdateEvent(
taskId="test-task",
contextId="test-context",
task_id="test-task",
context_id="test-context",
status=TaskStatus(state=TaskState.failed, message=failed_message),
final=True,
)
@@ -216,8 +216,8 @@ class TestTaskResultAggregator:
# Then process working - should not override state and should not update message
# because the current task state is not working
working_event = TaskStatusUpdateEvent(
taskId="test-task",
contextId="test-context",
task_id="test-task",
context_id="test-context",
status=TaskStatus(state=TaskState.working),
final=False,
)
@@ -231,8 +231,8 @@ class TestTaskResultAggregator:
# Start with input_required
input_message = create_test_message("Input message")
input_event = TaskStatusUpdateEvent(
taskId="test-task",
contextId="test-context",
task_id="test-task",
context_id="test-context",
status=TaskStatus(
state=TaskState.input_required, message=input_message
),
@@ -244,8 +244,8 @@ class TestTaskResultAggregator:
# Override with auth_required
auth_message = create_test_message("Auth message")
auth_event = TaskStatusUpdateEvent(
taskId="test-task",
contextId="test-context",
task_id="test-task",
context_id="test-context",
status=TaskStatus(state=TaskState.auth_required, message=auth_message),
final=False,
)
@@ -255,8 +255,8 @@ class TestTaskResultAggregator:
# Override with failed
failed_message = create_test_message("Failed message")
failed_event = TaskStatusUpdateEvent(
taskId="test-task",
contextId="test-context",
task_id="test-task",
context_id="test-context",
status=TaskStatus(state=TaskState.failed, message=failed_message),
final=True,
)
@@ -266,8 +266,8 @@ class TestTaskResultAggregator:
# Working should not override failed message because current task state is failed
working_message = create_test_message("Working message")
working_event = TaskStatusUpdateEvent(
taskId="test-task",
contextId="test-context",
task_id="test-task",
context_id="test-context",
status=TaskStatus(state=TaskState.working, message=working_message),
final=False,
)
@@ -281,8 +281,8 @@ class TestTaskResultAggregator:
"""Test that working state events update the status message."""
working_message = create_test_message("Working on task")
event = TaskStatusUpdateEvent(
taskId="test-task",
contextId="test-context",
task_id="test-task",
context_id="test-context",
status=TaskStatus(state=TaskState.working, message=working_message),
final=False,
)
@@ -296,8 +296,8 @@ class TestTaskResultAggregator:
def test_working_event_with_none_message(self):
"""Test that working state events handle None message properly."""
event = TaskStatusUpdateEvent(
taskId="test-task",
contextId="test-context",
task_id="test-task",
context_id="test-context",
status=TaskStatus(state=TaskState.working, message=None),
final=False,
)
@@ -311,8 +311,8 @@ class TestTaskResultAggregator:
# First set auth_required state
auth_message = create_test_message("Auth required")
auth_event = TaskStatusUpdateEvent(
taskId="test-task",
contextId="test-context",
task_id="test-task",
context_id="test-context",
status=TaskStatus(state=TaskState.auth_required, message=auth_message),
final=False,
)
@@ -323,8 +323,8 @@ class TestTaskResultAggregator:
# Then process working - should not update message because task state is not working
working_message = create_test_message("Working on auth")
working_event = TaskStatusUpdateEvent(
taskId="test-task",
contextId="test-context",
task_id="test-task",
context_id="test-context",
status=TaskStatus(state=TaskState.working, message=working_message),
final=False,
)
+48 -46
View File
@@ -24,8 +24,11 @@ import pytest
try:
from a2a.types import DataPart as A2ADataPart
from a2a.types import Message as A2AMessage
from a2a.types import MessageSendConfiguration
from a2a.types import MessageSendParams
from a2a.types import Part as A2APart
from a2a.types import Role
from a2a.types import SendMessageRequest
from a2a.types import Task as A2ATask
from a2a.types import TaskState
from a2a.types import TaskStatus
@@ -137,32 +140,31 @@ class TestBuildA2ARequestLog:
from google.adk.a2a.logs.log_utils import build_a2a_request_log
# Create mock request with all components
req = Mock()
req.id = "req-123"
req.method = "sendMessage"
req.jsonrpc = "2.0"
# Mock message
req.params.message.messageId = "msg-456"
req.params.message.role = "user"
req.params.message.taskId = "task-789"
req.params.message.contextId = "ctx-101"
# Mock message parts - use simple mocks since the function will call build_message_part_log
part1 = Mock()
part2 = Mock()
req.params.message.parts = [part1, part2]
# Mock configuration
req.params.configuration.acceptedOutputModes = ["text", "image"]
req.params.configuration.blocking = True
req.params.configuration.historyLength = 10
req.params.configuration.pushNotificationConfig = Mock() # Non-None
# Mock metadata
req.params.metadata = {"key1": "value1"}
# Mock message metadata to avoid JSON serialization issues
req.params.message.metadata = {"msg_key": "msg_value"}
req = SendMessageRequest(
id="req-123",
method="message/send",
jsonrpc="2.0",
params=MessageSendParams(
message=A2AMessage(
message_id="msg-456",
role="user",
task_id="task-789",
context_id="ctx-101",
parts=[
A2APart(root=A2ATextPart(text="Part 1")),
A2APart(root=A2ATextPart(text="Part 2")),
],
metadata={"msg_key": "msg_value"},
),
configuration=MessageSendConfiguration(
accepted_output_modes=["text", "image"],
blocking=True,
history_length=10,
push_notification_config=None,
),
metadata={"key1": "value1"},
),
)
with patch(
"google.adk.a2a.logs.log_utils.build_message_part_log"
@@ -173,7 +175,7 @@ class TestBuildA2ARequestLog:
# Verify all components are present
assert "req-123" in result
assert "sendMessage" in result
assert "message/send" in result
assert "2.0" in result
assert "msg-456" in result
assert "user" in result
@@ -191,13 +193,13 @@ class TestBuildA2ARequestLog:
req = Mock()
req.id = "req-123"
req.method = "sendMessage"
req.method = "message/send"
req.jsonrpc = "2.0"
req.params.message.messageId = "msg-456"
req.params.message.message_id = "msg-456"
req.params.message.role = "user"
req.params.message.taskId = "task-789"
req.params.message.contextId = "ctx-101"
req.params.message.task_id = "task-789"
req.params.message.context_id = "ctx-101"
req.params.message.parts = None # No parts
req.params.message.metadata = None # No message metadata
@@ -220,10 +222,10 @@ class TestBuildA2ARequestLog:
req.method = "sendMessage"
req.jsonrpc = "2.0"
req.params.message.messageId = "msg-456"
req.params.message.message_id = "msg-456"
req.params.message.role = "user"
req.params.message.taskId = "task-789"
req.params.message.contextId = "ctx-101"
req.params.message.task_id = "task-789"
req.params.message.context_id = "ctx-101"
req.params.message.parts = [] # Empty parts list
req.params.message.metadata = None # No message metadata
@@ -283,7 +285,7 @@ class TestBuildA2AResponseLog:
from google.adk.a2a.logs.log_utils import build_a2a_response_log
task_status = TaskStatus(state=TaskState.working)
task = A2ATask(id="task-123", contextId="ctx-456", status=task_status)
task = A2ATask(id="task-123", context_id="ctx-456", status=task_status)
resp = Mock()
resp.root.result = task
@@ -314,7 +316,7 @@ class TestBuildA2AResponseLog:
# Create status message using module-level imported types
status_message = A2AMessage(
messageId="status-msg-123",
message_id="status-msg-123",
role=Role.agent,
parts=[
A2APart(root=A2ATextPart(text="Status part 1")),
@@ -325,7 +327,7 @@ class TestBuildA2AResponseLog:
task_status = TaskStatus(state=TaskState.working, message=status_message)
task = A2ATask(
id="task-123",
contextId="ctx-456",
context_id="ctx-456",
status=task_status,
history=[],
artifacts=None,
@@ -358,10 +360,10 @@ class TestBuildA2AResponseLog:
# Use module-level imported types consistently
message = A2AMessage(
messageId="msg-123",
message_id="msg-123",
role=Role.agent,
taskId="task-456",
contextId="ctx-789",
task_id="task-456",
context_id="ctx-789",
parts=[A2APart(root=A2ATextPart(text="Message part 1"))],
)
@@ -395,10 +397,10 @@ class TestBuildA2AResponseLog:
# Use mock for this case since we want to test empty parts handling
message = Mock()
message.__class__.__name__ = "Message"
message.messageId = "msg-empty"
message.message_id = "msg-empty"
message.role = "agent"
message.taskId = "task-empty"
message.contextId = "ctx-empty"
message.task_id = "task-empty"
message.context_id = "ctx-empty"
message.parts = None # No parts
message.model_dump_json.return_value = '{"message": "empty"}'
@@ -488,10 +490,10 @@ class TestBuildA2AResponseLog:
req.method = "sendMessage"
req.jsonrpc = "2.0"
req.params.message.messageId = "msg-with-metadata"
req.params.message.message_id = "msg-with-metadata"
req.params.message.role = "user"
req.params.message.taskId = "task-metadata"
req.params.message.contextId = "ctx-metadata"
req.params.message.task_id = "task-metadata"
req.params.message.context_id = "ctx-metadata"
req.params.message.parts = []
req.params.message.metadata = {"msg_type": "test", "priority": "high"}
@@ -181,15 +181,15 @@ class TestAgentCardBuilder:
assert isinstance(result, AgentCard)
assert result.name == "test_agent"
assert result.description == "Test agent description"
assert result.documentationUrl is None
assert result.documentation_url is None
assert result.url == "http://localhost:80/a2a"
assert result.version == "0.0.1"
assert result.skills == [mock_primary_skill, mock_sub_skill]
assert result.defaultInputModes == ["text/plain"]
assert result.defaultOutputModes == ["text/plain"]
assert result.supportsAuthenticatedExtendedCard is False
assert result.default_input_modes == ["text/plain"]
assert result.default_output_modes == ["text/plain"]
assert result.supports_authenticated_extended_card is False
assert result.provider is None
assert result.securitySchemes is None
assert result.security_schemes is None
@patch("google.adk.a2a.utils.agent_card_builder._build_primary_skills")
@patch("google.adk.a2a.utils.agent_card_builder._build_sub_agent_skills")
@@ -225,15 +225,15 @@ class TestAgentCardBuilder:
# Assert
assert result.name == "test_agent"
assert result.description == "An ADK Agent" # Default description
# The source code uses doc_url parameter but AgentCard expects documentationUrl
# Since the source code doesn't map doc_url to documentationUrl, it will be None
assert result.documentationUrl is None
# The source code uses doc_url parameter but AgentCard expects documentation_url
# Since the source code doesn't map doc_url to documentation_url, it will be None
assert result.documentation_url is None
assert (
result.url == "https://example.com/a2a"
) # Should strip trailing slash
assert result.version == "2.0.0"
assert result.provider == mock_provider
assert result.securitySchemes == mock_security_schemes
assert result.security_schemes == mock_security_schemes
@patch("google.adk.a2a.utils.agent_card_builder._build_primary_skills")
@patch("google.adk.a2a.utils.agent_card_builder._build_sub_agent_skills")
+13 -13
View File
@@ -73,8 +73,8 @@ def create_test_agent_card(
description=description,
version="1.0",
capabilities=AgentCapabilities(),
defaultInputModes=["text/plain"],
defaultOutputModes=["application/json"],
default_input_modes=["text/plain"],
default_output_modes=["application/json"],
skills=[
AgentSkill(
id="test-skill",
@@ -316,8 +316,8 @@ class TestRemoteA2aAgentResolution:
description="test",
version="1.0",
capabilities=AgentCapabilities(),
defaultInputModes=["text/plain"],
defaultOutputModes=["application/json"],
default_input_modes=["text/plain"],
default_output_modes=["application/json"],
skills=[
AgentSkill(
id="test-skill",
@@ -347,8 +347,8 @@ class TestRemoteA2aAgentResolution:
description="test",
version="1.0",
capabilities=AgentCapabilities(),
defaultInputModes=["text/plain"],
defaultOutputModes=["application/json"],
default_input_modes=["text/plain"],
default_output_modes=["application/json"],
skills=[
AgentSkill(
id="test-skill",
@@ -483,7 +483,7 @@ class TestRemoteA2aAgentMessageHandling:
) as mock_convert:
# Create a proper mock A2A message
mock_a2a_message = Mock(spec=A2AMessage)
mock_a2a_message.taskId = None # Will be set by the method
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(
@@ -492,7 +492,7 @@ class TestRemoteA2aAgentMessageHandling:
assert result is not None
assert result.params.message == mock_a2a_message
assert mock_a2a_message.taskId == "task-123"
assert mock_a2a_message.task_id == "task-123"
def test_construct_message_parts_from_session_success(self):
"""Test successful message parts construction from session."""
@@ -542,8 +542,8 @@ class TestRemoteA2aAgentMessageHandling:
async def test_handle_a2a_response_success_with_message(self):
"""Test successful A2A response handling with message."""
mock_a2a_message = Mock(spec=A2AMessage)
mock_a2a_message.taskId = "task-123"
mock_a2a_message.contextId = "context-123"
mock_a2a_message.task_id = "task-123"
mock_a2a_message.context_id = "context-123"
mock_success_response = Mock(spec=SendMessageSuccessResponse)
mock_success_response.result = mock_a2a_message
@@ -581,7 +581,7 @@ class TestRemoteA2aAgentMessageHandling:
"""Test successful A2A response handling with task."""
mock_a2a_task = Mock(spec=A2ATask)
mock_a2a_task.id = "task-123"
mock_a2a_task.contextId = "context-123"
mock_a2a_task.context_id = "context-123"
mock_success_response = Mock(spec=SendMessageSuccessResponse)
mock_success_response.result = mock_a2a_task
@@ -950,8 +950,8 @@ class TestRemoteA2aAgentIntegration:
mock_response = Mock()
mock_success_response = Mock(spec=SendMessageSuccessResponse)
mock_a2a_message = Mock(spec=A2AMessage)
mock_a2a_message.taskId = "task-123"
mock_a2a_message.contextId = "context-123"
mock_a2a_message.task_id = "task-123"
mock_a2a_message.context_id = "context-123"
mock_success_response.result = mock_a2a_message
mock_response.root = mock_success_response
mock_a2a_client.send_message.return_value = mock_response