From db0ad4ba525af34ddd6760567374298650281634 Mon Sep 17 00:00:00 2001 From: Mykyta Netipa Date: Wed, 30 Sep 2026 08:50:11 +0000 Subject: [PATCH 1/3] fix(server): use keyset cursors for ListTasks page tokens A ListTasks page token was the ID of the first task of the next page, and each store looked that task up again to decide where to resume. If the task was updated between requests, the listing restarted and returned duplicates; if it was deleted, the request failed with InvalidParamsError. The token now carries the sort position (timestamp, id) of the last task returned, and the next page starts strictly after it, so changes to that task no longer affect pagination. The database store no longer needs an extra query to resolve the token. Tokens issued by earlier versions are still accepted and resolved the old way. Fixes #1280 --- src/a2a/server/tasks/database_task_store.py | 144 +++++++++++++----- src/a2a/server/tasks/inmemory_task_store.py | 78 +++++++--- src/a2a/utils/task.py | 82 +++++++++- .../server/tasks/test_database_task_store.py | 132 +++++++++++++++- .../server/tasks/test_inmemory_task_store.py | 115 +++++++++++++- tests/utils/test_task.py | 57 +++++++ 6 files changed, 530 insertions(+), 78 deletions(-) diff --git a/src/a2a/server/tasks/database_task_store.py b/src/a2a/server/tasks/database_task_store.py index 409518d8d..f3f2e9dd0 100644 --- a/src/a2a/server/tasks/database_task_store.py +++ b/src/a2a/server/tasks/database_task_store.py @@ -1,6 +1,7 @@ import logging from collections.abc import Callable +from datetime import datetime, timedelta, timezone from typing import Any, cast @@ -34,11 +35,33 @@ from a2a.types.a2a_pb2 import Task from a2a.utils.constants import DEFAULT_LIST_TASKS_PAGE_SIZE from a2a.utils.errors import InvalidParamsError -from a2a.utils.task import decode_page_token, encode_page_token +from a2a.utils.task import ( + ListTasksCursor, + decode_list_tasks_cursor, + decode_page_token, + encode_list_tasks_cursor, +) logger = logging.getLogger(__name__) +_EPOCH = datetime(1970, 1, 1) # noqa: DTZ001 -- last_updated is naive UTC + + +def _datetime_to_ns(value: datetime) -> int: + """Nanoseconds since the epoch for a stored (naive UTC) `last_updated`.""" + if value.tzinfo is not None: + value = value.astimezone(timezone.utc).replace(tzinfo=None) + delta = value - _EPOCH + return ( + delta.days * 86_400 + delta.seconds + ) * 1_000_000_000 + delta.microseconds * 1_000 + + +def _ns_to_datetime(timestamp_ns: int) -> datetime: + """Inverse of `_datetime_to_ns`; exact for values it produced.""" + return _EPOCH + timedelta(microseconds=timestamp_ns // 1_000) + class DatabaseTaskStore(TaskStore): """SQLAlchemy-based implementation of TaskStore. @@ -257,64 +280,109 @@ async def list( self.task_model.id.desc(), ) - # Get paginated results + # Get paginated results. The page token carries the position of the + # last task returned, so it stays valid if that task is updated or + # deleted. A task updated mid-listing can move above the cursor and + # be skipped. if params.page_token: - start_task_id = decode_page_token(params.page_token) - start_task = ( - await session.execute( - select(self.task_model).where( - and_( - self.task_model.id == start_task_id, - self.task_model.owner == owner, - ) + cursor = decode_list_tasks_cursor(params.page_token) + if cursor is None: + stmt = stmt.where( + await self._legacy_page_clause( + session, owner, params.page_token ) ) - ).scalar_one_or_none() - if not start_task: - raise InvalidParamsError( - f'Invalid page token: {params.page_token}' - ) - - start_task_timestamp = start_task.last_updated - where_clauses = [] - if start_task_timestamp: - where_clauses.append( - and_( - timestamp_col == start_task_timestamp, - self.task_model.id <= start_task_id, - ) - ) - where_clauses.append(timestamp_col < start_task_timestamp) - where_clauses.append(timestamp_col.is_(None)) else: - where_clauses.append( - and_( - timestamp_col.is_(None), - self.task_model.id <= start_task_id, - ) - ) - stmt = stmt.where(or_(*where_clauses)) + stmt = stmt.where(self._after_cursor(cursor)) page_size = params.page_size or DEFAULT_LIST_TASKS_PAGE_SIZE stmt = stmt.limit(page_size + 1) # Add 1 for next page token result = await session.execute(stmt) tasks_models = result.scalars().all() - tasks = [self._from_orm(task_model) for task_model in tasks_models] + page_models = tasks_models[:page_size] next_page_token = ( - encode_page_token(tasks[-1].id) - if len(tasks) == page_size + 1 + encode_list_tasks_cursor(self._cursor_for(page_models[-1])) + if len(tasks_models) == page_size + 1 else None ) return a2a_pb2.ListTasksResponse( - tasks=tasks[:page_size], + tasks=[ + self._from_orm(task_model) for task_model in page_models + ], total_size=total_count, next_page_token=next_page_token, page_size=page_size, ) + @staticmethod + def _cursor_for(task_model: TaskModel) -> ListTasksCursor: + # From the stored column, not the proto, so the cursor compares equal + # at the database's own timestamp precision. + last_updated = task_model.last_updated + return ListTasksCursor( + timestamp_ns=_datetime_to_ns(last_updated) + if last_updated is not None + else None, + task_id=task_model.id, + ) + + def _after_cursor(self, cursor: ListTasksCursor) -> Any: + """Rows strictly after `cursor` in the `ListTasks` sort order.""" + timestamp_col = self.task_model.last_updated + if cursor.timestamp_ns is None: + return and_( + timestamp_col.is_(None), self.task_model.id < cursor.task_id + ) + try: + cursor_timestamp = _ns_to_datetime(cursor.timestamp_ns) + except OverflowError as e: + raise InvalidParamsError('Invalid page token') from e + return or_( + timestamp_col < cursor_timestamp, + and_( + timestamp_col == cursor_timestamp, + self.task_model.id < cursor.task_id, + ), + timestamp_col.is_(None), + ) + + async def _legacy_page_clause( + self, session: AsyncSession, owner: str, page_token: str + ) -> Any: + """Resolves a legacy page token, which names the first task of the page.""" + timestamp_col = self.task_model.last_updated + start_task_id = decode_page_token(page_token) + start_task = ( + await session.execute( + select(self.task_model).where( + and_( + self.task_model.id == start_task_id, + self.task_model.owner == owner, + ) + ) + ) + ).scalar_one_or_none() + if not start_task: + raise InvalidParamsError(f'Invalid page token: {page_token}') + + start_task_timestamp = start_task.last_updated + if start_task_timestamp: + return or_( + and_( + timestamp_col == start_task_timestamp, + self.task_model.id <= start_task_id, + ), + timestamp_col < start_task_timestamp, + timestamp_col.is_(None), + ) + return and_( + timestamp_col.is_(None), + self.task_model.id <= start_task_id, + ) + async def delete(self, task_id: str, context: ServerCallContext) -> None: """Deletes a task from the database by ID, for the given owner.""" await self._ensure_initialized() diff --git a/src/a2a/server/tasks/inmemory_task_store.py b/src/a2a/server/tasks/inmemory_task_store.py index 26e1cad0d..e264fe73a 100644 --- a/src/a2a/server/tasks/inmemory_task_store.py +++ b/src/a2a/server/tasks/inmemory_task_store.py @@ -1,6 +1,8 @@ import logging import threading +from collections.abc import Sequence + from a2a.server.context import ServerCallContext from a2a.server.owner_resolver import OwnerResolver, resolve_user_scope from a2a.server.tasks.copying_task_store import CopyingTaskStoreAdapter @@ -9,12 +11,36 @@ from a2a.types.a2a_pb2 import Task from a2a.utils.constants import DEFAULT_LIST_TASKS_PAGE_SIZE from a2a.utils.errors import InvalidParamsError -from a2a.utils.task import decode_page_token, encode_page_token +from a2a.utils.task import ( + ListTasksCursor, + decode_list_tasks_cursor, + decode_page_token, + encode_list_tasks_cursor, +) logger = logging.getLogger(__name__) +def _list_sort_key(task: Task) -> tuple[bool, int, str]: + """`ListTasks` sort key: `(has timestamp, timestamp, id)`, sorted descending.""" + has_timestamp = task.HasField('status') and task.status.HasField( + 'timestamp' + ) + return ( + has_timestamp, + task.status.timestamp.ToNanoseconds() if has_timestamp else 0, + task.id, + ) + + +def _cursor_for(task: Task) -> ListTasksCursor: + has_timestamp, timestamp_ns, task_id = _list_sort_key(task) + return ListTasksCursor( + timestamp_ns=timestamp_ns if has_timestamp else None, task_id=task_id + ) + + class _InMemoryTaskStoreImpl(TaskStore): """Internal In-memory implementation of TaskStore. @@ -108,38 +134,31 @@ async def list( ] # Order tasks by last update time. To ensure stable sorting, in cases where timestamps are null or not unique, do a second order comparison of IDs. - tasks.sort( - key=lambda task: ( - task.status.HasField('timestamp') - if task.HasField('status') - else False, - task.status.timestamp.ToNanoseconds() - if task.HasField('status') and task.status.HasField('timestamp') - else 0, - task.id, - ), - reverse=True, - ) + tasks.sort(key=_list_sort_key, reverse=True) - # Paginate tasks + # Paginate tasks. The page token carries the position of the last task + # returned, so it stays valid if that task is updated or deleted. A + # task updated mid-listing can move above the cursor and be skipped. total_size = len(tasks) start_idx = 0 if params.page_token: - start_task_id = decode_page_token(params.page_token) - valid_token = False - for i, task in enumerate(tasks): - if task.id == start_task_id: - start_idx = i - valid_token = True - break - if not valid_token: - raise InvalidParamsError( - f'Invalid page token: {params.page_token}' + cursor = decode_list_tasks_cursor(params.page_token) + if cursor is None: + start_idx = self._legacy_start_index(tasks, params.page_token) + else: + after = cursor.sort_key() + start_idx = next( + ( + i + for i, task in enumerate(tasks) + if _list_sort_key(task) < after + ), + total_size, ) page_size = params.page_size or DEFAULT_LIST_TASKS_PAGE_SIZE end_idx = start_idx + page_size next_page_token = ( - encode_page_token(tasks[end_idx].id) + encode_list_tasks_cursor(_cursor_for(tasks[end_idx - 1])) if end_idx < total_size else None ) @@ -152,6 +171,15 @@ async def list( page_size=page_size, ) + @staticmethod + def _legacy_start_index(tasks: Sequence[Task], page_token: str) -> int: + """Resolves a legacy page token, which names the first task of the page.""" + start_task_id = decode_page_token(page_token) + for i, task in enumerate(tasks): + if task.id == start_task_id: + return i + raise InvalidParamsError(f'Invalid page token: {page_token}') + async def delete(self, task_id: str, context: ServerCallContext) -> None: """Deletes a task from the in-memory store by ID, for the given owner.""" owner = self.owner_resolver(context) diff --git a/src/a2a/utils/task.py b/src/a2a/utils/task.py index 4acf54e46..a179e3da4 100644 --- a/src/a2a/utils/task.py +++ b/src/a2a/utils/task.py @@ -1,8 +1,10 @@ """Utility functions for creating A2A Task objects.""" import binascii +import json -from base64 import b64decode, b64encode +from base64 import b64decode, b64encode, urlsafe_b64decode, urlsafe_b64encode +from dataclasses import dataclass from typing import Literal, Protocol, runtime_checkable from a2a.types.a2a_pb2 import Task @@ -119,3 +121,81 @@ def decode_page_token(page_token: str) -> str: 'Token is not a valid base64-encoded cursor.' ) from e return decoded + + +_CURSOR_VERSION = 1 +# Issued tokens are around 100 characters; anything much longer is not ours. +_MAX_PAGE_TOKEN_LENGTH = 512 + + +@dataclass(frozen=True) +class ListTasksCursor: + """A position in the `ListTasks` sort order. + + Tasks are listed by `(has timestamp, timestamp, id)` in descending order, + so tasks without a timestamp come last. A cursor names the last task of a + page; the next page starts strictly after it. Because the position is + carried in the token rather than looked up again, a page token stays valid + when that task is later updated or deleted. + """ + + timestamp_ns: int | None + task_id: str + + def sort_key(self) -> tuple[bool, int, str]: + """The cursor's position as a `(has timestamp, timestamp, id)` key.""" + return ( + self.timestamp_ns is not None, + self.timestamp_ns or 0, + self.task_id, + ) + + +def encode_list_tasks_cursor(cursor: ListTasksCursor) -> str: + """Encodes a `ListTasksCursor` as an opaque, URL-safe page token.""" + payload = json.dumps( + { + 'v': _CURSOR_VERSION, + 'ts': cursor.timestamp_ns, + 'id': cursor.task_id, + }, + separators=(',', ':'), + ) + return ( + urlsafe_b64encode(payload.encode(_ENCODING)) + .decode(_ENCODING) + .rstrip('=') + ) + + +def decode_list_tasks_cursor(page_token: str) -> ListTasksCursor | None: + """Decodes a page token produced by `encode_list_tasks_cursor`. + + Args: + page_token: The page token from a previous `ListTasks` response. + + Returns: + The decoded cursor, or None if the token is not a valid cursor token. + Callers treat None as a legacy task-ID token (see + `decode_page_token`), which also rejects tampered or unknown tokens. + + Raises: + InvalidParamsError: If the token is far longer than any issued token. + """ + if len(page_token) > _MAX_PAGE_TOKEN_LENGTH: + raise InvalidParamsError(f'Invalid page token: {page_token[:64]}...') + padded = page_token + '=' * (-len(page_token) % 4) + try: + data = json.loads(urlsafe_b64decode(padded.encode(_ENCODING))) + except (binascii.Error, ValueError): + return None + if not isinstance(data, dict) or data.get('v') != _CURSOR_VERSION: + return None + timestamp_ns = data.get('ts') + task_id = data.get('id') + timestamp_is_valid = timestamp_ns is None or ( + isinstance(timestamp_ns, int) and not isinstance(timestamp_ns, bool) + ) + if not timestamp_is_valid or not isinstance(task_id, str) or not task_id: + return None + return ListTasksCursor(timestamp_ns=timestamp_ns, task_id=task_id) diff --git a/tests/server/tasks/test_database_task_store.py b/tests/server/tasks/test_database_task_store.py index 40fbf31c8..cc9ce2ebb 100644 --- a/tests/server/tasks/test_database_task_store.py +++ b/tests/server/tasks/test_database_task_store.py @@ -1,7 +1,8 @@ import os +from base64 import urlsafe_b64encode from collections.abc import AsyncGenerator -from datetime import datetime, timezone +from datetime import datetime, timedelta, timezone from typing import Any from unittest.mock import MagicMock @@ -35,6 +36,7 @@ ) from a2a.utils.constants import DEFAULT_LIST_TASKS_PAGE_SIZE from a2a.utils.errors import InvalidParamsError +from a2a.utils.task import decode_list_tasks_cursor from google.protobuf.json_format import MessageToDict from sqlalchemy.ext.asyncio import create_async_engine from sqlalchemy.inspection import inspect @@ -58,6 +60,21 @@ def user_name(self) -> str: TEST_CONTEXT = ServerCallContext(user=SampleUser('test_user')) +def _decoded_cursor(page_token: str) -> tuple[datetime | None, str] | None: + """The (timestamp, task ID) a page token resumes after; None on the last page.""" + if not page_token: + return None + cursor = decode_list_tasks_cursor(page_token) + assert cursor is not None, 'expected a cursor token, not a legacy one' + timestamp = ( + datetime(1970, 1, 1, tzinfo=timezone.utc) + + timedelta(microseconds=cursor.timestamp_ns // 1_000) + if cursor.timestamp_ns is not None + else None + ) + return (timestamp, cursor.task_id) + + # DSNs for different databases SQLITE_TEST_DSN = ( 'sqlite+aiosqlite:///file:testdb?mode=memory&cache=shared&uri=true' @@ -208,7 +225,7 @@ async def test_get_task(db_store_parameterized: DatabaseTaskStore) -> None: @pytest.mark.asyncio @pytest.mark.parametrize( - 'params, expected_ids, total_count, next_page_token', + 'params, expected_ids, total_count, next_cursor', [ # No parameters, should return all tasks ( @@ -229,7 +246,7 @@ async def test_get_task(db_store_parameterized: DatabaseTaskStore) -> None: ListTasksRequest(page_size=2), ['task-2', 'task-1'], 5, - 'dGFzay0w', # base64 for 'task-0' + (datetime(2025, 1, 1, tzinfo=timezone.utc), 'task-1'), ), # Pagination (same timestamp) ( @@ -239,7 +256,7 @@ async def test_get_task(db_store_parameterized: DatabaseTaskStore) -> None: ), ['task-1', 'task-0'], 5, - 'dGFzay00', # base64 for 'task-4' + (datetime(2025, 1, 1, tzinfo=timezone.utc), 'task-0'), ), # Pagination (final page) ( @@ -282,7 +299,7 @@ async def test_get_task(db_store_parameterized: DatabaseTaskStore) -> None: ), ['task-2'], 3, - 'dGFzay0w', # base64 for 'task-0' + (datetime(2025, 1, 2, tzinfo=timezone.utc), 'task-2'), ), ], ) @@ -291,7 +308,7 @@ async def test_list_tasks( params: ListTasksRequest, expected_ids: list[str], total_count: int, - next_page_token: str, + next_cursor: tuple[datetime, str] | None, ) -> None: """Test listing tasks with various filters and pagination.""" tasks_to_create = [ @@ -338,7 +355,7 @@ async def test_list_tasks( retrieved_ids = [task.id for task in page.tasks] assert retrieved_ids == expected_ids assert page.total_size == total_count - assert page.next_page_token == (next_page_token or '') + assert _decoded_cursor(page.next_page_token) == next_cursor assert page.page_size == (params.page_size or DEFAULT_LIST_TASKS_PAGE_SIZE) # Cleanup @@ -978,4 +995,105 @@ async def test_core_to_0_3_model_conversion( await store.delete('v03-persistence-task', TEST_CONTEXT) +def _task(task_id: str, seconds: int | None) -> Task: + task = Task( + id=task_id, + context_id='ctx', + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ) + if seconds is not None: + task.status.timestamp.FromSeconds(seconds) + return task + + +async def _list_all(store: DatabaseTaskStore, page_size: int) -> list[str]: + seen: list[str] = [] + params = ListTasksRequest(page_size=page_size) + while True: + page = await store.list(params, TEST_CONTEXT) + seen.extend(task.id for task in page.tasks) + if not page.next_page_token: + return seen + params.page_token = page.next_page_token + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + 'change, task_id, expected_second_page', + [ + # The last task returned moves to the top: the listing continues. + ('update', 't4', ['t3', 't2']), + # The next task moves above the cursor: skipped this pass, no repeats. + ('update', 't3', ['t2', 't1']), + # Deleting either task does not invalidate the token. + ('delete', 't4', ['t3', 't2']), + ('delete', 't3', ['t2', 't1']), + ], +) +async def test_list_tasks_page_token_survives_task_changes( + db_store_parameterized: DatabaseTaskStore, + change: str, + task_id: str, + expected_second_page: list[str], +) -> None: + """Regression test for #1280: the token is a position, not a task lookup.""" + store = db_store_parameterized + for i in range(1, 6): + await store.save(_task(f't{i}', 1_700_000_000 + i), TEST_CONTEXT) + first = await store.list(ListTasksRequest(page_size=2), TEST_CONTEXT) + assert [task.id for task in first.tasks] == ['t5', 't4'] + + if change == 'update': + await store.save(_task(task_id, 1_800_000_000), TEST_CONTEXT) + else: + await store.delete(task_id, TEST_CONTEXT) + second = await store.list( + ListTasksRequest(page_size=2, page_token=first.next_page_token), + TEST_CONTEXT, + ) + + assert [task.id for task in second.tasks] == expected_second_page + + +@pytest.mark.asyncio +@pytest.mark.parametrize('page_size', [1, 2, 20]) +async def test_list_tasks_pages_through_ties_and_missing_timestamps( + db_store_parameterized: DatabaseTaskStore, page_size: int +) -> None: + """Equal timestamps (as on second-precision backends) and missing ones.""" + store = db_store_parameterized + for task in ( + _task('a', 1_700_000_000), + _task('b', 1_700_000_000), + _task('c', 1_700_000_000), + _task('newest', 1_700_000_100), + _task('undated-a', None), + _task('undated-b', None), + ): + await store.save(task, TEST_CONTEXT) + + assert await _list_all(store, page_size) == [ + 'newest', + 'c', + 'b', + 'a', + 'undated-b', + 'undated-a', + ] + + +@pytest.mark.asyncio +async def test_list_tasks_rejects_malformed_cursor_token( + db_store_parameterized: DatabaseTaskStore, +) -> None: + token = ( + urlsafe_b64encode(b'{"v":1,"ts":"soon","id":"t1"}').decode().rstrip('=') + ) + + with pytest.raises(InvalidParamsError): + await db_store_parameterized.list( + ListTasksRequest(page_size=2, page_token=token), TEST_CONTEXT + ) + + # Ensure aiosqlite, asyncpg, and aiomysql are installed in the test environment (added to pyproject.toml). diff --git a/tests/server/tasks/test_inmemory_task_store.py b/tests/server/tasks/test_inmemory_task_store.py index e20b72b2f..92e566df4 100644 --- a/tests/server/tasks/test_inmemory_task_store.py +++ b/tests/server/tasks/test_inmemory_task_store.py @@ -2,7 +2,8 @@ import concurrent.futures import threading -from datetime import datetime, timezone +from base64 import urlsafe_b64encode +from datetime import datetime, timedelta, timezone import pytest @@ -12,6 +13,7 @@ from a2a.types.a2a_pb2 import ListTasksRequest, Task, TaskState, TaskStatus from a2a.utils.constants import DEFAULT_LIST_TASKS_PAGE_SIZE from a2a.utils.errors import InvalidParamsError +from a2a.utils.task import decode_list_tasks_cursor class SampleUser(User): @@ -32,6 +34,21 @@ def user_name(self) -> str: TEST_CONTEXT = ServerCallContext(user=SampleUser('test_user')) +def _decoded_cursor(page_token: str) -> tuple[datetime | None, str] | None: + """The (timestamp, task ID) a page token resumes after; None on the last page.""" + if not page_token: + return None + cursor = decode_list_tasks_cursor(page_token) + assert cursor is not None, 'expected a cursor token, not a legacy one' + timestamp = ( + datetime(1970, 1, 1, tzinfo=timezone.utc) + + timedelta(microseconds=cursor.timestamp_ns // 1_000) + if cursor.timestamp_ns is not None + else None + ) + return (timestamp, cursor.task_id) + + def create_minimal_task( task_id: str = 'task-abc', context_id: str = 'session-xyz' ) -> Task: @@ -63,7 +80,7 @@ async def test_in_memory_task_store_get_nonexistent() -> None: @pytest.mark.asyncio @pytest.mark.parametrize( - 'params, expected_ids, total_count, next_page_token', + 'params, expected_ids, total_count, next_cursor', [ # No parameters, should return all tasks ( @@ -84,7 +101,7 @@ async def test_in_memory_task_store_get_nonexistent() -> None: ListTasksRequest(page_size=2), ['task-2', 'task-1'], 5, - 'dGFzay0w', # base64 for 'task-0' + (datetime(2025, 1, 1, tzinfo=timezone.utc), 'task-1'), ), # Pagination (same timestamp) ( @@ -94,7 +111,7 @@ async def test_in_memory_task_store_get_nonexistent() -> None: ), ['task-1', 'task-0'], 5, - 'dGFzay00', # base64 for 'task-4' + (datetime(2025, 1, 1, tzinfo=timezone.utc), 'task-0'), ), # Pagination (final page) ( @@ -137,7 +154,7 @@ async def test_in_memory_task_store_get_nonexistent() -> None: ), ['task-2'], 3, - 'dGFzay0w', # base64 for 'task-0' + (datetime(2025, 1, 2, tzinfo=timezone.utc), 'task-2'), ), ], ) @@ -145,7 +162,7 @@ async def test_list_tasks( params: ListTasksRequest, expected_ids: list[str], total_count: int, - next_page_token: str, + next_cursor: tuple[datetime, str] | None, ) -> None: """Test listing tasks with various filters and pagination.""" store = InMemoryTaskStore() @@ -193,7 +210,7 @@ async def test_list_tasks( retrieved_ids = [task.id for task in page.tasks] assert retrieved_ids == expected_ids assert page.total_size == total_count - assert page.next_page_token == (next_page_token or '') + assert _decoded_cursor(page.next_page_token) == next_cursor assert page.page_size == (params.page_size or DEFAULT_LIST_TASKS_PAGE_SIZE) # Cleanup @@ -342,6 +359,90 @@ async def test_list_tasks_timestamp_ordering_and_pagination( params.page_token = page.next_page_token +async def _store_with_five_tasks() -> InMemoryTaskStore: + """t1..t5 with increasing timestamps, so they list as t5, t4, t3, t2, t1.""" + store = InMemoryTaskStore() + for i in range(1, 6): + task = create_minimal_task(f't{i}') + task.status.timestamp.FromSeconds(1_700_000_000 + i) + await store.save(task, TEST_CONTEXT) + return store + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + 'change, task_id, expected_second_page', + [ + # The last task returned moves to the top: the listing continues. + ('update', 't4', ['t3', 't2']), + # The next task moves above the cursor: skipped this pass, no repeats. + ('update', 't3', ['t2', 't1']), + # Deleting either task does not invalidate the token. + ('delete', 't4', ['t3', 't2']), + ('delete', 't3', ['t2', 't1']), + ], +) +async def test_list_tasks_page_token_survives_task_changes( + change: str, task_id: str, expected_second_page: list[str] +) -> None: + """Regression test for #1280: the token is a position, not a task lookup.""" + store = await _store_with_five_tasks() + first = await store.list(ListTasksRequest(page_size=2), TEST_CONTEXT) + assert [task.id for task in first.tasks] == ['t5', 't4'] + + if change == 'update': + task = create_minimal_task(task_id) + task.status.timestamp.FromSeconds(1_800_000_000) + await store.save(task, TEST_CONTEXT) + else: + await store.delete(task_id, TEST_CONTEXT) + second = await store.list( + ListTasksRequest(page_size=2, page_token=first.next_page_token), + TEST_CONTEXT, + ) + + assert [task.id for task in second.tasks] == expected_second_page + + +@pytest.mark.asyncio +async def test_list_tasks_pages_through_tasks_without_timestamps() -> None: + """Tasks without a timestamp sort last and are each returned once.""" + store = InMemoryTaskStore() + timestamped = create_minimal_task('dated') + timestamped.status.timestamp.FromSeconds(1_700_000_000) + for task in ( + timestamped, + create_minimal_task('undated-a'), + create_minimal_task('undated-b'), + Task(id='no-status'), + ): + await store.save(task, TEST_CONTEXT) + + seen: list[str] = [] + params = ListTasksRequest(page_size=1) + while True: + page = await store.list(params, TEST_CONTEXT) + seen.extend(task.id for task in page.tasks) + if not page.next_page_token: + break + params.page_token = page.next_page_token + + assert seen == ['dated', 'undated-b', 'undated-a', 'no-status'] + + +@pytest.mark.asyncio +async def test_list_tasks_rejects_malformed_cursor_token() -> None: + store = await _store_with_five_tasks() + token = ( + urlsafe_b64encode(b'{"v":1,"ts":"soon","id":"t1"}').decode().rstrip('=') + ) + + with pytest.raises(InvalidParamsError): + await store.list( + ListTasksRequest(page_size=2, page_token=token), TEST_CONTEXT + ) + + @pytest.mark.asyncio async def test_in_memory_task_store_delete() -> None: """Test deleting a task from the store.""" diff --git a/tests/utils/test_task.py b/tests/utils/test_task.py index 8124955d1..57eba136b 100644 --- a/tests/utils/test_task.py +++ b/tests/utils/test_task.py @@ -1,5 +1,7 @@ import unittest +from base64 import urlsafe_b64encode + import pytest from a2a.helpers.proto_helpers import new_task @@ -14,8 +16,11 @@ ) from a2a.utils.errors import InvalidParamsError from a2a.utils.task import ( + ListTasksCursor, apply_history_length, + decode_list_tasks_cursor, decode_page_token, + encode_list_tasks_cursor, encode_page_token, ) @@ -39,6 +44,58 @@ def test_decode_page_token_fails(self): ) +@pytest.mark.parametrize( + 'cursor', + [ + ListTasksCursor(timestamp_ns=1_735_689_600_123_000_001, task_id='t-1'), + ListTasksCursor(timestamp_ns=-1, task_id='before-epoch'), + ListTasksCursor(timestamp_ns=None, task_id='no-timestamp'), + ListTasksCursor(timestamp_ns=0, task_id='ünïcode/+='), + ], +) +def test_list_tasks_cursor_round_trips(cursor: ListTasksCursor) -> None: + token = encode_list_tasks_cursor(cursor) + + assert decode_list_tasks_cursor(token) == cursor + # Safe to put in a query string without escaping. + assert not set(token) & set('+/=') + + +def test_legacy_task_id_token_is_not_a_cursor() -> None: + assert decode_list_tasks_cursor(encode_page_token('task-1')) is None + assert decode_list_tasks_cursor(encode_page_token('{"v":1}')) is None + assert decode_list_tasks_cursor('invalid') is None + + +@pytest.mark.parametrize( + 'payload', + [ + b'{"v":1,"ts":1}', + b'{"v":1,"ts":1,"id":""}', + b'{"v":1,"ts":"1","id":"t"}', + b'{"v":1,"ts":true,"id":"t"}', + b'{"v":1,"ts":1.5,"id":"t"}', + ], +) +def test_incomplete_cursor_token_is_not_a_cursor(payload: bytes) -> None: + """Falls back to the legacy path, which rejects it as an unknown task.""" + token = urlsafe_b64encode(payload).decode().rstrip('=') + + assert decode_list_tasks_cursor(token) is None + + +def test_overlong_page_token_is_rejected() -> None: + with pytest.raises(InvalidParamsError): + decode_list_tasks_cursor('A' * 513) + + +def test_cursor_sort_key_orders_missing_timestamps_last() -> None: + dated = ListTasksCursor(timestamp_ns=0, task_id='a') + undated = ListTasksCursor(timestamp_ns=None, task_id='z') + + assert undated.sort_key() < dated.sort_key() + + class TestApplyHistoryLength(unittest.TestCase): def setUp(self): self.history = [ From aa14a7a29ba3c73584398947d60f4d3c682cd8e4 Mon Sep 17 00:00:00 2001 From: Mykyta Netipa Date: Wed, 30 Sep 2026 10:00:54 +0000 Subject: [PATCH 2/3] chore: allow DTZ and urlsafe in spell check --- .github/actions/spelling/allow.txt | 2 ++ 1 file changed, 2 insertions(+) diff --git a/.github/actions/spelling/allow.txt b/.github/actions/spelling/allow.txt index 6cd148b96..9a8cc1dc8 100644 --- a/.github/actions/spelling/allow.txt +++ b/.github/actions/spelling/allow.txt @@ -48,6 +48,7 @@ denormals drivername dsn DSNs +DTZ dunders ES256 euo @@ -167,6 +168,7 @@ TResponse typ typeerror UIDs +urlsafe versioned Versioned vulnz From 7e1d7c321e8d45ad324664195c8e88238e6aae57 Mon Sep 17 00:00:00 2001 From: Mykyta Netipa Date: Wed, 30 Sep 2026 10:11:12 +0000 Subject: [PATCH 3/3] refactor(server): simplify ListTasks cursor tokens Drop the version field and the length cap from cursor tokens: a cursor is recognized by its exact {ts, id} shape, and the cap could reject tokens issued for long task IDs. Convert last_updated with protobuf Timestamp instead of hand-rolled epoch arithmetic. --- .github/actions/spelling/allow.txt | 1 - src/a2a/server/tasks/database_task_store.py | 20 ++++++++--------- src/a2a/utils/task.py | 22 ++++--------------- .../server/tasks/test_database_task_store.py | 4 +--- .../server/tasks/test_inmemory_task_store.py | 4 +--- tests/utils/test_task.py | 20 ++++++++--------- 6 files changed, 24 insertions(+), 47 deletions(-) diff --git a/.github/actions/spelling/allow.txt b/.github/actions/spelling/allow.txt index 9a8cc1dc8..0bd580314 100644 --- a/.github/actions/spelling/allow.txt +++ b/.github/actions/spelling/allow.txt @@ -48,7 +48,6 @@ denormals drivername dsn DSNs -DTZ dunders ES256 euo diff --git a/src/a2a/server/tasks/database_task_store.py b/src/a2a/server/tasks/database_task_store.py index f3f2e9dd0..a13d2ffd9 100644 --- a/src/a2a/server/tasks/database_task_store.py +++ b/src/a2a/server/tasks/database_task_store.py @@ -1,7 +1,7 @@ import logging from collections.abc import Callable -from datetime import datetime, timedelta, timezone +from datetime import datetime from typing import Any, cast @@ -23,6 +23,7 @@ "or 'pip install a2a-sdk[sql]'" ) from e from google.protobuf.json_format import MessageToDict, ParseDict +from google.protobuf.timestamp_pb2 import Timestamp from a2a.compat.v0_3.model_conversions import ( compat_task_model_to_core, @@ -45,22 +46,19 @@ logger = logging.getLogger(__name__) -_EPOCH = datetime(1970, 1, 1) # noqa: DTZ001 -- last_updated is naive UTC - def _datetime_to_ns(value: datetime) -> int: """Nanoseconds since the epoch for a stored (naive UTC) `last_updated`.""" - if value.tzinfo is not None: - value = value.astimezone(timezone.utc).replace(tzinfo=None) - delta = value - _EPOCH - return ( - delta.days * 86_400 + delta.seconds - ) * 1_000_000_000 + delta.microseconds * 1_000 + timestamp = Timestamp() + timestamp.FromDatetime(value) + return timestamp.ToNanoseconds() def _ns_to_datetime(timestamp_ns: int) -> datetime: - """Inverse of `_datetime_to_ns`; exact for values it produced.""" - return _EPOCH + timedelta(microseconds=timestamp_ns // 1_000) + """Inverse of `_datetime_to_ns`, as a naive UTC datetime.""" + timestamp = Timestamp() + timestamp.FromNanoseconds(timestamp_ns) + return timestamp.ToDatetime() class DatabaseTaskStore(TaskStore): diff --git a/src/a2a/utils/task.py b/src/a2a/utils/task.py index a179e3da4..0e340d7c5 100644 --- a/src/a2a/utils/task.py +++ b/src/a2a/utils/task.py @@ -123,11 +123,6 @@ def decode_page_token(page_token: str) -> str: return decoded -_CURSOR_VERSION = 1 -# Issued tokens are around 100 characters; anything much longer is not ours. -_MAX_PAGE_TOKEN_LENGTH = 512 - - @dataclass(frozen=True) class ListTasksCursor: """A position in the `ListTasks` sort order. @@ -154,11 +149,7 @@ def sort_key(self) -> tuple[bool, int, str]: def encode_list_tasks_cursor(cursor: ListTasksCursor) -> str: """Encodes a `ListTasksCursor` as an opaque, URL-safe page token.""" payload = json.dumps( - { - 'v': _CURSOR_VERSION, - 'ts': cursor.timestamp_ns, - 'id': cursor.task_id, - }, + {'ts': cursor.timestamp_ns, 'id': cursor.task_id}, separators=(',', ':'), ) return ( @@ -178,21 +169,16 @@ def decode_list_tasks_cursor(page_token: str) -> ListTasksCursor | None: The decoded cursor, or None if the token is not a valid cursor token. Callers treat None as a legacy task-ID token (see `decode_page_token`), which also rejects tampered or unknown tokens. - - Raises: - InvalidParamsError: If the token is far longer than any issued token. """ - if len(page_token) > _MAX_PAGE_TOKEN_LENGTH: - raise InvalidParamsError(f'Invalid page token: {page_token[:64]}...') padded = page_token + '=' * (-len(page_token) % 4) try: data = json.loads(urlsafe_b64decode(padded.encode(_ENCODING))) except (binascii.Error, ValueError): return None - if not isinstance(data, dict) or data.get('v') != _CURSOR_VERSION: + if not isinstance(data, dict) or data.keys() != {'ts', 'id'}: return None - timestamp_ns = data.get('ts') - task_id = data.get('id') + timestamp_ns = data['ts'] + task_id = data['id'] timestamp_is_valid = timestamp_ns is None or ( isinstance(timestamp_ns, int) and not isinstance(timestamp_ns, bool) ) diff --git a/tests/server/tasks/test_database_task_store.py b/tests/server/tasks/test_database_task_store.py index cc9ce2ebb..5376eaa17 100644 --- a/tests/server/tasks/test_database_task_store.py +++ b/tests/server/tasks/test_database_task_store.py @@ -1086,9 +1086,7 @@ async def test_list_tasks_pages_through_ties_and_missing_timestamps( async def test_list_tasks_rejects_malformed_cursor_token( db_store_parameterized: DatabaseTaskStore, ) -> None: - token = ( - urlsafe_b64encode(b'{"v":1,"ts":"soon","id":"t1"}').decode().rstrip('=') - ) + token = urlsafe_b64encode(b'{"ts":"soon","id":"t1"}').decode().rstrip('=') with pytest.raises(InvalidParamsError): await db_store_parameterized.list( diff --git a/tests/server/tasks/test_inmemory_task_store.py b/tests/server/tasks/test_inmemory_task_store.py index 92e566df4..40f9ecf8d 100644 --- a/tests/server/tasks/test_inmemory_task_store.py +++ b/tests/server/tasks/test_inmemory_task_store.py @@ -433,9 +433,7 @@ async def test_list_tasks_pages_through_tasks_without_timestamps() -> None: @pytest.mark.asyncio async def test_list_tasks_rejects_malformed_cursor_token() -> None: store = await _store_with_five_tasks() - token = ( - urlsafe_b64encode(b'{"v":1,"ts":"soon","id":"t1"}').decode().rstrip('=') - ) + token = urlsafe_b64encode(b'{"ts":"soon","id":"t1"}').decode().rstrip('=') with pytest.raises(InvalidParamsError): await store.list( diff --git a/tests/utils/test_task.py b/tests/utils/test_task.py index 57eba136b..2ce2f9e27 100644 --- a/tests/utils/test_task.py +++ b/tests/utils/test_task.py @@ -51,6 +51,7 @@ def test_decode_page_token_fails(self): ListTasksCursor(timestamp_ns=-1, task_id='before-epoch'), ListTasksCursor(timestamp_ns=None, task_id='no-timestamp'), ListTasksCursor(timestamp_ns=0, task_id='ünïcode/+='), + ListTasksCursor(timestamp_ns=1, task_id='x' * 1_000), ], ) def test_list_tasks_cursor_round_trips(cursor: ListTasksCursor) -> None: @@ -63,18 +64,20 @@ def test_list_tasks_cursor_round_trips(cursor: ListTasksCursor) -> None: def test_legacy_task_id_token_is_not_a_cursor() -> None: assert decode_list_tasks_cursor(encode_page_token('task-1')) is None - assert decode_list_tasks_cursor(encode_page_token('{"v":1}')) is None + assert decode_list_tasks_cursor(encode_page_token('{"ts":1}')) is None assert decode_list_tasks_cursor('invalid') is None @pytest.mark.parametrize( 'payload', [ - b'{"v":1,"ts":1}', - b'{"v":1,"ts":1,"id":""}', - b'{"v":1,"ts":"1","id":"t"}', - b'{"v":1,"ts":true,"id":"t"}', - b'{"v":1,"ts":1.5,"id":"t"}', + b'{"ts":1}', + b'{"id":"t"}', + b'{"ts":1,"id":"t","extra":0}', + b'{"ts":1,"id":""}', + b'{"ts":"1","id":"t"}', + b'{"ts":true,"id":"t"}', + b'{"ts":1.5,"id":"t"}', ], ) def test_incomplete_cursor_token_is_not_a_cursor(payload: bytes) -> None: @@ -84,11 +87,6 @@ def test_incomplete_cursor_token_is_not_a_cursor(payload: bytes) -> None: assert decode_list_tasks_cursor(token) is None -def test_overlong_page_token_is_rejected() -> None: - with pytest.raises(InvalidParamsError): - decode_list_tasks_cursor('A' * 513) - - def test_cursor_sort_key_orders_missing_timestamps_last() -> None: dated = ListTasksCursor(timestamp_ns=0, task_id='a') undated = ListTasksCursor(timestamp_ns=None, task_id='z')