Source code for fsh_lib.pagination

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