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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions .github/actions/spelling/allow.txt
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
a2a

Check warning on line 1 in .github/actions/spelling/allow.txt

View workflow job for this annotation

GitHub Actions / Check Spelling

Ignoring entry because it contains non-alpha characters (non-alpha-in-dictionary)
A2A

Check warning on line 2 in .github/actions/spelling/allow.txt

View workflow job for this annotation

GitHub Actions / Check Spelling

Ignoring entry because it contains non-alpha characters (non-alpha-in-dictionary)
A2AFastAPI

Check warning on line 3 in .github/actions/spelling/allow.txt

View workflow job for this annotation

GitHub Actions / Check Spelling

Ignoring entry because it contains non-alpha characters (non-alpha-in-dictionary)
AAgent
Expand Down Expand Up @@ -29,7 +29,7 @@
AUser
autouse
backticks
base64url

Check warning on line 32 in .github/actions/spelling/allow.txt

View workflow job for this annotation

GitHub Actions / Check Spelling

Ignoring entry because it contains non-alpha characters (non-alpha-in-dictionary)
buf
bufbuild
cla
Expand All @@ -49,7 +49,7 @@
dsn
DSNs
dunders
ES256

Check warning on line 52 in .github/actions/spelling/allow.txt

View workflow job for this annotation

GitHub Actions / Check Spelling

Ignoring entry because it contains non-alpha characters (non-alpha-in-dictionary)
euo
EUR
evt
Expand All @@ -66,8 +66,8 @@
gowebpki
GVsb
hazmat
HS256

Check warning on line 69 in .github/actions/spelling/allow.txt

View workflow job for this annotation

GitHub Actions / Check Spelling

Ignoring entry because it contains non-alpha characters (non-alpha-in-dictionary)
HS384

Check warning on line 70 in .github/actions/spelling/allow.txt

View workflow job for this annotation

GitHub Actions / Check Spelling

Ignoring entry because it contains non-alpha characters (non-alpha-in-dictionary)
ietf
importlib
initdb
Expand Down Expand Up @@ -110,11 +110,11 @@
Oneof
OpenAPI
openapiv
openapiv2

Check warning on line 113 in .github/actions/spelling/allow.txt

View workflow job for this annotation

GitHub Actions / Check Spelling

Ignoring entry because it contains non-alpha characters (non-alpha-in-dictionary)
opensource
otherurl
outerjoin
pb2

Check warning on line 117 in .github/actions/spelling/allow.txt

View workflow job for this annotation

GitHub Actions / Check Spelling

Ignoring entry because it contains non-alpha characters (non-alpha-in-dictionary)
podman
Podman
poolclass
Expand All @@ -138,7 +138,7 @@
respx
resub
rmi
RS256

Check warning on line 141 in .github/actions/spelling/allow.txt

View workflow job for this annotation

GitHub Actions / Check Spelling

Ignoring entry because it contains non-alpha characters (non-alpha-in-dictionary)
RUF
Rundgren
SECP
Expand Down Expand Up @@ -167,6 +167,7 @@
typ
typeerror
UIDs
urlsafe
versioned
Versioned
vulnz
Expand Down
142 changes: 104 additions & 38 deletions src/a2a/server/tasks/database_task_store.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import logging

from collections.abc import Callable
from datetime import datetime
from typing import Any, cast


Expand All @@ -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,
Expand All @@ -29,17 +31,36 @@
from a2a.server.context import ServerCallContext
from a2a.server.models import Base, TaskModel, create_task_model
from a2a.server.owner_resolver import OwnerResolver, resolve_user_scope
from a2a.server.tasks.task_store import TaskStore
from a2a.types import a2a_pb2
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:

Check notice on line 50 in src/a2a/server/tasks/database_task_store.py

View workflow job for this annotation

GitHub Actions / Lint Code Base

Copy/pasted code

see src/a2a/server/tasks/inmemory_task_store.py (9-25)
"""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.

Expand Down Expand Up @@ -257,64 +278,109 @@
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()
Expand Down
78 changes: 53 additions & 25 deletions src/a2a/server/tasks/inmemory_task_store.py
Original file line number Diff line number Diff line change
@@ -1,20 +1,46 @@
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
from a2a.server.tasks.task_store import TaskStore
from a2a.types import a2a_pb2
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]:

Check notice on line 25 in src/a2a/server/tasks/inmemory_task_store.py

View workflow job for this annotation

GitHub Actions / Lint Code Base

Copy/pasted code

see src/a2a/server/tasks/database_task_store.py (34-50)
"""`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.

Expand Down Expand Up @@ -108,38 +134,31 @@
]

# 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
)
Expand All @@ -152,6 +171,15 @@
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)
Expand Down
Loading
Loading