diff --git a/src/a2a/server/agent_execution/active_task.py b/src/a2a/server/agent_execution/active_task.py index 720d18bff..4074ba12d 100644 --- a/src/a2a/server/agent_execution/active_task.py +++ b/src/a2a/server/agent_execution/active_task.py @@ -48,12 +48,15 @@ from collections.abc import AsyncGenerator, Callable from a2a.server.agent_execution.agent_executor import AgentExecutor + from a2a.server.cluster.event_stream import TaskEventStream from a2a.server.context import ServerCallContext from a2a.server.tasks.push_notification_sender import ( PushNotificationSender, ) from a2a.server.tasks.task_manager import TaskManager +from a2a.server.cluster.event_stream import VersionedEvent +from a2a.server.cluster.task_store import ConcurrentTaskModificationError from a2a.server.events.event_queue_v2 import ( AsyncQueue, Event, @@ -154,6 +157,29 @@ async def run(self) -> None: await self._enqueue_to_subscribers(cast('Event', e), updated_task) async def _process_event(self, event: Event) -> None: + try: + await self._process_event_inner(event) + except ConcurrentTaskModificationError: + # Another writer advanced this task + logger.info( + 'Consumer[%s]: concurrent modification, reloading', + self.active_task._task_id, + ) + self.active_task._task_manager.invalidate() + task = await self.active_task._task_manager.get_task() + if task is not None and task.status.state in TERMINAL_TASK_STATES: + await self._enqueue_to_subscribers(task, task) + producer = self.active_task._producer_task + if producer is not None and not producer.done(): + producer.cancel() + await self._handle_terminal_state(task) + await self.active_task._event_queue_subscribers.close( + immediate=False + ) + return + await self._process_event_inner(event) + + async def _process_event_inner(self, event: Event) -> None: updated_task = None handled_event: ( Task @@ -333,6 +359,16 @@ async def _enqueue_to_subscribers( await self.active_task._event_queue_subscribers.enqueue_event( cast('Any', (event, updated_task)) ) + + # Fan out to other replicas + stream = self.active_task._event_stream + if stream is not None and isinstance(event, Event): + version = self.active_task._task_manager.current_version + await stream.publish( + self.active_task._task_id, + VersionedEvent(event=event, version=version), + ) + self.active_task._event_queue_agent.task_done() @@ -354,13 +390,14 @@ class ActiveTask: permanently ceased execution and closed its queues. """ - def __init__( + def __init__( # noqa: PLR0913 self, agent_executor: AgentExecutor, task_id: str, task_manager: TaskManager, push_sender: PushNotificationSender | None = None, on_cleanup: Callable[[ActiveTask], None] | None = None, + event_stream: TaskEventStream | None = None, ) -> None: """Initializes the ActiveTask. @@ -372,6 +409,8 @@ def __init__( on_cleanup: Optional callback triggered when the task is fully finished and the last subscriber has disconnected. Used to prune the task from the ActiveTaskRegistry. + event_stream: Optional cross-replica stream; applied events are + published to it so other replicas can observe them. """ # --- Core Dependencies --- self._agent_executor = agent_executor @@ -383,6 +422,7 @@ def __init__( self._task_manager = task_manager self._push_sender = push_sender self._on_cleanup = on_cleanup + self._event_stream = event_stream # --- Synchronization Primitives --- # `_lock` protects structural lifecycle changes: start(), subscribe() counting, @@ -419,6 +459,15 @@ def task_id(self) -> str: """The ID of the task.""" return self._task_id + @property + def has_running_execution(self) -> bool: + """Whether this replica is actively executing the agent for this task.""" + return ( + self._producer_task is not None + and not self._producer_task.done() + and self._request_lock.locked() + ) + async def enqueue_request( self, request_context: RequestContext ) -> uuid.UUID: @@ -520,6 +569,9 @@ async def _run_producer(self) -> None: # TODO: Should we create task manager every time? self._task_manager._call_context = request_context.call_context + # Drop the cached snapshot and re-read to pick + # up state another replica may have advanced. + self._task_manager.invalidate() request_context.current_task = ( await self._task_manager.get_task() ) @@ -781,6 +833,7 @@ async def cancel(self, call_context: ServerCallContext) -> Task: ) await self._is_finished.wait() + self._task_manager.invalidate() task = await self._task_manager.get_task() if not task: raise RuntimeError('Task should have been created') diff --git a/src/a2a/server/agent_execution/active_task_registry.py b/src/a2a/server/agent_execution/active_task_registry.py index 15e9c6350..cae2a55ce 100644 --- a/src/a2a/server/agent_execution/active_task_registry.py +++ b/src/a2a/server/agent_execution/active_task_registry.py @@ -9,12 +9,15 @@ if TYPE_CHECKING: from a2a.server.agent_execution.agent_executor import AgentExecutor + from a2a.server.cluster.event_stream import TaskEventStream + from a2a.server.cluster.task_store import VersionedTaskStore 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.cluster.task_store import StoredTask from a2a.server.tasks.task_manager import TaskManager from a2a.utils.errors import TaskNotFoundError @@ -28,12 +31,14 @@ class ActiveTaskRegistry: def __init__( self, agent_executor: AgentExecutor, - task_store: TaskStore, + task_store: TaskStore | VersionedTaskStore, push_sender: PushNotificationSender | None = None, + event_stream: TaskEventStream | None = None, ): self._agent_executor = agent_executor self._task_store = task_store self._push_sender = push_sender + self._event_stream = event_stream self._active_tasks: dict[str, ActiveTask] = {} self._lock = threading.RLock() self._cleanup_tasks: set[asyncio.Task[None]] = set() @@ -67,6 +72,7 @@ async def get_or_create( task_manager=task_manager, push_sender=self._push_sender, on_cleanup=self._on_active_task_cleanup, + event_stream=self._event_stream, ) self._active_tasks[task_id] = active_task @@ -84,10 +90,16 @@ async def get_or_create( # ownership. Masked as not-found so existence is not leaked. Done # outside _lock because the store read is I/O and the miss-path # check runs outside the lock too. - if not create_task_if_missing and not await self._task_store.get( - task_id, call_context - ): - raise TaskNotFoundError + if not create_task_if_missing: + existing_task = await self._task_store.get( + task_id, call_context + ) + # A VersionedTaskStore returns a StoredTask; a plain store + # returns the task (or None). Normalize before the check. + if isinstance(existing_task, StoredTask): + existing_task = existing_task.task + if not existing_task: + raise TaskNotFoundError return existing await active_task.start( 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 59996236e..cd9bb9e24 100644 --- a/src/a2a/server/request_handlers/default_request_handler_v2.py +++ b/src/a2a/server/request_handlers/default_request_handler_v2.py @@ -4,6 +4,7 @@ import logging import warnings +from contextlib import aclosing from typing import TYPE_CHECKING, Any, cast from a2a.server.agent_execution import ( @@ -17,6 +18,12 @@ TERMINAL_TASK_STATES, ) from a2a.server.agent_execution.active_task_registry import ActiveTaskRegistry +from a2a.server.cluster.event_stream import VersionedEvent +from a2a.server.cluster.task_store import ( + ConcurrentTaskModificationError, + LegacyTaskStoreAdapter, + VersionedTaskStore, +) from a2a.server.request_handlers.request_handler import ( RequestHandler, validate, @@ -39,6 +46,7 @@ Task, TaskPushNotificationConfig, TaskState, + TaskStatusUpdateEvent, ) from a2a.utils.errors import ( ExtendedAgentCardNotConfiguredError, @@ -60,6 +68,8 @@ from collections.abc import AsyncGenerator, Awaitable, Callable from a2a.server.agent_execution.active_task import ActiveTask + from a2a.server.cluster.event_stream import TaskEventStream + from a2a.server.cluster.version import TaskVersion from a2a.server.context import ServerCallContext from a2a.server.events import Event from a2a.server.tasks import ( @@ -87,7 +97,7 @@ class DefaultRequestHandlerV2(RequestHandler): def __init__( # noqa: PLR0913 self, agent_executor: AgentExecutor, - task_store: TaskStore, + task_store: TaskStore | VersionedTaskStore, agent_card: AgentCard, queue_manager: Any | None = None, # Accepted for signature compat; ignored in v2 (warns) @@ -100,6 +110,7 @@ def __init__( # noqa: PLR0913 ] | None = None, push_url_validator: Callable[[str], Awaitable[bool]] | None = None, + event_stream: TaskEventStream | None = None, ) -> None: if queue_manager is not None: message = ( @@ -113,23 +124,38 @@ def __init__( # noqa: PLR0913 warnings.warn(message, DeprecationWarning, stacklevel=2) logger.warning(message) self.agent_executor = agent_executor - self.task_store = task_store + self.task_store: TaskStore | VersionedTaskStore = task_store + self._versioned_store: VersionedTaskStore = ( + task_store + if isinstance(task_store, VersionedTaskStore) + else LegacyTaskStoreAdapter(task_store) + ) self._agent_card = agent_card self._push_config_store = push_config_store self._push_sender = push_sender self._push_url_validator = push_url_validator self.extended_agent_card = extended_agent_card self.extended_card_modifier = extended_card_modifier + self._event_stream = event_stream + if isinstance(task_store, VersionedTaskStore) and event_stream is None: + message = ( + 'A VersionedTaskStore was configured without an event_stream, ' + 'so cross-replica streaming is disabled' + ) + warnings.warn(message, stacklevel=2) + logger.warning(message) self._request_context_builder = ( request_context_builder or SimpleRequestContextBuilder( - should_populate_referred_tasks=False, task_store=self.task_store + should_populate_referred_tasks=False, + task_store=None, ) ) self._active_task_registry = ActiveTaskRegistry( agent_executor=self.agent_executor, - task_store=self.task_store, + task_store=self._versioned_store, push_sender=self._push_sender, + event_stream=self._event_stream, ) self._background_tasks = set() @@ -158,11 +184,11 @@ async def on_get_task( # noqa: D102 validate_history_length(params) task_id = params.id - task: Task | None = await self.task_store.get(task_id, context) - if not task: + stored = await self._versioned_store.get(task_id, context) + if stored is None: raise TaskNotFoundError - return apply_history_length(task, params) + return apply_history_length(stored.task, params) @validate_request_params async def on_list_tasks( # noqa: D102 @@ -174,7 +200,7 @@ async def on_list_tasks( # noqa: D102 if params.HasField('page_size'): validate_page_size(params.page_size) - page = await self.task_store.list(params, context) + page = await self._versioned_store.list(params, context) for task in page.tasks: if not params.include_artifacts: task.ClearField('artifacts') @@ -193,20 +219,82 @@ async def on_cancel_task( # noqa: D102 ) -> Task | None: task_id = params.id + # Owner-scoped read of the current state + stored = await self._versioned_store.get(task_id, context) + if stored is None: + raise TaskNotFoundError + task, version = stored.task, stored.version + + # Check whether already cancelled + if task.status.state == TaskState.TASK_STATE_CANCELED: + return task + if task.status.state in TERMINAL_TASK_STATES: + raise TaskNotCancelableError + + # Fast path: this replica is running the agent -> stop it directly. + local = await self._active_task_registry.get(task_id) + if local is not None and local.has_running_execution: + try: + result = await local.cancel(context) + except InvalidParamsError as e: + raise TaskNotCancelableError from e + if isinstance(result, Message): + raise InternalError( + message='Cancellation returned a message instead of a task.' + ) + return result + + # Running on another replica (or nowhere): record CANCELED in shared + # state; the owner sees the version bump on its next save and aborts. + return await self._cancel_remote(task_id, task, version, context) + + async def _cancel_remote( + self, + task_id: str, + task: Task, + version: TaskVersion, + context: ServerCallContext, + ) -> Task: try: - active_task = await self._active_task_registry.get_or_create( - task_id, call_context=context, create_task_if_missing=False - ) - result = await active_task.cancel(context) - except InvalidParamsError as e: - raise TaskNotCancelableError from e + return await self._write_cancel(task_id, task, version, context) + except ConcurrentTaskModificationError: + reloaded_stored = await self._versioned_store.get(task_id, context) + if reloaded_stored is None: + raise TaskNotFoundError from None + if ( + reloaded_stored.task.status.state + == TaskState.TASK_STATE_CANCELED + ): + return reloaded_stored.task + raise TaskNotCancelableError from None - if isinstance(result, Message): - raise InternalError( - message='Cancellation returned a message instead of a task.' + async def _write_cancel( + self, + task_id: str, + current: Task, + version: TaskVersion, + context: ServerCallContext, + ) -> Task: + cancelled = Task() + cancelled.CopyFrom(current) + cancelled.status.state = TaskState.TASK_STATE_CANCELED + event = TaskStatusUpdateEvent( + task_id=task_id, + context_id=current.context_id, + status=cancelled.status, + ) + new_version = await self._versioned_store.save( + cancelled, + event=event, + prev=current, + prev_version=version, + context=context, + ) + if self._event_stream is not None: + await self._event_stream.publish( + task_id, VersionedEvent(event=event, version=new_version) ) - - return result + return cancelled def _validate_task_id_match(self, task_id: str, event_task_id: str) -> None: if task_id != event_task_id: @@ -227,10 +315,11 @@ async def _setup_active_task( original_task_id = params.message.task_id or None original_context_id = params.message.context_id or None - if original_task_id: - task = await self.task_store.get(original_task_id, call_context) - if not task: - raise TaskNotFoundError(f'Task {original_task_id} not found') + if original_task_id and ( + await self._versioned_store.get(original_task_id, call_context) + is None + ): + raise TaskNotFoundError(f'Task {original_task_id} not found') # Build context to resolve or generate missing IDs request_context = await self._request_context_builder.build( @@ -379,8 +468,7 @@ async def on_create_task_push_notification_config( # noqa: D102 raise PushNotificationNotSupportedError task_id = params.task_id - task: Task | None = await self.task_store.get(task_id, context) - if not task: + if await self._versioned_store.get(task_id, context) is None: raise TaskNotFoundError await self._reject_unsafe_push_url(params.url) @@ -409,8 +497,7 @@ async def on_get_task_push_notification_config( # noqa: D102 task_id = params.task_id config_id = params.id - task: Task | None = await self.task_store.get(task_id, context) - if not task: + if await self._versioned_store.get(task_id, context) is None: raise TaskNotFoundError push_notification_configs: list[TaskPushNotificationConfig] = ( @@ -435,15 +522,68 @@ async def on_subscribe_to_task( # noqa: D102 ) -> AsyncGenerator[Event, None]: task_id = params.id - active_task = await self._active_task_registry.get_or_create( - task_id, - call_context=context, - create_task_if_missing=False, - ) + stored = await self._versioned_store.get(task_id, context) + if stored is None: + raise TaskNotFoundError + task, snapshot_version = stored.task, stored.version + + # A terminal task cannot be resubscribed to (nothing further to stream). + if task.status.state in TERMINAL_TASK_STATES: + raise InvalidParamsError( + message=f'Task {task_id} is in terminal state: ' + f'{task.status.state}' + ) + + if self._event_stream is None: + # Single-process mode + active = await self._active_task_registry.get_or_create( + task_id, call_context=context, create_task_if_missing=False + ) + async for event in active.subscribe(include_initial_task=True): + yield event + return - async for event in active_task.subscribe(include_initial_task=True): + # Shared-stream mode. Fast path: this replica runs the agent -> tap it. + stream = self._event_stream + local = await self._active_task_registry.get(task_id) + if local is not None: + async for event in local.subscribe(include_initial_task=True): + yield event + return + + # Not running here: serve the snapshot and tail the remote stream. + async for event in self._subscribe_remote( + task_id, task, snapshot_version, stream + ): yield event + async def _subscribe_remote( + self, + task_id: str, + task: Task, + snapshot_version: TaskVersion, + stream: TaskEventStream, + ) -> AsyncGenerator[Event, None]: + """Serves a resubscription for a task running on another replica.""" + yield task + + async with aclosing( + stream.subscribe(task_id, after=snapshot_version) + ) as subscription: + async for versioned in subscription: + if not versioned.version.is_after(snapshot_version): + continue + event = versioned.event + yield event + # Stop tailing once the stream ends: a Message, or a status + # update / Task in a terminal or interrupted state. + if isinstance(event, Message) or ( + isinstance(event, TaskStatusUpdateEvent | Task) + and event.status.state + in (TERMINAL_TASK_STATES | INTERRUPTED_TASK_STATES) + ): + return + @validate_request_params @validate( lambda self: self._agent_card.capabilities.push_notifications, @@ -459,8 +599,7 @@ async def on_list_task_push_notification_configs( # noqa: D102 raise PushNotificationNotSupportedError task_id = params.task_id - task: Task | None = await self.task_store.get(task_id, context) - if not task: + if await self._versioned_store.get(task_id, context) is None: raise TaskNotFoundError push_notification_config_list = await self._push_config_store.get_info( @@ -487,8 +626,7 @@ async def on_delete_task_push_notification_config( # noqa: D102 task_id = params.task_id config_id = params.id - task: Task | None = await self.task_store.get(task_id, context) - if not task: + if await self._versioned_store.get(task_id, context) is None: raise TaskNotFoundError await self._push_config_store.delete_info(task_id, context, config_id) diff --git a/src/a2a/server/tasks/task_manager.py b/src/a2a/server/tasks/task_manager.py index c9dfc879f..152083459 100644 --- a/src/a2a/server/tasks/task_manager.py +++ b/src/a2a/server/tasks/task_manager.py @@ -1,8 +1,9 @@ +from __future__ import annotations + import logging -from a2a.server.context import ServerCallContext -from a2a.server.events.event_queue import Event -from a2a.server.tasks.task_store import TaskStore +from typing import TYPE_CHECKING + from a2a.types.a2a_pb2 import ( Artifact, Message, @@ -16,6 +17,14 @@ from a2a.utils.telemetry import trace_function +if TYPE_CHECKING: + from a2a.server.cluster.task_store import VersionedTaskStore + from a2a.server.cluster.version import TaskVersion + from a2a.server.context import ServerCallContext + from a2a.server.events.event_queue import Event + from a2a.server.tasks.task_store import TaskStore + + logger = logging.getLogger(__name__) @@ -96,7 +105,7 @@ class TaskManager: def __init__( self, - task_store: TaskStore, + task_store: TaskStore | VersionedTaskStore, context: ServerCallContext, task_id: str | None, context_id: str | None, @@ -105,7 +114,8 @@ def __init__( """Initializes the TaskManager. Args: - task_store: The `TaskStore` instance for persistence. + task_store: The `TaskStore` (or `VersionedTaskStore`) for + persistence. context: The `ServerCallContext` that this task is produced under. task_id: The ID of the task, if known from the request. context_id: The ID of the context, if known from the request. @@ -115,18 +125,52 @@ def __init__( if task_id is not None and not (isinstance(task_id, str) and task_id): raise ValueError('Task ID must be a non-empty string') + # Imported lazily to avoid an import cycle: the cluster package imports + # from a2a.server.tasks, which imports this module. + from a2a.server.cluster.task_store import ( # noqa: PLC0415 + LegacyTaskStoreAdapter, + VersionedTaskStore, + ) + from a2a.server.cluster.version import TaskVersion # noqa: PLC0415 + self.task_store = task_store + self._versioned_store: VersionedTaskStore = ( + task_store + if isinstance(task_store, VersionedTaskStore) + else LegacyTaskStoreAdapter(task_store) + ) self._call_context: ServerCallContext = context self.task_id = task_id self.context_id = context_id self._initial_message = initial_message self._current_task: Task | None = None + # Version of `_current_task` as last read/written. MISSING until a + # versioned store reports a real version. + self._current_version: TaskVersion = TaskVersion.MISSING + # Transient: the event currently being persisted, read by _save_task. + self._pending_event: Event | None = None logger.debug( 'TaskManager initialized with task_id: %s, context_id: %s', task_id, context_id, ) + def invalidate(self) -> None: + """Drops the cached snapshot so the next read hits the store. + + Only safe at a request boundary; used to pick up state another replica + may have advanced. + """ + from a2a.server.cluster.version import TaskVersion # noqa: PLC0415 + + self._current_task = None + self._current_version = TaskVersion.MISSING + + @property + def current_version(self) -> TaskVersion: + """The version of the most recently read/written task snapshot.""" + return self._current_version + async def get_task(self) -> Task | None: """Retrieves the current task object, either from memory or the store. @@ -146,12 +190,18 @@ async def get_task(self) -> Task | None: logger.debug( 'Attempting to get task from store with id: %s', self.task_id ) - self._current_task = await self.task_store.get( + from a2a.server.cluster.version import TaskVersion # noqa: PLC0415 + + stored = await self._versioned_store.get( self.task_id, self._call_context ) - if self._current_task: + if stored is not None: + self._current_task = stored.task + self._current_version = stored.version logger.debug('Task %s retrieved successfully.', self.task_id) else: + self._current_task = None + self._current_version = TaskVersion.MISSING logger.debug('Task %s not found.', self.task_id) return self._current_task @@ -195,7 +245,11 @@ async def save_task_event( task_id_from_event, ) if isinstance(event, Task): - await self._save_task(event) + self._pending_event = event + try: + await self._save_task(event) + finally: + self._pending_event = None return event task: Task = await self.ensure_task(event) @@ -213,7 +267,11 @@ async def save_task_event( logger.debug('Appending artifact to task %s', task.id) append_artifact_to_task(task, event) - await self._save_task(task) + self._pending_event = event + try: + await self._save_task(task) + finally: + self._pending_event = None return task async def ensure_task_id(self, task_id: str, context_id: str) -> Task: @@ -226,12 +284,21 @@ async def ensure_task_id(self, task_id: str, context_id: str) -> Task: Returns: An existing or newly created `Task` object. """ + from a2a.server.cluster.version import TaskVersion # noqa: PLC0415 + task: Task | None = self._current_task if not task and self.task_id: logger.debug( 'Attempting to retrieve existing task with id: %s', self.task_id ) - task = await self.task_store.get(self.task_id, self._call_context) + stored = await self._versioned_store.get( + self.task_id, self._call_context + ) + if stored is not None: + task = stored.task + self._current_version = stored.version + else: + self._current_version = TaskVersion.MISSING if not task: logger.info( @@ -302,13 +369,20 @@ def _init_task_obj(self, task_id: str, context_id: str) -> Task: ) async def _save_task(self, task: Task) -> None: - """Saves the given task to the task store and updates the in-memory `_current_task`. + """Saves the task and updates the snapshot. - Args: - task: The `Task` object to save. + Threads `_current_version` so a VersionedTaskStore can compare-and-swap; + raises ConcurrentTaskModificationError on a stale write. """ logger.debug('Saving task with id: %s', task.id) - await self.task_store.save(task, self._call_context) + prev = self._current_task + self._current_version = await self._versioned_store.save( + task, + event=self._pending_event, + prev=prev, + prev_version=self._current_version, + context=self._call_context, + ) self._current_task = task if not self.task_id: logger.info('New task created with id: %s', task.id) diff --git a/tests/server/cluster/conftest.py b/tests/server/cluster/conftest.py new file mode 100644 index 000000000..97da1eaf6 --- /dev/null +++ b/tests/server/cluster/conftest.py @@ -0,0 +1,421 @@ +import asyncio +import contextlib +import threading + +from collections.abc import AsyncGenerator + +import pytest + +from a2a.auth.user import User +from a2a.helpers.proto_helpers import new_task_from_user_message +from a2a.server.agent_execution.agent_executor import AgentExecutor +from a2a.server.cluster import ( + ConcurrentTaskModificationError, + StoredTask, + TaskEventStream, + TaskVersion, + VersionedEvent, + VersionedTaskStore, +) +from a2a.server.context import ServerCallContext +from a2a.server.events.event_queue import Event +from a2a.server.owner_resolver import OwnerResolver, resolve_user_scope +from a2a.server.request_handlers.default_request_handler_v2 import ( + DefaultRequestHandlerV2, +) +from a2a.server.tasks.inmemory_task_store import InMemoryTaskStore +from a2a.server.tasks.task_updater import TaskUpdater +from a2a.types.a2a_pb2 import ( + AgentCapabilities, + AgentCard, + ListTasksRequest, + ListTasksResponse, + Message, + Part, + Role, + SendMessageRequest, + Task, + TaskState, +) + + +# --- In-memory cluster doubles ---------------------------------------------- + +_TERMINAL_STATES = frozenset( + { + TaskState.TASK_STATE_COMPLETED, + TaskState.TASK_STATE_CANCELED, + TaskState.TASK_STATE_FAILED, + TaskState.TASK_STATE_REJECTED, + } +) + + +class VersionedInMemoryTaskStore(VersionedTaskStore): + """`VersionedTaskStore` backed by in-process dictionaries. + + Holds one integer version per (owner, task_id). `save` performs a + compare-and-swap against `prev_version` and raises + `ConcurrentTaskModificationError` on mismatch. Reads and list operations + delegate to a wrapped `InMemoryTaskStore` for the task data itself. + """ + + def __init__( + self, owner_resolver: OwnerResolver = resolve_user_scope + ) -> None: + self._owner_resolver = owner_resolver + self._store = InMemoryTaskStore() + # Maps owner to a mapping of task_id to its current integer version. + self._versions: dict[str, dict[str, int]] = {} + self._lock = threading.RLock() + + def _current_version_locked(self, owner: str, task_id: str) -> int: + return self._versions.get(owner, {}).get(task_id, 0) + + async def save( + self, + task: Task, + *, + event: Event | None, + prev: Task | None, + prev_version: TaskVersion, + context: ServerCallContext, + ) -> TaskVersion: + """Persists `task`, bumping its version, with a compare-and-swap. + + Raises `ConcurrentTaskModificationError` if `prev_version` does not + match the currently stored version. `TaskVersion.MISSING` skips the + check; a write moving `task` to CANCELED overwrites a non-terminal + stored task without a version check, and raises when it is terminal + or absent. + """ + del event, prev # in-memory store keeps no event log + owner = self._owner_resolver(context) + with self._lock: + stored = self._current_version_locked(owner, task_id=task.id) + if task.status.state == TaskState.TASK_STATE_CANCELED: + current = await self._store.get(task.id, context) + if current is None or current.status.state in _TERMINAL_STATES: + raise ConcurrentTaskModificationError(task.id) + elif not prev_version.is_missing: + if stored == 0: + raise ConcurrentTaskModificationError(task.id) + if TaskVersion(stored) != prev_version: + raise ConcurrentTaskModificationError(task.id) + new_version = stored + 1 + await self._store.save(task, context) + self._versions.setdefault(owner, {})[task.id] = new_version + return TaskVersion(new_version) + + async def get( + self, task_id: str, context: ServerCallContext + ) -> StoredTask | None: + """Returns the task with its current version, or None if absent.""" + owner = self._owner_resolver(context) + with self._lock: + task = await self._store.get(task_id, context) + if task is None: + return None + stored = self._current_version_locked(owner, task_id) + return StoredTask(task, TaskVersion(stored)) + + async def list( + self, + params: ListTasksRequest, + context: ServerCallContext, + ) -> ListTasksResponse: + """Lists tasks via the wrapped in-memory store.""" + return await self._store.list(params, context) + + async def delete(self, task_id: str, context: ServerCallContext) -> None: + """Deletes a task and forgets its version.""" + owner = self._owner_resolver(context) + with self._lock: + await self._store.delete(task_id, context) + owner_versions = self._versions.get(owner) + if owner_versions is not None: + owner_versions.pop(task_id, None) + if not owner_versions: + del self._versions[owner] + + +class InMemoryTaskEventStream(TaskEventStream): + """`TaskEventStream` that fans out to in-process subscribers. + + `publish` delivers to every current subscriber of the task; `subscribe` + returns an async iterator backed by a per-subscriber queue and filters on + the caller's `after` version. + """ + + def __init__(self) -> None: + # {task_id: set of subscriber queues} + self._subscribers: dict[str, set[asyncio.Queue[VersionedEvent]]] = {} + self._lock = asyncio.Lock() + + async def publish(self, task_id: str, event: VersionedEvent) -> None: + """Delivers `event` to all current subscribers of `task_id`.""" + async with self._lock: + queues = list(self._subscribers.get(task_id, ())) + for queue in queues: + await queue.put(event) + + async def subscribe( # type: ignore[override] + self, task_id: str, *, after: TaskVersion + ) -> AsyncGenerator[VersionedEvent, None]: + """Yields events for `task_id` published after this call, newer than `after`.""" + queue: asyncio.Queue[VersionedEvent] = asyncio.Queue() + async with self._lock: + self._subscribers.setdefault(task_id, set()).add(queue) + try: + while True: + versioned = await queue.get() + if versioned.version.is_after(after): + yield versioned + finally: + async with self._lock: + subs = self._subscribers.get(task_id) + if subs is not None: + subs.discard(queue) + if not subs: + del self._subscribers[task_id] + + async def destroy(self, task_id: str) -> None: + """Drops all subscribers for `task_id`.""" + async with self._lock: + self._subscribers.pop(task_id, None) + + +# --- Identity / context helpers --------------------------------------------- + + +class SampleUser(User): + """Minimal authenticated User for tests.""" + + def __init__(self, user_name: str = 'test_user') -> None: + self._user_name = user_name + + @property + def is_authenticated(self) -> bool: + return True + + @property + def user_name(self) -> str: + return self._user_name + + +def make_context(user: str = 'test_user') -> ServerCallContext: + """Builds a ServerCallContext for the given user name.""" + return ServerCallContext(user=SampleUser(user)) + + +def streaming_agent_card() -> AgentCard: + """An AgentCard that advertises streaming support.""" + return AgentCard(capabilities=AgentCapabilities(streaming=True)) + + +def build_send_request( + text: str, + task_id: str = '', + context_id: str = '', + message_id: str = 'm', +) -> SendMessageRequest: + """Builds a SendMessageRequest, optionally continuing an existing task.""" + msg = Message( + role=Role.ROLE_USER, + message_id=message_id, + parts=[Part(text=text)], + ) + if task_id: + msg.task_id = task_id + if context_id: + msg.context_id = context_id + return SendMessageRequest(message=msg) + + +async def drain(agen) -> None: # noqa: ANN001 + """Consumes an async generator to completion, swallowing exceptions.""" + with contextlib.suppress(Exception): + async for _ in agen: + pass + + +async def wait_for_state( + store: VersionedTaskStore, + task_id: str, + state: TaskState, + context: ServerCallContext, + *, + tries: int = 100, + delay: float = 0.02, +) -> None: + """Polls the store until the task reaches `state` (for async persistence).""" + for _ in range(tries): + stored = await store.get(task_id, context) + if stored is not None and stored.task.status.state == state: + return + await asyncio.sleep(delay) + raise AssertionError(f'Task {task_id} did not reach {state} in time') + + +# --- Reusable agent executors ----------------------------------------------- + + +class CompletingAgent(AgentExecutor): + """Creates the task if needed, goes WORKING, then completes.""" + + async def execute(self, context, event_queue) -> None: # noqa: ANN001 + if context.current_task is None: + await event_queue.enqueue_event( + new_task_from_user_message(context.message) + ) + updater = TaskUpdater( + event_queue, + str(context.task_id or ''), + str(context.context_id or ''), + ) + await updater.start_work() + await updater.complete() + + async def cancel(self, context, event_queue) -> None: # noqa: ANN001 + pass + + +class ControlledAgent(AgentExecutor): + """Creates a task, goes WORKING, then completes only when released. + + `working` is set once execution reaches WORKING; the agent then waits on + `release` before emitting an artifact and completing. Useful for observing a + task mid-flight from another replica. + """ + + def __init__(self) -> None: + self.working = asyncio.Event() + self.release = asyncio.Event() + + async def execute(self, context, event_queue) -> None: # noqa: ANN001 + if context.current_task is None: + await event_queue.enqueue_event( + new_task_from_user_message(context.message) + ) + updater = TaskUpdater( + event_queue, + str(context.task_id or ''), + str(context.context_id or ''), + ) + await updater.start_work() + self.working.set() + await self.release.wait() + await updater.add_artifact([Part(text='result')]) + await updater.complete() + + async def cancel(self, context, event_queue) -> None: # noqa: ANN001 + pass + + +class LongRunningAgent(AgentExecutor): + """Goes WORKING then loops, saving periodically until aborted. + + Periodic saves let a remote cancel be observed as a CAS conflict, which + aborts the execution (`aborted` is set on CancelledError). + """ + + def __init__(self) -> None: + self.working = asyncio.Event() + self.aborted = asyncio.Event() + + async def execute(self, context, event_queue) -> None: # noqa: ANN001 + if context.current_task is None: + await event_queue.enqueue_event( + new_task_from_user_message(context.message) + ) + updater = TaskUpdater( + event_queue, + str(context.task_id or ''), + str(context.context_id or ''), + ) + await updater.start_work() + self.working.set() + try: + for _ in range(200): + await asyncio.sleep(0.05) + await updater.update_status( + TaskState.TASK_STATE_WORKING, + message=updater.new_agent_message([Part(text='tick')]), + ) + except asyncio.CancelledError: + self.aborted.set() + raise + + async def cancel(self, context, event_queue) -> None: # noqa: ANN001 + pass + + +class InputRequiredThenCompleteAgent(AgentExecutor): + """Asks for input until two user turns exist, then completes. + + Derives the turn purely from durable task history, so it behaves correctly + regardless of which replica runs each turn. + """ + + async def execute(self, context, event_queue) -> None: # noqa: ANN001 + task = context.current_task + if task is None: + await event_queue.enqueue_event( + new_task_from_user_message(context.message) + ) + updater = TaskUpdater( + event_queue, + str(context.task_id or ''), + str(context.context_id or ''), + ) + user_msgs = 0 + if task is not None: + user_msgs = sum(1 for m in task.history if m.role == Role.ROLE_USER) + + if task is None or user_msgs <= 1: + await updater.requires_input( + message=updater.new_agent_message([Part(text='need input')]) + ) + else: + await updater.complete() + + async def cancel(self, context, event_queue) -> None: # noqa: ANN001 + pass + + +# --- Replica / infra factories ---------------------------------------------- + + +def make_replica( + store: VersionedTaskStore, + stream: TaskEventStream, + agent: AgentExecutor, +) -> DefaultRequestHandlerV2: + """Builds a handler ('replica') wired to a shared store and stream.""" + return DefaultRequestHandlerV2( + agent_executor=agent, + task_store=store, + agent_card=streaming_agent_card(), + event_stream=stream, + ) + + +# --- Fixtures ---------------------------------------------------------------- + + +@pytest.fixture +def shared_store() -> VersionedInMemoryTaskStore: + """A versioned in-memory store shared across replicas in a test.""" + return VersionedInMemoryTaskStore() + + +@pytest.fixture +def shared_stream() -> InMemoryTaskEventStream: + """An in-memory event stream shared across replicas in a test.""" + return InMemoryTaskEventStream() + + +@pytest.fixture +def context() -> ServerCallContext: + """A default server call context for 'test_user'.""" + return make_context() diff --git a/tests/server/cluster/test_cancel_multireplica.py b/tests/server/cluster/test_cancel_multireplica.py new file mode 100644 index 000000000..8d028f980 --- /dev/null +++ b/tests/server/cluster/test_cancel_multireplica.py @@ -0,0 +1,144 @@ +import asyncio +import contextlib + +import pytest + +from a2a.types.a2a_pb2 import CancelTaskRequest, TaskState +from a2a.utils.errors import TaskNotCancelableError, TaskNotFoundError + +from .conftest import ( + CompletingAgent, + LongRunningAgent, + build_send_request, + drain, + make_context, + make_replica, + wait_for_state, +) + + +@pytest.mark.asyncio +@pytest.mark.timeout(15) +async def test_cancel_on_non_owning_replica_stops_remote_agent( + shared_store, shared_stream, context +) -> None: + agent = LongRunningAgent() + replica_a = make_replica(shared_store, shared_stream, agent) + replica_b = make_replica(shared_store, shared_stream, agent) + try: + a_task = asyncio.create_task( + drain( + replica_a.on_message_send_stream( + build_send_request('go', message_id='m1'), context + ) + ) + ) + await asyncio.wait_for(agent.working.wait(), timeout=5) + + task_id = next( + iter(replica_a._active_task_registry._active_tasks) # noqa: SLF001 + ) + await wait_for_state( + shared_store, task_id, TaskState.TASK_STATE_WORKING, context + ) + + # Cancel from B (not running the agent). + result = await replica_b.on_cancel_task( + CancelTaskRequest(id=task_id), context + ) + assert result.status.state == TaskState.TASK_STATE_CANCELED + + # A's agent observes the CAS conflict and aborts; its stream ends. + await asyncio.wait_for(agent.aborted.wait(), timeout=8) + await asyncio.wait_for(a_task, timeout=8) + + final = await shared_store.get(task_id, context) + assert final is not None + assert final.task.status.state == TaskState.TASK_STATE_CANCELED + finally: + await replica_a.aclose() + await replica_b.aclose() + + +@pytest.mark.asyncio +@pytest.mark.timeout(10) +async def test_cancel_already_cancelled_is_idempotent( + shared_store, shared_stream, context +) -> None: + agent = LongRunningAgent() + handler = make_replica(shared_store, shared_stream, agent) + other = make_replica(shared_store, shared_stream, agent) + try: + a_task = asyncio.create_task( + drain( + handler.on_message_send_stream( + build_send_request('go', message_id='m1'), context + ) + ) + ) + await asyncio.wait_for(agent.working.wait(), timeout=5) + task_id = next( + iter(handler._active_task_registry._active_tasks) # noqa: SLF001 + ) + await wait_for_state( + shared_store, task_id, TaskState.TASK_STATE_WORKING, context + ) + + r1 = await other.on_cancel_task(CancelTaskRequest(id=task_id), context) + assert r1.status.state == TaskState.TASK_STATE_CANCELED + # Second cancel returns the cancelled task, not an error. + r2 = await other.on_cancel_task(CancelTaskRequest(id=task_id), context) + assert r2.status.state == TaskState.TASK_STATE_CANCELED + with contextlib.suppress(Exception): + await asyncio.wait_for(a_task, timeout=8) + finally: + await handler.aclose() + await other.aclose() + + +@pytest.mark.asyncio +@pytest.mark.timeout(10) +async def test_cancel_completed_task_raises_not_cancelable( + shared_store, shared_stream, context +) -> None: + handler = make_replica(shared_store, shared_stream, CompletingAgent()) + try: + result = await handler.on_message_send( + build_send_request('go', message_id='m1'), context + ) + assert result.status.state == TaskState.TASK_STATE_COMPLETED + with pytest.raises(TaskNotCancelableError): + await handler.on_cancel_task( + CancelTaskRequest(id=result.id), context + ) + finally: + await handler.aclose() + + +@pytest.mark.asyncio +@pytest.mark.timeout(10) +async def test_cancel_absent_task_raises_not_found( + shared_store, shared_stream, context +) -> None: + handler = make_replica(shared_store, shared_stream, CompletingAgent()) + try: + with pytest.raises(TaskNotFoundError): + await handler.on_cancel_task(CancelTaskRequest(id='nope'), context) + finally: + await handler.aclose() + + +@pytest.mark.asyncio +@pytest.mark.timeout(10) +async def test_cancel_non_owner_rejected(shared_store, shared_stream) -> None: + handler = make_replica(shared_store, shared_stream, CompletingAgent()) + try: + result = await handler.on_message_send( + build_send_request('go', message_id='m1'), make_context('alice') + ) + with pytest.raises(TaskNotFoundError): + await handler.on_cancel_task( + CancelTaskRequest(id=result.id), make_context('bob') + ) + finally: + await handler.aclose() diff --git a/tests/server/cluster/test_handler_wiring.py b/tests/server/cluster/test_handler_wiring.py new file mode 100644 index 000000000..8451aad98 --- /dev/null +++ b/tests/server/cluster/test_handler_wiring.py @@ -0,0 +1,165 @@ +import pytest + +from a2a.server.agent_execution.agent_executor import AgentExecutor +from a2a.server.cluster import ( + LegacyTaskStoreAdapter, + TaskEventStream, +) +from a2a.server.request_handlers import DefaultRequestHandler +from a2a.server.request_handlers.default_request_handler_v2 import ( + DefaultRequestHandlerV2, +) +from a2a.server.tasks import InMemoryTaskStore +from a2a.types.a2a_pb2 import AgentCard + +from .conftest import ( + InMemoryTaskEventStream, + VersionedInMemoryTaskStore, + make_context, +) + + +class NoopExecutor(AgentExecutor): + async def execute(self, context, event_queue) -> None: # noqa: ANN001 + pass + + async def cancel(self, context, event_queue) -> None: # noqa: ANN001 + pass + + +def make_handler(event_stream: TaskEventStream | None = None): + return DefaultRequestHandlerV2( + agent_executor=NoopExecutor(), + task_store=InMemoryTaskStore(), + agent_card=AgentCard(), + event_stream=event_stream, + ) + + +def test_plain_store_kept_as_is_and_adapted_for_versioned() -> None: + # A plain TaskStore is exposed unchanged as `task_store`; internal cluster + # code sees it wrapped in a LegacyTaskStoreAdapter. + store = InMemoryTaskStore() + handler = DefaultRequestHandlerV2( + agent_executor=NoopExecutor(), + task_store=store, + agent_card=AgentCard(), + ) + assert handler.task_store is store + assert isinstance(handler._versioned_store, LegacyTaskStoreAdapter) # noqa: SLF001 + assert handler._versioned_store.store is store # noqa: SLF001 + + +def test_versioned_store_used_directly_and_not_swapped() -> None: + # LegacyTaskStoreAdapter is itself a VersionedTaskStore, so this asserts the + # no-swap contract without a test double. + store = LegacyTaskStoreAdapter(InMemoryTaskStore()) + handler = DefaultRequestHandlerV2( + agent_executor=NoopExecutor(), + task_store=store, + agent_card=AgentCard(), + ) + assert handler.task_store is store + assert handler._versioned_store is store # noqa: SLF001 + + +def test_default_handler_has_no_stream() -> None: + # No stream supplied: single-process mode, the handler holds no event stream. + handler = make_handler() + assert handler._event_stream is None # noqa: SLF001 + + +def test_explicit_stream_is_stored() -> None: + stream = InMemoryTaskEventStream() + handler = make_handler(event_stream=stream) + assert handler._event_stream is stream # noqa: SLF001 + + +def test_stream_is_threaded_to_registry() -> None: + stream = InMemoryTaskEventStream() + handler = make_handler(event_stream=stream) + assert handler._active_task_registry._event_stream is stream # noqa: SLF001 + + +def test_default_alias_accepts_event_stream() -> None: + # DefaultRequestHandler is an alias for v2; ensure the kwarg flows through. + stream = InMemoryTaskEventStream() + handler = DefaultRequestHandler( + NoopExecutor(), + InMemoryTaskStore(), + AgentCard(), + event_stream=stream, + ) + assert handler._event_stream is stream # noqa: SLF001 + + +@pytest.mark.asyncio +async def test_active_task_receives_stream_from_registry() -> None: + stream = InMemoryTaskEventStream() + handler = make_handler(event_stream=stream) + context = make_context() + + active_task = await handler._active_task_registry.get_or_create( # noqa: SLF001 + 'task-1', + call_context=context, + create_task_if_missing=True, + initial_message=None, + ) + try: + assert active_task._event_stream is stream # noqa: SLF001 + finally: + await handler.aclose() + + +@pytest.mark.asyncio +async def test_none_event_stream_propagates_none_to_active_task() -> None: + # When no stream is passed, the ActiveTask receives None and uses only its + # in-process queues (single-process behaviour). + handler = make_handler(event_stream=None) + context = make_context() + active_task = await handler._active_task_registry.get_or_create( # noqa: SLF001 + 'task-1', + call_context=context, + create_task_if_missing=True, + initial_message=None, + ) + try: + assert active_task._event_stream is None # noqa: SLF001 + finally: + await handler.aclose() + + +def test_versioned_store_without_stream_warns() -> None: + # A versioned store signals cluster intent; without a shared stream, + # cross-replica streaming is disabled, so construction warns once. + with pytest.warns(UserWarning, match='cross-replica streaming is disabled'): + DefaultRequestHandlerV2( + agent_executor=NoopExecutor(), + task_store=VersionedInMemoryTaskStore(), + agent_card=AgentCard(), + ) + + +def test_versioned_store_with_stream_does_not_warn(recwarn) -> None: # noqa: ANN001 + # Supplying a shared stream is the multi-replica configuration: no warning. + DefaultRequestHandlerV2( + agent_executor=NoopExecutor(), + task_store=VersionedInMemoryTaskStore(), + agent_card=AgentCard(), + event_stream=InMemoryTaskEventStream(), + ) + assert not [ + w + for w in recwarn + if 'cross-replica streaming is disabled' in str(w.message) + ] + + +def test_plain_store_without_stream_does_not_warn(recwarn) -> None: # noqa: ANN001 + # A plain store is the single-process default: no cluster intent, no warning. + make_handler() + assert not [ + w + for w in recwarn + if 'cross-replica streaming is disabled' in str(w.message) + ] diff --git a/tests/server/cluster/test_send_multireplica.py b/tests/server/cluster/test_send_multireplica.py new file mode 100644 index 000000000..92fa8f70b --- /dev/null +++ b/tests/server/cluster/test_send_multireplica.py @@ -0,0 +1,110 @@ +import pytest + +from a2a.server.cluster import ConcurrentTaskModificationError +from a2a.types.a2a_pb2 import Role, TaskState + +from .conftest import ( + CompletingAgent, + InputRequiredThenCompleteAgent, + build_send_request, + make_replica, +) + + +@pytest.mark.asyncio +@pytest.mark.timeout(10) +async def test_single_send_completes_on_shared_store( + shared_store, shared_stream, context +) -> None: + handler = make_replica(shared_store, shared_stream, CompletingAgent()) + try: + result = await handler.on_message_send( + build_send_request('hi'), context + ) + assert result.status.state == TaskState.TASK_STATE_COMPLETED + # Task persisted with a real version. + stored = await shared_store.get(result.id, context) + assert stored is not None + assert stored.task.status.state == TaskState.TASK_STATE_COMPLETED + assert not stored.version.is_missing + finally: + await handler.aclose() + + +@pytest.mark.asyncio +@pytest.mark.timeout(15) +async def test_multiturn_input_required_across_replicas( + shared_store, shared_stream, context +) -> None: + """turn 1 -> replica A, turn 2 -> replica B, turn 3 -> replica A. + + The pre-fix bug: replica A reuses a stale cached snapshot on turn 3 and + clobbers replica B's turn-2 write. With request-boundary invalidation + + CAS, A re-reads and the task advances correctly. + """ + agent = InputRequiredThenCompleteAgent() + replica_a = make_replica(shared_store, shared_stream, agent) + replica_b = make_replica(shared_store, shared_stream, agent) + try: + # Turn 1 on A: creates task, asks for input. + r1 = await replica_a.on_message_send( + build_send_request('q1', message_id='m1'), context + ) + assert r1.status.state == TaskState.TASK_STATE_INPUT_REQUIRED + task_id = r1.id + context_id = r1.context_id + + # Turn 2 on B: provides input; agent still needs more (2 user msgs). + r2 = await replica_b.on_message_send( + build_send_request( + 'a1', task_id=task_id, context_id=context_id, message_id='m2' + ), + context, + ) + assert r2.status.state == TaskState.TASK_STATE_INPUT_REQUIRED + # B's write is durable and visible. + stored = await shared_store.get(task_id, context) + assert stored is not None + assert ( + sum(1 for m in stored.task.history if m.role == Role.ROLE_USER) == 2 + ) + + # Turn 3 back on A: must see B's state (not the stale turn-1 snapshot) + # and complete. + r3 = await replica_a.on_message_send( + build_send_request( + 'a2', task_id=task_id, context_id=context_id, message_id='m3' + ), + context, + ) + assert r3.status.state == TaskState.TASK_STATE_COMPLETED + finally: + await replica_a.aclose() + await replica_b.aclose() + + +@pytest.mark.asyncio +@pytest.mark.timeout(10) +async def test_stale_version_save_raises_conflict_directly( + shared_store, shared_stream, context +) -> None: + """The store-level CAS that underpins the handler fix.""" + handler = make_replica(shared_store, shared_stream, CompletingAgent()) + try: + result = await handler.on_message_send( + build_send_request('hi'), context + ) + stored = await shared_store.get(result.id, context) + assert stored is not None + task, v = stored.task, stored.version + # Simulate a stale writer: another write advances the version. + await shared_store.save( + task, event=None, prev=None, prev_version=v, context=context + ) + # Now the original version is stale. + with pytest.raises(ConcurrentTaskModificationError): + await shared_store.save( + task, event=None, prev=None, prev_version=v, context=context + ) + finally: + await handler.aclose() diff --git a/tests/server/cluster/test_subscribe_multireplica.py b/tests/server/cluster/test_subscribe_multireplica.py new file mode 100644 index 000000000..c15a56007 --- /dev/null +++ b/tests/server/cluster/test_subscribe_multireplica.py @@ -0,0 +1,139 @@ +import asyncio + +import pytest + +from a2a.types.a2a_pb2 import SubscribeToTaskRequest, TaskState +from a2a.utils.errors import InvalidParamsError, TaskNotFoundError + +from .conftest import ( + ControlledAgent, + build_send_request, + make_context, + make_replica, + wait_for_state, +) + + +@pytest.mark.asyncio +@pytest.mark.timeout(15) +async def test_resubscribe_on_non_owning_replica_streams_events( + shared_store, shared_stream, context +) -> None: + """Replica A runs the agent; a resubscribe on replica B streams its events. + + Pre-fix, B would build a local ActiveTask, emit the snapshot, then hang. + """ + agent = ControlledAgent() + replica_a = make_replica(shared_store, shared_stream, agent) + replica_b = make_replica(shared_store, shared_stream, agent) + try: + a_events: list = [] + + async def run_a() -> None: + async for ev in replica_a.on_message_send_stream( + build_send_request('go', message_id='m1'), context + ): + a_events.append(ev) + + a_task = asyncio.create_task(run_a()) + await asyncio.wait_for(agent.working.wait(), timeout=5) + + task_id = next( + iter(replica_a._active_task_registry._active_tasks) # noqa: SLF001 + ) + await wait_for_state( + shared_store, task_id, TaskState.TASK_STATE_WORKING, context + ) + + # Resubscribe on B (not running the agent). Collect until COMPLETED. + b_states: list = [] + + async def run_b() -> None: + async for ev in replica_b.on_subscribe_to_task( + SubscribeToTaskRequest(id=task_id), context + ): + if getattr(ev, 'status', None): + b_states.append(ev.status.state) + if ev.status.state == TaskState.TASK_STATE_COMPLETED: + return + + b_task = asyncio.create_task(run_b()) + await asyncio.sleep(0.1) # let B read snapshot + start tailing + + # Release the agent on A -> it completes; B must observe via the stream. + agent.release.set() + + await asyncio.wait_for(b_task, timeout=8) + await asyncio.wait_for(a_task, timeout=8) + + assert b_states[0] == TaskState.TASK_STATE_WORKING + assert b_states[-1] == TaskState.TASK_STATE_COMPLETED + finally: + await replica_a.aclose() + await replica_b.aclose() + + +@pytest.mark.asyncio +@pytest.mark.timeout(10) +async def test_resubscribe_to_terminal_task_rejected( + shared_store, shared_stream, context +) -> None: + """A resubscribe to an already-finished task is rejected (no hang).""" + agent = ControlledAgent() + agent.release.set() # completes immediately + handler = make_replica(shared_store, shared_stream, agent) + other = make_replica(shared_store, shared_stream, agent) + try: + result = await handler.on_message_send( + build_send_request('go', message_id='m1'), context + ) + assert result.status.state == TaskState.TASK_STATE_COMPLETED + task_id = result.id + + # Terminal task: the non-owning replica rejects rather than hanging. + with pytest.raises(InvalidParamsError, match='terminal state'): + async for _ in other.on_subscribe_to_task( + SubscribeToTaskRequest(id=task_id), context + ): + pass + finally: + await handler.aclose() + await other.aclose() + + +@pytest.mark.asyncio +@pytest.mark.timeout(10) +async def test_resubscribe_absent_task_raises_not_found( + shared_store, shared_stream, context +) -> None: + handler = make_replica(shared_store, shared_stream, ControlledAgent()) + try: + with pytest.raises(TaskNotFoundError): + async for _ in handler.on_subscribe_to_task( + SubscribeToTaskRequest(id='nope'), context + ): + pass + finally: + await handler.aclose() + + +@pytest.mark.asyncio +@pytest.mark.timeout(10) +async def test_resubscribe_non_owner_rejected( + shared_store, shared_stream +) -> None: + """Issue #1159: a non-owner cannot resubscribe to another user's task.""" + agent = ControlledAgent() + agent.release.set() + handler = make_replica(shared_store, shared_stream, agent) + try: + result = await handler.on_message_send( + build_send_request('go', message_id='m1'), make_context('alice') + ) + with pytest.raises(TaskNotFoundError): + async for _ in handler.on_subscribe_to_task( + SubscribeToTaskRequest(id=result.id), make_context('bob') + ): + pass + finally: + await handler.aclose() 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 2eb7e4725..1a69ab80d 100644 --- a/tests/server/request_handlers/test_default_request_handler_v2.py +++ b/tests/server/request_handlers/test_default_request_handler_v2.py @@ -138,7 +138,7 @@ def test_init_default_dependencies(): handler._request_context_builder._should_populate_referred_tasks is False ) - assert handler._request_context_builder._task_store == task_store + assert handler._request_context_builder._task_store is None def test_init_warns_when_queue_manager_passed(caplog):