Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
55 changes: 54 additions & 1 deletion src/a2a/server/agent_execution/active_task.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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()


Expand All @@ -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.

Expand All @@ -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
Expand All @@ -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,
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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()
)
Expand Down Expand Up @@ -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')
Expand Down
22 changes: 17 additions & 5 deletions src/a2a/server/agent_execution/active_task_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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()
Expand Down Expand Up @@ -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

Expand All @@ -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(
Expand Down
Loading
Loading