diff --git a/src/a2a/server/request_handlers/default_request_handler.py b/src/a2a/server/request_handlers/default_request_handler.py index 85a82ceb7..8f35372c8 100644 --- a/src/a2a/server/request_handlers/default_request_handler.py +++ b/src/a2a/server/request_handlers/default_request_handler.py @@ -32,6 +32,9 @@ TaskManager, TaskStore, ) +from a2a.server.tasks.push_notification_config_store import ( + normalize_push_notification_config, +) from a2a.types.a2a_pb2 import ( AgentCard, CancelTaskRequest, @@ -560,13 +563,15 @@ async def on_create_task_push_notification_config( await self._reject_unsafe_push_url(params.url) - await self._push_config_store.set_info( + stored = await self._push_config_store.set_info( task_id, params, context, ) - - return params + if stored is not None: + return stored + # Custom stores written before set_info returned the stored config. + return normalize_push_notification_config(task_id, params) @validate_request_params @validate( diff --git a/src/a2a/server/request_handlers/default_request_handler_v2.py b/src/a2a/server/request_handlers/default_request_handler_v2.py index e78980e96..a13ddb522 100644 --- a/src/a2a/server/request_handlers/default_request_handler_v2.py +++ b/src/a2a/server/request_handlers/default_request_handler_v2.py @@ -29,6 +29,9 @@ validate, validate_request_params, ) +from a2a.server.tasks.push_notification_config_store import ( + normalize_push_notification_config, +) from a2a.types.a2a_pb2 import ( AgentCard, CancelTaskRequest, @@ -497,13 +500,15 @@ async def on_create_task_push_notification_config( # noqa: D102 await self._reject_unsafe_push_url(params.url) - await self._push_config_store.set_info( + stored = await self._push_config_store.set_info( task_id, params, context, ) - - return params + if stored is not None: + return stored + # Custom stores written before set_info returned the stored config. + return normalize_push_notification_config(task_id, params) @validate_request_params @validate( diff --git a/src/a2a/server/tasks/database_push_notification_config_store.py b/src/a2a/server/tasks/database_push_notification_config_store.py index 697b75542..ce5979095 100644 --- a/src/a2a/server/tasks/database_push_notification_config_store.py +++ b/src/a2a/server/tasks/database_push_notification_config_store.py @@ -38,6 +38,7 @@ from a2a.server.owner_resolver import OwnerResolver, resolve_user_scope from a2a.server.tasks.push_notification_config_store import ( PushNotificationConfigStore, + normalize_push_notification_config, ) from a2a.types.a2a_pb2 import TaskPushNotificationConfig @@ -283,16 +284,14 @@ 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) - if not config_to_save.id: - config_to_save.id = task_id + config_to_save = normalize_push_notification_config( + task_id, notification_config + ) db_config = self._to_orm(task_id, config_to_save, owner) async with self.async_session_maker.begin() as session: @@ -303,6 +302,7 @@ async def set_info( config_to_save.id, owner, ) + return config_to_save async def _select_configs( self, diff --git a/src/a2a/server/tasks/inmemory_push_notification_config_store.py b/src/a2a/server/tasks/inmemory_push_notification_config_store.py index f8b0b151b..29cbcfc34 100644 --- a/src/a2a/server/tasks/inmemory_push_notification_config_store.py +++ b/src/a2a/server/tasks/inmemory_push_notification_config_store.py @@ -5,6 +5,7 @@ from a2a.server.owner_resolver import OwnerResolver, resolve_user_scope from a2a.server.tasks.push_notification_config_store import ( PushNotificationConfigStore, + normalize_push_notification_config, ) from a2a.types.a2a_pb2 import TaskPushNotificationConfig @@ -40,31 +41,36 @@ 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 = normalize_push_notification_config( + task_id, notification_config + ) + 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, diff --git a/src/a2a/server/tasks/push_notification_config_store.py b/src/a2a/server/tasks/push_notification_config_store.py index e1e65c3fb..06a2726e5 100644 --- a/src/a2a/server/tasks/push_notification_config_store.py +++ b/src/a2a/server/tasks/push_notification_config_store.py @@ -9,6 +9,18 @@ logger = logging.getLogger(__name__) +def normalize_push_notification_config( + task_id: str, notification_config: TaskPushNotificationConfig +) -> TaskPushNotificationConfig: + """Returns a copy with task_id set and an empty id defaulted to the task id.""" + normalized = TaskPushNotificationConfig() + normalized.CopyFrom(notification_config) + normalized.task_id = task_id + if not normalized.id: + normalized.id = task_id + return normalized + + class PushNotificationConfigStore(ABC): """Interface for storing and retrieving push notification configurations for tasks.""" @@ -18,8 +30,14 @@ async def set_info( task_id: str, notification_config: TaskPushNotificationConfig, context: ServerCallContext, - ) -> None: - """Sets or updates the push notification configuration for a task.""" + ) -> TaskPushNotificationConfig | None: + """Sets or updates the push notification configuration for a task. + + Implementations should not mutate notification_config. They store a + copy normalized by normalize_push_notification_config and return it. + Returning None is still accepted for existing implementations; the + request handlers then normalize the request themselves. + """ @abstractmethod async def get_info( diff --git a/tests/server/request_handlers/test_default_request_handler.py b/tests/server/request_handlers/test_default_request_handler.py index 15622e69e..30ad37fff 100644 --- a/tests/server/request_handlers/test_default_request_handler.py +++ b/tests/server/request_handlers/test_default_request_handler.py @@ -1977,6 +1977,87 @@ 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_create_task_push_notification_config_normalizes_when_store_returns_none( + agent_card, +): + """A custom store that still returns None from set_info keeps working.""" + push_config_store = AsyncMock(spec=PushNotificationConfigStore) + push_config_store.set_info.return_value = None + + 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 + ) + + assert (response.task_id, response.id) == (task.id, task.id) + 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.""" diff --git a/tests/server/request_handlers/test_default_request_handler_v2.py b/tests/server/request_handlers/test_default_request_handler_v2.py index a4e5a0e60..e0442f4c0 100644 --- a/tests/server/request_handlers/test_default_request_handler_v2.py +++ b/tests/server/request_handlers/test_default_request_handler_v2.py @@ -513,6 +513,84 @@ 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 ( + 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_create_task_push_notification_config_normalizes_when_store_returns_none(): + """A custom store that still returns None from set_info keeps working.""" + push_config_store = AsyncMock(spec=PushNotificationConfigStore) + push_config_store.set_info.return_value = None + + 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 + ) + + assert (response.task_id, response.id) == (task.id, task.id) + 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.""" @@ -1710,6 +1788,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() diff --git a/tests/server/tasks/test_database_push_notification_config_store.py b/tests/server/tasks/test_database_push_notification_config_store.py index 72cf39c99..5ae9aa012 100644 --- a/tests/server/tasks/test_database_push_notification_config_store.py +++ b/tests/server/tasks/test_database_push_notification_config_store.py @@ -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( @@ -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, @@ -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 @@ -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) @@ -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 @@ -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) @@ -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) @@ -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 diff --git a/tests/server/tasks/test_inmemory_push_notifications.py b/tests/server/tasks/test_inmemory_push_notifications.py index 0a53352f8..8acae7d9d 100644 --- a/tests/server/tasks/test_inmemory_push_notifications.py +++ b/tests/server/tasks/test_inmemory_push_notifications.py @@ -44,8 +44,11 @@ def _create_sample_push_config( url: str = 'http://example.com/callback', config_id: str = 'cfg1', token: str | None = None, + task_id: str = '', ) -> TaskPushNotificationConfig: - return TaskPushNotificationConfig(id=config_id, url=url, token=token) + return TaskPushNotificationConfig( + id=config_id, url=url, token=token, task_id=task_id + ) class SampleUser(User): @@ -103,7 +106,9 @@ def test_constructor_stores_client(self) -> None: async def test_set_info_adds_new_config(self) -> None: task_id = 'task_new' - config = _create_sample_push_config(url='http://new.url/callback') + config = _create_sample_push_config( + task_id=task_id, url='http://new.url/callback' + ) await self.config_store.set_info(task_id, config, MINIMAL_CALL_CONTEXT) @@ -112,17 +117,38 @@ async def test_set_info_adds_new_config(self) -> None: ) self.assertEqual(retrieved, [config]) + async def test_set_info_returns_the_stored_config_without_mutating_input( + self, + ) -> None: + task_id = 'task_normalize' + config = TaskPushNotificationConfig(url='http://normalize.url/callback') + + stored = await self.config_store.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 = await self.config_store.get_info( + task_id, MINIMAL_CALL_CONTEXT + ) + self.assertEqual(retrieved, [stored]) + async def test_set_info_appends_to_existing_config(self) -> None: task_id = 'task_update' initial_config = _create_sample_push_config( - url='http://initial.url/callback', config_id='cfg_initial' + task_id=task_id, + url='http://initial.url/callback', + config_id='cfg_initial', ) await self.config_store.set_info( task_id, initial_config, MINIMAL_CALL_CONTEXT ) updated_config = _create_sample_push_config( - url='http://updated.url/callback', config_id='cfg_updated' + task_id=task_id, + url='http://updated.url/callback', + config_id='cfg_updated', ) await self.config_store.set_info( task_id, updated_config, MINIMAL_CALL_CONTEXT @@ -164,7 +190,9 @@ async def test_set_info_without_config_id(self) -> None: async def test_get_info_existing_config(self) -> None: task_id = 'task_get_exist' - config = _create_sample_push_config(url='http://get.this/callback') + config = _create_sample_push_config( + task_id=task_id, url='http://get.this/callback' + ) await self.config_store.set_info(task_id, config, MINIMAL_CALL_CONTEXT) retrieved_config = await self.config_store.get_info(