Coverage for polar/kit/pagination.py: 91%

62 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-07 12:42 +0000

1import math 

2from collections.abc import Sequence 

3from typing import Annotated, Any, NamedTuple, Self, overload 

4 

5from fastapi import Depends, Query 

6from pydantic import BaseModel, GetCoreSchemaHandler 

7from pydantic._internal._repr import display_as_type 

8from pydantic_core import CoreSchema 

9from sqlalchemy import Select, func, over 

10from sqlalchemy.sql._typing import _ColumnsClauseArgument 

11 

12from polar.config import settings 

13from polar.kit.db.models import RecordModel 

14from polar.kit.db.models.base import Model 

15from polar.kit.db.postgres import AsyncReadSession 

16from polar.kit.schemas import ClassName, Schema 

17 

18 

19class PaginationParams(NamedTuple): 

20 page: int 

21 limit: int 

22 

23 

24@overload 

25async def paginate[RM: RecordModel]( 25 ↛ 34line 25 didn't jump to line 34 because

26 session: AsyncReadSession, 

27 statement: Select[tuple[RM]], 

28 *, 

29 pagination: PaginationParams, 

30 count_clause: _ColumnsClauseArgument[Any] | None = None, 

31) -> tuple[Sequence[RM], int]: ... 

32 

33 

34@overload 

35async def paginate[M: Model]( 35 ↛ 44line 35 didn't jump to line 44 because

36 session: AsyncReadSession, 

37 statement: Select[tuple[M]], 

38 *, 

39 pagination: PaginationParams, 

40 count_clause: _ColumnsClauseArgument[Any] | None = None, 

41) -> tuple[Sequence[M], int]: ... 

42 

43 

44@overload 

45async def paginate[T: Any]( 45 ↛ 54line 45 didn't jump to line 54 because

46 session: AsyncReadSession, 

47 statement: Select[T], 

48 *, 

49 pagination: PaginationParams, 

50 count_clause: _ColumnsClauseArgument[Any] | None = None, 

51) -> tuple[Sequence[T], int]: ... 

52 

53 

54async def paginate( 

55 session: AsyncReadSession, 

56 statement: Select[Any], 

57 *, 

58 pagination: PaginationParams, 

59 count_clause: _ColumnsClauseArgument[Any] | None = None, 

60) -> tuple[Sequence[Any], int]: 

61 page, limit = pagination 

62 offset = limit * (page - 1) 

63 statement = statement.offset(offset).limit(limit) 

64 

65 if count_clause is not None: 65 ↛ 66line 65 didn't jump to line 66 because the condition on line 65 was never true

66 statement = statement.add_columns(count_clause) 

67 else: 

68 statement = statement.add_columns(over(func.count())) 

69 

70 result = await session.execute(statement) 

71 

72 results: list[Any] = [] 

73 count = 0 

74 for row in result.unique().all(): 

75 (*queried_data, c) = row._tuple() 

76 count = int(c) 

77 if len(queried_data) == 1: 77 ↛ 80line 77 didn't jump to line 80 because the condition on line 77 was always true

78 results.append(queried_data[0]) 

79 else: 

80 results.append(queried_data) 

81 

82 return results, count 

83 

84 

85async def get_pagination_params( 

86 page: int = Query(1, description="Page number, defaults to 1.", gt=0), 

87 limit: int = Query( 

88 10, 

89 description=( 

90 f"Size of a page, defaults to 10. " 

91 f"Maximum is {settings.API_PAGINATION_MAX_LIMIT}." 

92 ), 

93 gt=0, 

94 ), 

95) -> PaginationParams: 

96 return PaginationParams(page, min(settings.API_PAGINATION_MAX_LIMIT, limit)) 

97 

98 

99PaginationParamsQuery = Annotated[PaginationParams, Depends(get_pagination_params)] 

100 

101 

102class Pagination(Schema): 

103 total_count: int 

104 max_page: int 

105 

106 

107class ListResource[T: Any](BaseModel): 

108 items: list[T] 

109 pagination: Pagination 

110 

111 @classmethod 

112 def from_paginated_results( 

113 cls, items: Sequence[T], total_count: int, pagination_params: PaginationParams 

114 ) -> Self: 

115 return cls( 

116 items=list(items), 

117 pagination=Pagination( 

118 total_count=total_count, 

119 max_page=math.ceil(total_count / pagination_params.limit), 

120 ), 

121 ) 

122 

123 @classmethod 

124 def model_parametrized_name(cls, params: tuple[type[Any], ...]) -> str: 

125 """ 

126 Override default model name implementation to detect `ClassName` metadata. 

127 

128 It's useful to shorten the name when a long union type is used. 

129 """ 

130 param_names = [] 

131 for param in params: 

132 if hasattr(param, "__metadata__"): 

133 for metadata in param.__metadata__: 

134 if isinstance(metadata, ClassName): 

135 param_names.append(metadata.name) 

136 else: 

137 param_names.append(display_as_type(param)) 

138 

139 params_component = ", ".join(param_names) 

140 return f"{cls.__name__}[{params_component}]" 

141 

142 @classmethod 

143 def __get_pydantic_core_schema__( 

144 cls, source: type[BaseModel], handler: GetCoreSchemaHandler, / 

145 ) -> CoreSchema: 

146 """ 

147 Override the schema to set the `ref` field to the overridden class name. 

148 """ 

149 result = handler(source) 

150 result["ref"] = cls.__name__ # type: ignore 

151 return result