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

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. 

19 

20:meta private: 

21""" 

22 

23from __future__ import annotations 

24 

25import base64 

26import uuid as uuid_mod 

27from typing import Any 

28 

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 

35 

36from airflow.api_fastapi.common.parameters import SortParam 

37 

38 

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")) 

44 

45 

46def _dialect_nulls_last(is_desc: bool, dialect: str) -> bool: 

47 """ 

48 Where a plain ``ORDER BY col`` puts NULLs on this backend. 

49 

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 

54 

55 

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. 

61 

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 

76 

77 

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. 

83 

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 

98 

99 

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 

111 

112 

113_BACKWARD_PREFIX = "~" 

114 

115 

116def encode_cursor(row: Any, sort_param: SortParam) -> str: 

117 """ 

118 Encode cursor token from the boundary row of a result set. 

119 

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.") 

126 

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("=") 

130 

131 

132def make_backward_cursor(token: str) -> str: 

133 """Prefix a cursor token with the backward direction marker (``~``).""" 

134 return f"{_BACKWARD_PREFIX}{token}" 

135 

136 

137def parse_cursor(cursor: str) -> tuple[str, bool]: 

138 """ 

139 Parse a raw cursor string into ``(token, is_backward)``. 

140 

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 

147 

148 

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") 

155 

156 try: 

157 data: Any = msgspec.msgpack.decode(raw) 

158 except Exception: 

159 raise HTTPException(status.HTTP_400_BAD_REQUEST, "Invalid cursor token") 

160 

161 if not isinstance(data, list): 

162 raise HTTPException(status.HTTP_400_BAD_REQUEST, "Invalid cursor token structure") 

163 

164 return data 

165 

166 

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. 

172 

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. 

178 

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) 

184 

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") 

188 

189 parsed_values = [_coerce_value(col, val) for (_, col, _), val in zip(resolved, raw_values, strict=True)] 

190 

191 if is_backward: 

192 resolved = [(name, col, not is_desc) for name, col, is_desc in resolved] 

193 

194 return statement.where(_nested_keyset_predicate(resolved, parsed_values, dialect))