Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 15 additions & 3 deletions src/a2a/server/agent_execution/active_task.py
Original file line number Diff line number Diff line change
Expand Up @@ -346,9 +346,21 @@ async def _update_task_state(
'Consumer[%s]: Sending push notification',
self.active_task._task_id,
)
await self.active_task._push_sender.send_notification(
self.active_task._task_id, event
)
try:
await self.active_task._push_sender.send_notification(
self.active_task._task_id, event
)
except Exception:
# Push delivery is best-effort: a sender failure must not
# fail the task or disturb the event stream. This guard is
# defense-in-depth on top of the boundary inside
# BasePushNotificationSender, and also protects against
# custom PushNotificationSender implementations.
logger.exception(
'Consumer[%s]: Push notification sender raised; '
'ignoring the failure.',
self.active_task._task_id,
)

async def _handle_terminal_state(self, updated_task: Task) -> None:
logger.debug(
Expand Down
11 changes: 10 additions & 1 deletion src/a2a/server/request_handlers/default_request_handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -382,7 +382,16 @@ async def _send_push_notification_if_needed(
and task_id
and isinstance(event, PushNotificationEvent)
):
await self._push_sender.send_notification(task_id, event)
try:
await self._push_sender.send_notification(task_id, event)
except Exception:
# Push delivery is best-effort: a sender failure must not
# fail the task or disturb the event stream.
logger.exception(
'Push notification sender raised for task_id=%s; '
'ignoring the failure.',
task_id,
)

@validate_request_params
async def on_message_send(
Expand Down
39 changes: 32 additions & 7 deletions src/a2a/server/tasks/base_push_notification_sender.py
Original file line number Diff line number Diff line change
Expand Up @@ -68,8 +68,25 @@ def __init__(
async def send_notification(
self, task_id: str, event: PushNotificationEvent
) -> None:
"""Sends a push notification for an event if configuration exists."""
push_configs = await self._config_store.get_info_for_dispatch(task_id)
"""Sends a push notification for an event if configuration exists.

Best-effort by design: failures to read the configuration store or
to deliver a notification are logged and do not raise.
"""
try:
push_configs = await self._config_store.get_info_for_dispatch(
task_id
)
except Exception:
# Push delivery is best-effort: an infrastructure failure while
# reading the config store must not propagate into the task
# lifecycle.
logger.exception(
'Failed to read push notification configs for task_id=%s; '
'skipping push notification delivery.',
task_id,
)
return
if not push_configs:
return

Expand All @@ -91,11 +108,19 @@ async def _dispatch_notification(
task_id: str,
) -> bool:
url = push_info.url
if (
self._push_url_validator is not None
and not await self._push_url_validator(url)
):
return False
if self._push_url_validator is not None:
try:
accepted = await self._push_url_validator(url)
except Exception:
logger.exception(
'Push URL validator raised for task_id=%s, URL: %s. '
'Failing closed and treating the URL as rejected.',
task_id,
url,
)
return False
if not accepted:
return False
try:
headers: dict[str, str] = {}
if push_info.token:
Expand Down
7 changes: 6 additions & 1 deletion src/a2a/server/tasks/push_notification_sender.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,4 +17,9 @@ class PushNotificationSender(ABC):
async def send_notification(
self, task_id: str, event: PushNotificationEvent
) -> None:
"""Sends a push notification containing the latest task state."""
"""Sends a push notification containing the latest task state.

Implementations are treated as best-effort: the framework catches
and logs any exception raised here, so a failure never affects the
task lifecycle or the event stream.
"""
24 changes: 24 additions & 0 deletions tests/server/request_handlers/test_default_request_handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -3451,3 +3451,27 @@ def test_init_does_not_warn_for_complete_agent_cards(
extended_agent_card=_complete_agent_card(),
)
assert caplog.records == []


@pytest.mark.asyncio
async def test_push_sender_failure_does_not_propagate(agent_card):
"""A push-notification infrastructure failure raised by the sender must
not propagate out of _send_push_notification_if_needed into the task
lifecycle (#1313)."""
mock_push_sender = AsyncMock(spec=PushNotificationSender)
mock_push_sender.send_notification.side_effect = RuntimeError(
'transient push infrastructure error'
)
request_handler = DefaultRequestHandler(
agent_executor=AsyncMock(spec=AgentExecutor),
task_store=AsyncMock(spec=TaskStore),
push_config_store=AsyncMock(spec=PushNotificationConfigStore),
push_sender=mock_push_sender,
agent_card=agent_card,
)
event = create_sample_task()

# Must not raise: push delivery is best-effort.
await request_handler._send_push_notification_if_needed(event.id, event)

mock_push_sender.send_notification.assert_awaited_once_with(event.id, event)
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@

from unittest.mock import AsyncMock, MagicMock, patch

import httpx
import pytest

from a2a.auth.user import UnauthenticatedUser, User
Expand All @@ -33,6 +34,9 @@
TaskStore,
TaskUpdater,
)
from a2a.server.tasks.base_push_notification_sender import (
BasePushNotificationSender,
)
from a2a.server.tasks.task_manager import TaskManager
from a2a.types import (
ContentTypeNotSupportedError,
Expand Down Expand Up @@ -2591,3 +2595,98 @@ def test_init_does_not_warn_for_complete_agent_cards(
extended_agent_card=_complete_agent_card(),
)
assert caplog.records == []


class FlakyPushConfigStore(InMemoryPushNotificationConfigStore):
"""Config store whose dispatch read fails transiently (issue #1313)."""

async def get_info_for_dispatch(
self, task_id: str
) -> list[TaskPushNotificationConfig]:
raise RuntimeError('transient DB error while reading push configs')


class RaisingPushSender(PushNotificationSender):
"""A custom (public-interface) sender that always fails (issue #1313)."""

async def send_notification(self, task_id, event) -> None:
raise RuntimeError('sender exploded')


@pytest.mark.asyncio
async def test_push_config_store_failure_does_not_fail_task():
"""A push-notification infrastructure failure must not rewrite a
successfully completed task as FAILED, and message/send must still
return the task result (#1313)."""
task_store = InMemoryTaskStore()
request_handler = DefaultRequestHandlerV2(
agent_executor=HelloAgentExecutor(),
task_store=task_store,
push_config_store=FlakyPushConfigStore(),
push_sender=BasePushNotificationSender(
httpx_client=AsyncMock(spec=httpx.AsyncClient),
config_store=FlakyPushConfigStore(),
),
agent_card=create_default_agent_card(),
)
params = SendMessageRequest(
message=Message(
role=Role.ROLE_USER,
message_id='msg_push_store_fail',
parts=[Part(text='Hi')],
),
configuration=SendMessageConfiguration(
accepted_output_modes=['text/plain']
),
)

result = await request_handler.on_message_send(
params, create_server_call_context()
)

assert isinstance(result, Task)
assert result.status.state == TaskState.TASK_STATE_COMPLETED
get_task_result = await request_handler.on_get_task(
GetTaskRequest(id=result.id), create_server_call_context()
)
assert get_task_result is not None
assert isinstance(get_task_result, Task)
assert get_task_result.status.state == TaskState.TASK_STATE_COMPLETED


@pytest.mark.asyncio
async def test_custom_push_sender_failure_does_not_fail_task():
"""A raising custom PushNotificationSender must not rewrite a
completed task as FAILED — the consumer-level guard is
defense-in-depth behind the BasePushNotificationSender boundary
(#1313)."""
task_store = InMemoryTaskStore()
request_handler = DefaultRequestHandlerV2(
agent_executor=HelloAgentExecutor(),
task_store=task_store,
push_sender=RaisingPushSender(),
agent_card=create_default_agent_card(),
)
params = SendMessageRequest(
message=Message(
role=Role.ROLE_USER,
message_id='msg_push_sender_fail',
parts=[Part(text='Hi')],
),
configuration=SendMessageConfiguration(
accepted_output_modes=['text/plain']
),
)

result = await request_handler.on_message_send(
params, create_server_call_context()
)

assert isinstance(result, Task)
assert result.status.state == TaskState.TASK_STATE_COMPLETED
get_task_result = await request_handler.on_get_task(
GetTaskRequest(id=result.id), create_server_call_context()
)
assert get_task_result is not None
assert isinstance(get_task_result, Task)
assert get_task_result.status.state == TaskState.TASK_STATE_COMPLETED
40 changes: 40 additions & 0 deletions tests/server/tasks/test_push_notification_sender.py
Original file line number Diff line number Diff line change
Expand Up @@ -313,6 +313,46 @@ async def test_send_notification_artifact_update_event(self) -> None:
headers={},
)

@patch('a2a.server.tasks.base_push_notification_sender.logger')
async def test_send_notification_config_store_failure_is_swallowed(
self, mock_logger: MagicMock
) -> None:
"""A config store read failure is logged and delivery is skipped
instead of propagating into the task lifecycle (issue #1313)."""
task_id = 'task_store_error'
task_data = _create_sample_task(task_id=task_id)
self.mock_config_store.get_info_for_dispatch.side_effect = RuntimeError(
'transient DB error'
)

await self.sender.send_notification(task_id, task_data)

self.mock_config_store.get_info_for_dispatch.assert_awaited_once_with(
task_id
)
self.mock_httpx_client.post.assert_not_called()
mock_logger.exception.assert_called_once()

@patch('a2a.server.tasks.base_push_notification_sender.logger')
async def test_push_url_validator_raising_rejects_url(
self, mock_logger: MagicMock
) -> None:
"""A raising push_url_validator is treated as a rejected URL rather
than propagating into the task lifecycle (issue #1313)."""
sender = BasePushNotificationSender(
httpx_client=self.mock_httpx_client,
config_store=self.mock_config_store,
push_url_validator=AsyncMock(side_effect=RuntimeError('boom')),
)
task_data = _create_sample_task(task_id='task_validator_error')
config = _create_sample_push_config(url='http://notify.me/here')
self.mock_config_store.get_info_for_dispatch.return_value = [config]

await sender.send_notification(task_data.id, task_data)

self.mock_httpx_client.post.assert_not_called()
mock_logger.exception.assert_called_once()


def _gai_result(ip: str, port: int = 80):
return [(2, 1, 6, '', (ip, port))]
Expand Down
Loading