diff --git a/.github/actions/spelling/allow.txt b/.github/actions/spelling/allow.txt index 6cd148b96..0bd580314 100644 --- a/.github/actions/spelling/allow.txt +++ b/.github/actions/spelling/allow.txt @@ -167,6 +167,7 @@ TResponse typ typeerror UIDs +urlsafe versioned Versioned vulnz diff --git a/src/a2a/server/tasks/database_task_store.py b/src/a2a/server/tasks/database_task_store.py index 409518d8d..a13d2ffd9 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 from typing import Any, cast @@ -22,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, @@ -34,12 +36,31 @@ 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 _datetime_to_ns(value: datetime) -> int: + """Nanoseconds since the epoch for a stored (naive UTC) `last_updated`.""" + timestamp = Timestamp() + timestamp.FromDatetime(value) + return timestamp.ToNanoseconds() + + +def _ns_to_datetime(timestamp_ns: int) -> datetime: + """Inverse of `_datetime_to_ns`, as a naive UTC datetime.""" + timestamp = Timestamp() + timestamp.FromNanoseconds(timestamp_ns) + return timestamp.ToDatetime() + + class DatabaseTaskStore(TaskStore): """SQLAlchemy-based implementation of TaskStore. @@ -257,64 +278,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..0e340d7c5 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,67 @@ def decode_page_token(page_token: str) -> str: 'Token is not a valid base64-encoded cursor.' ) from e return decoded + + +@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( + {'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. + """ + 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.keys() != {'ts', 'id'}: + return None + 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) + ) + 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..5376eaa17 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,103 @@ 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'{"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..40f9ecf8d 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,88 @@ 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'{"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..2ce2f9e27 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,56 @@ 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/+='), + ListTasksCursor(timestamp_ns=1, task_id='x' * 1_000), + ], +) +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('{"ts":1}')) is None + assert decode_list_tasks_cursor('invalid') is None + + +@pytest.mark.parametrize( + 'payload', + [ + 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: + """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_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 = [