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
26 changes: 26 additions & 0 deletions tests/providers/test_context_resources.py
Original file line number Diff line number Diff line change
Expand Up @@ -1067,6 +1067,32 @@ class _Container(BaseContainer):
assert await _Container.p_app.resolve() is not None


async def test_context_resource_context_async_cleans_up_after_exception() -> None:
events: list[str] = []

async def create_resource() -> typing.AsyncIterator[int]:
events.append("enter")
try:
yield 1
finally:
events.append("exit")

resource = providers.ContextResource(create_resource)

async def raise_in_resource_context() -> None:
async with resource.context_async():
assert await resource.resolve() == 1
msg = "expected"
raise RuntimeError(msg)

with pytest.raises(RuntimeError, match="expected"):
await raise_in_resource_context()

assert events == ["enter", "exit"]
with pytest.raises(RuntimeError, match="Context is not set"):
await resource.resolve()


def test_sync_force_enter_context_for_scoped_resource() -> None:
class _Container(BaseContainer):
p_app = providers.ContextResource(create_sync_context_resource).with_config(scope=ContextScopes.APP)
Expand Down
32 changes: 32 additions & 0 deletions tests/providers/test_state.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,38 @@ class _Container(BaseContainer):
assert _Container.state.resolve_sync() == state_value


def test_state_restores_context_after_exception() -> None:
state = State[str]()

def raise_in_inner_context() -> None:
with state.init("inner"):
msg = "expected"
raise RuntimeError(msg)

with state.init("outer"):
with pytest.raises(RuntimeError, match="expected"):
raise_in_inner_context()
assert state.resolve_sync() == "outer"

with pytest.raises(StateNotInitializedError):
state.resolve_sync()


async def test_state_restores_context_after_async_exception() -> None:
state = State[str]()

async def raise_in_context() -> None:
with state.init("value"):
msg = "expected"
raise RuntimeError(msg)

with pytest.raises(RuntimeError, match="expected"):
await raise_in_context()

with pytest.raises(StateNotInitializedError):
await state.resolve()


async def test_state_correctly_manages_its_context() -> None:
async def _async_creator(x: int) -> int:
await asyncio.sleep(random.random())
Expand Down
14 changes: 9 additions & 5 deletions that_depends/providers/context_resources.py
Original file line number Diff line number Diff line change
Expand Up @@ -480,11 +480,15 @@ async def context_async(self, force: bool = False) -> typing.AsyncIterator[Resou
async with self._async_lock:
val = await self._enter_context_async(force=force)
temp_token = self._token
yield val
async with self._async_lock:
self._token = temp_token
await self._exit_context_async()
self._token = token
try:
yield val
finally:
async with self._async_lock:
self._token = temp_token
try:
await self._exit_context_async()
finally:
self._token = token

def _fetch_context(self) -> ResourceContext[T_co]:
try:
Expand Down
6 changes: 4 additions & 2 deletions that_depends/providers/state.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,8 +40,10 @@ def init(self, state: T) -> typing.Iterator[T]:

"""
token = self._state.set(state)
yield state
self._state.reset(token)
try:
yield state
finally:
self._state.reset(token)

@override
async def resolve(self) -> T:
Expand Down
Loading