diff --git a/.release-please-manifest.json b/.release-please-manifest.json index c4ddc7481..53b7bc9c4 100644 --- a/.release-please-manifest.json +++ b/.release-please-manifest.json @@ -1,3 +1,3 @@ { - ".": "1.1.1" + ".": "1.1.2" } diff --git a/CHANGELOG.md b/CHANGELOG.md index 6f064c21a..30ea7908b 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,18 @@ # Changelog +## [1.1.2](https://github.com/a2aproject/a2a-python/compare/v1.1.1...v1.1.2) (2026-07-20) + + +### Features + +* **server:** add aclose() to drain ActiveTask background tasks ([#1101](https://github.com/a2aproject/a2a-python/issues/1101)) ([#1105](https://github.com/a2aproject/a2a-python/issues/1105)) ([9801f46](https://github.com/a2aproject/a2a-python/commit/9801f4637fc0689b461dd3db00968bec496cd7ac)) + + +### Bug Fixes + +* **agent_execution:** resolve ActiveTask 'destroyed but pending' warning during teardown ([#1122](https://github.com/a2aproject/a2a-python/issues/1122)) ([d19c4d2](https://github.com/a2aproject/a2a-python/commit/d19c4d260e375d52175eec4bbd9018edb3eb270b)) +* persist early producer failure as FAILED with originating message ([#1106](https://github.com/a2aproject/a2a-python/issues/1106)) ([4e3d724](https://github.com/a2aproject/a2a-python/commit/4e3d7249e2316859fe8e9b85fd688d6d4525d690)) + ## [1.1.1](https://github.com/a2aproject/a2a-python/compare/v1.1.0...v1.1.1) (2026-07-15) diff --git a/src/a2a/server/agent_execution/active_task.py b/src/a2a/server/agent_execution/active_task.py index b0154c8d6..ea1955fdd 100644 --- a/src/a2a/server/agent_execution/active_task.py +++ b/src/a2a/server/agent_execution/active_task.py @@ -141,7 +141,7 @@ async def run(self) -> None: updated_task = None task = await self.active_task._task_manager.get_task() - if task: + if task and task.status.state not in TERMINAL_TASK_STATES: handled_event = TaskStatusUpdateEvent( task_id=task.id, context_id=task.context_id, @@ -260,11 +260,20 @@ async def _handle_task_modification_event( ) if self.message_to_save is not None: - updated_task = self.active_task._task_manager.update_with_message( - self.message_to_save, - updated_task, + message_already_saved = any( + message.message_id == self.message_to_save.message_id + for message in updated_task.history ) - await self.active_task._task_manager.save_task_event(updated_task) + if not message_already_saved: + updated_task = ( + self.active_task._task_manager.update_with_message( + self.message_to_save, + updated_task, + ) + ) + await self.active_task._task_manager.save_task_event( + updated_task + ) self.message_to_save = None self.active_task._task_manager.context_id = event.context_id @@ -551,12 +560,16 @@ async def _run_producer(self) -> None: 'Producer[%s]: Execution failed', self._task_id, ) - # Create task and mark as failed. + # Persist the failure directly instead of relying on the closing + # event queue to carry a final status update. if request_context: - await self._task_manager.ensure_task_id( + task = await self._task_manager.ensure_task_id( self._task_id, request_context.context_id or '', ) + if task.status.state not in TERMINAL_TASK_STATES: + task.status.state = TaskState.TASK_STATE_FAILED + await self._task_manager.save_task_event(task) self._task_created.set() await self._event_queue_agent.enqueue_event(cast('Event', e)) @@ -574,6 +587,19 @@ async def _run_consumer(self) -> None: self._is_finished.set() self._request_queue.shutdown(immediate=True) await self._event_queue_agent.close(immediate=True) + + if self._producer_task and not self._producer_task.done(): + try: + await self._producer_task + except asyncio.CancelledError: + pass + except Exception as e: # noqa: BLE001 + logger.debug( + 'Consumer[%s]: Awaited producer_task raised %r', + self._task_id, + e, + ) + async with self._lock: self._reference_count -= 1 logger.debug('Consumer[%s]: Finishing', self._task_id) @@ -729,6 +755,48 @@ async def cancel(self, call_context: ServerCallContext) -> Task: raise RuntimeError('Task should have been created') return task + async def aclose(self) -> None: + """Force-closes the task's queues and drains its background tasks. + + Provides a bounded, public teardown for the producer and consumer + ``asyncio.Task``s spawned in ``start()``. Without it, a producer + wedged in its ``finally`` closing an abandoned subscriber sink can + survive until event-loop shutdown and surface as + ``Task was destroyed but it is pending!``. + + Always forces: the queues are closed with ``immediate=True`` and the + background tasks are cancelled, so teardown is bounded even when a + subscriber sink was never drained. It is safe to call multiple times. + """ + # Shut down the request queue first, mirroring the producer's + # ``finally``. If `start()` was never called, the producer/consumer + # ``finally`` blocks never run, and without this a caller parked in + # `enqueue_request()` would wait forever. + self._request_queue.shutdown(immediate=True) + await self._event_queue_agent.close(immediate=True) + await self._event_queue_subscribers.close(immediate=True) + # Set `_is_finished` and collect the background tasks under `_lock` so + # this is mutually exclusive with `start()`, which refuses to spawn + # once `_is_finished` is set. The lock is released before awaiting the + # tasks, because their teardown re-acquires it. + async with self._lock: + self._is_finished.set() + background_tasks = [ + task + for task in (self._producer_task, self._consumer_task) + if task is not None + ] + for task in background_tasks: + task.cancel() + if background_tasks: + results = await asyncio.gather( + *background_tasks, return_exceptions=True + ) + for result in results: + # CancelledError is a BaseException, so it is excluded here. + if isinstance(result, Exception): + logger.error('Error during aclose', exc_info=result) + async def _maybe_cleanup(self) -> None: """Triggers cleanup if task is finished and has no subscribers. diff --git a/src/a2a/server/agent_execution/active_task_registry.py b/src/a2a/server/agent_execution/active_task_registry.py index 9c1299ab3..ab7d6a11c 100644 --- a/src/a2a/server/agent_execution/active_task_registry.py +++ b/src/a2a/server/agent_execution/active_task_registry.py @@ -11,6 +11,7 @@ from a2a.server.context import ServerCallContext from a2a.server.tasks.push_notification_sender import PushNotificationSender from a2a.server.tasks.task_store import TaskStore + from a2a.types.a2a_pb2 import Message from a2a.server.agent_execution.active_task import ActiveTask from a2a.server.tasks.task_manager import TaskManager @@ -34,6 +35,7 @@ def __init__( self._active_tasks: dict[str, ActiveTask] = {} self._lock = asyncio.Lock() self._cleanup_tasks: set[asyncio.Task[None]] = set() + self._closed = False async def get_or_create( self, @@ -41,9 +43,12 @@ async def get_or_create( call_context: ServerCallContext, context_id: str | None = None, create_task_if_missing: bool = False, + initial_message: Message | None = None, ) -> ActiveTask: """Retrieves an existing ActiveTask or creates a new one.""" async with self._lock: + if self._closed: + raise RuntimeError('ActiveTaskRegistry is closed') if task_id in self._active_tasks: return self._active_tasks[task_id] @@ -51,7 +56,7 @@ async def get_or_create( task_id=task_id, context_id=context_id, task_store=self._task_store, - initial_message=None, + initial_message=initial_message, context=call_context, ) @@ -86,3 +91,39 @@ async def get(self, task_id: str) -> ActiveTask | None: """Retrieves an existing task.""" async with self._lock: return self._active_tasks.get(task_id) + + async def aclose(self) -> None: + """Closes the registry and drains all active tasks. + + Marks the registry closed so ``get_or_create`` refuses new work, then + force-closes every registered ``ActiveTask`` and awaits the in-flight + ``_remove_task`` cleanup tasks they schedule, so no SDK-owned + ``asyncio.Task`` is left pending at event-loop shutdown. Safe to call + multiple times. + + The close flag is set and the active-task snapshot is taken under + ``_lock``, and the lock is then released before awaiting, because + ``_remove_task`` re-acquires ``_lock``; holding it while draining + would deadlock. Marking closed under the same lock prevents a + concurrent ``get_or_create`` from registering a task that the drain + would miss. + """ + async with self._lock: + self._closed = True + active_tasks = list(self._active_tasks.values()) + + if active_tasks: + results = await asyncio.gather( + *(task.aclose() for task in active_tasks), + return_exceptions=True, + ) + for result in results: + if isinstance(result, Exception): + logger.error('Error draining active task', exc_info=result) + + cleanup_tasks = list(self._cleanup_tasks) + if cleanup_tasks: + await asyncio.gather(*cleanup_tasks, return_exceptions=True) + + async with self._lock: + self._active_tasks.clear() diff --git a/src/a2a/server/events/event_queue_v2.py b/src/a2a/server/events/event_queue_v2.py index 3386a732e..7f9a53e87 100644 --- a/src/a2a/server/events/event_queue_v2.py +++ b/src/a2a/server/events/event_queue_v2.py @@ -239,6 +239,21 @@ async def close(self, immediate: bool = False) -> None: *(sink.close(immediate=False) for sink in sinks_to_close) ) + # Both branches above cancel the dispatcher task, but the sink gather + # only awaits the sinks. Await the cancelled dispatcher here so it is + # never left pending when close() returns -- e.g. a source created with + # create_default_sink=False and no taps has no sinks for the gather to + # await, which would otherwise leave the cancelled dispatcher pending + # and trigger "Task was destroyed but it is pending!". + # + # Suppress the expected CancelledError and any error the dispatcher + # raised while shutting down: close() is teardown (reachable from + # __aexit__) and must not raise, and _dispatch_loop already logs its own + # exceptions, so a dispatcher failure stays observable without + # propagating out of close(). + with contextlib.suppress(asyncio.CancelledError, Exception): + await self._dispatcher_task + def is_closed(self) -> bool: """[DEPRECATED] Checks if the queue is closed. 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 30304609a..872a3bfa2 100644 --- a/src/a2a/server/request_handlers/default_request_handler_v2.py +++ b/src/a2a/server/request_handlers/default_request_handler_v2.py @@ -112,6 +112,15 @@ def __init__( # noqa: PLR0913 ) self._background_tasks = set() + async def aclose(self) -> None: + """Shuts down the handler, draining all active tasks. + + Drains the ``ActiveTaskRegistry`` so a server shutdown leaves no + pending ``asyncio.Task``. Intended to be wired into an ASGI + ``lifespan`` / ``on_shutdown`` hook. Safe to call multiple times. + """ + await self._active_task_registry.aclose() + @validate_request_params async def on_get_task( # noqa: D102 self, @@ -222,6 +231,7 @@ async def _setup_active_task( context_id=context_id, call_context=call_context, create_task_if_missing=True, + initial_message=request_context.message, ) return active_task, request_context diff --git a/tests/server/agent_execution/test_active_task.py b/tests/server/agent_execution/test_active_task.py index ce9e2c068..1be233ee1 100644 --- a/tests/server/agent_execution/test_active_task.py +++ b/tests/server/agent_execution/test_active_task.py @@ -20,6 +20,7 @@ TaskStatus, TaskStatusUpdateEvent, ) +from a2a.utils._async_queue_compat import QueueShutDown, create_async_queue from a2a.utils.errors import InvalidParamsError @@ -895,3 +896,209 @@ async def execute_mock(req, q): assert len(events) == 0 await active_task.cancel(request_context) + + +@pytest.mark.timeout(5) +@pytest.mark.asyncio +async def test_producer_awaited_on_normal_completion(): + """Verify producer is awaited before cleanup to prevent GC warnings. + + Regression test for #1121: when the consumer finishes first (normal + completion path), the producer must be awaited before + _maybe_cleanup() releases the ActiveTask reference, otherwise asyncio + logs "Task was destroyed but it is pending!". + """ + agent_executor = Mock() + task_manager = Mock() + cleanup_called = False + + def on_cleanup(_task: ActiveTask) -> None: + nonlocal cleanup_called + cleanup_called = True + + active_task = ActiveTask( + agent_executor=agent_executor, + task_id='test-task-id', + task_manager=task_manager, + on_cleanup=on_cleanup, + ) + + execute_started = asyncio.Event() + execute_barrier = asyncio.Event() + + async def execute_mock(req, q): + execute_started.set() + await execute_barrier.wait() + + agent_executor.execute = AsyncMock(side_effect=execute_mock) + agent_executor.cancel = AsyncMock() + task_manager.get_task = AsyncMock( + return_value=Task( + id='test-task-id', + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ) + ) + task_manager.save_task_event = AsyncMock() + task_manager.ensure_task_id = AsyncMock( + return_value=Task( + id='test-task-id', + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ) + ) + task_manager.process = AsyncMock(side_effect=lambda x: x) + + request_context = Mock(spec=RequestContext) + request_context.call_context = ServerCallContext() + request_context.context_id = 'test-context' + request_context.message = None + + await active_task.enqueue_request(request_context) + await active_task.start( + call_context=ServerCallContext(), create_task_if_missing=True + ) + + await execute_started.wait() + + if active_task._consumer_task: + active_task._consumer_task.cancel() + + # Release the producer so it can finish its work and loop back + # to get(), where it will receive QueueShutDown. + execute_barrier.set() + + try: + await active_task._consumer_task + except asyncio.CancelledError: + pass + + await active_task._is_finished.wait() + + assert active_task._producer_task is not None + assert active_task._producer_task.done(), ( + 'Producer task should be done after consumer teardown' + ) + assert cleanup_called, ( + 'on_cleanup should be called after producer is drained' + ) + + +@pytest.mark.timeout(5) +@pytest.mark.asyncio +async def test_active_task_aclose_reaps_background_tasks(): + """aclose() drains a live producer and consumer.""" + agent_executor = Mock() + task_manager = Mock() + request_context = Mock(spec=RequestContext) + + active_task = ActiveTask( + agent_executor=agent_executor, + task_id='test-task-id', + task_manager=task_manager, + push_sender=Mock(), + ) + + async def slow_execute(req, q): + await asyncio.sleep(10) + + agent_executor.execute = AsyncMock(side_effect=slow_execute) + task_manager.get_task = AsyncMock( + return_value=Task( + id='test-task-id', + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ) + ) + + await active_task.enqueue_request(request_context) + await active_task.start( + call_context=ServerCallContext(), create_task_if_missing=True + ) + + await active_task.aclose() + + assert active_task._producer_task is not None + assert active_task._producer_task.done() + assert active_task._consumer_task is not None + assert active_task._consumer_task.done() + assert active_task._is_finished.is_set() + + +@pytest.mark.timeout(5) +@pytest.mark.asyncio +async def test_active_task_aclose_force_closes_undrained_subscriber(): + """aclose() unblocks past an undrained subscriber sink. + + Reproduces issue #1101: a graceful close(immediate=False) would block + forever on the leaked sink's join(). + """ + agent_executor = Mock() + task_manager = Mock() + request_context = Mock(spec=RequestContext) + + active_task = ActiveTask( + agent_executor=agent_executor, + task_id='test-task-id', + task_manager=task_manager, + push_sender=Mock(), + ) + + async def slow_execute(req, q): + await asyncio.sleep(10) + + agent_executor.execute = AsyncMock(side_effect=slow_execute) + task_manager.get_task = AsyncMock( + return_value=Task( + id='test-task-id', + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ) + ) + + await active_task.enqueue_request(request_context) + await active_task.start( + call_context=ServerCallContext(), create_task_if_missing=True + ) + + # Leak a subscriber sink and push an event into it without draining it. + leaked = await active_task._event_queue_subscribers.tap() + await active_task._event_queue_subscribers.enqueue_event(Message()) + await asyncio.sleep(0.05) + + await active_task.aclose() + + assert active_task._producer_task is not None + assert active_task._producer_task.done() + assert leaked.is_closed() + + +@pytest.mark.timeout(5) +@pytest.mark.asyncio +async def test_active_task_aclose_never_started_shuts_down_request_queue(): + """aclose() on a never-started task shuts down the request queue. + + If start() was never called, the producer and consumer `finally` blocks + that normally shut down `_request_queue` never run, so without this a + caller parked in `enqueue_request()` would never be released. + """ + active_task = ActiveTask( + agent_executor=Mock(), + task_id='test-task-id', + task_manager=Mock(), + push_sender=Mock(), + ) + + # Swap in a bounded queue and fill it so the next enqueue_request parks. + active_task._request_queue = create_async_queue(maxsize=1) + await active_task.enqueue_request(Mock(spec=RequestContext)) + waiter = asyncio.create_task( + active_task.enqueue_request(Mock(spec=RequestContext)) + ) + await asyncio.sleep(0.05) + assert not waiter.done() + + await active_task.aclose() + + # The parked waiter is released with the shutdown error, and any + # subsequent enqueue fails fast instead of hanging. + with pytest.raises(QueueShutDown): + await waiter + with pytest.raises(QueueShutDown): + await active_task.enqueue_request(Mock(spec=RequestContext)) diff --git a/tests/server/agent_execution/test_active_task_registry.py b/tests/server/agent_execution/test_active_task_registry.py new file mode 100644 index 000000000..16d9c8797 --- /dev/null +++ b/tests/server/agent_execution/test_active_task_registry.py @@ -0,0 +1,107 @@ +import asyncio +import logging + +from unittest.mock import AsyncMock + +import pytest + +from a2a.server.agent_execution.active_task_registry import ActiveTaskRegistry +from a2a.server.agent_execution.agent_executor import AgentExecutor +from a2a.server.agent_execution.context import RequestContext +from a2a.server.context import ServerCallContext +from a2a.server.events.event_queue_v2 import EventQueue +from a2a.server.tasks import InMemoryTaskStore + + +class _SlowExecutor(AgentExecutor): + """An executor whose execute() blocks until cancelled.""" + + async def execute( + self, context: RequestContext, event_queue: EventQueue + ) -> None: + await asyncio.sleep(10) + + async def cancel( + self, context: RequestContext, event_queue: EventQueue + ) -> None: + return None + + +def _make_registry() -> ActiveTaskRegistry: + return ActiveTaskRegistry( + agent_executor=_SlowExecutor(), + task_store=InMemoryTaskStore(), + ) + + +@pytest.mark.timeout(5) +@pytest.mark.asyncio +async def test_aclose_reaps_active_tasks_and_empties_registry(): + """aclose() reaps background tasks and removes them.""" + registry = _make_registry() + active = await registry.get_or_create( + 'task-1', + call_context=ServerCallContext(), + create_task_if_missing=True, + ) + + await registry.aclose() + + assert active._producer_task is not None + assert active._producer_task.done() + assert active._consumer_task is not None + assert active._consumer_task.done() + assert await registry.get('task-1') is None + + +@pytest.mark.timeout(5) +@pytest.mark.asyncio +async def test_aclose_is_idempotent(): + """Calling aclose() repeatedly is a safe no-op.""" + registry = _make_registry() + await registry.get_or_create( + 'task-1', + call_context=ServerCallContext(), + create_task_if_missing=True, + ) + + await registry.aclose() + await registry.aclose() + + +@pytest.mark.timeout(5) +@pytest.mark.asyncio +async def test_aclose_on_empty_registry(): + """aclose() with no active tasks returns immediately.""" + registry = _make_registry() + await registry.aclose() + + +@pytest.mark.timeout(5) +@pytest.mark.asyncio +async def test_get_or_create_rejected_after_aclose(): + """A closed registry refuses to create new tasks (no orphan race).""" + registry = _make_registry() + await registry.aclose() + + with pytest.raises(RuntimeError): + await registry.get_or_create( + 'task-1', + call_context=ServerCallContext(), + create_task_if_missing=True, + ) + + +@pytest.mark.timeout(5) +@pytest.mark.asyncio +async def test_aclose_logs_and_swallows_task_errors(caplog): + """A failing ActiveTask.aclose is logged, not propagated.""" + registry = _make_registry() + failing = AsyncMock() + failing.aclose = AsyncMock(side_effect=ValueError('boom')) + registry._active_tasks['bad'] = failing + + with caplog.at_level(logging.ERROR): + await registry.aclose() + + assert 'Error draining active task' in caplog.text diff --git a/tests/server/events/test_event_queue_v2.py b/tests/server/events/test_event_queue_v2.py index 27bceea4c..5e93ab870 100644 --- a/tests/server/events/test_event_queue_v2.py +++ b/tests/server/events/test_event_queue_v2.py @@ -592,6 +592,54 @@ async def test_dispatch_task_failed(event_queue: EventQueueSource) -> None: await asyncio.wait_for(event_queue.close(immediate=False), timeout=0.1) +@pytest.mark.asyncio +async def test_close_immediate_awaits_dispatcher_without_sinks() -> None: + """A subscriber-less source must not leave its dispatcher pending on close. + + Regression: a source created with create_default_sink=False and no taps has + no sinks for close()'s gather to await. close(immediate=True) cancels the + dispatcher task, so with nothing else to await it used to return while the + cancelled dispatcher was still pending, producing "Task was destroyed but it + is pending!". close() must await its own cancelled dispatcher task. + """ + source = EventQueueSource(create_default_sink=False) + assert not source._dispatcher_task.done() + + await source.close(immediate=True) + + assert source._dispatcher_task.done() + + +@pytest.mark.asyncio +async def test_close_does_not_propagate_dispatcher_crash() -> None: + """close() awaits its dispatcher without surfacing a non-cancel crash. + + close() is teardown (reachable from __aexit__) and must not raise, and + _dispatch_loop already logs its own exceptions. If the dispatcher died with a + non-cancel error, awaiting it in close() must not re-raise that error. + """ + source = EventQueueSource(create_default_sink=False) + + # Retire the real dispatcher, then stand in a task that has already failed + # with a non-cancel error, simulating a _dispatch_loop crash. + real = source._dispatcher_task + real.cancel() + try: + await real + except asyncio.CancelledError: + pass + + async def crashed() -> None: + raise RuntimeError('dispatcher boom') + + source._dispatcher_task = asyncio.ensure_future(crashed()) + await asyncio.sleep(0.01) # let it fail so the exception is stored + + # Must return cleanly, not re-raise the dispatcher's RuntimeError. + await source.close(immediate=True) + assert source._dispatcher_task.done() + + @pytest.mark.asyncio async def test_concurrent_close_immediate_false() -> None: """Test that concurrent close(immediate=False) calls both wait for join() deterministically.""" 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 caaa4f88e..b276fb77a 100644 --- a/tests/server/request_handlers/test_default_request_handler_v2.py +++ b/tests/server/request_handlers/test_default_request_handler_v2.py @@ -21,6 +21,7 @@ from a2a.server.agent_execution.active_task_registry import ActiveTaskRegistry from a2a.server.context import ServerCallContext from a2a.server.events import EventQueue +from a2a.server.events.event_queue_v2 import EventQueueSource from a2a.server.request_handlers import DefaultRequestHandlerV2 from a2a.server.tasks import ( InMemoryPushNotificationConfigStore, @@ -30,6 +31,7 @@ TaskStore, TaskUpdater, ) +from a2a.server.tasks.task_manager import TaskManager from a2a.types import ( InternalError, InvalidAgentResponseError, @@ -306,6 +308,48 @@ async def cancel(self, context: RequestContext, event_queue: EventQueue): pass +class EarlyFailingAgentExecutor(AgentExecutor): + async def execute(self, context: RequestContext, event_queue: EventQueue): + raise RuntimeError('early producer failure') + + async def cancel(self, context: RequestContext, event_queue: EventQueue): + pass + + +class LateFailingTerminalAgentExecutor(AgentExecutor): + def __init__( + self, terminal_state: TaskState, terminal_state_persisted: asyncio.Event + ) -> None: + self.terminal_state = terminal_state + self.terminal_state_persisted = terminal_state_persisted + self.raised = asyncio.Event() + + async def execute(self, context: RequestContext, event_queue: EventQueue): + assert context.message is not None + task = new_task_from_user_message(context.message) + await event_queue.enqueue_event(task) + task_updater = TaskUpdater(event_queue, task.id, task.context_id) + await task_updater.update_status(self.terminal_state) + await self.terminal_state_persisted.wait() + self.raised.set() + raise RuntimeError('late producer failure') + + async def cancel(self, context: RequestContext, event_queue: EventQueue): + pass + + +async def send_message_with_early_failure( + request_handler: DefaultRequestHandlerV2, + params: SendMessageRequest, + context: ServerCallContext, +) -> Message | Task | None: + try: + return await request_handler.on_message_send(params, context) + except RuntimeError as e: + assert str(e) == 'early producer failure' + return None + + @pytest.mark.asyncio async def test_on_get_task_limit_history(): task_store = InMemoryTaskStore() @@ -1126,6 +1170,117 @@ async def test_on_message_send_limit_history(): assert task.history is not None and len(task.history) > 1 +@pytest.mark.asyncio +async def test_on_message_send_early_producer_exception_marks_task_failed_and_preserves_originating_message(): + task_store = InMemoryTaskStore() + request_handler = DefaultRequestHandlerV2( + agent_executor=HelloAgentExecutor(), + task_store=task_store, + agent_card=create_default_agent_card(), + ) + params = SendMessageRequest( + message=Message( + role=Role.ROLE_USER, + message_id='msg_early_failure_state', + parts=[Part(text='Hi')], + ) + ) + context = create_server_call_context() + original_enqueue_event = EventQueueSource.enqueue_event + + async def fail_before_request_started(self, event): + if type(event).__name__ == '_RequestStarted': + raise RuntimeError('early producer failure') + return await original_enqueue_event(self, event) + + with patch.object( + EventQueueSource, 'enqueue_event', fail_before_request_started + ): + await send_message_with_early_failure(request_handler, params, context) + + stored_task = await task_store.get(params.message.task_id, context) + assert stored_task is not None + assert stored_task.status.state == TaskState.TASK_STATE_FAILED + assert len(stored_task.history) == 1 + assert stored_task.history[0].message_id == 'msg_early_failure_state' + assert stored_task.history[0].parts[0].text == 'Hi' + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + 'terminal_state', + [ + TaskState.TASK_STATE_COMPLETED, + TaskState.TASK_STATE_CANCELED, + ], +) +async def test_on_message_send_late_producer_exception_preserves_persisted_terminal_state( + terminal_state: TaskState, +): + task_store = InMemoryTaskStore() + terminal_state_persisted = asyncio.Event() + agent_executor = LateFailingTerminalAgentExecutor( + terminal_state, terminal_state_persisted + ) + request_handler = DefaultRequestHandlerV2( + agent_executor=agent_executor, + task_store=task_store, + agent_card=create_default_agent_card(), + ) + params = SendMessageRequest( + message=Message( + role=Role.ROLE_USER, + message_id='msg_late_failure_terminal_state', + parts=[Part(text='Hi')], + ) + ) + context = create_server_call_context() + original_save_task = TaskManager._save_task + + async def save_task_and_signal_terminal_state(self, task): + await original_save_task(self, task) + if task.status.state == terminal_state: + terminal_state_persisted.set() + + with patch.object( + TaskManager, '_save_task', save_task_and_signal_terminal_state + ): + await request_handler.on_message_send(params, context) + await agent_executor.raised.wait() + await asyncio.sleep(0) + + stored_task = await task_store.get(params.message.task_id, context) + assert stored_task is not None + assert stored_task.status.state == terminal_state + + +@pytest.mark.asyncio +async def test_on_message_send_early_producer_exception_preserves_originating_message(): + task_store = InMemoryTaskStore() + request_handler = DefaultRequestHandlerV2( + agent_executor=EarlyFailingAgentExecutor(), + task_store=task_store, + agent_card=create_default_agent_card(), + ) + params = SendMessageRequest( + message=Message( + role=Role.ROLE_USER, + message_id='msg_early_failure_history', + parts=[Part(text='Hi')], + ) + ) + context = create_server_call_context() + + await send_message_with_early_failure(request_handler, params, context) + + stored_task = await task_store.get(params.message.task_id, context) + assert stored_task is not None + assert stored_task.history is not None + assert len(stored_task.history) == 1 + assert stored_task.history[0].message_id == 'msg_early_failure_history' + assert stored_task.history[0].parts[0].text == 'Hi' + + @pytest.mark.asyncio async def test_on_message_send_stream_task_id_mismatch(): mock_task_store = AsyncMock(spec=TaskStore) @@ -1557,3 +1712,37 @@ async def test_on_get_task_push_notification_config_is_owner_scoped(): ), _ctx('bob'), ) + + +@pytest.mark.timeout(5) +@pytest.mark.asyncio +async def test_aclose_drains_registry(): + """aclose() drains the active-task registry on shutdown.""" + handler = DefaultRequestHandlerV2( + agent_executor=MockAgentExecutor(), + task_store=InMemoryTaskStore(), + agent_card=create_default_agent_card(), + ) + await handler._active_task_registry.get_or_create( + 'task-1', + call_context=ServerCallContext(user=UnauthenticatedUser()), + create_task_if_missing=True, + ) + + await handler.aclose() + + assert await handler._active_task_registry.get('task-1') is None + + +@pytest.mark.timeout(5) +@pytest.mark.asyncio +async def test_aclose_is_idempotent_and_handles_empty(): + """aclose() is safe with no active tasks and when called twice.""" + handler = DefaultRequestHandlerV2( + agent_executor=MockAgentExecutor(), + task_store=InMemoryTaskStore(), + agent_card=create_default_agent_card(), + ) + + await handler.aclose() + await handler.aclose()