From 52df4555b5e546f74769be9804fceb5a3b4af07e Mon Sep 17 00:00:00 2001 From: Mykyta Netipa Date: Mon, 14 Sep 2026 22:14:31 +0000 Subject: [PATCH] feat(server/cluster): add clustered task store and event stream building blocks Introduces the a2a.server.cluster package: the extension points for running the A2A server as multiple replicas behind a load balancer with no task affinity. The package is inert at this commit -- nothing in the runtime imports it yet -- but it also adds a `version` column to the shared tasks schema that database-backed deployments must provision (see the schema bullet). * Optimistic concurrency: TaskVersion, VersionedTaskStore (compare-and-set) whose `get` returns a `StoredTask` (task + version), with ConcurrentTaskModificationError and a LegacyTaskStoreAdapter that runs an existing TaskStore under the versioned interface (last-writer-wins). * Versioned store: VersionedDatabaseTaskStore (CAS on a version column) with a configurable bounded retry (max_attempts / retry_delay_s) for transient DB contention. * Cross-replica event stream: TaskEventStream / VersionedEvent and DatabaseTaskEventStream tailing a transactional task_events outbox. * Schema: nullable `version` column on TaskMixin plus the `task_events` outbox table. Their VALUES are inert for the non-versioned stores (version stays NULL, task_events empty), but DatabaseTaskStore selects and writes every mapped column, so the `version` column must exist in the database. Any database-backed deployment must run `a2a-db` after upgrading -- versioned or not; create_table=True (dev/tests) picks it up automatically. --- .github/actions/spelling/allow.txt | 3 + src/a2a/server/cluster/__init__.py | 57 ++ .../server/cluster/database_event_stream.py | 181 +++++++ src/a2a/server/cluster/database_task_store.py | 308 +++++++++++ src/a2a/server/cluster/event_stream.py | 32 ++ src/a2a/server/cluster/task_store.py | 128 +++++ src/a2a/server/cluster/version.py | 54 ++ src/a2a/server/models.py | 73 ++- tests/server/cluster/__init__.py | 0 .../cluster/test_database_event_stream.py | 278 ++++++++++ .../test_database_versioned_task_store.py | 505 ++++++++++++++++++ .../server/cluster/test_task_store_adapter.py | 165 ++++++ tests/server/cluster/test_version.py | 102 ++++ 13 files changed, 1885 insertions(+), 1 deletion(-) create mode 100644 src/a2a/server/cluster/__init__.py create mode 100644 src/a2a/server/cluster/database_event_stream.py create mode 100644 src/a2a/server/cluster/database_task_store.py create mode 100644 src/a2a/server/cluster/event_stream.py create mode 100644 src/a2a/server/cluster/task_store.py create mode 100644 src/a2a/server/cluster/version.py create mode 100644 tests/server/cluster/__init__.py create mode 100644 tests/server/cluster/test_database_event_stream.py create mode 100644 tests/server/cluster/test_database_versioned_task_store.py create mode 100644 tests/server/cluster/test_task_store_adapter.py create mode 100644 tests/server/cluster/test_version.py diff --git a/.github/actions/spelling/allow.txt b/.github/actions/spelling/allow.txt index cba8a880d..1bef2eb69 100644 --- a/.github/actions/spelling/allow.txt +++ b/.github/actions/spelling/allow.txt @@ -33,6 +33,7 @@ buf bufbuild cla cls +clustermode coc codegen coro @@ -158,5 +159,7 @@ TResponse typ typeerror UIDs +versioned +Versioned vulnz whl diff --git a/src/a2a/server/cluster/__init__.py b/src/a2a/server/cluster/__init__.py new file mode 100644 index 000000000..f7c4f1ac9 --- /dev/null +++ b/src/a2a/server/cluster/__init__.py @@ -0,0 +1,57 @@ +import logging + +from a2a.server.cluster.event_stream import TaskEventStream, VersionedEvent +from a2a.server.cluster.task_store import ( + ConcurrentTaskModificationError, + LegacyTaskStoreAdapter, + StoredTask, + VersionedTaskStore, +) +from a2a.server.cluster.version import TaskVersion + + +logger = logging.getLogger(__name__) + +try: + from a2a.server.cluster.database_event_stream import DatabaseTaskEventStream + from a2a.server.cluster.database_task_store import ( + VersionedDatabaseTaskStore, + ) +except ImportError as e: + _original_error = e + logger.debug( + 'Database-backed cluster stores not loaded. This is expected if ' + 'database dependencies are not installed. Error: %s', + _original_error, + ) + + class VersionedDatabaseTaskStore: # type: ignore[no-redef] + """Placeholder when database dependencies are not installed.""" + + def __init__(self, *args: object, **kwargs: object) -> None: + raise ImportError( + 'To use VersionedDatabaseTaskStore, its dependencies must be ' + "installed. Install with 'pip install a2a-sdk[sql]'." + ) from _original_error + + class DatabaseTaskEventStream: # type: ignore[no-redef] + """Placeholder when database dependencies are not installed.""" + + def __init__(self, *args: object, **kwargs: object) -> None: + raise ImportError( + 'To use DatabaseTaskEventStream, its dependencies must be ' + "installed. Install with 'pip install a2a-sdk[sql]'." + ) from _original_error + + +__all__ = [ + 'ConcurrentTaskModificationError', + 'DatabaseTaskEventStream', + 'LegacyTaskStoreAdapter', + 'StoredTask', + 'TaskEventStream', + 'TaskVersion', + 'VersionedDatabaseTaskStore', + 'VersionedEvent', + 'VersionedTaskStore', +] diff --git a/src/a2a/server/cluster/database_event_stream.py b/src/a2a/server/cluster/database_event_stream.py new file mode 100644 index 000000000..d77db1971 --- /dev/null +++ b/src/a2a/server/cluster/database_event_stream.py @@ -0,0 +1,181 @@ +import asyncio +import logging + +from collections.abc import AsyncGenerator + + +try: + from sqlalchemy import Table, and_, select + from sqlalchemy.exc import OperationalError + from sqlalchemy.ext.asyncio import ( + AsyncEngine, + AsyncSession, + async_sessionmaker, + ) + from sqlalchemy.orm import class_mapper +except ImportError as e: + raise ImportError( + 'DatabaseTaskEventStream requires SQLAlchemy and a database driver. ' + 'Install with one of: ' + "'pip install a2a-sdk[postgresql]', " + "'pip install a2a-sdk[mysql]', " + "'pip install a2a-sdk[sqlite]', " + "or 'pip install a2a-sdk[sql]'" + ) from e + +from a2a.server.cluster.event_stream import TaskEventStream, VersionedEvent +from a2a.server.cluster.version import TaskVersion +from a2a.server.events.event_queue import Event +from a2a.server.models import Base, TaskEventModel, create_task_event_model +from a2a.types.a2a_pb2 import StreamResponse + + +logger = logging.getLogger(__name__) + +_DEFAULT_POLL_INTERVAL_S = 0.5 + + +def _as_int(version: TaskVersion) -> int: + """Extracts the integer value of a version, or raises if not an int.""" + value = version._value # noqa: SLF001 + if not isinstance(value, int): + raise TypeError( + 'DatabaseTaskEventStream requires integer versions, got ' + f'{type(value).__name__}' + ) + return value + + +def stream_response_to_event(response: StreamResponse) -> Event: + """Converts a `StreamResponse` proto back to an internal `Event`.""" + which = response.WhichOneof('payload') + if which == 'task': + return response.task + if which == 'message': + return response.message + if which == 'status_update': + return response.status_update + if which == 'artifact_update': + return response.artifact_update + raise ValueError(f'StreamResponse has no known payload set: {which!r}') + + +class DatabaseTaskEventStream(TaskEventStream): + """`TaskEventStream` backed by polling the shared ``task_events`` table.""" + + _event_model: type[TaskEventModel] + + def __init__( + self, + engine: AsyncEngine, + create_table: bool = True, + table_name: str = 'task_events', + poll_interval_s: float = _DEFAULT_POLL_INTERVAL_S, + ) -> None: + """Initializes the stream over an existing SQLAlchemy AsyncEngine.""" + self._engine = engine + self._session_maker = async_sessionmaker(engine, expire_on_commit=False) + self._create_table = create_table + self._poll_interval_s = poll_interval_s + self._initialized = False + self._event_model = ( # ty:ignore[invalid-assignment] + TaskEventModel + if table_name == 'task_events' + else create_task_event_model(table_name) + ) + + async def initialize(self) -> None: + """Creates the ``task_events`` table if requested.""" + if self._initialized: + return + if self._create_table: + async with self._engine.begin() as conn: + mapper = class_mapper(self._event_model) + tables = [t for t in mapper.tables if isinstance(t, Table)] + await conn.run_sync(Base.metadata.create_all, tables=tables) + self._initialized = True + + async def _ensure_initialized(self) -> None: + if not self._initialized: + await self.initialize() + + async def publish(self, task_id: str, event: VersionedEvent) -> None: + """No-op: events are persisted transactionally by the task store.""" + del task_id, event + + async def _seq_at_or_before_version( + self, session: AsyncSession, task_id: str, after: TaskVersion + ) -> int: + """Seq to start polling from so nothing after `after` is missed.""" + if after.is_missing: + return 0 + stmt = ( + select(self._event_model.seq) + .where( + and_( + self._event_model.task_id == task_id, + self._event_model.task_version <= _as_int(after), + ) + ) + .order_by(self._event_model.seq.desc()) + .limit(1) + ) + result = (await session.execute(stmt)).scalar_one_or_none() + return result or 0 + + async def subscribe( # type: ignore[override] + self, task_id: str, *, after: TaskVersion + ) -> AsyncGenerator[VersionedEvent, None]: + """Polls the log for events of `task_id` newer than `after`.""" + await self._ensure_initialized() + async with self._session_maker() as session: + cursor = await self._seq_at_or_before_version( + session, task_id, after + ) + + while True: + try: + async with self._session_maker() as session: + stmt = ( + select( + self._event_model.seq, + self._event_model.task_version, + self._event_model.event_data, + ) + .where( + and_( + self._event_model.task_id == task_id, + self._event_model.seq > cursor, + ) + ) + .order_by(self._event_model.seq.asc()) + ) + rows = (await session.execute(stmt)).all() + except OperationalError: + logger.debug( + 'Transient DB error polling events for %s; retrying', + task_id, + exc_info=True, + ) + await asyncio.sleep(self._poll_interval_s) + continue + + for row in rows: + seq, task_version, event_data = row[0], row[1], row[2] + cursor = seq + version = TaskVersion(task_version) + if not version.is_after(after): + continue + response = StreamResponse() + response.ParseFromString(event_data) + yield VersionedEvent( + event=stream_response_to_event(response), + version=version, + ) + + if not rows: + await asyncio.sleep(self._poll_interval_s) + + async def destroy(self, task_id: str) -> None: + """No-op: the append-only log is retained; subscribers stop on their own.""" + del task_id diff --git a/src/a2a/server/cluster/database_task_store.py b/src/a2a/server/cluster/database_task_store.py new file mode 100644 index 000000000..620f34b47 --- /dev/null +++ b/src/a2a/server/cluster/database_task_store.py @@ -0,0 +1,308 @@ +import asyncio +import logging +import time + +from collections.abc import Callable + + +try: + from sqlalchemy import Table, and_, insert, inspect, select, update + from sqlalchemy.exc import IntegrityError, OperationalError + from sqlalchemy.ext.asyncio import AsyncEngine + from sqlalchemy.orm import class_mapper +except ImportError as e: + raise ImportError( + 'VersionedDatabaseTaskStore requires SQLAlchemy and a database driver. ' + 'Install with one of: ' + "'pip install a2a-sdk[postgresql]', " + "'pip install a2a-sdk[mysql]', " + "'pip install a2a-sdk[sqlite]', " + "or 'pip install a2a-sdk[sql]'" + ) from e + +from a2a.server.cluster.task_store import ( + ConcurrentTaskModificationError, + StoredTask, + 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.models import ( + Base, + TaskEventModel, + TaskModel, + create_task_event_model, +) +from a2a.server.owner_resolver import OwnerResolver, resolve_user_scope +from a2a.server.tasks.database_task_store import DatabaseTaskStore +from a2a.server.tasks.task_store import TaskStore +from a2a.types.a2a_pb2 import ( + ListTasksRequest, + ListTasksResponse, + Task, + TaskState, +) +from a2a.utils.proto_utils import to_stream_response + + +logger = logging.getLogger(__name__) + +_TERMINAL_STATES = frozenset( + { + TaskState.TASK_STATE_COMPLETED, + TaskState.TASK_STATE_CANCELED, + TaskState.TASK_STATE_FAILED, + TaskState.TASK_STATE_REJECTED, + } +) + +# DB contention retry +_MAX_ATTEMPTS = 5 +_RETRY_DELAY_S = 0.02 + + +class VersionedDatabaseTaskStore(VersionedTaskStore): + """`VersionedTaskStore` backed by SQLAlchemy.""" + + _event_model: type[TaskEventModel] + + def __init__( # noqa: PLR0913 + self, + engine: AsyncEngine, + create_table: bool = True, + table_name: str = 'tasks', + owner_resolver: OwnerResolver = resolve_user_scope, + core_to_model_conversion: Callable[[Task, str], TaskModel] + | None = None, + model_to_core_conversion: Callable[[TaskModel], Task] | None = None, + event_table_name: str = 'task_events', + max_attempts: int = _MAX_ATTEMPTS, + retry_delay_s: float = _RETRY_DELAY_S, + ) -> None: + """Initializes the store, delegating schema to `DatabaseTaskStore`.""" + self._db = DatabaseTaskStore( + engine=engine, + create_table=create_table, + table_name=table_name, + owner_resolver=owner_resolver, + core_to_model_conversion=core_to_model_conversion, + model_to_core_conversion=model_to_core_conversion, + ) + self._create_table = create_table + self._max_attempts = max_attempts + self._retry_delay_s = retry_delay_s + self._event_model = ( # ty:ignore[invalid-assignment] + TaskEventModel + if event_table_name == 'task_events' + else create_task_event_model(event_table_name) + ) + self._event_table_ready = False + + @property + def as_task_store(self) -> TaskStore: + """The underlying non-versioned `TaskStore`.""" + return self._db + + async def initialize(self) -> None: + """Initializes the database schema (task table and event log).""" + await self._db.initialize() + await self._ensure_event_table() + + async def _ensure_event_table(self) -> None: + if self._event_table_ready: + return + if self._create_table: + async with self._db.engine.begin() as conn: + mapper = class_mapper(self._event_model) + tables = [t for t in mapper.tables if isinstance(t, Table)] + await conn.run_sync(Base.metadata.create_all, tables=tables) + self._event_table_ready = True + + async def save( + self, + task: Task, + *, + event: Event | None = None, + prev: Task | None = None, + prev_version: TaskVersion, + context: ServerCallContext, + ) -> TaskVersion: + """Persists `task` with a compare-and-swap on the version column. + + On the first write (``prev_version`` MISSING) this INSERTs and lets a + primary-key collision surface as a conflict. On updates it runs a + conditional UPDATE keyed on the previous version and treats zero + affected rows as a concurrent modification. + + When `event` is provided it is appended to the ``task_events`` log in the + same transaction, so the log is a consistent source for cross-replica + replay via `DatabaseTaskEventStream`. + """ + del prev + await self._db._ensure_initialized() # noqa: SLF001 + await self._ensure_event_table() + owner = self._db.owner_resolver(context) + + # Retry only transient DB contention (e.g. lock/serialization errors). + # A ConcurrentTaskModificationError is the real signal and is never + # retried - it propagates so the caller can reload and decide. + attempts = 0 + while True: + attempts += 1 + try: + return await self._save_once( + task, event=event, prev_version=prev_version, owner=owner + ) + except OperationalError: + if attempts >= self._max_attempts: + raise + await asyncio.sleep(self._retry_delay_s * attempts) + + async def _save_once( + self, + task: Task, + *, + event: Event | None, + prev_version: TaskVersion, + owner: str, + ) -> TaskVersion: + new_version = time.time_ns() + model = self._db._to_orm(task, owner) # noqa: SLF001 + model.version = new_version + task_model = self._db.task_model + values = _column_values(model, task_model) + + async with self._db.async_session_maker.begin() as session: + if prev_version.is_missing: + try: + await session.execute(insert(task_model).values(**values)) + except IntegrityError as e: + raise ConcurrentTaskModificationError(task.id) from e + elif task.status.state == TaskState.TASK_STATE_CANCELED: + row = ( + await session.execute( + select(task_model) + .where( + and_( + task_model.id == task.id, + task_model.owner == owner, + ) + ) + .with_for_update() + ) + ).scalar_one_or_none() + current = ( + self._db._from_orm(row) # noqa: SLF001 + if row is not None + else None + ) + if current is None or current.status.state in _TERMINAL_STATES: + raise ConcurrentTaskModificationError(task.id) + await session.execute( + update(task_model) + .where( + and_( + task_model.id == task.id, + task_model.owner == owner, + ) + ) + .values(**values) + ) + else: + result = await session.execute( + update(task_model) + .where( + and_( + task_model.id == task.id, + task_model.owner == owner, + task_model.version == _as_int(prev_version), + ) + ) + .values(**values) + ) + if result.rowcount == 0: # ty:ignore[unresolved-attribute] + raise ConcurrentTaskModificationError(task.id) + + if event is not None: + await session.execute( + insert(self._event_model).values( + task_id=task.id, + owner=owner, + task_version=new_version, + event_data=to_stream_response( + event + ).SerializeToString(), + ) + ) + + return TaskVersion(new_version) + + async def get( + self, task_id: str, context: ServerCallContext + ) -> StoredTask | None: + """Returns the task with its stored version, or None if absent. + + Retries transient contention (e.g. a lock held while another writer + commits) rather than failing the read. + """ + await self._db._ensure_initialized() # noqa: SLF001 + owner = self._db.owner_resolver(context) + attempts = 0 + while True: + attempts += 1 + try: + return await self._get_once(task_id, owner) + except OperationalError: + if attempts >= self._max_attempts: + raise + await asyncio.sleep(self._retry_delay_s * attempts) + + async def _get_once(self, task_id: str, owner: str) -> StoredTask | None: + task_model = self._db.task_model + async with self._db.async_session_maker() as session: + stmt = select(task_model).where( + and_( + task_model.id == task_id, + task_model.owner == owner, + ) + ) + row = (await session.execute(stmt)).scalar_one_or_none() + if row is None: + return None + task = self._db._from_orm(row) # noqa: SLF001 + version = ( + TaskVersion(row.version) + if row.version is not None + else TaskVersion.MISSING + ) + return StoredTask(task, version) + + async def list( + self, + params: ListTasksRequest, + context: ServerCallContext, + ) -> ListTasksResponse: + """Lists tasks via the underlying store.""" + return await self._db.list(params, context) + + async def delete(self, task_id: str, context: ServerCallContext) -> None: + """Deletes a task via the underlying store.""" + await self._db.delete(task_id, context) + + +def _as_int(version: TaskVersion) -> int: + """Extracts the integer value of a version, or raises if not an int.""" + value = version._value # noqa: SLF001 + if not isinstance(value, int): + raise TypeError( + 'VersionedDatabaseTaskStore requires integer versions, got ' + f'{type(value).__name__}' + ) + return value + + +def _column_values(model: object, task_model: type) -> dict[str, object]: + """Extracts mapped column values from an ORM instance as a dict.""" + mapper = inspect(task_model) + return {col.key: getattr(model, col.key) for col in mapper.column_attrs} diff --git a/src/a2a/server/cluster/event_stream.py b/src/a2a/server/cluster/event_stream.py new file mode 100644 index 000000000..d65838a39 --- /dev/null +++ b/src/a2a/server/cluster/event_stream.py @@ -0,0 +1,32 @@ +from abc import ABC, abstractmethod +from collections.abc import AsyncGenerator +from dataclasses import dataclass + +from a2a.server.cluster.version import TaskVersion +from a2a.server.events.event_queue import Event + + +@dataclass(frozen=True) +class VersionedEvent: + """An event together with the task version produced by applying it.""" + + event: Event + version: TaskVersion + + +class TaskEventStream(ABC): + """Delivers task events across replicas.""" + + @abstractmethod + async def publish(self, task_id: str, event: VersionedEvent) -> None: + """Publishes one event for `task_id` to all replicas.""" + + @abstractmethod + def subscribe( + self, task_id: str, *, after: TaskVersion + ) -> AsyncGenerator[VersionedEvent, None]: + """Yields events for `task_id` newer than `after`.""" + + @abstractmethod + async def destroy(self, task_id: str) -> None: + """Releases resources for a task that has reached a terminal state.""" diff --git a/src/a2a/server/cluster/task_store.py b/src/a2a/server/cluster/task_store.py new file mode 100644 index 000000000..aed8777fe --- /dev/null +++ b/src/a2a/server/cluster/task_store.py @@ -0,0 +1,128 @@ +from abc import ABC, abstractmethod +from dataclasses import dataclass + +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 +from a2a.types.a2a_pb2 import ListTasksRequest, ListTasksResponse, Task + + +@dataclass(frozen=True) +class StoredTask: + """A task together with the version it was read at.""" + + task: Task + version: TaskVersion + + +class ConcurrentTaskModificationError(Exception): + """Raised by `VersionedTaskStore.save` when `prev_version` is stale.""" + + def __init__(self, task_id: str) -> None: + super().__init__( + f'Task {task_id} was modified concurrently by another writer' + ) + self.task_id = task_id + + +class VersionedTaskStore(ABC): + """A `TaskStore` variant with snapshot versioning to prevent concurrent re-writes.""" + + @abstractmethod + async def save( + self, + task: Task, + *, + event: Event | None, + prev: Task | None, + prev_version: TaskVersion, + context: ServerCallContext, + ) -> TaskVersion: + """Persists `task` and returns its new version. + + Args: + task: The task state to persist. + event: The event that produced this state, or `None` for a direct + write. + prev: The task as previously read, for implementations that diff. + prev_version: The version `task` was derived from. Implementations + MUST raise `ConcurrentTaskModificationError` if the currently + stored version differs. `TaskVersion.MISSING` marks a first + write. A write moving `task` to CANCELED overwrites a + non-terminal stored task without a version check, and raises + `ConcurrentTaskModificationError` if the stored task is already + terminal or absent. + context: The server call context (used to resolve the owner). + + Returns: + The new `TaskVersion` for the persisted task. + + Raises: + ConcurrentTaskModificationError: If `prev_version` is stale. + """ + + @abstractmethod + async def get( + self, task_id: str, context: ServerCallContext + ) -> StoredTask | None: + """Retrieves a task with its version, or `None` if it does not exist.""" + + @abstractmethod + async def list( + self, + params: ListTasksRequest, + context: ServerCallContext, + ) -> ListTasksResponse: + """Retrieves a list of tasks from the store.""" + + @abstractmethod + async def delete(self, task_id: str, context: ServerCallContext) -> None: + """Deletes a task from the store by ID.""" + + +class LegacyTaskStoreAdapter(VersionedTaskStore): + """Runs an unversioned `TaskStore` under the `VersionedTaskStore` interface.""" + + def __init__(self, store: TaskStore) -> None: + self._store = store + + @property + def store(self) -> TaskStore: + """The wrapped task store.""" + return self._store + + async def save( + self, + task: Task, + *, + event: Event | None, + prev: Task | None, + prev_version: TaskVersion, + context: ServerCallContext, + ) -> TaskVersion: + """Saves via the wrapped store; ignores version args, returns MISSING.""" + del event, prev, prev_version # unversioned store ignores these + await self._store.save(task, context) + return TaskVersion.MISSING + + async def get( + self, task_id: str, context: ServerCallContext + ) -> StoredTask | None: + """Gets from the wrapped store, pairing the result with MISSING.""" + task = await self._store.get(task_id, context) + if task is None: + return None + return StoredTask(task, TaskVersion.MISSING) + + async def list( + self, + params: ListTasksRequest, + context: ServerCallContext, + ) -> ListTasksResponse: + """Lists tasks via the wrapped store.""" + return await self._store.list(params, context) + + async def delete(self, task_id: str, context: ServerCallContext) -> None: + """Deletes a task via the wrapped store.""" + await self._store.delete(task_id, context) diff --git a/src/a2a/server/cluster/version.py b/src/a2a/server/cluster/version.py new file mode 100644 index 000000000..5a2f05062 --- /dev/null +++ b/src/a2a/server/cluster/version.py @@ -0,0 +1,54 @@ +from typing import ClassVar + + +class TaskVersion: + """A version marker a `VersionedTaskStore` assigns to a stored `Task`. + + Prevents concurrent state re-writes. The wrapped value is the store's + choice - a counter, a commit timestamp, etc. Callers do not read it or do + arithmetic on it; they only pass it back to `save` and order two versions + with `is_after`. + """ + + __slots__ = ('_value',) + + MISSING: 'ClassVar[TaskVersion]' + + def __init__(self, value: int | str) -> None: + self._value = value + + def __eq__(self, other: object) -> bool: + """Two versions are equal when they wrap equal values.""" + return isinstance(other, TaskVersion) and self._value == other._value + + def __hash__(self) -> int: + """Hash by the wrapped value so versions are usable as keys.""" + return hash(self._value) + + def __repr__(self) -> str: + """Render `MISSING` specially, otherwise show the wrapped value.""" + if self.is_missing: + return 'TaskVersion.MISSING' + return f'TaskVersion({self._value!r})' + + @property + def is_missing(self) -> bool: + """Whether this token means "versioning is not tracked".""" + return self._value == 0 + + def is_after(self, other: 'TaskVersion') -> bool: + """Whether `self` is a strictly later version than `other`.""" + if other.is_missing: + return True + if self.is_missing: + return False + if type(self._value) is not type(other._value): # noqa: SLF001 + raise TypeError( + 'Cannot compare TaskVersions with different value types: ' + f'{type(other._value).__name__} and ' # noqa: SLF001 + f'{type(self._value).__name__}' + ) + return other._value < self._value # ty:ignore[unsupported-operator] # noqa: SLF001 + + +TaskVersion.MISSING = TaskVersion(0) diff --git a/src/a2a/server/models.py b/src/a2a/server/models.py index b3ae1a389..6f43e23cc 100644 --- a/src/a2a/server/models.py +++ b/src/a2a/server/models.py @@ -15,7 +15,15 @@ def override(func): # noqa: ANN001, ANN201 try: - from sqlalchemy import JSON, DateTime, Index, LargeBinary, String + from sqlalchemy import ( + JSON, + BigInteger, + DateTime, + Index, + Integer, + LargeBinary, + String, + ) from sqlalchemy.orm import ( DeclarativeBase, Mapped, @@ -59,6 +67,16 @@ class TaskMixin: protocol_version: Mapped[str | None] = mapped_column( String(16), nullable=True ) + # Optimistic-concurrency version for the clustered VersionedTaskStore + # (a2a.server.cluster). Its VALUE is inert for the non-versioned + # DatabaseTaskStore and for existing rows (both leave it NULL). The COLUMN, + # however, is part of the tasks schema that DatabaseTaskStore reads and + # writes for every task, so any database-backed deployment must provision it + # (run `a2a-db`) after upgrading - versioned or not. See migration + # b5e3d1c8a2f7. + version: Mapped[int | None] = mapped_column( + BigInteger, nullable=True, default=None + ) # Using declared_attr to avoid conflict with Pydantic's metadata @declared_attr @@ -192,3 +210,56 @@ class PushNotificationConfigModel(PushNotificationConfigMixin, Base): """Default push notification config model with standard table name.""" __tablename__ = 'push_notification_configs' + + +# TaskEventMixin: append-only log of task events for the clustered event stream. +class TaskEventMixin: + """Mixin providing columns for an append-only task-event log. + + Written transactionally with the task row by the clustered + `VersionedTaskStore` and read by `DatabaseTaskEventStream`. Only used by + `a2a.server.cluster`; unused by the default single-process stores. + """ + + # Monotonic sequence id, primary source of ordering for the poller. + # BigInteger everywhere except SQLite, whose AUTOINCREMENT requires the + # column to be a plain INTEGER. + seq: Mapped[int] = mapped_column( + BigInteger().with_variant(Integer, 'sqlite'), + primary_key=True, + autoincrement=True, + ) + task_id: Mapped[str] = mapped_column(String(36), nullable=False, index=True) + owner: Mapped[str] = mapped_column(String(255), nullable=True) + # Task version produced by applying this event. + task_version: Mapped[int] = mapped_column(BigInteger, nullable=False) + # Serialized StreamResponse proto for the event. + event_data: Mapped[bytes] = mapped_column(LargeBinary, nullable=False) + + @override + def __repr__(self) -> str: + """Return a string representation of the task event.""" + return ( + f'<{self.__class__.__name__}(seq={getattr(self, "seq", None)}, ' + f'task_id="{self.task_id}", task_version={self.task_version})>' + ) + + +def create_task_event_model( + table_name: str = 'task_events', base: type[DeclarativeBase] = Base +) -> type: + """Create a TaskEventModel class with a configurable table name.""" + + class TaskEventModel(TaskEventMixin, base): # type: ignore + __tablename__ = table_name + + TaskEventModel.__name__ = f'TaskEventModel_{table_name}' + TaskEventModel.__qualname__ = f'TaskEventModel_{table_name}' + return TaskEventModel + + +# Default TaskEventModel for backward compatibility. +class TaskEventModel(TaskEventMixin, Base): + """Default task-event model with standard table name.""" + + __tablename__ = 'task_events' diff --git a/tests/server/cluster/__init__.py b/tests/server/cluster/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests/server/cluster/test_database_event_stream.py b/tests/server/cluster/test_database_event_stream.py new file mode 100644 index 000000000..f2a71ca86 --- /dev/null +++ b/tests/server/cluster/test_database_event_stream.py @@ -0,0 +1,278 @@ +"""Tests for `DatabaseTaskEventStream` including a multi-replica simulation.""" + +import asyncio +import contextlib +import os +import uuid + +from collections.abc import AsyncGenerator + +import pytest +import pytest_asyncio + +from _pytest.mark.structures import ParameterSet + + +pytest.importorskip('sqlalchemy', reason='Database tests require SQLAlchemy') + +from a2a.auth.user import User +from a2a.server.cluster import TaskVersion, VersionedEvent +from a2a.server.cluster.database_event_stream import ( + DatabaseTaskEventStream, + stream_response_to_event, +) +from a2a.server.cluster.database_task_store import VersionedDatabaseTaskStore +from a2a.server.context import ServerCallContext +from a2a.server.models import Base +from a2a.types.a2a_pb2 import ( + Message, + Part, + Role, + StreamResponse, + Task, + TaskArtifactUpdateEvent, + TaskState, + TaskStatus, + TaskStatusUpdateEvent, +) +from sqlalchemy.ext.asyncio import create_async_engine + + +class SampleUser(User): + def __init__(self, user_name: str): + self._user_name = user_name + + @property + def is_authenticated(self) -> bool: + return True + + @property + def user_name(self) -> str: + return self._user_name + + +TEST_CONTEXT = ServerCallContext(user=SampleUser('test_user')) + +POSTGRES_TEST_DSN = os.environ.get('POSTGRES_TEST_DSN') +MYSQL_TEST_DSN = os.environ.get('MYSQL_TEST_DSN') + + +def _sqlite_dsn() -> str: + # Unique shared-cache in-memory DB per test so multiple engines see the + # same data while different tests never collide. + name = f'eventstream_{uuid.uuid4().hex}' + return f'sqlite+aiosqlite:///file:{name}?mode=memory&cache=shared&uri=true' + + +DB_CONFIGS: list[ParameterSet | tuple[str | None, str]] = [ + pytest.param(('sqlite', 'sqlite'), id='sqlite') +] +if POSTGRES_TEST_DSN: + DB_CONFIGS.append( + pytest.param((POSTGRES_TEST_DSN, 'postgresql'), id='postgresql') + ) +else: + DB_CONFIGS.append( + pytest.param( + (None, 'postgresql'), + marks=pytest.mark.skip(reason='POSTGRES_TEST_DSN not set'), + id='postgresql_skipped', + ) + ) +if MYSQL_TEST_DSN: + DB_CONFIGS.append(pytest.param((MYSQL_TEST_DSN, 'mysql'), id='mysql')) +else: + DB_CONFIGS.append( + pytest.param( + (None, 'mysql'), + marks=pytest.mark.skip(reason='MYSQL_TEST_DSN not set'), + id='mysql_skipped', + ) + ) + + +def create_task(task_id: str = 'task-abc') -> Task: + return Task( + id=task_id, + context_id='ctx', + status=TaskStatus(state=TaskState.TASK_STATE_SUBMITTED), + ) + + +def status_event( + task_id: str = 'task-abc', + state: TaskState = TaskState.TASK_STATE_WORKING, +) -> TaskStatusUpdateEvent: + return TaskStatusUpdateEvent( + task_id=task_id, + context_id='ctx', + status=TaskStatus(state=state), + ) + + +@pytest_asyncio.fixture(params=DB_CONFIGS) +async def db_url(request) -> AsyncGenerator[str, None]: + param, dialect = request.param + if param is None: + pytest.skip(f'DSN for {dialect} not set.') + url = _sqlite_dsn() if param == 'sqlite' else param + + engine = create_async_engine(url) + # Keep one connection open for the whole test so a shared-cache in-memory + # SQLite database is not torn down while other engines connect to it. + keepalive = await engine.connect() + try: + async with engine.begin() as conn: + await conn.run_sync(Base.metadata.create_all) + yield url + finally: + await keepalive.close() + await engine.dispose() + + +# --- Pure serialization helper tests (no DB needed) --- + + +def test_stream_response_to_event_roundtrips_all_kinds() -> None: + task = create_task() + msg = Message( + message_id='m1', role=Role.ROLE_AGENT, parts=[Part(text='hi')] + ) + status = status_event() + artifact = TaskArtifactUpdateEvent(task_id='task-abc', context_id='ctx') + + for event in (task, msg, status, artifact): + response = StreamResponse() + # Round-trip through serialize to mimic DB storage. + from a2a.utils.proto_utils import to_stream_response + + response = to_stream_response(event) + raw = response.SerializeToString() + parsed = StreamResponse() + parsed.ParseFromString(raw) + assert stream_response_to_event(parsed) == event + + +def test_stream_response_to_event_raises_on_empty() -> None: + with pytest.raises(ValueError, match='no known payload'): + stream_response_to_event(StreamResponse()) + + +def test_as_int_rejects_non_integer_version() -> None: + from a2a.server.cluster.database_event_stream import _as_int + + with pytest.raises(TypeError, match='integer versions'): + _as_int(TaskVersion('etag')) + + +@pytest.mark.asyncio +@pytest.mark.timeout(10) +async def test_initialize_creates_table_and_is_idempotent( + tmp_path, +) -> None: + # A dedicated file-backed SQLite DB (not shared-cache in-memory) so + # create_table=True owns the schema with no cross-engine contention. + url = f'sqlite+aiosqlite:///{tmp_path / "events.db"}' + engine = create_async_engine(url) + try: + stream = DatabaseTaskEventStream(engine=engine, create_table=True) + await stream.initialize() + # Second call is a no-op (already initialized) -> early-return path. + await stream.initialize() + # Subscribing seeds the cursor from an empty table (0) without error. + agen = stream.subscribe('task-abc', after=TaskVersion.MISSING) + consumer = asyncio.ensure_future(agen.__anext__()) + await asyncio.sleep(0.05) + consumer.cancel() + with contextlib.suppress(asyncio.CancelledError): + await consumer + await agen.aclose() + finally: + await engine.dispose() + + +# --- Integration tests (shared DB across store + stream) --- + + +@pytest.mark.asyncio +@pytest.mark.timeout(10) +async def test_stream_delivers_events_written_by_store(db_url: str) -> None: + """A store writes events; a stream on a SEPARATE engine reads them. + + This simulates the producing replica (store) and a subscribing replica + (stream) sharing one database. + """ + store_engine = create_async_engine(db_url) + stream_engine = create_async_engine(db_url) + try: + store = VersionedDatabaseTaskStore( + engine=store_engine, create_table=False + ) + stream = DatabaseTaskEventStream( + engine=stream_engine, create_table=False, poll_interval_s=0.05 + ) + await store.initialize() + await stream.initialize() + + # Snapshot version (as a subscriber would have after reading the task). + v1 = await store.save( + create_task(), + event=None, + prev=None, + prev_version=TaskVersion.MISSING, + context=TEST_CONTEXT, + ) + + received: list[TaskState] = [] + + async def consume() -> None: + async for ve in stream.subscribe('task-abc', after=v1): + received.append(ve.event.status.state) + if ve.event.status.state == TaskState.TASK_STATE_COMPLETED: + return + + consumer = asyncio.create_task(consume()) + await asyncio.sleep(0.1) # subscriber captures the tail cursor + + # Producer writes two events with the task, in transaction. + v2 = await store.save( + create_task(), + event=status_event(state=TaskState.TASK_STATE_WORKING), + prev=None, + prev_version=v1, + context=TEST_CONTEXT, + ) + await store.save( + create_task(), + event=status_event(state=TaskState.TASK_STATE_COMPLETED), + prev=None, + prev_version=v2, + context=TEST_CONTEXT, + ) + + await asyncio.wait_for(consumer, timeout=8) + assert received == [ + TaskState.TASK_STATE_WORKING, + TaskState.TASK_STATE_COMPLETED, + ] + finally: + await store_engine.dispose() + await stream_engine.dispose() + + +@pytest.mark.asyncio +@pytest.mark.timeout(10) +async def test_publish_is_noop(db_url: str) -> None: + stream_engine = create_async_engine(db_url) + try: + stream = DatabaseTaskEventStream( + engine=stream_engine, create_table=False + ) + await stream.initialize() + # publish is a no-op for the DB stream; must not raise. + await stream.publish( + 'task-abc', VersionedEvent(create_task(), TaskVersion(1)) + ) + await stream.destroy('task-abc') # also a no-op + finally: + await stream_engine.dispose() diff --git a/tests/server/cluster/test_database_versioned_task_store.py b/tests/server/cluster/test_database_versioned_task_store.py new file mode 100644 index 000000000..1b852e4b6 --- /dev/null +++ b/tests/server/cluster/test_database_versioned_task_store.py @@ -0,0 +1,505 @@ +"""Tests for `VersionedDatabaseTaskStore` (CAS over SQLAlchemy).""" + +import os + +from collections.abc import AsyncGenerator + +import pytest +import pytest_asyncio + +from _pytest.mark.structures import ParameterSet + + +pytest.importorskip('sqlalchemy', reason='Database tests require SQLAlchemy') + +from a2a.auth.user import User +from a2a.server.cluster import ( + ConcurrentTaskModificationError, + TaskVersion, + VersionedTaskStore, +) +from a2a.server.cluster.database_task_store import VersionedDatabaseTaskStore +from a2a.server.context import ServerCallContext +from a2a.server.models import Base +from a2a.types.a2a_pb2 import ( + ListTasksRequest, + Task, + TaskState, + TaskStatus, +) +from sqlalchemy.ext.asyncio import create_async_engine + + +class SampleUser(User): + """A test implementation of the User interface.""" + + def __init__(self, user_name: str): + self._user_name = user_name + + @property + def is_authenticated(self) -> bool: + return True + + @property + def user_name(self) -> str: + return self._user_name + + +TEST_CONTEXT = ServerCallContext(user=SampleUser('test_user')) + + +SQLITE_TEST_DSN = 'sqlite+aiosqlite:///file:testdb_versioned?mode=memory&cache=shared&uri=true' +POSTGRES_TEST_DSN = os.environ.get('POSTGRES_TEST_DSN') +MYSQL_TEST_DSN = os.environ.get('MYSQL_TEST_DSN') + +DB_CONFIGS: list[ParameterSet | tuple[str | None, str]] = [ + pytest.param((SQLITE_TEST_DSN, 'sqlite'), id='sqlite') +] + +if POSTGRES_TEST_DSN: + DB_CONFIGS.append( + pytest.param((POSTGRES_TEST_DSN, 'postgresql'), id='postgresql') + ) +else: + DB_CONFIGS.append( + pytest.param( + (None, 'postgresql'), + marks=pytest.mark.skip(reason='POSTGRES_TEST_DSN not set'), + id='postgresql_skipped', + ) + ) + +if MYSQL_TEST_DSN: + DB_CONFIGS.append(pytest.param((MYSQL_TEST_DSN, 'mysql'), id='mysql')) +else: + DB_CONFIGS.append( + pytest.param( + (None, 'mysql'), + marks=pytest.mark.skip(reason='MYSQL_TEST_DSN not set'), + id='mysql_skipped', + ) + ) + + +def create_task( + task_id: str = 'task-abc', + context_id: str = 'session-xyz', + state: TaskState = TaskState.TASK_STATE_SUBMITTED, +) -> Task: + return Task( + id=task_id, + context_id=context_id, + status=TaskStatus(state=state), + ) + + +@pytest_asyncio.fixture(params=DB_CONFIGS) +async def versioned_store( + request, +) -> AsyncGenerator[VersionedDatabaseTaskStore, None]: + db_url, dialect_name = request.param + if db_url is None: + pytest.skip(f'DSN for {dialect_name} not set in environment variables.') + + engine = create_async_engine(db_url) + try: + async with engine.begin() as conn: + await conn.run_sync(Base.metadata.create_all) + store = VersionedDatabaseTaskStore(engine=engine, create_table=False) + await store.initialize() + yield store + finally: + async with engine.begin() as conn: + await conn.run_sync(Base.metadata.drop_all) + await engine.dispose() + + +@pytest.mark.asyncio +async def test_is_a_versioned_task_store( + versioned_store: VersionedDatabaseTaskStore, +) -> None: + assert isinstance(versioned_store, VersionedTaskStore) + + +@pytest.mark.asyncio +async def test_as_task_store_exposes_underlying_store( + versioned_store: VersionedDatabaseTaskStore, +) -> None: + from a2a.server.tasks.task_store import TaskStore + + assert isinstance(versioned_store.as_task_store, TaskStore) + + +@pytest.mark.asyncio +async def test_non_integer_version_on_update_raises_type_error( + versioned_store: VersionedDatabaseTaskStore, +) -> None: + await versioned_store.save( + create_task(), + event=None, + prev=None, + prev_version=TaskVersion.MISSING, + context=TEST_CONTEXT, + ) + # A string-valued version can't back the integer `version` column. + with pytest.raises(TypeError, match='integer versions'): + await versioned_store.save( + create_task(state=TaskState.TASK_STATE_WORKING), + event=None, + prev=None, + prev_version=TaskVersion('not-an-int'), + context=TEST_CONTEXT, + ) + + +@pytest.mark.asyncio +async def test_get_missing_returns_none( + versioned_store: VersionedDatabaseTaskStore, +) -> None: + assert await versioned_store.get('nope', TEST_CONTEXT) is None + + +@pytest.mark.asyncio +async def test_first_save_then_get( + versioned_store: VersionedDatabaseTaskStore, +) -> None: + task = create_task() + v1 = await versioned_store.save( + task, + event=None, + prev=None, + prev_version=TaskVersion.MISSING, + context=TEST_CONTEXT, + ) + assert not v1.is_missing + + stored = await versioned_store.get('task-abc', TEST_CONTEXT) + assert stored is not None + assert stored.task.id == 'task-abc' + assert stored.version == v1 + + +@pytest.mark.asyncio +async def test_update_with_matching_version_succeeds( + versioned_store: VersionedDatabaseTaskStore, +) -> None: + v1 = await versioned_store.save( + create_task(), + event=None, + prev=None, + prev_version=TaskVersion.MISSING, + context=TEST_CONTEXT, + ) + v2 = await versioned_store.save( + create_task(state=TaskState.TASK_STATE_WORKING), + event=None, + prev=None, + prev_version=v1, + context=TEST_CONTEXT, + ) + assert v2.is_after(v1) + stored = await versioned_store.get('task-abc', TEST_CONTEXT) + assert stored is not None + assert stored.task.status.state == TaskState.TASK_STATE_WORKING + assert stored.version == v2 + + +@pytest.mark.asyncio +async def test_stale_update_raises_conflict( + versioned_store: VersionedDatabaseTaskStore, +) -> None: + v1 = await versioned_store.save( + create_task(), + event=None, + prev=None, + prev_version=TaskVersion.MISSING, + context=TEST_CONTEXT, + ) + # Winner advances the task. + await versioned_store.save( + create_task(state=TaskState.TASK_STATE_WORKING), + event=None, + prev=None, + prev_version=v1, + context=TEST_CONTEXT, + ) + # Loser still holds v1 -> CAS fails. + with pytest.raises(ConcurrentTaskModificationError): + await versioned_store.save( + create_task(state=TaskState.TASK_STATE_COMPLETED), + event=None, + prev=None, + prev_version=v1, + context=TEST_CONTEXT, + ) + # Store still reflects the winner's write. + stored = await versioned_store.get('task-abc', TEST_CONTEXT) + assert stored is not None + assert stored.task.status.state == TaskState.TASK_STATE_WORKING + + +@pytest.mark.asyncio +async def test_concurrent_first_insert_one_wins( + versioned_store: VersionedDatabaseTaskStore, +) -> None: + # First insert succeeds. + await versioned_store.save( + create_task(), + event=None, + prev=None, + prev_version=TaskVersion.MISSING, + context=TEST_CONTEXT, + ) + # A second "first write" (MISSING) for the same id must not silently + # clobber; it collides on the primary key -> conflict. + with pytest.raises(ConcurrentTaskModificationError): + await versioned_store.save( + create_task(state=TaskState.TASK_STATE_WORKING), + event=None, + prev=None, + prev_version=TaskVersion.MISSING, + context=TEST_CONTEXT, + ) + + +@pytest.mark.asyncio +async def test_delete_removes_task( + versioned_store: VersionedDatabaseTaskStore, +) -> None: + v1 = await versioned_store.save( + create_task(), + event=None, + prev=None, + prev_version=TaskVersion.MISSING, + context=TEST_CONTEXT, + ) + assert not v1.is_missing + await versioned_store.delete('task-abc', TEST_CONTEXT) + assert await versioned_store.get('task-abc', TEST_CONTEXT) is None + + +@pytest.mark.asyncio +async def test_list_returns_saved_tasks( + versioned_store: VersionedDatabaseTaskStore, +) -> None: + for i in range(3): + await versioned_store.save( + create_task(task_id=f'task-{i}'), + event=None, + prev=None, + prev_version=TaskVersion.MISSING, + context=TEST_CONTEXT, + ) + resp = await versioned_store.list(ListTasksRequest(), TEST_CONTEXT) + assert {t.id for t in resp.tasks} == {'task-0', 'task-1', 'task-2'} + + +@pytest.mark.asyncio +async def test_save_retries_transient_operational_error( + versioned_store: VersionedDatabaseTaskStore, +) -> None: + """A transient OperationalError on write is retried, not surfaced.""" + from unittest import mock + + from sqlalchemy.exc import OperationalError + + calls = {'n': 0} + real_save_once = versioned_store._save_once # noqa: SLF001 + + async def flaky_save_once(*args, **kwargs): # noqa: ANN002, ANN003 + calls['n'] += 1 + if calls['n'] == 1: + raise OperationalError('stmt', {}, Exception('database is locked')) + return await real_save_once(*args, **kwargs) + + with mock.patch.object( + versioned_store, '_save_once', side_effect=flaky_save_once + ): + version = await versioned_store.save( + create_task(), + event=None, + prev=None, + prev_version=TaskVersion.MISSING, + context=TEST_CONTEXT, + ) + assert calls['n'] == 2 # first failed, second succeeded + assert not version.is_missing + + +@pytest.mark.asyncio +async def test_save_raises_after_exhausting_retries( + versioned_store: VersionedDatabaseTaskStore, +) -> None: + """Persistent OperationalError surfaces after exactly `max_attempts` tries.""" + from unittest import mock + + from sqlalchemy.exc import OperationalError + + # A store tuned to two attempts with no backoff, so the loop is bounded by + # max_attempts (not retries after the first) and the test does not sleep. + store = VersionedDatabaseTaskStore( + engine=versioned_store._db.engine, # noqa: SLF001 + create_table=False, + max_attempts=2, + retry_delay_s=0.0, + ) + + calls = {'n': 0} + + async def always_locked(*args, **kwargs): # noqa: ANN002, ANN003 + calls['n'] += 1 + raise OperationalError('stmt', {}, Exception('database is locked')) + + with ( + mock.patch.object(store, '_save_once', side_effect=always_locked), + pytest.raises(OperationalError), + ): + await store.save( + create_task(), + event=None, + prev=None, + prev_version=TaskVersion.MISSING, + context=TEST_CONTEXT, + ) + assert calls['n'] == 2 # initial try + one retry, then surfaced + + +@pytest.mark.asyncio +async def test_get_retries_transient_operational_error( + versioned_store: VersionedDatabaseTaskStore, +) -> None: + """A transient OperationalError on read is retried, not surfaced.""" + from unittest import mock + + from sqlalchemy.exc import OperationalError + + await versioned_store.save( + create_task(), + event=None, + prev=None, + prev_version=TaskVersion.MISSING, + context=TEST_CONTEXT, + ) + + calls = {'n': 0} + real_get_once = versioned_store._get_once # noqa: SLF001 + + async def flaky_get_once(*args, **kwargs): # noqa: ANN002, ANN003 + calls['n'] += 1 + if calls['n'] == 1: + raise OperationalError('stmt', {}, Exception('database is locked')) + return await real_get_once(*args, **kwargs) + + with mock.patch.object( + versioned_store, '_get_once', side_effect=flaky_get_once + ): + stored = await versioned_store.get('task-abc', TEST_CONTEXT) + assert calls['n'] == 2 + assert stored is not None + assert not stored.version.is_missing + + +@pytest.mark.asyncio +async def test_get_raises_after_exhausting_retries( + versioned_store: VersionedDatabaseTaskStore, +) -> None: + """Read surfaces the error after exactly `max_attempts` tries.""" + from unittest import mock + + from sqlalchemy.exc import OperationalError + + # As with the write path, bound the read loop to two attempts with no + # backoff so the count is asserted and no real sleep occurs. + store = VersionedDatabaseTaskStore( + engine=versioned_store._db.engine, # noqa: SLF001 + create_table=False, + max_attempts=2, + retry_delay_s=0.0, + ) + + calls = {'n': 0} + + async def always_locked(*args, **kwargs): # noqa: ANN002, ANN003 + calls['n'] += 1 + raise OperationalError('stmt', {}, Exception('database is locked')) + + with ( + mock.patch.object(store, '_get_once', side_effect=always_locked), + pytest.raises(OperationalError), + ): + await store.get('task-abc', TEST_CONTEXT) + assert calls['n'] == 2 # initial try + one retry, then surfaced + + +@pytest.mark.asyncio +async def test_retry_defaults( + versioned_store: VersionedDatabaseTaskStore, +) -> None: + """The shipped retry budget defaults are 5 attempts / 0.02s backoff.""" + assert versioned_store._max_attempts == 5 # noqa: SLF001 + assert versioned_store._retry_delay_s == 0.02 # noqa: SLF001 + + +@pytest.mark.asyncio +async def test_cancel_overwrites_without_version_match( + versioned_store: VersionedDatabaseTaskStore, +) -> None: + v1 = await versioned_store.save( + create_task(state=TaskState.TASK_STATE_SUBMITTED), + event=None, + prev=None, + prev_version=TaskVersion.MISSING, + context=TEST_CONTEXT, + ) + await versioned_store.save( + create_task(state=TaskState.TASK_STATE_WORKING), + event=None, + prev=None, + prev_version=v1, + context=TEST_CONTEXT, + ) + await versioned_store.save( + create_task(state=TaskState.TASK_STATE_CANCELED), + event=None, + prev=None, + prev_version=v1, + context=TEST_CONTEXT, + ) + stored = await versioned_store.get('task-abc', TEST_CONTEXT) + assert stored is not None + assert stored.task.status.state == TaskState.TASK_STATE_CANCELED + + +@pytest.mark.asyncio +async def test_cancel_on_terminal_task_raises( + versioned_store: VersionedDatabaseTaskStore, +) -> None: + v1 = await versioned_store.save( + create_task(state=TaskState.TASK_STATE_COMPLETED), + event=None, + prev=None, + prev_version=TaskVersion.MISSING, + context=TEST_CONTEXT, + ) + with pytest.raises(ConcurrentTaskModificationError): + await versioned_store.save( + create_task(state=TaskState.TASK_STATE_CANCELED), + event=None, + prev=None, + prev_version=v1, + context=TEST_CONTEXT, + ) + + +@pytest.mark.asyncio +async def test_cancel_absent_task_raises( + versioned_store: VersionedDatabaseTaskStore, +) -> None: + with pytest.raises(ConcurrentTaskModificationError): + await versioned_store.save( + create_task(state=TaskState.TASK_STATE_CANCELED), + event=None, + prev=None, + prev_version=TaskVersion(123), + context=TEST_CONTEXT, + ) diff --git a/tests/server/cluster/test_task_store_adapter.py b/tests/server/cluster/test_task_store_adapter.py new file mode 100644 index 000000000..3124c0879 --- /dev/null +++ b/tests/server/cluster/test_task_store_adapter.py @@ -0,0 +1,165 @@ +"""Tests for `LegacyTaskStoreAdapter` and the `VersionedTaskStore` contract.""" + +import pytest + +from a2a.auth.user import User +from a2a.server.cluster import ( + ConcurrentTaskModificationError, + LegacyTaskStoreAdapter, + TaskVersion, + VersionedTaskStore, +) +from a2a.server.context import ServerCallContext +from a2a.server.tasks import InMemoryTaskStore +from a2a.types.a2a_pb2 import ( + ListTasksRequest, + Task, + TaskState, + TaskStatus, +) + + +class SampleUser(User): + """A test implementation of the User interface.""" + + def __init__(self, user_name: str): + self._user_name = user_name + + @property + def is_authenticated(self) -> bool: + return True + + @property + def user_name(self) -> str: + return self._user_name + + +TEST_CONTEXT = ServerCallContext(user=SampleUser('test_user')) + + +def create_minimal_task( + task_id: str = 'task-abc', context_id: str = 'session-xyz' +) -> Task: + return Task( + id=task_id, + context_id=context_id, + status=TaskStatus(state=TaskState.TASK_STATE_SUBMITTED), + ) + + +def make_adapter() -> LegacyTaskStoreAdapter: + return LegacyTaskStoreAdapter(InMemoryTaskStore()) + + +def test_concurrent_modification_error_carries_task_id() -> None: + err = ConcurrentTaskModificationError('task-abc') + assert err.task_id == 'task-abc' + assert 'task-abc' in str(err) + assert isinstance(err, Exception) + + +def test_adapter_is_a_versioned_task_store() -> None: + assert isinstance(make_adapter(), VersionedTaskStore) + + +def test_store_property_exposes_wrapped_store() -> None: + store = InMemoryTaskStore() + adapter = LegacyTaskStoreAdapter(store) + assert adapter.store is store + + +@pytest.mark.asyncio +async def test_save_returns_missing_version() -> None: + adapter = make_adapter() + task = create_minimal_task() + version = await adapter.save( + task, + event=None, + prev=None, + prev_version=TaskVersion.MISSING, + context=TEST_CONTEXT, + ) + assert version.is_missing + + +@pytest.mark.asyncio +async def test_get_returns_task_and_missing_version() -> None: + adapter = make_adapter() + task = create_minimal_task() + await adapter.save( + task, + event=None, + prev=None, + prev_version=TaskVersion.MISSING, + context=TEST_CONTEXT, + ) + + stored = await adapter.get('task-abc', TEST_CONTEXT) + assert stored is not None + assert stored.task == task + assert stored.version.is_missing + + +@pytest.mark.asyncio +async def test_get_missing_task_returns_none() -> None: + adapter = make_adapter() + assert await adapter.get('does-not-exist', TEST_CONTEXT) is None + + +@pytest.mark.asyncio +async def test_save_never_raises_on_stale_prev_version() -> None: + # An unversioned store performs no CAS: a "stale" prev_version is ignored + # and the write succeeds (last-writer-wins), matching today's behaviour. + adapter = make_adapter() + task = create_minimal_task() + await adapter.save( + task, + event=None, + prev=None, + prev_version=TaskVersion.MISSING, + context=TEST_CONTEXT, + ) + # Save again with a deliberately non-matching, non-missing prev_version. + updated = create_minimal_task() + updated.status.state = TaskState.TASK_STATE_WORKING + version = await adapter.save( + updated, + event=None, + prev=task, + prev_version=TaskVersion(999), + context=TEST_CONTEXT, + ) + assert version.is_missing + stored = await adapter.get('task-abc', TEST_CONTEXT) + assert stored is not None + assert stored.task.status.state == TaskState.TASK_STATE_WORKING + + +@pytest.mark.asyncio +async def test_delete_delegates_to_inner() -> None: + adapter = make_adapter() + task = create_minimal_task() + await adapter.save( + task, + event=None, + prev=None, + prev_version=TaskVersion.MISSING, + context=TEST_CONTEXT, + ) + await adapter.delete('task-abc', TEST_CONTEXT) + assert await adapter.get('task-abc', TEST_CONTEXT) is None + + +@pytest.mark.asyncio +async def test_list_delegates_to_inner() -> None: + adapter = make_adapter() + for i in range(3): + await adapter.save( + create_minimal_task(task_id=f'task-{i}'), + event=None, + prev=None, + prev_version=TaskVersion.MISSING, + context=TEST_CONTEXT, + ) + response = await adapter.list(ListTasksRequest(), TEST_CONTEXT) + assert {t.id for t in response.tasks} == {'task-0', 'task-1', 'task-2'} diff --git a/tests/server/cluster/test_version.py b/tests/server/cluster/test_version.py new file mode 100644 index 000000000..65bc0e67a --- /dev/null +++ b/tests/server/cluster/test_version.py @@ -0,0 +1,102 @@ +"""Tests for `TaskVersion`.""" + +import pytest + +from a2a.server.cluster import TaskVersion + + +def test_missing_is_missing() -> None: + assert TaskVersion.MISSING.is_missing is True + + +def test_real_version_is_not_missing() -> None: + assert TaskVersion(1).is_missing is False + assert TaskVersion('etag-abc').is_missing is False + + +def test_zero_int_is_missing() -> None: + # 0 is the reserved "not tracked" value. + assert TaskVersion(0).is_missing is True + assert TaskVersion(0) == TaskVersion.MISSING + + +def test_string_zero_is_not_missing() -> None: + # '0' != 0 in Python, so a string ETag of '0' is a real version. + assert TaskVersion('0').is_missing is False + + +def test_equality_by_value() -> None: + assert TaskVersion(5) == TaskVersion(5) + assert TaskVersion(5) != TaskVersion(6) + assert TaskVersion('a') == TaskVersion('a') + assert TaskVersion('a') != TaskVersion('b') + + +def test_equality_with_non_taskversion() -> None: + assert TaskVersion(5) != 5 + assert TaskVersion(5) != 'TaskVersion(5)' + assert (TaskVersion(5) == object()) is False + + +def test_hashable_and_usable_in_sets_and_dicts() -> None: + versions = {TaskVersion(1), TaskVersion(1), TaskVersion(2)} + assert versions == {TaskVersion(1), TaskVersion(2)} + + mapping = {TaskVersion('x'): 'value'} + assert mapping[TaskVersion('x')] == 'value' + + +def test_repr() -> None: + assert repr(TaskVersion.MISSING) == 'TaskVersion.MISSING' + assert repr(TaskVersion(0)) == 'TaskVersion.MISSING' + assert repr(TaskVersion(7)) == 'TaskVersion(7)' + assert repr(TaskVersion('etag')) == "TaskVersion('etag')" + + +def test_is_after_ordering_integers() -> None: + assert TaskVersion(2).is_after(TaskVersion(1)) is True + assert TaskVersion(1).is_after(TaskVersion(2)) is False + + +def test_is_after_equal_versions_is_false() -> None: + # "strictly later" -> equal is not after. + assert TaskVersion(1).is_after(TaskVersion(1)) is False + + +def test_is_after_missing_asymmetry() -> None: + # Anything real is "after" MISSING (untracked baseline is oldest). + assert TaskVersion(1).is_after(TaskVersion.MISSING) is True + # MISSING is never after a real version. + assert TaskVersion.MISSING.is_after(TaskVersion(1)) is False + + +def test_is_after_both_missing_is_false() -> None: + # other.is_missing short-circuits to True per the contract. + assert TaskVersion.MISSING.is_after(TaskVersion.MISSING) is True + + +def test_is_after_string_versions() -> None: + # Works for orderable non-int values too (e.g. zero-padded ETags). + assert TaskVersion('0002').is_after(TaskVersion('0001')) is True + assert TaskVersion('0001').is_after(TaskVersion('0002')) is False + + +def test_missing_is_a_singleton_value() -> None: + # Not necessarily the same object, but value-equal and both missing. + assert TaskVersion(0) == TaskVersion.MISSING + assert TaskVersion(0).is_missing and TaskVersion.MISSING.is_missing + + +def test_is_after_mismatched_value_types_raises() -> None: + # A store uses one value type consistently; comparing an int-backed + # version against a str-backed one is a programming error. + with pytest.raises(TypeError, match='different value types'): + TaskVersion(1).is_after(TaskVersion('a')) + with pytest.raises(TypeError, match='different value types'): + TaskVersion('a').is_after(TaskVersion(1)) + + +def test_is_after_mismatch_guard_not_triggered_when_either_missing() -> None: + # MISSING short-circuits before the type check, so no TypeError. + assert TaskVersion('a').is_after(TaskVersion.MISSING) is True + assert TaskVersion.MISSING.is_after(TaskVersion(1)) is False