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
Original file line number Diff line number Diff line change
Expand Up @@ -555,27 +555,25 @@

task_id = params.task_id
task: Task | None = await self.task_store.get(task_id, context)
if not task:
raise TaskNotFoundError

await self._reject_unsafe_push_url(params.url)

await self._push_config_store.set_info(
return await self._push_config_store.set_info(
task_id,
params,
context,
)

return params

@validate_request_params
@validate(
lambda self: self._agent_card.capabilities.push_notifications,
error_message='Push notifications are not supported by the agent',
error_type=PushNotificationNotSupportedError,
)
async def on_get_task_push_notification_config(
self,

Check notice on line 576 in src/a2a/server/request_handlers/default_request_handler.py

View workflow job for this annotation

GitHub Actions / Lint Code Base

Copy/pasted code

see src/a2a/server/request_handlers/default_request_handler_v2.py (495-512)
params: GetTaskPushNotificationConfigRequest,
context: ServerCallContext,
) -> TaskPushNotificationConfig:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -492,26 +492,24 @@
raise PushNotificationNotSupportedError

task_id = params.task_id
if await self._versioned_store.get(task_id, context) is None:
raise TaskNotFoundError

await self._reject_unsafe_push_url(params.url)

await self._push_config_store.set_info(
return await self._push_config_store.set_info(

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

to avoid breaking changes for existing custom implementations: here and in legacy handler lets implement a fallback for None case (see typing comment)

smth like this:

stored = await self._push_config_store.set_info(
   task_id,
   params,
   context,
)
if stored is not None:
   return stored
fallback = TaskPushNotificationConfig()
fallback.CopyFrom(params)
fallback.task_id = task_id
if not fallback.id:
   fallback.id = task_id
return fallback

task_id,
params,
context,
)

return params

@validate_request_params
@validate(
lambda self: self._agent_card.capabilities.push_notifications,
error_message='Push notifications are not supported by the agent',
error_type=PushNotificationNotSupportedError,
)
async def on_get_task_push_notification_config( # noqa: D102

Check notice on line 512 in src/a2a/server/request_handlers/default_request_handler_v2.py

View workflow job for this annotation

GitHub Actions / Lint Code Base

Copy/pasted code

see src/a2a/server/request_handlers/default_request_handler.py (558-576)
self,
params: GetTaskPushNotificationConfigRequest,
context: ServerCallContext,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -283,14 +283,15 @@ async def set_info(
task_id: str,
notification_config: TaskPushNotificationConfig,
context: ServerCallContext,
) -> None:
) -> TaskPushNotificationConfig:
"""Sets or updates the push notification configuration for a task."""
await self._ensure_initialized()
owner = self.owner_resolver(context)

# Create a copy of the config using proto CopyFrom
config_to_save = TaskPushNotificationConfig()
config_to_save.CopyFrom(notification_config)
config_to_save.task_id = task_id
if not config_to_save.id:
config_to_save.id = task_id

Expand All @@ -303,6 +304,7 @@ async def set_info(
config_to_save.id,
owner,
)
return config_to_save

async def _select_configs(
self,
Expand Down
21 changes: 14 additions & 7 deletions src/a2a/server/tasks/inmemory_push_notification_config_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,31 +40,38 @@ async def set_info(
task_id: str,
notification_config: TaskPushNotificationConfig,
context: ServerCallContext,
) -> None:
) -> TaskPushNotificationConfig:
"""Sets or updates the push notification configuration for a task in memory."""
owner = self.owner_resolver(context)
stored = TaskPushNotificationConfig()
stored.CopyFrom(notification_config)
stored.task_id = task_id
if not stored.id:
stored.id = task_id

with self.lock:
owner_infos = self._push_notification_infos.setdefault(owner, {})
if task_id not in owner_infos:
owner_infos[task_id] = []

if not notification_config.id:
notification_config.id = task_id

# Remove existing config with the same ID
for config in owner_infos[task_id]:
if config.id == notification_config.id:
if config.id == stored.id:
owner_infos[task_id].remove(config)
break

owner_infos[task_id].append(notification_config)
owner_infos[task_id].append(stored)
logger.debug(
'Push notification config for task %s with config id %s for owner %s saved/updated.',
task_id,
notification_config.id,
stored.id,
owner,
)

result = TaskPushNotificationConfig()
result.CopyFrom(stored)
return result

async def get_info(
self,
task_id: str,
Expand Down
9 changes: 7 additions & 2 deletions src/a2a/server/tasks/push_notification_config_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,8 +18,13 @@ async def set_info(
task_id: str,
notification_config: TaskPushNotificationConfig,
context: ServerCallContext,
) -> None:
"""Sets or updates the push notification configuration for a task."""
) -> TaskPushNotificationConfig:

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

to avoid breaking changes for existing custom implementations: lets make return type TaskPushNotificationConfig | None.

inmemory and database stores can stay just TaskPushNotificationConfig

"""Sets or updates the push notification configuration for a task.

Implementations MUST NOT mutate notification_config. They store a
copy with task_id set to the given task and an empty id defaulted to
the task id, and return that stored configuration.
"""

@abstractmethod
async def get_info(
Expand Down
50 changes: 50 additions & 0 deletions tests/server/request_handlers/test_default_request_handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -1977,6 +1977,56 @@ async def test_set_task_push_notification_config_task_not_found(agent_card):
mock_push_store.set_info.assert_not_awaited()


@pytest.mark.asyncio
@pytest.mark.parametrize('store_kind', ['inmemory', 'database'])
async def test_create_task_push_notification_config_returns_stored_id(
agent_card, store_kind
):
"""Test on_create_task_push_notification_config returns the id that was stored."""
if store_kind == 'database':
pytest.importorskip('sqlalchemy')
pytest.importorskip('aiosqlite')
from a2a.server.tasks.database_push_notification_config_store import (
DatabasePushNotificationConfigStore,
)
from sqlalchemy.ext.asyncio import create_async_engine

engine = create_async_engine(
'sqlite+aiosqlite:///file:pushid?mode=memory&cache=shared&uri=true'
)
push_config_store = DatabasePushNotificationConfigStore(engine=engine)
else:
push_config_store = InMemoryPushNotificationConfigStore()

task = create_sample_task()
task_store = InMemoryTaskStore()
context = create_server_call_context()
await task_store.save(task, context)

request_handler = DefaultRequestHandler(
agent_executor=MockAgentExecutor(),
task_store=task_store,
push_config_store=push_config_store,
agent_card=agent_card,
)
params = TaskPushNotificationConfig(
task_id=task.id,
url='http://example.com',
)

response = await request_handler.on_create_task_push_notification_config(
params, context
)

stored = await push_config_store.get_info(task.id, context)
if store_kind == 'database':
await engine.dispose()

assert response.id == task.id
assert list(stored) == [response]
assert params.id == '', 'the request object must not be mutated'


@pytest.mark.asyncio
async def test_get_task_push_notification_config_no_store(agent_card):
"""Test on_get_task_push_notification_config when _push_config_store is None."""
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -513,6 +513,55 @@ async def test_set_task_push_notification_config_task_not_found():
mock_push_store.set_info.assert_not_awaited()


@pytest.mark.asyncio
@pytest.mark.parametrize('store_kind', ['inmemory', 'database'])
async def test_create_task_push_notification_config_returns_stored_id(
store_kind,
):
"""Test on_create_task_push_notification_config returns the id that was stored."""
if store_kind == 'database':
pytest.importorskip('sqlalchemy')
pytest.importorskip('aiosqlite')
from a2a.server.tasks.database_push_notification_config_store import (

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

sqlalchemy and aiosqlite are optional extras (a2a-sdk[sqlite]), not core dependencies in pyproject.toml.

Unlike test_database_push_notification_config_store.py, these core handler test modules run in environments without optional SQL extras.

Guard the database branch with pytest.importorskip('sqlalchemy') and pytest.importorskip('aiosqlite')

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Guarded the database branch in both handler test modules with pytest.importorskip('sqlalchemy') and pytest.importorskip('aiosqlite').

DatabasePushNotificationConfigStore,
)
from sqlalchemy.ext.asyncio import create_async_engine

engine = create_async_engine(
'sqlite+aiosqlite:///file:pushidv2?mode=memory&cache=shared&uri=true'
)
push_config_store = DatabasePushNotificationConfigStore(engine=engine)
else:
push_config_store = InMemoryPushNotificationConfigStore()

task = create_sample_task()
task_store = InMemoryTaskStore()
context = create_server_call_context()
await task_store.save(task, context)

request_handler = DefaultRequestHandlerV2(
agent_executor=MockAgentExecutor(),
task_store=task_store,
push_config_store=push_config_store,
agent_card=create_default_agent_card(),
)
params = TaskPushNotificationConfig(
task_id=task.id, url='http://example.com'
)

response = await request_handler.on_create_task_push_notification_config(
params, context
)

stored = await push_config_store.get_info(task.id, context)
if store_kind == 'database':
await engine.dispose()

assert response.id == task.id
assert list(stored) == [response]
assert params.id == '', 'the request object must not be mutated'


@pytest.mark.asyncio
async def test_get_task_push_notification_config_no_store():
"""Test on_get_task_push_notification_config when _push_config_store is None."""
Expand Down Expand Up @@ -1710,6 +1759,40 @@ async def test_on_message_send_with_push_notification():
)


@pytest.mark.asyncio
async def test_on_message_send_stores_inline_push_config_under_its_task():
task_store = InMemoryTaskStore()
push_store = InMemoryPushNotificationConfigStore()

request_handler = DefaultRequestHandlerV2(
agent_executor=HelloAgentExecutor(),
task_store=task_store,
push_config_store=push_store,
agent_card=create_default_agent_card(),
)
# SendMessageConfiguration carries neither task_id nor id for the config.
params = SendMessageRequest(
message=Message(
role=Role.ROLE_USER,
message_id='msg_push_inline',
parts=[Part(text='Hi')],
),
configuration=SendMessageConfiguration(
task_push_notification_config=TaskPushNotificationConfig(
url='http://example.com/webhook'
)
),
)

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

stored = await push_store.get_info(result.id, context)
assert [(config.task_id, config.id) for config in stored] == [
(result.id, result.id)
]


@pytest.mark.asyncio
async def test_on_message_send_with_empty_push_notification_config_does_not_call_set_info():
task_store = InMemoryTaskStore()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -204,7 +204,9 @@ async def test_set_and_get_info_single_config(
):
"""Test setting and retrieving a single configuration."""
task_id = 'task-1'
config = TaskPushNotificationConfig(id='config-1', url='http://example.com')
config = TaskPushNotificationConfig(
task_id=task_id, id='config-1', url='http://example.com'
)

await db_store_parameterized.set_info(task_id, config, MINIMAL_CALL_CONTEXT)
retrieved_configs = await db_store_parameterized.get_info(
Expand All @@ -215,6 +217,26 @@ async def test_set_and_get_info_single_config(
assert retrieved_configs[0] == config


@pytest.mark.asyncio
async def test_set_info_returns_the_stored_config_without_mutating_input(
db_store_parameterized: DatabasePushNotificationConfigStore,
):
"""set_info returns what it stored and leaves the caller's config alone."""
task_id = 'task-normalize'
config = TaskPushNotificationConfig(url='http://example.com')

stored = await db_store_parameterized.set_info(
task_id, config, MINIMAL_CALL_CONTEXT
)

assert (stored.task_id, stored.id) == (task_id, task_id)
assert (config.task_id, config.id) == ('', '')
retrieved_configs = await db_store_parameterized.get_info(
task_id, MINIMAL_CALL_CONTEXT
)
assert retrieved_configs == [stored]


@pytest.mark.asyncio
async def test_set_and_get_info_multiple_configs(
db_store_parameterized: DatabasePushNotificationConfigStore,
Expand Down Expand Up @@ -306,8 +328,12 @@ async def test_delete_info_specific_config(
):
"""Test deleting a single, specific configuration."""
task_id = 'task-1'
config1 = TaskPushNotificationConfig(id='config-1', url='http://a.com')
config2 = TaskPushNotificationConfig(id='config-2', url='http://b.com')
config1 = TaskPushNotificationConfig(
task_id=task_id, id='config-1', url='http://a.com'
)
config2 = TaskPushNotificationConfig(
task_id=task_id, id='config-2', url='http://b.com'
)

await db_store_parameterized.set_info(
task_id, config1, MINIMAL_CALL_CONTEXT
Expand Down Expand Up @@ -372,7 +398,10 @@ async def test_data_is_encrypted_in_db(
"""Verify that the data stored in the database is actually encrypted."""
task_id = 'encrypted-task'
config = TaskPushNotificationConfig(
id='config-1', url='http://secret.url', token='secret-token'
task_id=task_id,
id='config-1',
url='http://secret.url',
token='secret-token',
)
plain_json = MessageToJson(config)

Expand Down Expand Up @@ -481,7 +510,7 @@ async def test_custom_table_name(

task_id = 'custom-table-task'
config = TaskPushNotificationConfig(
id='config-1', url='http://custom.url'
task_id=task_id, id='config-1', url='http://custom.url'
)

# This will create the table on first use
Expand Down Expand Up @@ -530,10 +559,10 @@ async def test_set_and_get_info_multiple_configs_no_key(

task_id = 'task-1'
config1 = TaskPushNotificationConfig(
id='config-1', url='http://example.com/1'
task_id=task_id, id='config-1', url='http://example.com/1'
)
config2 = TaskPushNotificationConfig(
id='config-2', url='http://example.com/2'
task_id=task_id, id='config-2', url='http://example.com/2'
)

await store.set_info(task_id, config1, MINIMAL_CALL_CONTEXT)
Expand All @@ -560,7 +589,7 @@ async def test_data_is_not_encrypted_in_db_if_no_key_is_set(

task_id = 'task-1'
config = TaskPushNotificationConfig(
id='config-1', url='http://example.com/1'
task_id=task_id, id='config-1', url='http://example.com/1'
)
plain_json = MessageToJson(config)

Expand Down Expand Up @@ -594,7 +623,9 @@ async def test_decryption_fallback_for_unencrypted_data(
await unencrypted_store.initialize()

task_id = 'mixed-encryption-task'
config = TaskPushNotificationConfig(id='config-1', url='http://plain.url')
config = TaskPushNotificationConfig(
task_id=task_id, id='config-1', url='http://plain.url'
)
await unencrypted_store.set_info(task_id, config, MINIMAL_CALL_CONTEXT)

# 2. Try to read with the encryption-enabled store from the fixture
Expand Down
Loading
Loading