Skip to content
136 changes: 130 additions & 6 deletions python/packages/a2a/agent_framework_a2a/_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,7 @@
)
from agent_framework._telemetry import mark_feature_used
from agent_framework._types import AgentRunInputs
from agent_framework.exceptions import AgentInvalidRequestException
from agent_framework.observability import AgentTelemetryLayer
from google.protobuf.json_format import MessageToDict

Expand Down Expand Up @@ -484,10 +485,21 @@ def run(
a2a_stream: AsyncIterable[A2AStreamItem] = self.client.subscribe(
SubscribeToTaskRequest(id=continuation_token["task_id"])
)
input_request_occurrence_id = continuation_token["task_id"]
else:
if not normalized_messages:
raise ValueError("At least one message is required when starting a new task (no continuation_token).")
context_id, task_id, task_state = self._extract_a2a_session_state(session)
message = f"A2A agent {self.name!r} requires a real message or an explicit continuation token."
if context_id is not None or task_id is not None or task_state is not None:
task_context = [
f"context_id={context_id!r}",
f"task_id={task_id!r}",
f"task_state={TaskState.Name(task_state) if task_state is not None else None}",
]
message = f"{message} Session context: {', '.join(task_context)}."
raise AgentInvalidRequestException(message)
a2a_message = self._prepare_message_for_a2a(normalized_messages[-1], session=session)
input_request_occurrence_id = a2a_message.message_id
request = SendMessageRequest(message=a2a_message)
if background and not stream:
# return_immediately only applies to non-streaming (message/send)
Expand All @@ -510,6 +522,7 @@ def run(
a2a_stream,
background=background,
emit_intermediate=stream,
input_request_occurrence_id=input_request_occurrence_id,
session=provider_session,
session_context=session_context,
),
Expand All @@ -525,6 +538,7 @@ async def _map_a2a_stream(
*,
background: bool = False,
emit_intermediate: bool = False,
input_request_occurrence_id: str,
session: AgentSession | None = None,
session_context: SessionContext | None = None,
) -> AsyncIterable[AgentResponseUpdate]:
Expand All @@ -541,6 +555,7 @@ async def _map_a2a_stream(
carry message content are yielded to the caller. Typically
set for streaming callers so non-streaming consumers only
receive terminal task outputs.
input_request_occurrence_id: Stable identity for message-less input requests observed in this run.
session: The agent session for context providers.
session_context: The session context for context providers.
"""
Expand All @@ -563,6 +578,7 @@ async def _map_a2a_stream(
)

all_updates: list[AgentResponseUpdate] = []
seen_user_input_request_ids: set[str] = set()
streamed_artifact_ids_by_task: dict[str, set[str]] = {}
last_task_id: str | None = None
last_context_id: str | None = None
Expand Down Expand Up @@ -600,8 +616,10 @@ async def _map_a2a_stream(
task,
background=background,
emit_intermediate=emit_intermediate,
input_request_occurrence_id=input_request_occurrence_id,
streamed_artifact_ids=streamed_artifact_ids_by_task.get(task.id),
)
updates = self._deduplicate_user_input_request_updates(updates, seen_user_input_request_ids)
if task.status.state in TERMINAL_TASK_STATES:
streamed_artifact_ids_by_task.pop(task.id, None)
# If the terminal Task has no content, flush accumulated updates
Expand All @@ -621,7 +639,12 @@ async def _map_a2a_stream(
if status_event.context_id:
last_context_id = status_event.context_id
last_task_state = status_event.status.state
updates = self._updates_from_task_update_event(status_event)
updates = self._updates_from_task_update_event(
status_event,
background=background,
input_request_occurrence_id=input_request_occurrence_id,
)
updates = self._deduplicate_user_input_request_updates(updates, seen_user_input_request_ids)
is_terminal = status_event.status.state in TERMINAL_TASK_STATES
is_input_required = status_event.status.state == TaskState.TASK_STATE_INPUT_REQUIRED
if emit_intermediate:
Expand Down Expand Up @@ -699,12 +722,56 @@ async def _map_a2a_stream(
# Task helpers
# ------------------------------------------------------------------

def _user_input_request_id(self, task_id: str, occurrence_id: str) -> str:
"""Return a workflow request ID scoped to one remote prompt occurrence."""
scoped_occurrence = f"{len(self.id)}:{self.id}{len(task_id)}:{task_id}{occurrence_id}"
return f"a2a-input-{uuid.uuid5(uuid.NAMESPACE_URL, scoped_occurrence)}"

def _input_required_request(
self,
task_id: str,
message: A2AMessage | None,
*,
fallback_occurrence_id: str,
) -> Content:
"""Normalize an A2A input requirement into one durable caller request."""
contents = self._parse_contents_from_a2a(message.parts) if message is not None else []
prompt_parts = [
value for content in contents if (value := content.text or (content.uri if content.type == "uri" else None))
]
request = Content.from_text(
text="\n".join(prompt_parts) or "Remote A2A task requires input.",
additional_properties=(
{"a2a_input_required_message": MessageToDict(message)} if message is not None else None
),
)
occurrence_id = message.message_id if message is not None and message.message_id else fallback_occurrence_id
request.id = self._user_input_request_id(task_id, occurrence_id)
request.user_input_request = True
return request

@staticmethod
def _deduplicate_user_input_request_updates(
updates: list[AgentResponseUpdate],
seen_request_ids: set[str],
) -> list[AgentResponseUpdate]:
"""Drop duplicate representations of a user-input request within one agent run."""
deduplicated: list[AgentResponseUpdate] = []
for update in updates:
request_ids = {request.id for request in update.user_input_requests if request.id}
if request_ids and request_ids.issubset(seen_request_ids):
continue
seen_request_ids.update(request_ids)
deduplicated.append(update)
return deduplicated

def _updates_from_task(
self,
task: Task,
*,
background: bool = False,
emit_intermediate: bool = False,
input_request_occurrence_id: str | None = None,
streamed_artifact_ids: set[str] | None = None,
) -> list[AgentResponseUpdate]:
"""Convert an A2A Task into AgentResponseUpdate(s).
Expand All @@ -720,6 +787,26 @@ def _updates_from_task(
status = task.status
task_metadata = MessageToDict(task.metadata) if task.metadata else None

if status.state == TaskState.TASK_STATE_INPUT_REQUIRED:
message = status.message if status.HasField("message") and status.message.parts else None
occurrence_id = input_request_occurrence_id or task.id
return [
AgentResponseUpdate(
contents=[
self._input_required_request(
task.id,
message,
fallback_occurrence_id=occurrence_id,
)
],
role="assistant" if message is None or message.role == A2ARole.ROLE_AGENT else "user",
response_id=task.id,
continuation_token=self._build_continuation_token(task) if background else None,
additional_properties={"a2a_metadata": task_metadata} if task_metadata else None,
raw_representation=task,
)
]

if status.state in TERMINAL_TASK_STATES:
task_messages = self._parse_messages_from_task(task)
if task.artifacts and streamed_artifact_ids:
Expand Down Expand Up @@ -789,7 +876,11 @@ def _updates_from_task(
return []

def _updates_from_task_update_event(
self, update_event: TaskStatusUpdateEvent | TaskArtifactUpdateEvent
self,
update_event: TaskStatusUpdateEvent | TaskArtifactUpdateEvent,
*,
background: bool = False,
input_request_occurrence_id: str | None = None,
) -> list[AgentResponseUpdate]:
"""Convert A2A task update events into streaming AgentResponseUpdates."""
if isinstance(update_event, TaskArtifactUpdateEvent):
Expand All @@ -813,18 +904,50 @@ def _updates_from_task_update_event(
if not isinstance(update_event, TaskStatusUpdateEvent):
return []

state = update_event.status.state
continuation_token = (
A2AContinuationToken(task_id=update_event.task_id, context_id=update_event.context_id)
if background and state in IN_PROGRESS_TASK_STATES
else None
)
if state == TaskState.TASK_STATE_INPUT_REQUIRED:
message = (
update_event.status.message
if update_event.status.HasField("message") and update_event.status.message.parts
else None
)
message_meta = MessageToDict(message.metadata) if message is not None and message.metadata else {}
event_meta = MessageToDict(update_event.metadata) if update_event.metadata else {}
merged_metadata = {**message_meta, **event_meta} or None
occurrence_id = input_request_occurrence_id or update_event.task_id
return [
AgentResponseUpdate(
contents=[
self._input_required_request(
update_event.task_id,
message,
fallback_occurrence_id=occurrence_id,
)
],
role="assistant" if message is None or message.role == A2ARole.ROLE_AGENT else "user",
response_id=update_event.task_id,
message_id=message.message_id if message is not None else None,
continuation_token=continuation_token,
additional_properties={"a2a_metadata": merged_metadata} if merged_metadata else None,
raw_representation=update_event,
)
]

if not update_event.status.HasField("message") or not update_event.status.message.parts:
return []

state = update_event.status.state
if state not in TERMINAL_TASK_STATES and state != TaskState.TASK_STATE_INPUT_REQUIRED:
if state not in TERMINAL_TASK_STATES:
return []

message = update_event.status.message
contents = self._parse_contents_from_a2a(message.parts)
if not contents:
return []

msg_meta = MessageToDict(message.metadata) if message.metadata else {}
event_meta = MessageToDict(update_event.metadata) if update_event.metadata else {}
merged_metadata = {**msg_meta, **event_meta} or None
Expand All @@ -835,6 +958,7 @@ def _updates_from_task_update_event(
role="assistant" if message.role == A2ARole.ROLE_AGENT else "user",
response_id=update_event.task_id,
message_id=message.message_id,
continuation_token=continuation_token,
additional_properties={"a2a_metadata": merged_metadata} if merged_metadata else None,
raw_representation=update_event,
)
Expand Down
Loading
Loading