Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
43 changes: 43 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -407,6 +407,49 @@ run_coroutine_sync(cleanup_all_exit_stacks()) # accepts raise_exception=True too
assert machine.db.closed is True
```

#### Unwinding with an in-flight exception (rollback on error)

In a FastAPI request, when your endpoint raises, generator dependencies receive the exception at their `yield` — so `except` branches (e.g. `rollback()`) run before `finally`. Outside a request you get the same behavior by passing the caught exception to the cleanup helpers via `exc=`, or by using `injectable_scope()`, which forwards it automatically:

```python
from collections.abc import AsyncGenerator
from typing import Annotated

from fastapi import Depends
from fastapi_injectable import cleanup_all_exit_stacks, injectable, injectable_scope

async def get_connection() -> AsyncGenerator[Connection, None]:
conn = Connection()
try:
yield conn
except Exception:
conn.rollback() # runs when the decorated function raised
raise
else:
conn.commit()
finally:
conn.close()

@injectable
async def process(conn: Annotated[Connection, Depends(get_connection)]) -> None:
raise ValueError("boom")

# Option #1: pass the exception to the cleanup helpers explicitly
try:
await process()
except ValueError as exc:
# conn.rollback() and conn.close() run; without exc= only conn.commit()/close() would
await cleanup_all_exit_stacks(exc=exc)
# cleanup_exit_stack_of_func(process, exc=exc) also accepts it

# Option #2: injectable_scope() forwards the exception automatically,
# exactly like a failing FastAPI request unwinding its exit stack
async with injectable_scope():
await process() # raises -> conn.rollback() + conn.close() run during unwind
```

Just like in FastAPI, teardown that must **always** run belongs in `finally`: when an exception is delivered to the generator, code placed after a bare `yield` (with no `try`/`finally`) is skipped.

### Async Support

`fastapi-injectable` provides full support for both synchronous and asynchronous dependencies, allowing you to mix and match them as needed. You can freely use async dependencies in sync functions and vice versa. For cases where you need to run async code in a synchronous context, we provide the `run_coroutine_sync` utility function.
Expand Down
4 changes: 2 additions & 2 deletions noxfile.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,12 +12,12 @@
python_versions_without_free_threaded = [version for version in python_versions if not version.endswith("t")]
latest_python_version = python_versions_without_free_threaded[-1]
nox.needs_version = ">= 2025.05.01"
nox.options.sessions = (
nox.options.sessions = [
"pre-commit",
"mypy",
"tests",
"docs-build",
)
]
nox.options.default_venv_backend = "uv"


Expand Down
51 changes: 44 additions & 7 deletions src/fastapi_injectable/async_exit_stack.py
Original file line number Diff line number Diff line change
Expand Up @@ -80,7 +80,30 @@ async def get_stack(self, func: Callable[..., Any]) -> AsyncExitStack:
per_loop[func] = stack
return stack

async def _close_stack(self, loop: _Loop, stack: AsyncExitStack) -> None:
@staticmethod
async def _unwind_stack(stack: AsyncExitStack, exc: BaseException | None) -> None:
"""Unwind one stack, mirroring how a context manager would exit.

With ``exc`` the stack exits as if the exception propagated out of an
``async with`` block: every teardown (generator ``except``/rollback
branches included) sees the original exception. Without it the stack is
closed normally (``aclose()``, i.e. ``__aexit__(None, None, None)``).
"""
if exc is None:
await stack.aclose()
return
try:
await stack.__aexit__(type(exc), exc, exc.__traceback__)
except BaseException as unwind_exc:
# A generator dependency that re-raises the in-flight exception
# ("except: rollback(); raise") follows the context-manager protocol
# (exception not suppressed) -- AsyncExitStack.__aexit__ then re-raises
# it here. The caller already caught and handled ``exc``, so only
# genuine teardown failures (a *different* exception) propagate.
if unwind_exc is not exc:
raise

async def _close_stack(self, loop: _Loop, stack: AsyncExitStack, exc: BaseException | None = None) -> None:
"""Close one stack on its owning loop.

- owning loop is the current running loop -> close it here directly;
Expand All @@ -92,19 +115,29 @@ async def _close_stack(self, loop: _Loop, stack: AsyncExitStack) -> None:
"""
running = self._running_loop()
if loop is running:
await stack.aclose()
await self._unwind_stack(stack, exc)
return
if loop.is_closed() or not loop.is_running():
return
future = asyncio.run_coroutine_threadsafe(stack.aclose(), loop)
future = asyncio.run_coroutine_threadsafe(self._unwind_stack(stack, exc), loop)
await asyncio.wrap_future(future)

async def cleanup_stack(self, func: Callable[..., Any], *, raise_exception: bool = False) -> None:
async def cleanup_stack(
self,
func: Callable[..., Any],
*,
raise_exception: bool = False,
exc: BaseException | None = None,
) -> None:
"""Clean up the stack(s) associated with the given function.

Args:
func: The function whose exit stack should be cleaned up
raise_exception: If True, raises DependencyCleanupError when cleanup fails
exc: The in-flight exception, if any. When provided, each stack is unwound
with the exception details (``__aexit__(type(exc), exc, exc.__traceback__)``)
exactly as if it had propagated out of an ``async with`` block, so
generator dependencies can run their ``except``/rollback branches.

Raises:
DependencyCleanupError: When cleanup fails and raise_exception is True
Expand All @@ -125,7 +158,7 @@ async def cleanup_stack(self, func: Callable[..., Any], *, raise_exception: bool
exception_: Exception | None = None
msg = ""
try:
await asyncio.gather(*(self._close_stack(loop, stack) for loop, stack in entries))
await asyncio.gather(*(self._close_stack(loop, stack, exc) for loop, stack in entries))
except RuntimeError as e:
msg = f"Failed to cleanup stack for {func.__name__} during teardown: {e}"
logger.warning(msg)
Expand All @@ -138,11 +171,15 @@ async def cleanup_stack(self, func: Callable[..., Any], *, raise_exception: bool
if exception_ is not None and raise_exception:
raise DependencyCleanupError(msg) from exception_

async def cleanup_all_stacks(self, *, raise_exception: bool = False) -> None:
async def cleanup_all_stacks(self, *, raise_exception: bool = False, exc: BaseException | None = None) -> None:
"""Clean up all stacks, each on its own owning loop.

Args:
raise_exception: If True, raises DependencyCleanupError when any cleanup fails
exc: The in-flight exception, if any. When provided, every stack is unwound
with the exception details (``__aexit__(type(exc), exc, exc.__traceback__)``)
exactly as if it had propagated out of an ``async with`` block, so
generator dependencies can run their ``except``/rollback branches.

Raises:
DependencyCleanupError: When any cleanup fails and raise_exception is True
Expand All @@ -157,7 +194,7 @@ async def cleanup_all_stacks(self, *, raise_exception: bool = False) -> None:
exception_: Exception | None = None
msg = ""
try:
await asyncio.gather(*(self._close_stack(loop, stack) for loop, stack in entries))
await asyncio.gather(*(self._close_stack(loop, stack, exc) for loop, stack in entries))
except RuntimeError as e:
msg = f"Failed to cleanup one or more dependency stacks during teardown: {e}"
logger.warning(msg)
Expand Down
18 changes: 16 additions & 2 deletions src/fastapi_injectable/concurrency.py
Original file line number Diff line number Diff line change
Expand Up @@ -83,13 +83,27 @@ def _get_or_create_current_loop(self) -> asyncio.AbstractEventLoop:
Compatible with Python 3.12+ and 3.14+. Attempts to get loop via policy,
falls back to creating new loop if RuntimeError is raised.

A newly created loop is registered as the thread's current loop
(``set_event_loop``) so that consecutive synchronous calls keep landing on
the SAME loop. Without this, each call would run on a fresh throwaway loop:
dependencies would be resolved on loop A while their cleanup ran on loop B,
and the per-loop exit-stack registry (see ``AsyncExitStackManager``) would
rightly refuse to close A's stacks from B -- silently skipping generator
teardown. This mirrors the pre-3.12 auto-create semantics of
``asyncio.get_event_loop()``.

Returns:
Event loop instance.
"""
policy = asyncio.get_event_loop_policy()
try:
return asyncio.get_event_loop_policy().get_event_loop()
loop = policy.get_event_loop()
except RuntimeError:
return asyncio.get_event_loop_policy().new_event_loop()
loop = None
if loop is None or loop.is_closed():
loop = policy.new_event_loop()
asyncio.set_event_loop(loop)
return loop

@property
def loop_strategy(self) -> Literal["current", "isolated", "background_thread"]:
Expand Down
10 changes: 8 additions & 2 deletions src/fastapi_injectable/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -170,8 +170,14 @@ async def resolve_dependencies(
fastapi_inner_astack = AsyncExitStack()
fastapi_function_astack = AsyncExitStack()

async_exit_stack.push_async_callback(fastapi_inner_astack.aclose)
async_exit_stack.push_async_callback(fastapi_function_astack.aclose)
# Registered via push_async_exit (NOT push_async_callback(stack.aclose)) so that
# when the owning stack unwinds with an exception, the exception details reach
# these inner stacks. On fastapi>=0.121 generator dependencies are entered into
# them (see the fake_request_scope note below), and an ``aclose()`` callback
# would strip the exception -- generators would see a bare GeneratorExit and
# their ``except``/rollback branches could never run (issue #255).
async_exit_stack.push_async_exit(fastapi_inner_astack)
async_exit_stack.push_async_exit(fastapi_function_astack)

fake_request_scope: dict[str, Any] = {
"type": "http",
Expand Down
6 changes: 6 additions & 0 deletions src/fastapi_injectable/scope.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,12 @@ class InjectableScope:
cached values into this scope's cache. Leaving the scope closes the exit
stack (running all cleanup) and drops the cache.

When an exception propagates out of the ``async with`` block, it is forwarded
to the exit stack -- generator dependencies receive it at their ``yield`` and
can run their ``except``/rollback branches, exactly as FastAPI unwinds a
failing request (issue #255). As in FastAPI, teardown that must always run
belongs in a ``finally`` block.

The cache is partitioned by event loop. A scope object reused across event
loops -- resolved on loop A then on loop B while both are alive -- must never
serve a loop-A resource to loop B: a cached value commonly holds a loop-bound
Expand Down
23 changes: 18 additions & 5 deletions src/fastapi_injectable/util.py
Original file line number Diff line number Diff line change
Expand Up @@ -142,7 +142,7 @@ async def inner(dep: T2) -> T2:
# If it's an async generator, get the first value
if inspect.isasyncgen(dep):
async for value in dep: # pragma: no cover
return value # type: ignore[no-any-return]
return cast("T2", value)
return dep

# Nice signature for docs/inspection
Expand Down Expand Up @@ -446,13 +446,22 @@ async def get_db() -> AsyncGenerator[Database, None]:
return await coro


async def cleanup_exit_stack_of_func(func: Callable[..., Any], *, raise_exception: bool = False) -> None:
async def cleanup_exit_stack_of_func(
func: Callable[..., Any],
*,
raise_exception: bool = False,
exc: BaseException | None = None,
) -> None:
"""Clean up the exit stack associated with a specific function.

Args:
func: The function whose exit stack should be cleaned up.
raise_exception: Whether to raise exceptions during cleanup.
If False, exceptions are logged as warnings. Defaults to False.
exc: The in-flight exception, if any. When provided, the exit stack is unwound
with the exception details -- exactly as FastAPI does when a request fails --
so generator dependencies see the original exception and can run their
``except``/rollback branches instead of a plain close.

Notes:
- This ensures that resources such as context managers or other async cleanup routines
Expand All @@ -462,15 +471,19 @@ async def cleanup_exit_stack_of_func(func: Callable[..., Any], *, raise_exceptio
DependencyCleanupError: When cleanup fails and raise_exception is True
"""
for wrapper in PROVIDER_TO_WRAPPER_FUNC_MAP.get(func, [func]):
await async_exit_stack_manager.cleanup_stack(wrapper, raise_exception=raise_exception)
await async_exit_stack_manager.cleanup_stack(wrapper, raise_exception=raise_exception, exc=exc)


async def cleanup_all_exit_stacks(*, raise_exception: bool = False) -> None:
async def cleanup_all_exit_stacks(*, raise_exception: bool = False, exc: BaseException | None = None) -> None:
"""Clean up all active exit stacks.

Args:
raise_exception: Whether to raise exceptions during cleanup.
If False, exceptions are logged as warnings. Defaults to False.
exc: The in-flight exception, if any. When provided, every exit stack is unwound
with the exception details -- exactly as FastAPI does when a request fails --
so generator dependencies see the original exception and can run their
``except``/rollback branches instead of a plain close.

Notes:
- This method iterates through all registered exit stacks and ensures they are properly closed.
Expand All @@ -479,7 +492,7 @@ async def cleanup_all_exit_stacks(*, raise_exception: bool = False) -> None:
Raises:
DependencyCleanupError: When cleanup fails and raise_exception is True
"""
await async_exit_stack_manager.cleanup_all_stacks(raise_exception=raise_exception)
await async_exit_stack_manager.cleanup_all_stacks(raise_exception=raise_exception, exc=exc)


async def clear_dependency_cache() -> None:
Expand Down
78 changes: 78 additions & 0 deletions test/test_async_exit_stack.py
Original file line number Diff line number Diff line change
Expand Up @@ -186,6 +186,84 @@ async def test_cleanup_all_stacks_with_empty_stacks(manager: AsyncExitStackManag
assert len(manager._stacks) == 0


def _make_exc() -> ValueError:
try:
msg = "boom"
raise ValueError(msg) # noqa: TRY301
except ValueError as e:
return e


async def test_cleanup_stack_with_exc_unwinds_with_exception_details(
manager: AsyncExitStackManager, mock_func: Mock, mock_stack: AsyncMock
) -> None:
"""With ``exc``, the stack exits via __aexit__(type, exc, tb) instead of aclose()."""
_register(manager, mock_func, mock_stack)
exc = _make_exc()

await manager.cleanup_stack(mock_func, exc=exc)

mock_stack.__aexit__.assert_awaited_once_with(ValueError, exc, exc.__traceback__)
mock_stack.aclose.assert_not_awaited()
assert not _contains(manager, mock_func)


async def test_cleanup_all_stacks_with_exc_unwinds_with_exception_details(
manager: AsyncExitStackManager, mock_func: Mock, mock_stack: AsyncMock
) -> None:
"""cleanup_all_stacks(exc=...) unwinds every stack with the exception details."""
other_func = Mock()
other_func.__name__ = "other_func"
other_stack = AsyncMock(spec=AsyncExitStack)
_register(manager, mock_func, mock_stack)
_register(manager, other_func, other_stack)
exc = _make_exc()

await manager.cleanup_all_stacks(exc=exc)

mock_stack.__aexit__.assert_awaited_once_with(ValueError, exc, exc.__traceback__)
other_stack.__aexit__.assert_awaited_once_with(ValueError, exc, exc.__traceback__)
mock_stack.aclose.assert_not_awaited()
other_stack.aclose.assert_not_awaited()
assert len(manager._stacks) == 0


async def test_cleanup_with_exc_swallows_reraised_in_flight_exception(
manager: AsyncExitStackManager, mock_func: Mock, mock_stack: AsyncMock
) -> None:
"""A dependency re-raising the in-flight exception is CM protocol, not a cleanup failure.

AsyncExitStack.__aexit__ re-raises the passed-in exception when a generator's
``except: ...; raise`` branch doesn't suppress it. The caller already handled that
exception, so cleanup must not report it as a DependencyCleanupError.
"""
exc = _make_exc()
tb = exc.__traceback__ # re-raising below appends frames to exc.__traceback__
mock_stack.__aexit__.side_effect = exc
_register(manager, mock_func, mock_stack)

await manager.cleanup_stack(mock_func, raise_exception=True, exc=exc)

mock_stack.__aexit__.assert_awaited_once_with(ValueError, exc, tb)
assert not _contains(manager, mock_func)


async def test_cleanup_with_exc_still_raises_on_distinct_teardown_failure(
manager: AsyncExitStackManager, mock_func: Mock, mock_stack: AsyncMock
) -> None:
"""A teardown failure distinct from the in-flight exception still surfaces."""
exc = _make_exc()
mock_stack.__aexit__.side_effect = RuntimeError("teardown boom")
_register(manager, mock_func, mock_stack)

with pytest.raises(Exception, match="Failed to cleanup stack for mock_func") as exc_info:
await manager.cleanup_stack(mock_func, raise_exception=True, exc=exc)

assert isinstance(exc_info.value, DependencyCleanupError)
assert isinstance(exc_info.value.__cause__, RuntimeError)
assert not _contains(manager, mock_func)


async def test_cleanup_stack_runtime_error_names_func_not_loop(
manager: AsyncExitStackManager, mock_func: Mock, mock_stack: AsyncMock
) -> None:
Expand Down
5 changes: 5 additions & 0 deletions test/test_concurrency.py
Original file line number Diff line number Diff line change
Expand Up @@ -58,6 +58,7 @@ def test_get_loop_current_strategy(mock_get_running_loop: Mock, loop_manager_ins
def test_get_loop_current_strategy_fallback(loop_manager_instance: LoopManager) -> None:
"""Test get_loop with 'current' strategy when no running loop exists (Python 3.14+ compatibility)."""
mock_loop = Mock()
mock_loop.is_closed.return_value = False
mock_policy = Mock()
mock_policy.get_event_loop.return_value = mock_loop

Expand All @@ -83,12 +84,16 @@ def test_get_loop_current_strategy_fallback_create_new(loop_manager_instance: Lo
with (
patch("src.fastapi_injectable.concurrency.asyncio.get_running_loop", side_effect=RuntimeError),
patch("src.fastapi_injectable.concurrency.asyncio.get_event_loop_policy", return_value=mock_policy),
patch("src.fastapi_injectable.concurrency.asyncio.set_event_loop") as mock_set_event_loop,
):
loop_manager_instance.set_loop_strategy("current")

result = loop_manager_instance.get_loop()

assert result == mock_loop
# The freshly created loop is registered as the thread's current loop so that
# subsequent synchronous calls (resolution + cleanup) land on the SAME loop.
mock_set_event_loop.assert_called_once_with(mock_loop)
mock_policy.get_event_loop.assert_called_once()
mock_policy.new_event_loop.assert_called_once()

Expand Down
Loading
Loading