fastapi-toolsets 5.1.0__tar.gz → 5.1.2__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (44) hide show
  1. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/PKG-INFO +1 -1
  2. fastapi_toolsets-5.1.2/pyproject.toml +133 -0
  3. fastapi_toolsets-5.1.0/pyproject.toml → fastapi_toolsets-5.1.2/pyproject.toml.orig +3 -3
  4. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/__init__.py +1 -1
  5. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/crud/factory.py +173 -67
  6. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/crud/search.py +4 -2
  7. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/db/core.py +9 -6
  8. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/dependencies.py +74 -55
  9. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/models/watched.py +36 -7
  10. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/LICENSE +0 -0
  11. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/README.md +0 -0
  12. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/_imports.py +0 -0
  13. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/cli/__init__.py +0 -0
  14. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/cli/app.py +0 -0
  15. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/cli/commands/__init__.py +0 -0
  16. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/cli/commands/fixtures.py +0 -0
  17. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/cli/config.py +0 -0
  18. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/cli/pyproject.py +0 -0
  19. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/cli/utils.py +0 -0
  20. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/crud/__init__.py +0 -0
  21. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/db/__init__.py +0 -0
  22. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/db/locks.py +0 -0
  23. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/db/m2m.py +0 -0
  24. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/db/testing.py +0 -0
  25. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/db/watch.py +0 -0
  26. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/exceptions/__init__.py +0 -0
  27. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/exceptions/exceptions.py +0 -0
  28. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/exceptions/handler.py +0 -0
  29. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/fixtures/__init__.py +0 -0
  30. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/fixtures/enum.py +0 -0
  31. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/fixtures/registry.py +0 -0
  32. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/fixtures/utils.py +0 -0
  33. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/logger.py +0 -0
  34. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/metrics/__init__.py +0 -0
  35. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/metrics/handler.py +0 -0
  36. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/metrics/registry.py +0 -0
  37. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/models/__init__.py +0 -0
  38. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/models/columns.py +0 -0
  39. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/py.typed +0 -0
  40. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/pytest/__init__.py +0 -0
  41. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/pytest/plugin.py +0 -0
  42. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/pytest/utils.py +0 -0
  43. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/schemas.py +0 -0
  44. {fastapi_toolsets-5.1.0 → fastapi_toolsets-5.1.2}/src/fastapi_toolsets/types.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: fastapi-toolsets
3
- Version: 5.1.0
3
+ Version: 5.1.2
4
4
  Summary: Production-ready utilities for FastAPI applications
5
5
  Keywords: fastapi,sqlalchemy,postgresql
6
6
  Author: d3vyce
@@ -0,0 +1,133 @@
1
+ [project]
2
+ name = "fastapi-toolsets"
3
+ version = "5.1.2"
4
+ description = "Production-ready utilities for FastAPI applications"
5
+ readme = "README.md"
6
+ license = "MIT"
7
+ license-files = ["LICENSE"]
8
+ requires-python = ">=3.11"
9
+ keywords = [
10
+ "fastapi",
11
+ "sqlalchemy",
12
+ "postgresql",
13
+ ]
14
+ classifiers = [
15
+ "Development Status :: 5 - Production/Stable",
16
+ "Framework :: AsyncIO",
17
+ "Framework :: FastAPI",
18
+ "Framework :: Pydantic",
19
+ "Intended Audience :: Developers",
20
+ "Intended Audience :: Information Technology",
21
+ "Intended Audience :: System Administrators",
22
+ "License :: OSI Approved :: MIT License",
23
+ "Operating System :: OS Independent",
24
+ "Programming Language :: Python :: 3 :: Only",
25
+ "Programming Language :: Python :: 3.11",
26
+ "Programming Language :: Python :: 3.12",
27
+ "Programming Language :: Python :: 3.13",
28
+ "Programming Language :: Python :: 3.14",
29
+ "Topic :: Software Development :: Libraries :: Python Modules",
30
+ "Topic :: Software Development :: Libraries",
31
+ "Topic :: Software Development",
32
+ "Typing :: Typed",
33
+ ]
34
+ dependencies = [
35
+ "asyncpg>=0.29.0",
36
+ "fastapi>=0.100.0",
37
+ "pydantic>=2.0",
38
+ "sqlalchemy[asyncio]>=2.0",
39
+ ]
40
+
41
+ [[project.authors]]
42
+ name = "d3vyce"
43
+ email = "contact@d3vyce.fr"
44
+
45
+ [project.urls]
46
+ Homepage = "https://github.com/d3vyce/fastapi-toolsets"
47
+ Documentation = "https://fastapi-toolsets.d3vyce.fr/"
48
+ Repository = "https://github.com/d3vyce/fastapi-toolsets"
49
+ Issues = "https://github.com/d3vyce/fastapi-toolsets/issues"
50
+
51
+ [project.optional-dependencies]
52
+ cli = ["typer>=0.9.0"]
53
+ metrics = ["prometheus_client>=0.20.0"]
54
+ pytest = [
55
+ "httpx>=0.25.0",
56
+ "pytest-xdist>=3.0.0",
57
+ "pytest>=8.0.0",
58
+ ]
59
+ all = ["fastapi-toolsets[cli,metrics,pytest]"]
60
+
61
+ [project.scripts]
62
+ manager = "fastapi_toolsets.cli.app:cli"
63
+
64
+ [dependency-groups]
65
+ dev = [
66
+ { include-group = "tests" },
67
+ { include-group = "docs" },
68
+ { include-group = "docs-src" },
69
+ "fastapi-toolsets[all]",
70
+ "prek>=0.3.8",
71
+ "ruff>=0.1.0",
72
+ "ty>=0.0.1a0",
73
+ ]
74
+ tests = [
75
+ "async-lru>=1.0",
76
+ "coverage>=7.0.0",
77
+ "httpx>=0.25.0",
78
+ "pytest-anyio>=0.0.0",
79
+ "pytest-cov>=4.0.0",
80
+ "pytest-xdist>=3.0.0",
81
+ "pytest>=8.0.0",
82
+ ]
83
+ docs = [
84
+ "mike",
85
+ "mkdocstrings-python>=2.0.2",
86
+ "zensical>=0.0.30",
87
+ ]
88
+ docs-src = ["bcrypt>=4.0.0"]
89
+
90
+ [build-system]
91
+ requires = ["uv_build>=0.10,<0.13.0"]
92
+ build-backend = "uv_build"
93
+
94
+ [tool.ruff.format]
95
+ exclude = ["*.md"]
96
+
97
+ [tool.ruff.lint]
98
+ extend-select = ["E712"]
99
+
100
+ [tool.ruff.lint.flake8-bugbear]
101
+ extend-immutable-calls = [
102
+ "fastapi.Depends",
103
+ "fastapi.Security",
104
+ ]
105
+
106
+ [tool.ruff.lint.per-file-ignores]
107
+ "tests/**" = [
108
+ "RUF012",
109
+ "RUF059",
110
+ "SIM117",
111
+ "DTZ001",
112
+ "S110",
113
+ "BLE001",
114
+ ]
115
+
116
+ [tool.pytest.ini_options]
117
+ testpaths = ["tests"]
118
+ filterwarnings = ["ignore::DeprecationWarning"]
119
+
120
+ [tool.coverage.run]
121
+ source = ["src/fastapi_toolsets"]
122
+ branch = true
123
+
124
+ [tool.coverage.report]
125
+ exclude_lines = [
126
+ "pragma: no cover",
127
+ "if TYPE_CHECKING:",
128
+ "raise NotImplementedError",
129
+ ]
130
+
131
+ [tool.uv.sources.mike]
132
+ git = "https://github.com/squidfunk/mike.git"
133
+ tag = "2.2.0+zensical-0.1.0"
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "fastapi-toolsets"
3
- version = "5.1.0"
3
+ version = "5.1.2"
4
4
  description = "Production-ready utilities for FastAPI applications"
5
5
  readme = "README.md"
6
6
  license = "MIT"
@@ -91,7 +91,7 @@ docs-src = [
91
91
  ]
92
92
 
93
93
  [build-system]
94
- requires = ["uv_build>=0.10,<0.12.0"]
94
+ requires = ["uv_build>=0.10,<0.13.0"]
95
95
  build-backend = "uv_build"
96
96
 
97
97
  [tool.ruff.format]
@@ -101,7 +101,7 @@ exclude = ["*.md"]
101
101
  extend-select = ["E712"]
102
102
 
103
103
  [tool.ruff.lint.flake8-bugbear]
104
- extend-immutable-calls = ["fastapi.Depends"]
104
+ extend-immutable-calls = ["fastapi.Depends", "fastapi.Security"]
105
105
 
106
106
  [tool.ruff.lint.per-file-ignores]
107
107
  "tests/**" = ["RUF012", "RUF059", "SIM117", "DTZ001", "S110", "BLE001"]
@@ -24,4 +24,4 @@ Example usage:
24
24
  return Response(data={"user": user.username}, message="Success")
25
25
  """
26
26
 
27
- __version__ = "5.1.0"
27
+ __version__ = "5.1.2"
@@ -14,12 +14,25 @@ from typing import Any, ClassVar, Generic, Literal, Self, TypeAlias, cast, overl
14
14
 
15
15
  from fastapi import Query
16
16
  from pydantic import BaseModel
17
- from sqlalchemy import Date, DateTime, Float, Integer, Numeric, Uuid, and_, func, select
17
+ from sqlalchemy import (
18
+ Date,
19
+ DateTime,
20
+ Float,
21
+ Integer,
22
+ Numeric,
23
+ Uuid,
24
+ and_,
25
+ func,
26
+ select,
27
+ tuple_,
28
+ )
18
29
  from sqlalchemy.dialects.postgresql import insert
19
30
  from sqlalchemy.exc import NoResultFound
20
31
  from sqlalchemy.ext.asyncio import AsyncSession
21
32
  from sqlalchemy.orm import DeclarativeBase, QueryableAttribute, selectinload
33
+ from sqlalchemy.sql import operators
22
34
  from sqlalchemy.sql.base import ExecutableOption
35
+ from sqlalchemy.sql.elements import UnaryExpression
23
36
  from sqlalchemy.sql.roles import WhereHavingRole
24
37
 
25
38
  from ..db import transaction
@@ -128,6 +141,32 @@ def _apply_joins(q: Any, joins: JoinType | None, outer_join: bool) -> Any:
128
141
  return q
129
142
 
130
143
 
144
+ def _fans_out(
145
+ search_joins: Sequence[Any] | None, order_joins: Sequence[Any] | None
146
+ ) -> bool:
147
+ """True if any relationship join yields a collection."""
148
+ return any(
149
+ rel.property.uselist for rel in (*(search_joins or ()), *(order_joins or ()))
150
+ )
151
+
152
+
153
+ def _grouped_order(clause: Any, table: Any) -> Any:
154
+ """Recast an order clause for a query grouped by the entity's key."""
155
+ inner = clause.element if isinstance(clause, UnaryExpression) else clause
156
+ expr = inner.__clause_element__() if hasattr(inner, "__clause_element__") else inner
157
+ tables = {
158
+ t
159
+ for c in getattr(expr, "base_columns", ())
160
+ if (t := getattr(c, "table", None)) is not None
161
+ }
162
+ if tables and tables <= {table}:
163
+ return clause
164
+ agg = func.min(inner)
165
+ if isinstance(clause, UnaryExpression) and clause.modifier is operators.desc_op:
166
+ return agg.desc()
167
+ return agg.asc()
168
+
169
+
131
170
  class AsyncCrud(Generic[ModelType]):
132
171
  """Generic async CRUD operations for SQLAlchemy models.
133
172
 
@@ -163,6 +202,64 @@ class AsyncCrud(Generic[ModelType]):
163
202
  ):
164
203
  cls.searchable_fields = [pk_col, *raw_fields]
165
204
 
205
+ @classmethod
206
+ def _pk_attrs(cls: type[Self]) -> list[QueryableAttribute[Any]]:
207
+ """The model's primary key columns as instrumented attributes."""
208
+ return [
209
+ getattr(cls.model, cast(str, col.key))
210
+ for col in cls.model.__mapper__.primary_key
211
+ ]
212
+
213
+ @classmethod
214
+ async def _page_entities(
215
+ cls: type[Self],
216
+ session: AsyncSession,
217
+ q: Any,
218
+ *,
219
+ order_clauses: Sequence[Any],
220
+ limit: int,
221
+ offset: int | None = None,
222
+ load_options: Sequence[ExecutableOption] | None = None,
223
+ with_for_update: _ForUpdateMode = False,
224
+ ) -> list[ModelType]:
225
+ """Return up to *limit* entities, paging over distinct primary keys."""
226
+ pk_attrs = cls._pk_attrs()
227
+ table = cls.model.__table__
228
+ # Fall back to the primary key so the page boundary is deterministic.
229
+ grouped = [_grouped_order(c, table) for c in order_clauses] or [pk_attrs[0]]
230
+ id_q = (
231
+ q.order_by(None)
232
+ .with_only_columns(*pk_attrs)
233
+ .group_by(*pk_attrs)
234
+ .order_by(*grouped)
235
+ .limit(limit)
236
+ )
237
+ if offset:
238
+ id_q = id_q.offset(offset)
239
+ rows = (await session.execute(id_q)).all()
240
+ ids = [row[0] if len(pk_attrs) == 1 else tuple(row) for row in rows]
241
+ if not ids:
242
+ return []
243
+
244
+ where = (
245
+ pk_attrs[0].in_(ids) if len(pk_attrs) == 1 else tuple_(*pk_attrs).in_(ids)
246
+ )
247
+ item_q = select(cls.model).where(where)
248
+ if resolved := cls._resolve_load_options(load_options):
249
+ item_q = item_q.options(*resolved)
250
+ item_q = _apply_for_update(item_q, with_for_update)
251
+ found = (await session.execute(item_q)).unique().scalars().all()
252
+
253
+ rank = {pk: n for n, pk in enumerate(ids)}
254
+
255
+ def _key(obj: Any) -> Any:
256
+ values = tuple(getattr(obj, a.key) for a in pk_attrs)
257
+ return values[0] if len(pk_attrs) == 1 else values
258
+
259
+ return cast(
260
+ list[ModelType], sorted(found, key=lambda o: rank.get(_key(o), len(ids)))
261
+ )
262
+
166
263
  @classmethod
167
264
  def _resolve_load_options(
168
265
  cls, load_options: Sequence[ExecutableOption] | None
@@ -173,30 +270,17 @@ class AsyncCrud(Generic[ModelType]):
173
270
  return cls.default_load_options
174
271
 
175
272
  @classmethod
176
- def _capture_pk_values(cls: type[Self], instance: ModelType) -> dict[str, Any]:
177
- """Capture PK values off instance — call before commit expires attributes."""
178
- return {
179
- cast(str, col.key): getattr(instance, cast(str, col.key))
180
- for col in cls.model.__mapper__.primary_key
181
- }
182
-
183
- @classmethod
184
- async def _reload_with_options_by_pk(
185
- cls: type[Self], session: AsyncSession, pk_values: dict[str, Any]
273
+ async def _reload_with_options(
274
+ cls: type[Self], session: AsyncSession, instance: DeclarativeBase
186
275
  ) -> ModelType:
187
- """Re-query by previously captured PK values, with default_load_options applied."""
188
- # Only called when cls.default_load_options is set (see call sites).
276
+ """Re-query instance by PK with default_load_options applied."""
277
+ mapper = cls.model.__mapper__
189
278
  pk_filters = [
190
- getattr(cls.model, key) == value for key, value in pk_values.items()
279
+ getattr(cls.model, cast(str, col.key))
280
+ == getattr(instance, cast(str, col.key))
281
+ for col in mapper.primary_key
191
282
  ]
192
- q = select(cls.model).where(and_(*pk_filters))
193
- q = q.execution_options(populate_existing=True)
194
- q = q.options(*cast(Sequence[ExecutableOption], cls.default_load_options))
195
- result = await session.execute(q)
196
- item = result.unique().scalar_one_or_none()
197
- if item is None: # pragma: no cover — row was just flushed in this transaction
198
- raise NotFoundError()
199
- return cast(ModelType, item)
283
+ return await cls.get(session, filters=pk_filters)
200
284
 
201
285
  @classmethod
202
286
  async def _resolve_m2m(
@@ -750,14 +834,9 @@ class AsyncCrud(Generic[ModelType]):
750
834
  setattr(db_model, rel_attr, related_instances)
751
835
 
752
836
  session.add(db_model)
753
- pk_values: dict[str, Any] | None = None
754
- if cls.default_load_options:
755
- await session.flush()
756
- pk_values = cls._capture_pk_values(db_model)
757
- if pk_values is not None:
758
- db_model = await cls._reload_with_options_by_pk(session, pk_values)
759
- else:
760
- await session.refresh(db_model)
837
+ await session.refresh(db_model)
838
+ if cls.default_load_options:
839
+ db_model = await cls._reload_with_options(session, db_model)
761
840
  result = cast(ModelType, db_model)
762
841
  if schema:
763
842
  return Response(data=schema.model_validate(result))
@@ -1023,9 +1102,21 @@ class AsyncCrud(Generic[ModelType]):
1023
1102
  q = q.where(and_(*filters))
1024
1103
  if resolved := cls._resolve_load_options(load_options):
1025
1104
  q = q.options(*resolved)
1026
- q = _apply_for_update(q, with_for_update)
1027
1105
  if order_by is not None:
1028
1106
  q = q.order_by(order_by)
1107
+
1108
+ if limit is not None and joins:
1109
+ return await cls._page_entities(
1110
+ session,
1111
+ q,
1112
+ order_clauses=[] if order_by is None else [order_by],
1113
+ limit=limit,
1114
+ offset=offset,
1115
+ load_options=load_options,
1116
+ with_for_update=with_for_update,
1117
+ )
1118
+
1119
+ q = _apply_for_update(q, with_for_update)
1029
1120
  if offset is not None:
1030
1121
  q = q.offset(offset)
1031
1122
  if limit is not None:
@@ -1122,15 +1213,9 @@ class AsyncCrud(Generic[ModelType]):
1122
1213
  m2m_resolved = await cls._resolve_m2m(session, obj, only_set=True)
1123
1214
  for rel_attr, related_instances in m2m_resolved.items():
1124
1215
  setattr(db_model, rel_attr, related_instances)
1125
-
1126
- pk_values: dict[str, Any] | None = None
1127
- if cls.default_load_options:
1128
- await session.flush()
1129
- pk_values = cls._capture_pk_values(db_model)
1130
- if pk_values is not None:
1131
- db_model = await cls._reload_with_options_by_pk(session, pk_values)
1132
- else:
1133
- await session.refresh(db_model)
1216
+ await session.refresh(db_model)
1217
+ if cls.default_load_options:
1218
+ db_model = await cls._reload_with_options(session, db_model)
1134
1219
  if schema:
1135
1220
  return Response(data=schema.model_validate(db_model))
1136
1221
  return db_model
@@ -1374,17 +1459,31 @@ class AsyncCrud(Generic[ModelType]):
1374
1459
  q = q.where(and_(*filters))
1375
1460
  if resolved := cls._resolve_load_options(load_options):
1376
1461
  q = q.options(*resolved)
1377
- if order_by is not None:
1378
- q = q.order_by(order_by)
1379
-
1380
- if include_total:
1381
- q = q.offset(offset).limit(items_per_page)
1382
- result = await session.execute(q)
1462
+ order_clauses: list[Any] = [] if order_by is None else [order_by]
1463
+ q = q.order_by(*order_clauses)
1464
+
1465
+ fetch_limit = items_per_page if include_total else items_per_page + 1
1466
+ total_count: int | None = None
1467
+ # A to-many join repeats each entity, so LIMIT would slice joined rows
1468
+ # and `.unique()` would shrink the page after the fact.
1469
+ if _fans_out(search_joins, order_joins):
1470
+ raw_items = await cls._page_entities(
1471
+ session,
1472
+ q,
1473
+ order_clauses=order_clauses,
1474
+ limit=fetch_limit,
1475
+ offset=offset,
1476
+ load_options=load_options,
1477
+ )
1478
+ else:
1479
+ result = await session.execute(q.offset(offset).limit(fetch_limit))
1383
1480
  raw_items = cast(list[ModelType], result.unique().scalars().all())
1481
+ fetched = len(raw_items)
1482
+ raw_items = raw_items[:items_per_page]
1384
1483
 
1484
+ if include_total:
1385
1485
  # Count query (with same joins and filters)
1386
- pk_col = cls.model.__mapper__.primary_key[0]
1387
- count_q = select(func.count(func.distinct(getattr(cls.model, pk_col.name))))
1486
+ count_q = select(func.count(func.distinct(cls._pk_attrs()[0])))
1388
1487
  count_q = count_q.select_from(cls.model)
1389
1488
 
1390
1489
  # Apply explicit joins to count query
@@ -1397,16 +1496,11 @@ class AsyncCrud(Generic[ModelType]):
1397
1496
  count_q = count_q.where(and_(*filters))
1398
1497
 
1399
1498
  count_result = await session.execute(count_q)
1400
- total_count: int = count_result.scalar_one()
1499
+ total_count = count_result.scalar_one()
1401
1500
  has_more = page * items_per_page < total_count
1402
1501
  else:
1403
- # Fetch one extra row to detect if a next page exists without COUNT
1404
- q = q.offset(offset).limit(items_per_page + 1)
1405
- result = await session.execute(q)
1406
- raw_items = cast(list[ModelType], result.unique().scalars().all())
1407
- has_more = len(raw_items) > items_per_page
1408
- raw_items = raw_items[:items_per_page]
1409
- total_count = None
1502
+ # One extra row was fetched to detect a next page without COUNT
1503
+ has_more = fetched > items_per_page
1410
1504
 
1411
1505
  items: list[Any] = [schema.model_validate(item) for item in raw_items]
1412
1506
 
@@ -1543,18 +1637,30 @@ class AsyncCrud(Generic[ModelType]):
1543
1637
  q = q.options(*resolved)
1544
1638
 
1545
1639
  # Cursor column is always the primary sort; reverse direction for prev traversal
1546
- if direction is _CursorDirection.PREV:
1547
- q = q.order_by(cursor_column.desc())
1548
- else:
1549
- q = q.order_by(cursor_column)
1640
+ cursor_clause = (
1641
+ cursor_column.desc()
1642
+ if direction is _CursorDirection.PREV
1643
+ else cursor_column
1644
+ )
1645
+ order_clauses: list[Any] = [cursor_clause]
1550
1646
  if order_by is not None:
1551
- q = q.order_by(order_by)
1552
-
1553
- # Fetch one extra to detect whether another page exists in this direction
1554
- q = q.limit(items_per_page + 1)
1555
- result = await session.execute(q)
1556
- raw_items = cast(list[ModelType], result.unique().scalars().all())
1557
-
1647
+ order_clauses.append(order_by)
1648
+ q = q.order_by(*order_clauses)
1649
+
1650
+ # One extra row detects whether another page exists in this direction.
1651
+ # Under a to-many join that extra row may be a duplicate of one already
1652
+ # on the page, which reads as "no next page" and ends traversal early.
1653
+ if _fans_out(search_joins, order_joins):
1654
+ raw_items = await cls._page_entities(
1655
+ session,
1656
+ q,
1657
+ order_clauses=order_clauses,
1658
+ limit=items_per_page + 1,
1659
+ load_options=load_options,
1660
+ )
1661
+ else:
1662
+ result = await session.execute(q.limit(items_per_page + 1))
1663
+ raw_items = cast(list[ModelType], result.unique().scalars().all())
1558
1664
  has_more = len(raw_items) > items_per_page
1559
1665
  items_page = raw_items[:items_per_page]
1560
1666
 
@@ -160,9 +160,11 @@ def build_search_filters(
160
160
  column = field
161
161
 
162
162
  # Build the filter (cast to String only when needed, to preserve
163
- # pg_trgm GIN index usability on already-String columns)
163
+ # pg_trgm GIN index usability on already-String columns).
164
164
  column_as_string = (
165
- column if isinstance(column.type, String) else column.cast(String)
165
+ column
166
+ if isinstance(column.type, String) and not isinstance(column.type, Enum)
167
+ else column.cast(String)
166
168
  )
167
169
  if config.case_sensitive:
168
170
  filters.append(column_as_string.like(f"%{query}%"))
@@ -66,10 +66,8 @@ class _CommitOnResponseMiddleware:
66
66
 
67
67
  async def send_wrapper(message: Message) -> None:
68
68
  if message["type"] == "http.response.start":
69
- # ``scope["state"]`` is the same dict ``request.state`` writes
70
- # to, so this is the session stashed by the dependency.
71
69
  state = scope.get("state")
72
- session = state.get(self.state_attr) if state else None
70
+ session = state.pop(self.state_attr, None) if state else None
73
71
  if session is not None and session.in_transaction():
74
72
  await session.commit()
75
73
  await send(message)
@@ -158,7 +156,6 @@ class Database:
158
156
  # Private, per-instance state attribute; cannot collide with another
159
157
  # Database or be mismatched against the middleware.
160
158
  self._state_attr = f"_ft_db_session_{id(self):x}"
161
- self._middleware_installed = False
162
159
  self._disposed = False
163
160
 
164
161
  async def _dispose(self) -> None:
@@ -206,7 +203,6 @@ class Database:
206
203
  ```
207
204
  """
208
205
  app.add_middleware(_CommitOnResponseMiddleware, state_attr=self._state_attr)
209
- self._middleware_installed = True
210
206
 
211
207
  inner_lifespan = app.router.lifespan_context
212
208
 
@@ -243,10 +239,17 @@ class Database:
243
239
  return await UserCrud.get(session, [User.id == user_id])
244
240
  ```
245
241
  """
242
+ borrowed = getattr(request.state, self._state_attr, None)
243
+ if borrowed is not None:
244
+ yield borrowed
245
+ return
246
246
  async with self._open() as session:
247
247
  setattr(request.state, self._state_attr, session)
248
248
  yield session
249
- if not self._middleware_installed and session.in_transaction():
249
+ if (
250
+ getattr(request.state, self._state_attr, None) is session
251
+ and session.in_transaction()
252
+ ):
250
253
  await session.commit()
251
254
 
252
255
  @asynccontextmanager
@@ -2,14 +2,15 @@
2
2
 
3
3
  import inspect
4
4
  import typing
5
- from collections.abc import Callable
5
+ from collections.abc import Callable, Sequence
6
6
  from typing import Any, cast
7
7
 
8
8
  from fastapi import Depends
9
9
  from fastapi.params import Depends as DependsClass
10
10
  from sqlalchemy.ext.asyncio import AsyncSession
11
+ from sqlalchemy.sql.base import ExecutableOption
11
12
 
12
- from .crud import CrudFactory
13
+ from .crud import AsyncCrud, CrudFactory
13
14
  from .types import ModelType, SessionDependency
14
15
 
15
16
  __all__ = ["BodyDependency", "PathDependency"]
@@ -24,12 +25,59 @@ def _unwrap_session_dep(session_dep: SessionDependency) -> Callable[..., Any]:
24
25
  return session_dep
25
26
 
26
27
 
28
+ def _fetch_dependency(
29
+ model: type[ModelType],
30
+ field: Any,
31
+ *,
32
+ session_dep: SessionDependency,
33
+ param_name: str,
34
+ crud: type[AsyncCrud[ModelType]] | None,
35
+ load_options: Sequence[ExecutableOption] | None,
36
+ ) -> ModelType:
37
+ """Build a Depends() that fetches one row by ``field == <param_name>``."""
38
+ session_callable = _unwrap_session_dep(session_dep)
39
+ if crud is not None and crud.model is not model:
40
+ raise ValueError(
41
+ f"crud is bound to {crud.model.__name__}, not {model.__name__}"
42
+ )
43
+ crud = crud or CrudFactory(model)
44
+
45
+ # `session` has no default here: the __signature__ override below is what
46
+ # FastAPI reads, and it always passes `session` explicitly.
47
+ async def dependency(session: AsyncSession, **kwargs: Any) -> ModelType:
48
+ return await crud.get(
49
+ session,
50
+ filters=[field == kwargs[param_name]],
51
+ load_options=load_options,
52
+ )
53
+
54
+ dependency.__signature__ = inspect.Signature( # ty:ignore[unresolved-attribute]
55
+ parameters=[
56
+ inspect.Parameter(
57
+ param_name,
58
+ inspect.Parameter.KEYWORD_ONLY,
59
+ annotation=field.type.python_type,
60
+ ),
61
+ inspect.Parameter(
62
+ "session",
63
+ inspect.Parameter.KEYWORD_ONLY,
64
+ annotation=AsyncSession,
65
+ default=Depends(session_callable),
66
+ ),
67
+ ]
68
+ )
69
+
70
+ return cast(ModelType, Depends(cast(Callable[..., ModelType], dependency)))
71
+
72
+
27
73
  def PathDependency(
28
74
  model: type[ModelType],
29
75
  field: Any,
30
76
  *,
31
77
  session_dep: SessionDependency,
32
78
  param_name: str | None = None,
79
+ crud: type[AsyncCrud[ModelType]] | None = None,
80
+ load_options: Sequence[ExecutableOption] | None = None,
33
81
  ) -> ModelType:
34
82
  """Create a dependency that fetches a DB object from a path parameter.
35
83
 
@@ -38,6 +86,10 @@ def PathDependency(
38
86
  field: Model field to filter by (e.g., User.id)
39
87
  session_dep: Session dependency function (e.g., get_db)
40
88
  param_name: Path parameter name (defaults to model_field, e.g., user_id)
89
+ crud: Existing CRUD class to fetch with, so its ``default_load_options``
90
+ apply. Defaults to a bare ``CrudFactory(model)``.
91
+ load_options: SQLAlchemy loader options for the fetch. Overrides the CRUD's
92
+ ``default_load_options`` entirely rather than merging with them.
41
93
 
42
94
  Returns:
43
95
  A Depends() instance that resolves to the model instance
@@ -55,36 +107,14 @@ def PathDependency(
55
107
  ): ...
56
108
  ```
57
109
  """
58
- session_callable = _unwrap_session_dep(session_dep)
59
- crud = CrudFactory(model)
60
- name = (
61
- param_name
62
- if param_name is not None
63
- else f"{model.__name__.lower()}_{field.key}"
110
+ return _fetch_dependency(
111
+ model,
112
+ field,
113
+ session_dep=session_dep,
114
+ param_name=param_name or f"{model.__name__.lower()}_{field.key}",
115
+ crud=crud,
116
+ load_options=load_options,
64
117
  )
65
- python_type = field.type.python_type
66
-
67
- async def dependency(
68
- session: AsyncSession = Depends(session_callable), **kwargs: Any
69
- ) -> ModelType:
70
- value = kwargs[name]
71
- return await crud.get(session, filters=[field == value])
72
-
73
- dependency.__signature__ = inspect.Signature( # ty:ignore[unresolved-attribute]
74
- parameters=[
75
- inspect.Parameter(
76
- name, inspect.Parameter.KEYWORD_ONLY, annotation=python_type
77
- ),
78
- inspect.Parameter(
79
- "session",
80
- inspect.Parameter.KEYWORD_ONLY,
81
- annotation=AsyncSession,
82
- default=Depends(session_callable),
83
- ),
84
- ]
85
- )
86
-
87
- return cast(ModelType, Depends(cast(Callable[..., ModelType], dependency)))
88
118
 
89
119
 
90
120
  def BodyDependency(
@@ -93,6 +123,8 @@ def BodyDependency(
93
123
  *,
94
124
  session_dep: SessionDependency,
95
125
  body_field: str,
126
+ crud: type[AsyncCrud[ModelType]] | None = None,
127
+ load_options: Sequence[ExecutableOption] | None = None,
96
128
  ) -> ModelType:
97
129
  """Create a dependency that fetches a DB object from a body field.
98
130
 
@@ -101,6 +133,10 @@ def BodyDependency(
101
133
  field: Model field to filter by (e.g., User.id)
102
134
  session_dep: Session dependency function (e.g., get_db)
103
135
  body_field: Name of the field in the request body
136
+ crud: Existing CRUD class to fetch with, so its ``default_load_options``
137
+ apply. Defaults to a bare ``CrudFactory(model)``.
138
+ load_options: SQLAlchemy loader options for the fetch. Overrides the CRUD's
139
+ ``default_load_options`` entirely rather than merging with them.
104
140
 
105
141
  Returns:
106
142
  A Depends() instance that resolves to the model instance
@@ -120,28 +156,11 @@ def BodyDependency(
120
156
  ): ...
121
157
  ```
122
158
  """
123
- session_callable = _unwrap_session_dep(session_dep)
124
- crud = CrudFactory(model)
125
- python_type = field.type.python_type
126
-
127
- async def dependency(
128
- session: AsyncSession = Depends(session_callable), **kwargs: Any
129
- ) -> ModelType:
130
- value = kwargs[body_field]
131
- return await crud.get(session, filters=[field == value])
132
-
133
- dependency.__signature__ = inspect.Signature( # ty:ignore[unresolved-attribute]
134
- parameters=[
135
- inspect.Parameter(
136
- body_field, inspect.Parameter.KEYWORD_ONLY, annotation=python_type
137
- ),
138
- inspect.Parameter(
139
- "session",
140
- inspect.Parameter.KEYWORD_ONLY,
141
- annotation=AsyncSession,
142
- default=Depends(session_callable),
143
- ),
144
- ]
159
+ return _fetch_dependency(
160
+ model,
161
+ field,
162
+ session_dep=session_dep,
163
+ param_name=body_field,
164
+ crud=crud,
165
+ load_options=load_options,
145
166
  )
146
-
147
- return cast(ModelType, Depends(cast(Callable[..., ModelType], dependency)))
@@ -8,6 +8,7 @@ from typing import Any
8
8
  from sqlalchemy import event, select, tuple_
9
9
  from sqlalchemy import inspect as sa_inspect
10
10
  from sqlalchemy.ext.asyncio import AsyncSession
11
+ from sqlalchemy.orm import selectinload
11
12
  from sqlalchemy.orm.attributes import set_committed_value as _sa_set_committed_value
12
13
 
13
14
  from ..logger import get_logger
@@ -189,17 +190,44 @@ async def _invoke_callback(
189
190
  await result
190
191
 
191
192
 
193
+ def _loaded_relationships(obj: Any) -> set[str]:
194
+ """Relationship keys currently loaded on *obj*."""
195
+ state = sa_inspect(obj)
196
+ unloaded = state.unloaded
197
+ return {
198
+ rel.key
199
+ for rel in state.mapper.relationships
200
+ if rel.key not in unloaded and rel.lazy not in ("dynamic", "write_only")
201
+ }
202
+
203
+
204
+ def _snapshot_loaded_relationships(session: Any) -> dict[int, set[str]]:
205
+ """Record loaded relationships for the tracked objects, keyed by ``id``."""
206
+ objs = list(session.info.get(_SESSION_CREATES, []))
207
+ objs += [obj for obj, _ in session.info.get(_SESSION_UPDATES, {}).values()]
208
+ return {id(obj): _loaded_relationships(obj) for obj in objs}
209
+
210
+
192
211
  async def _batch_reload(
193
- session: AsyncSession, model: type, pk_tuples: list[tuple[Any, ...]]
212
+ session: AsyncSession,
213
+ model: type,
214
+ objs: list[Any],
215
+ preloaded: dict[int, set[str]],
194
216
  ) -> None:
195
- """Re-populate all rows of *model* identified by *pk_tuples* in one round trip."""
217
+ """Re-populate all rows of *model* in one round trip."""
196
218
  pk_cols = sa_inspect(model, raiseerr=True).primary_key
219
+ pk_tuples = [sa_inspect(obj).key[1] for obj in objs]
197
220
  where = (
198
221
  pk_cols[0].in_([pk[0] for pk in pk_tuples])
199
222
  if len(pk_cols) == 1
200
223
  else tuple_(*pk_cols).in_(pk_tuples)
201
224
  )
202
225
  q = select(model).where(where).execution_options(populate_existing=True)
226
+ loaded: set[str] = set()
227
+ for obj in objs:
228
+ loaded |= preloaded.get(id(obj), set())
229
+ if loaded:
230
+ q = q.options(*(selectinload(getattr(model, key)) for key in loaded))
203
231
  await session.execute(q)
204
232
 
205
233
 
@@ -207,6 +235,7 @@ class EventSession(AsyncSession):
207
235
  """AsyncSession subclass that dispatches lifecycle callbacks after commit."""
208
236
 
209
237
  async def commit(self) -> None:
238
+ preloaded = _snapshot_loaded_relationships(self)
210
239
  await super().commit()
211
240
 
212
241
  creates: list[Any] = self.info.pop(_SESSION_CREATES, [])
@@ -249,25 +278,25 @@ class EventSession(AsyncSession):
249
278
  # session.get() per object.
250
279
  create_items: list[Any] = []
251
280
  update_items: list[tuple[Any, dict[str, dict[str, Any]]]] = []
252
- pk_by_type: dict[type, list[tuple[Any, ...]]] = {}
281
+ objs_by_type: dict[type, list[Any]] = {}
253
282
 
254
283
  for obj in creates:
255
284
  state = sa_inspect(obj, raiseerr=False)
256
285
  if state is None or state.detached or state.transient: # pragma: no cover
257
286
  continue
258
287
  create_items.append(obj)
259
- pk_by_type.setdefault(type(obj), []).append(state.key[1])
288
+ objs_by_type.setdefault(type(obj), []).append(obj)
260
289
 
261
290
  for obj, changes in field_changes.values():
262
291
  state = sa_inspect(obj, raiseerr=False)
263
292
  if state is None or state.detached or state.transient: # pragma: no cover
264
293
  continue
265
294
  update_items.append((obj, changes))
266
- pk_by_type.setdefault(type(obj), []).append(state.key[1])
295
+ objs_by_type.setdefault(type(obj), []).append(obj)
267
296
 
268
- for model, pk_tuples in pk_by_type.items():
297
+ for model, objs in objs_by_type.items():
269
298
  try:
270
- await _batch_reload(self, model, pk_tuples)
299
+ await _batch_reload(self, model, objs, preloaded)
271
300
  except Exception as exc:
272
301
  _logger.error(_CALLBACK_ERROR_MSG, exc_info=exc)
273
302