fastapi-toolsets 5.1.2__tar.gz → 5.1.3__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.2 → fastapi_toolsets-5.1.3}/PKG-INFO +1 -1
  2. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/pyproject.toml +1 -1
  3. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/pyproject.toml.orig +1 -1
  4. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/__init__.py +1 -1
  5. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/crud/factory.py +33 -6
  6. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/db/core.py +23 -5
  7. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/db/locks.py +67 -28
  8. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/models/watched.py +25 -17
  9. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/pytest/utils.py +49 -6
  10. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/LICENSE +0 -0
  11. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/README.md +0 -0
  12. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/_imports.py +0 -0
  13. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/cli/__init__.py +0 -0
  14. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/cli/app.py +0 -0
  15. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/cli/commands/__init__.py +0 -0
  16. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/cli/commands/fixtures.py +0 -0
  17. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/cli/config.py +0 -0
  18. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/cli/pyproject.py +0 -0
  19. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/cli/utils.py +0 -0
  20. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/crud/__init__.py +0 -0
  21. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/crud/search.py +0 -0
  22. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/db/__init__.py +0 -0
  23. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/db/m2m.py +0 -0
  24. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/db/testing.py +0 -0
  25. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/db/watch.py +0 -0
  26. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/dependencies.py +0 -0
  27. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/exceptions/__init__.py +0 -0
  28. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/exceptions/exceptions.py +0 -0
  29. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/exceptions/handler.py +0 -0
  30. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/fixtures/__init__.py +0 -0
  31. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/fixtures/enum.py +0 -0
  32. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/fixtures/registry.py +0 -0
  33. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/fixtures/utils.py +0 -0
  34. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/logger.py +0 -0
  35. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/metrics/__init__.py +0 -0
  36. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/metrics/handler.py +0 -0
  37. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/metrics/registry.py +0 -0
  38. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/models/__init__.py +0 -0
  39. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/models/columns.py +0 -0
  40. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/py.typed +0 -0
  41. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/pytest/__init__.py +0 -0
  42. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/pytest/plugin.py +0 -0
  43. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/src/fastapi_toolsets/schemas.py +0 -0
  44. {fastapi_toolsets-5.1.2 → fastapi_toolsets-5.1.3}/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.2
3
+ Version: 5.1.3
4
4
  Summary: Production-ready utilities for FastAPI applications
5
5
  Keywords: fastapi,sqlalchemy,postgresql
6
6
  Author: d3vyce
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "fastapi-toolsets"
3
- version = "5.1.2"
3
+ version = "5.1.3"
4
4
  description = "Production-ready utilities for FastAPI applications"
5
5
  readme = "README.md"
6
6
  license = "MIT"
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "fastapi-toolsets"
3
- version = "5.1.2"
3
+ version = "5.1.3"
4
4
  description = "Production-ready utilities for FastAPI applications"
5
5
  readme = "README.md"
6
6
  license = "MIT"
@@ -24,4 +24,4 @@ Example usage:
24
24
  return Response(data={"user": user.username}, message="Success")
25
25
  """
26
26
 
27
- __version__ = "5.1.2"
27
+ __version__ = "5.1.3"
@@ -386,13 +386,29 @@ class AsyncCrud(Generic[ModelType]):
386
386
  own_filters=own_filters,
387
387
  )
388
388
 
389
+ @classmethod
390
+ def _resolve_search_fields(
391
+ cls: type[Self],
392
+ search_fields: Sequence[SearchFieldType] | None,
393
+ ) -> Sequence[SearchFieldType] | None:
394
+ """Return search_fields if given, otherwise fall back to the class-level default."""
395
+ return search_fields if search_fields is not None else cls.searchable_fields
396
+
397
+ @classmethod
398
+ def _resolve_order_fields(
399
+ cls: type[Self],
400
+ order_fields: Sequence[OrderFieldType] | None,
401
+ ) -> Sequence[OrderFieldType] | None:
402
+ """Return order_fields if given, otherwise fall back to the class-level default."""
403
+ return order_fields if order_fields is not None else cls.order_fields
404
+
389
405
  @classmethod
390
406
  def _resolve_search_columns(
391
407
  cls: type[Self],
392
408
  search_fields: Sequence[SearchFieldType] | None,
393
409
  ) -> list[str] | None:
394
410
  """Return search column keys, or None if no searchable fields configured."""
395
- fields = search_fields if search_fields is not None else cls.searchable_fields
411
+ fields = cls._resolve_search_fields(search_fields)
396
412
  if not fields:
397
413
  return None
398
414
  return search_field_keys(fields)
@@ -403,7 +419,7 @@ class AsyncCrud(Generic[ModelType]):
403
419
  order_fields: Sequence[OrderFieldType] | None,
404
420
  ) -> list[str] | None:
405
421
  """Return sort column keys, or None if no order fields configured."""
406
- fields = order_fields if order_fields is not None else cls.order_fields
422
+ fields = cls._resolve_order_fields(order_fields)
407
423
  if not fields:
408
424
  return None
409
425
  return sorted(facet_keys(fields))
@@ -482,9 +498,7 @@ class AsyncCrud(Generic[ModelType]):
482
498
  order_field_map: dict[str, OrderFieldType] | None = None
483
499
  order_valid_keys: list[str] | None = None
484
500
  if order:
485
- resolved_order = (
486
- order_fields if order_fields is not None else cls.order_fields
487
- )
501
+ resolved_order = cls._resolve_order_fields(order_fields)
488
502
  if resolved_order:
489
503
  keys = facet_keys(resolved_order)
490
504
  order_field_map = dict(zip(keys, resolved_order))
@@ -510,8 +524,21 @@ class AsyncCrud(Generic[ModelType]):
510
524
  ]
511
525
  )
512
526
 
527
+ fixed: dict[str, Any] = {
528
+ **pagination_fixed,
529
+ "search_fields": (cls._resolve_search_fields(search_fields) or [])
530
+ if search
531
+ else [],
532
+ "facet_fields": (cls._resolve_facet_fields(facet_fields) or [])
533
+ if filter
534
+ else [],
535
+ "order_fields": (cls._resolve_order_fields(order_fields) or [])
536
+ if order
537
+ else [],
538
+ }
539
+
513
540
  async def dependency(**kwargs: Any) -> dict[str, Any]:
514
- result: dict[str, Any] = dict(pagination_fixed)
541
+ result: dict[str, Any] = dict(fixed)
515
542
  for name in pagination_param_names:
516
543
  result[name] = kwargs[name]
517
544
 
@@ -295,21 +295,27 @@ class Database:
295
295
  self,
296
296
  tables: list[type[DeclarativeBase]],
297
297
  *,
298
+ session: AsyncSession | None = None,
298
299
  mode: LockMode = LockMode.SHARE_UPDATE_EXCLUSIVE,
299
300
  timeout: str = "5s",
300
301
  ) -> AbstractAsyncContextManager[AsyncSession]:
301
- """Lock PostgreSQL tables for the duration of a dedicated transaction.
302
+ """Lock PostgreSQL tables for the duration of a transaction.
302
303
 
303
- Opens its own session from the facade's sessionmaker, changes are
304
- committed when the context exits.
304
+ Without *session*, a dedicated session is opened from the facade's
305
+ sessionmaker and committed at block exit. That is a **second**
306
+ connection, and the block must not touch the request session under a
307
+ conflicting mode. Pass ``session=`` to lock on the request's own
308
+ transaction instead: one connection, and the block may use it freely.
305
309
 
306
310
  Args:
307
311
  tables: List of SQLAlchemy model classes to lock.
312
+ session: Existing session whose transaction takes the lock. That
313
+ transaction holds the lock until it ends.
308
314
  mode: Lock mode (default: ``SHARE UPDATE EXCLUSIVE``).
309
315
  timeout: Lock timeout (default: ``"5s"``).
310
316
 
311
317
  Yields:
312
- The dedicated session, open within the locked transaction.
318
+ The session holding the lock: the dedicated one, or *session*.
313
319
 
314
320
  Raises:
315
321
  LockTimeoutError: If the lock cannot be acquired within *timeout*.
@@ -320,6 +326,18 @@ class Database:
320
326
  async with db.lock_tables([User, Account]) as session:
321
327
  user = await UserCrud.get(session, [User.id == 1])
322
328
  user.balance += 100
329
+
330
+
331
+ @app.post("/transfer")
332
+ async def transfer(session=Depends(db)):
333
+ async with db.lock_tables([Account], session=session):
334
+ ... # same session, same connection
323
335
  ```
324
336
  """
325
- return lock_tables(self._sessionmaker, tables, mode=mode, timeout=timeout)
337
+ return lock_tables(
338
+ None if session is not None else self._sessionmaker,
339
+ tables,
340
+ session=session,
341
+ mode=mode,
342
+ timeout=timeout,
343
+ )
@@ -39,9 +39,10 @@ class LockMode(str, Enum):
39
39
 
40
40
 
41
41
  def lock_tables(
42
- session_maker: async_sessionmaker[_SessionT],
42
+ session_maker: async_sessionmaker[_SessionT] | None,
43
43
  tables: list[type[DeclarativeBase]],
44
44
  *,
45
+ session: _SessionT | None = None,
45
46
  mode: LockMode = LockMode.SHARE_UPDATE_EXCLUSIVE,
46
47
  timeout: str = "5s",
47
48
  ) -> AbstractAsyncContextManager[_SessionT]:
@@ -50,20 +51,31 @@ def lock_tables(
50
51
  Prefer the method on a :class:`Database` instance; use this
51
52
  directly only when you manage your own session factory.
52
53
 
54
+ Pass exactly one of *session_maker*, which opens a dedicated session and
55
+ commits it at block exit, or *session*, which locks on a transaction you
56
+ already own and so needs only one connection.
57
+
53
58
  Args:
54
- session_maker: Async session factory used to create the dedicated
55
- session.
59
+ session_maker: Async session factory for the dedicated session.
60
+ ``None`` when *session* is given.
56
61
  tables: List of SQLAlchemy model classes to lock.
62
+ session: Existing session whose transaction takes the lock.
57
63
  mode: Lock mode (default: SHARE UPDATE EXCLUSIVE).
58
64
  timeout: Lock timeout (default: "5s").
59
65
 
60
66
  Yields:
61
- The dedicated session, open within the locked transaction.
67
+ The session holding the lock: the dedicated one, or *session* itself.
62
68
 
63
69
  Raises:
70
+ TypeError: If neither or both of *session_maker* and *session* are given.
64
71
  LockTimeoutError: If the lock cannot be acquired within *timeout*.
65
72
  PoolExhaustedError: If the connection pool is exhausted.
66
73
 
74
+ Note:
75
+ With *session*, nothing is committed or rolled back for you: the lock
76
+ is held until the caller's transaction ends, pending ORM changes flush
77
+ as it is taken, and ``lock_timeout`` stays set on that transaction.
78
+
67
79
  Example:
68
80
  ```python
69
81
  from fastapi_toolsets.db import lock_tables
@@ -74,32 +86,53 @@ def lock_tables(
74
86
  ```
75
87
  """
76
88
  table_names = ",".join(table.__tablename__ for table in tables)
89
+ set_timeout = text(f"SET LOCAL lock_timeout='{timeout}'")
90
+ acquire = text(f"LOCK {table_names} IN {mode.value} MODE")
91
+
92
+ def _translate(e: BaseException) -> None:
93
+ if isinstance(e, sa_exc.TimeoutError):
94
+ raise PoolExhaustedError(
95
+ f"Connection pool exhausted while locking '{table_names}'. "
96
+ ) from e
97
+ if isinstance(e, sa_exc.DBAPIError) and _is_lock_not_available(e):
98
+ raise LockTimeoutError(
99
+ f"Lock on '{table_names}' could not be acquired within {timeout}."
100
+ ) from e
77
101
 
78
102
  @asynccontextmanager
79
- async def _lock() -> AsyncGenerator[_SessionT, None]:
80
- async with session_maker() as session:
103
+ async def _lock_dedicated(
104
+ maker: async_sessionmaker[_SessionT],
105
+ ) -> AsyncGenerator[_SessionT, None]:
106
+ async with maker() as dedicated:
81
107
  try:
82
- await session.execute(text(f"SET LOCAL lock_timeout='{timeout}'"))
83
- await session.execute(text(f"LOCK {table_names} IN {mode.value} MODE"))
84
- yield session
85
- await session.commit()
86
- except sa_exc.TimeoutError as e:
87
- await session.rollback()
88
- raise PoolExhaustedError(
89
- f"Connection pool exhausted while locking '{table_names}'. "
90
- ) from e
91
- except sa_exc.DBAPIError as e:
92
- await session.rollback()
93
- if _is_lock_not_available(e):
94
- raise LockTimeoutError(
95
- f"Lock on '{table_names}' could not be acquired within {timeout}."
96
- ) from e
97
- raise # pragma: no cover
98
- except BaseException:
99
- await session.rollback()
108
+ await dedicated.execute(set_timeout)
109
+ await dedicated.execute(acquire)
110
+ yield dedicated
111
+ await dedicated.commit()
112
+ except BaseException as e:
113
+ await dedicated.rollback()
114
+ _translate(e)
100
115
  raise
101
116
 
102
- return _lock()
117
+ @asynccontextmanager
118
+ async def _lock_caller(caller: _SessionT) -> AsyncGenerator[_SessionT, None]:
119
+ try:
120
+ with caller.no_autoflush:
121
+ await caller.execute(set_timeout)
122
+ async with caller.begin_nested():
123
+ await caller.execute(acquire)
124
+ except BaseException as e:
125
+ _translate(e)
126
+ raise
127
+ yield caller
128
+
129
+ if session_maker is not None and session is None:
130
+ return _lock_dedicated(session_maker)
131
+ if session is not None and session_maker is None:
132
+ return _lock_caller(session)
133
+ raise TypeError(
134
+ "lock_tables() requires exactly one of 'session_maker' or 'session'."
135
+ )
103
136
 
104
137
 
105
138
  @asynccontextmanager
@@ -109,15 +142,17 @@ async def advisory_lock(
109
142
  *,
110
143
  shared: bool = False,
111
144
  nowait: bool = False,
145
+ xact: bool = False,
112
146
  timeout: str | None = None,
113
147
  ) -> AsyncGenerator[bool, None]:
114
- """Acquire a PostgreSQL session-level advisory lock.
148
+ """Acquire a PostgreSQL advisory lock.
115
149
 
116
150
  Args:
117
151
  session: AsyncSession instance.
118
152
  key: Lock key, either a single ``int`` (bigint) or a ``(int, int)`` pair for namespacing.
119
153
  shared: Acquire a shared lock (multiple holders allowed). Default is exclusive.
120
154
  nowait: Return ``False`` immediately if the lock is unavailable instead of waiting.
155
+ xact: Hold the lock until the caller's transaction ends.
121
156
  timeout: Maximum wait time (e.g. ``"5s"``, ``"500ms"``). Raises ``DBAPIError``
122
157
  if exceeded. Ignored when *nowait* is ``True``.
123
158
 
@@ -144,10 +179,14 @@ async def advisory_lock(
144
179
 
145
180
  async with advisory_lock(session, (1, user_id), shared=True):
146
181
  ...
182
+
183
+ async with advisory_lock(session, (team_id, question_id), xact=True):
184
+ ... # held until the request's transaction commits
147
185
  ```
148
186
  """
149
187
  suffix = "_shared" if shared else ""
150
- acquire_fn = f"{'pg_try_advisory_lock' if nowait else 'pg_advisory_lock'}{suffix}"
188
+ scope = "_xact" if xact else ""
189
+ acquire_fn = f"pg_{'try_' if nowait else ''}advisory{scope}_lock{suffix}"
151
190
  release_fn = f"pg_advisory_unlock{suffix}"
152
191
 
153
192
  if isinstance(key, tuple):
@@ -180,6 +219,6 @@ async def advisory_lock(
180
219
  try:
181
220
  yield acquired
182
221
  finally:
183
- if acquired:
222
+ if acquired and not xact:
184
223
  with session.no_autoflush:
185
224
  await session.execute(release_sql, params)
@@ -155,11 +155,14 @@ def _after_flush(session: Any, flush_context: Any) -> None:
155
155
  )
156
156
  for field, attr_state in attrs:
157
157
  history = attr_state.history
158
- if history.has_changes() and history.deleted:
159
- changes[field] = {
160
- "old": history.deleted[0],
161
- "new": history.added[0] if history.added else None,
162
- }
158
+ if not history.has_changes():
159
+ continue
160
+ change: dict[str, Any] = {
161
+ "new": history.added[0] if history.added else None
162
+ }
163
+ if history.deleted:
164
+ change["old"] = history.deleted[0]
165
+ changes[field] = change
163
166
 
164
167
  if changes:
165
168
  _upsert_changes(
@@ -190,6 +193,19 @@ async def _invoke_callback(
190
193
  await result
191
194
 
192
195
 
196
+ async def _dispatch(
197
+ obj: Any,
198
+ event_type: ModelEvent,
199
+ changes: dict[str, dict[str, Any]] | None,
200
+ ) -> None:
201
+ """Run every handler for *obj*, isolating a failure to the handler that raised."""
202
+ for handler in _get_handlers(type(obj), event_type):
203
+ try:
204
+ await _invoke_callback(handler, obj, event_type, changes)
205
+ except Exception as exc:
206
+ _logger.error(_CALLBACK_ERROR_MSG, exc_info=exc)
207
+
208
+
193
209
  def _loaded_relationships(obj: Any) -> set[str]:
194
210
  """Relationship keys currently loaded on *obj*."""
195
211
  state = sa_inspect(obj)
@@ -302,29 +318,21 @@ class EventSession(AsyncSession):
302
318
 
303
319
  # Dispatch CREATE callbacks.
304
320
  for obj in create_items:
305
- try:
306
- for handler in _get_handlers(type(obj), ModelEvent.CREATE):
307
- await _invoke_callback(handler, obj, ModelEvent.CREATE, None)
308
- except Exception as exc:
309
- _logger.error(_CALLBACK_ERROR_MSG, exc_info=exc)
321
+ await _dispatch(obj, ModelEvent.CREATE, None)
310
322
 
311
323
  # Dispatch DELETE callbacks (restore snapshot; row is gone).
312
324
  for obj, snapshot in deletes:
313
325
  try:
314
326
  for key, value in snapshot.items():
315
327
  _sa_set_committed_value(obj, key, value)
316
- for handler in _get_handlers(type(obj), ModelEvent.DELETE):
317
- await _invoke_callback(handler, obj, ModelEvent.DELETE, None)
318
328
  except Exception as exc:
319
329
  _logger.error(_CALLBACK_ERROR_MSG, exc_info=exc)
330
+ continue
331
+ await _dispatch(obj, ModelEvent.DELETE, None)
320
332
 
321
333
  # Dispatch UPDATE callbacks.
322
334
  for obj, changes in update_items:
323
- try:
324
- for handler in _get_handlers(type(obj), ModelEvent.UPDATE):
325
- await _invoke_callback(handler, obj, ModelEvent.UPDATE, changes)
326
- except Exception as exc:
327
- _logger.error(_CALLBACK_ERROR_MSG, exc_info=exc)
335
+ await _dispatch(obj, ModelEvent.UPDATE, changes)
328
336
 
329
337
  async def rollback(self) -> None:
330
338
  await super().rollback()
@@ -4,6 +4,7 @@ import os
4
4
  from collections.abc import AsyncGenerator, Callable
5
5
  from contextlib import asynccontextmanager
6
6
  from typing import Any
7
+ from weakref import WeakKeyDictionary
7
8
 
8
9
  from httpx import ASGITransport, AsyncClient
9
10
  from sqlalchemy import text
@@ -18,6 +19,47 @@ from sqlalchemy.orm import DeclarativeBase
18
19
  from ..db.testing import cleanup_tables, create_database
19
20
  from ..models.watched import EventSession
20
21
 
22
+ _MISSING = object()
23
+ _BASE = object()
24
+
25
+ _override_layers: WeakKeyDictionary[Any, dict[Any, list[tuple[Any, Any]]]] = (
26
+ WeakKeyDictionary()
27
+ )
28
+
29
+
30
+ def _push_overrides(app: Any, overrides: dict[Any, Any], owner: Any) -> None:
31
+ """Register *overrides* for *owner* and apply them to *app*."""
32
+ layers = _override_layers.setdefault(app, {})
33
+ for key, override in overrides.items():
34
+ if key not in layers:
35
+ layers[key] = [(_BASE, app.dependency_overrides.get(key, _MISSING))]
36
+ layers[key].append((owner, override))
37
+ app.dependency_overrides[key] = override
38
+
39
+
40
+ def _pop_overrides(app: Any, overrides: dict[Any, Any], owner: Any) -> None:
41
+ """Drop *owner*'s layer and re-apply whichever layer is now on top."""
42
+ layers = _override_layers.get(app)
43
+ if layers is None:
44
+ return
45
+ for key in overrides:
46
+ stack = layers.get(key)
47
+ if stack is None:
48
+ continue
49
+ for index, (token, _) in enumerate(stack):
50
+ if token is owner:
51
+ del stack[index]
52
+ break
53
+ _, current = stack[-1]
54
+ if current is _MISSING:
55
+ app.dependency_overrides.pop(key, None)
56
+ else:
57
+ app.dependency_overrides[key] = current
58
+ if len(stack) == 1:
59
+ del layers[key]
60
+ if not layers:
61
+ del _override_layers[app]
62
+
21
63
 
22
64
  def _get_xdist_worker(default_test_db: str) -> str:
23
65
  """Return the pytest-xdist worker name, or *default_test_db* when not running under xdist.
@@ -168,7 +210,7 @@ async def create_async_client(
168
210
  base_url: Base URL for requests. Defaults to "http://test".
169
211
  dependency_overrides: Optional mapping of original dependencies to
170
212
  their test replacements. Applied via ``app.dependency_overrides``
171
- before yielding and cleaned up after.
213
+ before yielding and restored to their previous state after.
172
214
  **kwargs: Additional keyword arguments forwarded to
173
215
  :class:`httpx.AsyncClient` (e.g. ``headers``, ``cookies``,
174
216
  ``auth``, ``timeout``).
@@ -214,19 +256,20 @@ async def create_async_client(
214
256
  yield c
215
257
  ```
216
258
  """
217
- if dependency_overrides:
218
- app.dependency_overrides.update(dependency_overrides)
259
+ overrides = dependency_overrides or {}
260
+ owner = object()
219
261
 
220
262
  transport = ASGITransport(app=app)
221
263
  try:
264
+ if overrides:
265
+ _push_overrides(app, overrides, owner)
222
266
  async with AsyncClient(
223
267
  transport=transport, base_url=base_url, **kwargs
224
268
  ) as client:
225
269
  yield client
226
270
  finally:
227
- if dependency_overrides:
228
- for key in dependency_overrides:
229
- app.dependency_overrides.pop(key, None)
271
+ if overrides:
272
+ _pop_overrides(app, overrides, owner)
230
273
 
231
274
 
232
275
  @asynccontextmanager