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