Skip to content
Merged
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: 16 additions & 2 deletions src/a2a/server/tasks/base_push_notification_sender.py
Original file line number Diff line number Diff line change
Expand Up @@ -97,9 +97,23 @@ async def _dispatch_notification(
):
return False
try:
headers = None
headers: dict[str, str] = {}
if push_info.token:
headers = {'X-A2A-Notification-Token': push_info.token}
headers['X-A2A-Notification-Token'] = push_info.token
auth = push_info.authentication
if push_info.HasField('authentication'):
if auth.scheme and auth.credentials:
headers['Authorization'] = (
f'{auth.scheme} {auth.credentials}'
)
elif auth.scheme:
logger.warning(
'Push config %s sets an authentication scheme with no '
'credentials; sending no Authorization header for '
'task_id=%s',
push_info.id,
task_id,
)

response = await self._client.post(
url,
Expand Down
96 changes: 89 additions & 7 deletions tests/server/tasks/test_push_notification_sender.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
BasePushNotificationSender,
)
from a2a.types.a2a_pb2 import (
AuthenticationInfo,
StreamResponse,
Task,
TaskArtifactUpdateEvent,
Expand Down Expand Up @@ -36,8 +37,11 @@ def _create_sample_push_config(
url: str = 'http://example.com/callback',
config_id: str = 'cfg1',
token: str | None = None,
authentication: AuthenticationInfo | None = None,
) -> TaskPushNotificationConfig:
return TaskPushNotificationConfig(id=config_id, url=url, token=token)
return TaskPushNotificationConfig(
id=config_id, url=url, token=token, authentication=authentication
)


class TestBasePushNotificationSender(unittest.IsolatedAsyncioTestCase):
Expand Down Expand Up @@ -73,7 +77,7 @@ async def test_send_notification_success(self) -> None:
self.mock_httpx_client.post.assert_awaited_once_with(
config.url,
json=MessageToDict(StreamResponse(task=task_data)),
headers=None,
headers={},
)
mock_response.raise_for_status.assert_called_once()

Expand Down Expand Up @@ -103,6 +107,84 @@ async def test_send_notification_with_token_success(self) -> None:
)
mock_response.raise_for_status.assert_called_once()

async def _post_headers_for(self, **config_kwargs) -> dict[str, str] | None:
"""Sends one notification and returns the headers it posted with."""
task_data = _create_sample_task(task_id='task_auth')
self.mock_config_store.get_info_for_dispatch.return_value = [
_create_sample_push_config(**config_kwargs)
]
self.mock_httpx_client.post.return_value = AsyncMock(
spec=httpx.Response, status_code=200
)

await self.sender.send_notification(task_data.id, task_data)

return self.mock_httpx_client.post.await_args.kwargs['headers']

async def test_authentication_becomes_an_authorization_header(self) -> None:
"""Authentication becomes `Authorization: {scheme} {credentials}`."""
headers = await self._post_headers_for(
authentication=AuthenticationInfo(
scheme='Bearer', credentials='test-token'
)
)

assert headers == {'Authorization': 'Bearer test-token'}

async def test_authentication_and_token_are_sent_together(self) -> None:
headers = await self._post_headers_for(
token='notification-token',
authentication=AuthenticationInfo(
scheme='Bearer', credentials='test-token'
),
)

assert headers == {
'X-A2A-Notification-Token': 'notification-token',
'Authorization': 'Bearer test-token',
}

async def test_authentication_without_credentials_sends_no_header(
self,
) -> None:
"""A half-filled AuthenticationInfo yields no 'Bearer ' with nothing after it."""
headers = await self._post_headers_for(
authentication=AuthenticationInfo(scheme='Bearer')
)

assert headers == {}

async def test_authentication_without_credentials_warns(self) -> None:
"""A scheme with no credentials is valid to send, so say why it was dropped."""
with self.assertLogs(
'a2a.server.tasks.base_push_notification_sender', level='WARNING'
) as logs:
await self._post_headers_for(
config_id='cfg-half-auth',
authentication=AuthenticationInfo(scheme='Bearer'),
)

assert any(
'cfg-half-auth' in line and 'no ' in line for line in logs.output
)

async def test_complete_authentication_does_not_warn(self) -> None:
with self.assertNoLogs(
'a2a.server.tasks.base_push_notification_sender', level='WARNING'
):
await self._post_headers_for(
authentication=AuthenticationInfo(
scheme='Bearer', credentials='test-token'
)
)

async def test_no_authentication_sends_no_authorization_header(
self,
) -> None:
headers = await self._post_headers_for()

assert headers == {}

async def test_send_notification_no_config(self) -> None:
task_id = 'task_send_no_config'
task_data = _create_sample_task(task_id=task_id)
Expand Down Expand Up @@ -140,7 +222,7 @@ async def test_send_notification_http_status_error(
self.mock_httpx_client.post.assert_awaited_once_with(
config.url,
json=MessageToDict(StreamResponse(task=task_data)),
headers=None,
headers={},
)
mock_logger.exception.assert_called_once()

Expand Down Expand Up @@ -173,13 +255,13 @@ async def test_send_notification_multiple_configs(self) -> None:
self.mock_httpx_client.post.assert_any_call(
config1.url,
json=MessageToDict(StreamResponse(task=task_data)),
headers=None,
headers={},
)
# Check calls for config2
self.mock_httpx_client.post.assert_any_call(
config2.url,
json=MessageToDict(StreamResponse(task=task_data)),
headers=None,
headers={},
)
mock_response.raise_for_status.call_count = 2

Expand All @@ -204,7 +286,7 @@ async def test_send_notification_status_update_event(self) -> None:
self.mock_httpx_client.post.assert_awaited_once_with(
config.url,
json=MessageToDict(StreamResponse(status_update=event)),
headers=None,
headers={},
)

async def test_send_notification_artifact_update_event(self) -> None:
Expand All @@ -228,7 +310,7 @@ async def test_send_notification_artifact_update_event(self) -> None:
self.mock_httpx_client.post.assert_awaited_once_with(
config.url,
json=MessageToDict(StreamResponse(artifact_update=event)),
headers=None,
headers={},
)


Expand Down
Loading