Coverage for polar/kit/repository/base.py: 81%
108 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 12:42 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 12:42 +0000
1from collections.abc import AsyncGenerator, Sequence
2from datetime import datetime
3from enum import StrEnum
4from typing import Any, Protocol, Self, TypeAlias
6from sqlalchemy import Select, UnaryExpression, asc, desc, func, over, select
7from sqlalchemy.orm import Mapped
8from sqlalchemy.orm.attributes import flag_modified
9from sqlalchemy.sql.base import ExecutableOption
10from sqlalchemy.sql.expression import ColumnExpressionArgument
12from polar.config import settings
13from polar.kit.db.postgres import AsyncReadSession, AsyncSession
14from polar.kit.sorting import Sorting
15from polar.kit.utils import utc_now
18class ModelDeletedAtProtocol(Protocol):
19 deleted_at: Mapped[datetime | None]
22class ModelIDProtocol[ID_TYPE](Protocol):
23 id: Mapped[ID_TYPE]
26class ModelDeletedAtIDProtocol[ID_TYPE](Protocol):
27 id: Mapped[ID_TYPE]
28 deleted_at: Mapped[datetime | None]
31Options: TypeAlias = Sequence[ExecutableOption]
34class RepositoryProtocol[M](Protocol):
35 model: type[M]
37 async def get_one(self, statement: Select[tuple[M]]) -> M: ... 37 ↛ anywhereline 37 didn't jump anywhere: it always raised an exception.
39 async def get_one_or_none(self, statement: Select[tuple[M]]) -> M | None: ... 39 ↛ anywhereline 39 didn't jump anywhere: it always raised an exception.
41 async def get_all(self, statement: Select[tuple[M]]) -> Sequence[M]: ... 41 ↛ anywhereline 41 didn't jump anywhere: it always raised an exception.
43 async def paginate( 43 ↛ anywhereline 43 didn't jump anywhere: it always raised an exception.
44 self, statement: Select[tuple[M]], *, limit: int, page: int
45 ) -> tuple[list[M], int]: ...
47 def get_base_statement(self) -> Select[tuple[M]]: ... 47 ↛ anywhereline 47 didn't jump anywhere: it always raised an exception.
49 async def create(self, object: M, *, flush: bool = False) -> M: ... 49 ↛ anywhereline 49 didn't jump anywhere: it always raised an exception.
51 async def update( 51 ↛ anywhereline 51 didn't jump anywhere: it always raised an exception.
52 self,
53 object: M,
54 *,
55 update_dict: dict[str, Any] | None = None,
56 flush: bool = False,
57 ) -> M: ...
60class RepositoryBase[M]:
61 model: type[M]
63 def __init__(self, session: AsyncSession | AsyncReadSession) -> None:
64 self.session = session
66 async def get_one(self, statement: Select[tuple[M]]) -> M:
67 result = await self.session.execute(statement)
68 return result.unique().scalar_one()
70 async def get_one_or_none(self, statement: Select[tuple[M]]) -> M | None:
71 result = await self.session.execute(statement)
72 return result.unique().scalar_one_or_none()
74 async def get_all(self, statement: Select[tuple[M]]) -> Sequence[M]:
75 result = await self.session.execute(statement)
76 return result.scalars().unique().all()
78 async def stream(self, statement: Select[tuple[M]]) -> AsyncGenerator[M, None]:
79 """
80 Stream results from the database using the given statement.
82 This is useful for processing large datasets without loading everything
83 into memory at once.
85 The caveat is that your statement shouldn't join many-to-one or
86 many-to-many relationships, as we can't apply ORM's `unique()` method
87 to the results, which may lead to duplicates.
89 Args:
90 statement: The SQLAlchemy select statement to execute.
92 Yields:
93 Instances of the model `M` as they are fetched from the database.
94 """
95 results = await self.session.stream_scalars(
96 statement,
97 execution_options={"yield_per": settings.DATABASE_STREAM_YIELD_PER},
98 )
99 try:
100 async for result in results:
101 yield result
102 finally:
103 await results.close()
105 async def paginate(
106 self, statement: Select[tuple[M]], *, limit: int, page: int
107 ) -> tuple[list[M], int]:
108 offset = (page - 1) * limit
109 paginated_statement: Select[tuple[M, int]] = (
110 statement.add_columns(over(func.count())).limit(limit).offset(offset)
111 )
112 # Streaming can't be applied here, since we need to call ORM's unique()
113 results = await self.session.execute(paginated_statement)
115 items: list[M] = []
116 count = 0
117 for result in results.unique().all():
118 item, count = result._tuple()
119 items.append(item)
121 return items, count
123 def get_base_statement(self) -> Select[tuple[M]]:
124 return select(self.model)
126 async def create(self, object: M, *, flush: bool = False) -> M:
127 self.session.add(object)
129 if flush:
130 await self.session.flush()
132 return object
134 async def update(
135 self,
136 object: M,
137 *,
138 update_dict: dict[str, Any] | None = None,
139 flush: bool = False,
140 ) -> M:
141 if update_dict is not None: 141 ↛ 154line 141 didn't jump to line 154 because the condition on line 141 was always true
142 for attr, value in update_dict.items():
143 setattr(object, attr, value)
144 # Always consider that the attribute was modified if it's explictly set
145 # in the update_dict. This forces SQLAlchemy to include it in the
146 # UPDATE statement, even if the value is the same as before.
147 # Ref: https://docs.sqlalchemy.org/en/20/orm/session_api.html#sqlalchemy.orm.attributes.flag_modified
148 try:
149 flag_modified(object, attr)
150 # Don't fail if the attribute is not tracked by SQLAlchemy
151 except KeyError:
152 pass
154 self.session.add(object)
156 if flush: 156 ↛ 157line 156 didn't jump to line 157 because the condition on line 156 was never true
157 await self.session.flush()
159 return object
161 async def count(self, statement: Select[tuple[M]]) -> int:
162 count_statement = statement.with_only_columns(func.count())
163 result = await self.session.execute(count_statement)
164 return result.scalar_one()
166 @classmethod
167 def from_session(cls, session: AsyncSession | AsyncReadSession) -> Self:
168 return cls(session)
171class RepositorySoftDeletionProtocol[MODEL_DELETED_AT: ModelDeletedAtProtocol](
172 RepositoryProtocol[MODEL_DELETED_AT], Protocol
173):
174 def get_base_statement( 174 ↛ anywhereline 174 didn't jump anywhere: it always raised an exception.
175 self, *, include_deleted: bool = False
176 ) -> Select[tuple[MODEL_DELETED_AT]]: ...
178 async def soft_delete( 178 ↛ exitline 178 didn't return from function 'soft_delete' because
179 self, object: MODEL_DELETED_AT, *, flush: bool = False
180 ) -> MODEL_DELETED_AT: ...
183class RepositorySoftDeletionMixin[MODEL_DELETED_AT: ModelDeletedAtProtocol]:
184 def get_base_statement(
185 self: RepositoryProtocol[MODEL_DELETED_AT],
186 *,
187 include_deleted: bool = False,
188 ) -> Select[tuple[MODEL_DELETED_AT]]:
189 statement = super().get_base_statement() # type: ignore[safe-super]
190 if not include_deleted:
191 statement = statement.where(self.model.deleted_at.is_(None))
192 return statement
194 async def soft_delete(
195 self: RepositoryProtocol[MODEL_DELETED_AT],
196 object: MODEL_DELETED_AT,
197 *,
198 flush: bool = False,
199 ) -> MODEL_DELETED_AT:
200 return await self.update(
201 object, update_dict={"deleted_at": utc_now()}, flush=flush
202 )
205class RepositoryIDMixin[MODEL_ID: ModelIDProtocol, ID_TYPE]: # type: ignore[type-arg]
206 async def get_by_id(
207 self: RepositoryProtocol[MODEL_ID],
208 id: ID_TYPE,
209 *,
210 options: Options = (),
211 ) -> MODEL_ID | None:
212 statement = (
213 self.get_base_statement().where(self.model.id == id).options(*options)
214 )
215 return await self.get_one_or_none(statement)
218class RepositorySoftDeletionIDMixin[
219 MODEL_DELETED_AT_ID: ModelDeletedAtIDProtocol, # type: ignore[type-arg]
220 ID_TYPE,
221]:
222 async def get_by_id(
223 self: RepositorySoftDeletionProtocol[MODEL_DELETED_AT_ID],
224 id: ID_TYPE,
225 *,
226 options: Options = (),
227 include_deleted: bool = False,
228 ) -> MODEL_DELETED_AT_ID | None:
229 statement = (
230 self.get_base_statement(include_deleted=include_deleted)
231 .where(self.model.id == id)
232 .options(*options)
233 )
234 return await self.get_one_or_none(statement)
237SortingClause: TypeAlias = ColumnExpressionArgument[Any] | UnaryExpression[Any]
240class RepositorySortingMixin[M, PE: StrEnum]:
241 sorting_enum: type[PE]
243 def apply_sorting(
244 self,
245 statement: Select[tuple[M]],
246 sorting: list[Sorting[PE]],
247 ) -> Select[tuple[M]]:
248 order_by_clauses: list[UnaryExpression[Any]] = []
249 for criterion, is_desc in sorting:
250 clause_function = desc if is_desc else asc
251 order_by_clauses.append(clause_function(self.get_sorting_clause(criterion)))
252 return statement.order_by(*order_by_clauses)
254 def get_sorting_clause(self, property: PE) -> SortingClause:
255 raise NotImplementedError()