diff --git a/src/google/adk/a2a/executor/a2a_agent_executor.py b/src/google/adk/a2a/executor/a2a_agent_executor.py index 11303a8e..a4225a3d 100644 --- a/src/google/adk/a2a/executor/a2a_agent_executor.py +++ b/src/google/adk/a2a/executor/a2a_agent_executor.py @@ -40,6 +40,7 @@ from .a2a_agent_executor_impl import _A2aAgentExecutor as ExecutorImpl from .config import A2aAgentExecutorConfig from .executor_context import ExecutorContext from .task_result_aggregator import TaskResultAggregator +from .utils import _enqueue_canceled_task_event from .utils import execute_after_agent_interceptors from .utils import execute_after_event_interceptors from .utils import execute_before_agent_interceptors @@ -74,7 +75,7 @@ class A2aAgentExecutor(AgentExecutor): self._config = config or A2aAgentExecutorConfig() self._use_legacy = use_legacy self._force_new_version = force_new_version - self._executor_impl = None + self._executor_impl: ExecutorImpl | None = None async def _resolve_runner(self) -> Runner: """Resolve the runner, handling cases where it's a callable that returns a Runner.""" @@ -101,14 +102,11 @@ class A2aAgentExecutor(AgentExecutor): ) @override - async def cancel(self, context: RequestContext, event_queue: EventQueue): + async def cancel( + self, context: RequestContext, event_queue: EventQueue + ) -> None: """Cancel the execution.""" - if self._executor_impl: - await self._executor_impl.cancel(context, event_queue) - return - - # TODO: Implement proper cancellation logic if needed - raise NotImplementedError('Cancellation is not supported') + await _enqueue_canceled_task_event(context, event_queue) @override async def execute( diff --git a/src/google/adk/a2a/executor/a2a_agent_executor_impl.py b/src/google/adk/a2a/executor/a2a_agent_executor_impl.py index b2213de4..ad7f4250 100644 --- a/src/google/adk/a2a/executor/a2a_agent_executor_impl.py +++ b/src/google/adk/a2a/executor/a2a_agent_executor_impl.py @@ -41,6 +41,7 @@ from ..converters.utils import _get_adk_metadata_key from ..experimental import a2a_experimental from .config import A2aAgentExecutorConfig from .executor_context import ExecutorContext +from .utils import _enqueue_canceled_task_event from .utils import execute_after_agent_interceptors from .utils import execute_after_event_interceptors from .utils import execute_before_agent_interceptors @@ -66,10 +67,11 @@ class _A2aAgentExecutor(AgentExecutor): self._config = config or A2aAgentExecutorConfig() @override - async def cancel(self, context: RequestContext, event_queue: EventQueue): + async def cancel( + self, context: RequestContext, event_queue: EventQueue + ) -> None: """Cancel the execution.""" - # TODO: Implement proper cancellation logic if needed - raise NotImplementedError('Cancellation is not supported') + await _enqueue_canceled_task_event(context, event_queue) @override async def execute( diff --git a/src/google/adk/a2a/executor/utils.py b/src/google/adk/a2a/executor/utils.py index d7883c23..54859211 100644 --- a/src/google/adk/a2a/executor/utils.py +++ b/src/google/adk/a2a/executor/utils.py @@ -17,14 +17,36 @@ from typing import Optional from a2a.server.agent_execution.context import RequestContext from a2a.server.events import Event as A2AEvent +from a2a.server.events.event_queue import EventQueue from a2a.types import TaskStatusUpdateEvent +from .. import _compat from ...events.event import Event from ..converters.utils import _get_adk_metadata_key as _get_adk_metadata_key from .config import ExecuteInterceptor from .executor_context import ExecutorContext +async def _enqueue_canceled_task_event( + context: RequestContext, + event_queue: EventQueue, +) -> None: + """Publishes the terminal event required by the A2A cancellation contract.""" + if not context.task_id: + raise ValueError('A2A cancellation must have a task ID') + + # ``final`` is load-bearing on a2a-sdk 0.3.x, which keeps consuming the queue + # until it sees a status update with it set; 1.x has no such field. + await event_queue.enqueue_event( + _compat.make_task_status_update_event( + task_id=context.task_id, + context_id=context.context_id, + status=_compat.make_task_status(_compat.TS_CANCELED), + final=True, + ) + ) + + async def execute_before_agent_interceptors( context: RequestContext, execute_interceptors: Optional[list[ExecuteInterceptor]], diff --git a/tests/unittests/a2a/executor/test_a2a_agent_executor.py b/tests/unittests/a2a/executor/test_a2a_agent_executor.py index 70511dd1..36c4541a 100644 --- a/tests/unittests/a2a/executor/test_a2a_agent_executor.py +++ b/tests/unittests/a2a/executor/test_a2a_agent_executor.py @@ -12,6 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +import asyncio from unittest.mock import AsyncMock from unittest.mock import Mock from unittest.mock import patch @@ -729,22 +730,43 @@ class TestA2aAgentExecutor: """Test cancellation with a task ID.""" self.mock_context.task_id = "test-task-id" - # The current implementation raises NotImplementedError - with pytest.raises( - NotImplementedError, match="Cancellation is not supported" - ): - await self.executor.cancel(self.mock_context, self.mock_event_queue) + await self.executor.cancel(self.mock_context, self.mock_event_queue) + + self.mock_event_queue.enqueue_event.assert_awaited_once() + canceled_event = self.mock_event_queue.enqueue_event.await_args.args[0] + assert canceled_event.task_id == "test-task-id" + assert canceled_event.context_id == "test-context-id" + assert canceled_event.status.state == _compat.TS_CANCELED + _assert_final(canceled_event) @pytest.mark.asyncio async def test_cancel_without_task_id(self): """Test cancellation without a task ID.""" self.mock_context.task_id = None - # The current implementation raises NotImplementedError regardless of task_id - with pytest.raises( - NotImplementedError, match="Cancellation is not supported" - ): + with pytest.raises(ValueError, match="must have a task ID"): await self.executor.cancel(self.mock_context, self.mock_event_queue) + self.mock_event_queue.enqueue_event.assert_not_awaited() + + @pytest.mark.asyncio + async def test_execute_cancelled_does_not_publish_failure(self): + """Test that a cancelled execution is not reported as a failure.""" + self.mock_context.task_id = "test-task-id" + self.mock_context.current_task = None + + self.mock_request_converter.side_effect = asyncio.CancelledError() + + with pytest.raises(asyncio.CancelledError): + await self.executor.execute(self.mock_context, self.mock_event_queue) + + # The cancellation must have been raised inside the guarded region. + self.mock_request_converter.assert_called_once() + states = [ + call.args[0].status.state + for call in self.mock_event_queue.enqueue_event.call_args_list + if hasattr(call.args[0], "status") + ] + assert _compat.TS_FAILED not in states @pytest.mark.asyncio async def test_execute_with_exception_handling(self): diff --git a/tests/unittests/a2a/executor/test_a2a_agent_executor_impl.py b/tests/unittests/a2a/executor/test_a2a_agent_executor_impl.py index ef239149..b797c54f 100644 --- a/tests/unittests/a2a/executor/test_a2a_agent_executor_impl.py +++ b/tests/unittests/a2a/executor/test_a2a_agent_executor_impl.py @@ -14,6 +14,7 @@ from __future__ import annotations +import asyncio from unittest.mock import AsyncMock from unittest.mock import Mock from unittest.mock import patch @@ -424,10 +425,43 @@ class TestA2aAgentExecutor: """Test cancellation with a task ID.""" self.mock_context.task_id = "test-task-id" - with pytest.raises( - NotImplementedError, match="Cancellation is not supported" - ): + await self.executor.cancel(self.mock_context, self.mock_event_queue) + + self.mock_event_queue.enqueue_event.assert_awaited_once() + canceled_event = self.mock_event_queue.enqueue_event.await_args.args[0] + assert canceled_event.task_id == "test-task-id" + assert canceled_event.context_id == "test-context-id" + assert canceled_event.status.state == _compat.TS_CANCELED + _assert_final(canceled_event) + + @pytest.mark.asyncio + async def test_cancel_without_task_id(self): + """Test cancellation without a task ID.""" + self.mock_context.task_id = None + + with pytest.raises(ValueError, match="must have a task ID"): await self.executor.cancel(self.mock_context, self.mock_event_queue) + self.mock_event_queue.enqueue_event.assert_not_awaited() + + @pytest.mark.asyncio + async def test_execute_cancelled_does_not_publish_failure(self): + """Test that a cancelled execution is not reported as a failure.""" + self.mock_context.task_id = "test-task-id" + self.mock_context.current_task = None + + self.mock_request_converter.side_effect = asyncio.CancelledError() + + with pytest.raises(asyncio.CancelledError): + await self.executor.execute(self.mock_context, self.mock_event_queue) + + # The cancellation must have been raised inside the guarded region. + self.mock_request_converter.assert_called_once() + states = [ + call.args[0].status.state + for call in self.mock_event_queue.enqueue_event.call_args_list + if hasattr(call.args[0], "status") + ] + assert _compat.TS_FAILED not in states @pytest.mark.asyncio async def test_execute_with_exception_handling(self):