from __future__ import annotations
import base64
import binascii
import decimal
import enum
import json
import uuid
from datetime import date, datetime, time
from typing import TYPE_CHECKING, Any
from sqlalchemy import and_, false, or_
from fsh_lib.errors import FieldError
from fsh_lib.ordering import OrderKey
if TYPE_CHECKING:
from collections.abc import Sequence
from sqlalchemy import ColumnElement, Select
from sqlalchemy.ext.asyncio import AsyncSession
__all__ = [
"InvalidCursorError",
"decode_cursor",
"encode_cursor",
"keyset_predicate",
"run_keyset_query",
]
[docs]
class InvalidCursorError(FieldError):
def __init__(self) -> None:
super().__init__(("body", "cursor"), "Invalid pagination cursor.")
def _to_wire(value: Any) -> Any: # noqa: ANN401
if value is None or isinstance(value, (bool, int, float, str)):
return value
if isinstance(value, enum.Enum):
return _to_wire(value.value)
if isinstance(value, (datetime, date, time)):
return value.isoformat()
return str(value)
def encode_cursor(values: Sequence[Any]) -> str:
payload = json.dumps(
[_to_wire(v) for v in values],
separators=(",", ":"),
)
return base64.urlsafe_b64encode(payload.encode()).decode().rstrip("=")
_PARSERS: dict[type, Any] = {
datetime: datetime.fromisoformat,
date: date.fromisoformat,
time: time.fromisoformat,
uuid.UUID: uuid.UUID,
decimal.Decimal: decimal.Decimal,
}
def _coerce(value: Any, expr: ColumnElement) -> Any: # noqa: ANN401
if value is None:
return None
try:
py_type = expr.type.python_type
except NotImplementedError:
return value
if isinstance(value, py_type):
return value
parser = _PARSERS.get(py_type, py_type)
try:
return parser(value)
except (TypeError, ValueError, decimal.InvalidOperation) as err:
raise InvalidCursorError from err
def decode_cursor(cursor: str, keys: Sequence[OrderKey]) -> list[Any]:
try:
padded = cursor + "=" * (-len(cursor) % 4)
raw = json.loads(base64.urlsafe_b64decode(padded.encode()))
except (ValueError, binascii.Error) as err:
raise InvalidCursorError from err
if not isinstance(raw, list) or len(raw) != len(keys):
raise InvalidCursorError
return [
_coerce(value, key.expr) for value, key in zip(raw, keys, strict=True)
]
def _is_nullable(col: ColumnElement) -> bool:
try:
return bool(col.nullable)
except AttributeError:
return True
def _level_after(key: OrderKey, value: Any) -> ColumnElement: # noqa: ANN401
col = key.expr
if key.direction == "desc":
if value is None:
return col.is_not(None)
return col < value
if value is None:
return false()
if not _is_nullable(col):
return col > value
return or_(col > value, col.is_(None))
def _level_equal(key: OrderKey, value: Any) -> ColumnElement: # noqa: ANN401
col = key.expr
return col.is_(None) if value is None else col == value
def keyset_predicate(
keys: Sequence[OrderKey],
values: Sequence[Any],
) -> ColumnElement:
branches = []
for index, (key, value) in enumerate(zip(keys, values, strict=True)):
level = [
_level_equal(k, v)
for k, v in zip(keys[:index], values[:index], strict=True)
]
level.append(_level_after(key, value))
branches.append(and_(*level))
return or_(*branches)
async def run_keyset_query(
*,
db: AsyncSession,
stmt: Select,
model: type,
cursor: str | None,
cursor_field: str,
page_size: int,
max_page_size: int,
order_keys: Sequence[OrderKey] = (),
) -> tuple[list[Any], str | None]:
pk_col = getattr(model, cursor_field)
keys = [
*order_keys,
OrderKey(field=cursor_field, direction="asc", expr=pk_col),
]
effective_page_size = min(page_size, max_page_size)
stmt = stmt.order_by(pk_col.asc())
if cursor is not None:
stmt = stmt.where(keyset_predicate(keys, decode_cursor(cursor, keys)))
stmt = stmt.add_columns(*[k.expr for k in keys])
stmt = stmt.limit(effective_page_size + 1)
result = await db.execute(stmt)
rows = result.all()
has_more = len(rows) > effective_page_size
rows = rows[:effective_page_size]
items = [row[0] for row in rows]
next_cursor = (
encode_cursor(list(rows[-1])[1:]) if has_more and rows else None
)
return items, next_cursor