Coverage for /home/airflow/.local/lib/python3.12/site-packages/airflow/api_fastapi/common/cursors.py: 37%
83 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 14:22 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 14:22 +0000
1# Licensed to the Apache Software Foundation (ASF) under one
2# or more contributor license agreements. See the NOTICE file
3# distributed with this work for additional information
4# regarding copyright ownership. The ASF licenses this file
5# to you under the Apache License, Version 2.0 (the
6# "License"); you may not use this file except in compliance
7# with the License. You may obtain a copy of the License at
8#
9# http://www.apache.org/licenses/LICENSE-2.0
10#
11# Unless required by applicable law or agreed to in writing,
12# software distributed under the License is distributed on an
13# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14# KIND, either express or implied. See the License for the
15# specific language governing permissions and limitations
16# under the License.
17"""
18Cursor-based (keyset) pagination helpers.
20:meta private:
21"""
23from __future__ import annotations
25import base64
26import uuid as uuid_mod
27from typing import Any
29import msgspec
30from fastapi import HTTPException, status
31from sqlalchemy import and_, false, or_, true
32from sqlalchemy.sql import Select
33from sqlalchemy.sql.elements import ColumnElement
34from sqlalchemy.sql.sqltypes import Uuid
36from airflow.api_fastapi.common.parameters import SortParam
39def _b64url_decode_padded(token: str) -> bytes:
40 padding = 4 - (len(token) % 4)
41 if padding != 4:
42 token = token + ("=" * padding)
43 return base64.urlsafe_b64decode(token.encode("ascii"))
46def _dialect_nulls_last(is_desc: bool, dialect: str) -> bool:
47 """
48 Where a plain ``ORDER BY col`` puts NULLs on this backend.
50 PostgreSQL sorts NULLs last for ASC, first for DESC; MySQL and SQLite treat
51 NULL as the lowest value, so NULLs come first for ASC and last for DESC.
52 """
53 return (not is_desc) if dialect == "postgresql" else is_desc
56def _bounds(
57 col: ColumnElement, value: Any, is_desc: bool, dialect: str
58) -> tuple[ColumnElement[bool], ColumnElement[bool]]:
59 """
60 ``(non_strict, strict)`` keyset bounds matching the backend's native NULL placement.
62 A NULL *value* means the cursor sits in the NULL block; otherwise a trailing
63 NULL block (when NULLs sort last) is also admitted after a non-NULL cursor.
64 The ``col IS NULL`` terms are vacuous for non-nullable columns.
65 """
66 nulls_last = _dialect_nulls_last(is_desc, dialect)
67 if value is None:
68 if nulls_last:
69 return col.is_(None), false()
70 return true(), col.is_not(None)
71 ge = col <= value if is_desc else col >= value
72 gt = col < value if is_desc else col > value
73 if nulls_last:
74 return or_(ge, col.is_(None)), or_(gt, col.is_(None))
75 return ge, gt
78def _nested_keyset_predicate(
79 resolved: list[tuple[str, ColumnElement, bool]], values: list[Any], dialect: str
80) -> ColumnElement[bool]:
81 """
82 Keyset predicate for rows strictly after the cursor in ``ORDER BY`` order.
84 Uses nested ``and_(non-strict, or_(strict, ...))`` so leading sort keys use
85 inclusive range bounds and inner branches use strict inequalities—friendly
86 for composite index range scans. NULL placement follows each backend's
87 native ordering for a plain ``ORDER BY`` (see :func:`_bounds`), so the
88 ``ORDER BY`` stays a bare column and can still use an index.
89 """
90 n = len(resolved)
91 _, col, is_desc = resolved[n - 1]
92 _, inner = _bounds(col, values[n - 1], is_desc, dialect)
93 for i in range(n - 2, -1, -1):
94 _, col_i, is_desc_i = resolved[i]
95 non_strict, strict = _bounds(col_i, values[i], is_desc_i, dialect)
96 inner = and_(non_strict, or_(strict, inner))
97 return inner
100def _coerce_value(column: ColumnElement, value: Any) -> Any:
101 """Normalize decoded values for SQL bind parameters (e.g. UUID columns)."""
102 if value is None or not isinstance(value, str):
103 return value
104 ctype = getattr(column, "type", None)
105 if isinstance(ctype, Uuid):
106 try:
107 return uuid_mod.UUID(value)
108 except ValueError:
109 raise HTTPException(status.HTTP_400_BAD_REQUEST, "Invalid cursor token")
110 return value
113_BACKWARD_PREFIX = "~"
116def encode_cursor(row: Any, sort_param: SortParam) -> str:
117 """
118 Encode cursor token from the boundary row of a result set.
120 The token is a URL-safe base64 encoding of a MessagePack list of sort-key
121 values (no padding ``=``).
122 """
123 resolved = sort_param.get_resolved_columns()
124 if not resolved:
125 raise ValueError("SortParam has no resolved columns.")
127 parts = [sort_param.row_value(row, attr_name) for attr_name, _col, _desc in resolved]
128 payload = msgspec.msgpack.encode(parts)
129 return base64.urlsafe_b64encode(payload).decode("ascii").rstrip("=")
132def make_backward_cursor(token: str) -> str:
133 """Prefix a cursor token with the backward direction marker (``~``)."""
134 return f"{_BACKWARD_PREFIX}{token}"
137def parse_cursor(cursor: str) -> tuple[str, bool]:
138 """
139 Parse a raw cursor string into ``(token, is_backward)``.
141 Strips the ``~`` prefix if present and returns whether the cursor
142 represents a backward (previous-page) direction.
143 """
144 if cursor.startswith(_BACKWARD_PREFIX): 144 ↛ 145line 144 didn't jump to line 145 because the condition on line 144 was never true
145 return cursor[len(_BACKWARD_PREFIX) :], True
146 return cursor, False
149def decode_cursor(token: str) -> list[Any]:
150 """Decode a cursor token to the list of sort-key values."""
151 try:
152 raw = _b64url_decode_padded(token)
153 except Exception:
154 raise HTTPException(status.HTTP_400_BAD_REQUEST, "Invalid cursor token")
156 try:
157 data: Any = msgspec.msgpack.decode(raw)
158 except Exception:
159 raise HTTPException(status.HTTP_400_BAD_REQUEST, "Invalid cursor token")
161 if not isinstance(data, list):
162 raise HTTPException(status.HTTP_400_BAD_REQUEST, "Invalid cursor token structure")
164 return data
167def apply_cursor_filter(
168 statement: Select, token: str, sort_param: SortParam, dialect: str, *, is_backward: bool = False
169) -> Select:
170 """
171 Apply a keyset pagination WHERE clause from a cursor token.
173 For forward cursors the predicate selects rows strictly *after* the cursor
174 in ORDER BY order. When *is_backward* is True the ``is_desc`` flags are
175 flipped so the predicate selects rows strictly *before* the cursor in the
176 original sort order. The caller is responsible for reversing the ORDER BY
177 and the final result list when using a backward cursor.
179 *dialect* is the backend dialect name (``session.get_bind().dialect.name``);
180 the predicate matches that backend's native NULL placement, so the keyset
181 ``ORDER BY`` stays a bare column and keeps using indexes.
182 """
183 raw_values = decode_cursor(token)
185 resolved = sort_param.get_resolved_columns()
186 if len(raw_values) != len(resolved):
187 raise HTTPException(status.HTTP_400_BAD_REQUEST, "Cursor token does not match current query shape")
189 parsed_values = [_coerce_value(col, val) for (_, col, _), val in zip(resolved, raw_values, strict=True)]
191 if is_backward:
192 resolved = [(name, col, not is_desc) for name, col, is_desc in resolved]
194 return statement.where(_nested_keyset_predicate(resolved, parsed_values, dialect))