diff --git a/AGENTS.md b/AGENTS.md index ce6d2d0..1ac49e2 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -178,6 +178,7 @@ foundation: - asynchronous factories - generator factories - asynchronous generator factories +- context manager and asynchronous context manager factories - nested dependencies - callable objects - functools.partial diff --git a/docs/adr/0020-context-manager-factories-normalized-at-entry.md b/docs/adr/0020-context-manager-factories-normalized-at-entry.md new file mode 100644 index 0000000..69f85c7 --- /dev/null +++ b/docs/adr/0020-context-manager-factories-normalized-at-entry.md @@ -0,0 +1,52 @@ +# ADR 0020: Context manager factories are normalized at entry + +- Status: Accepted +- Date: 2026-08-01 + +## Context + +`@contextmanager` and `@asynccontextmanager` are the ordinary way to write a +resource in Python, and most real dependencies (database sessions, clients, +transactions) already exist in that form. FastDepends understands generator +and async-generator factories, but a decorated factory returns a context +manager object instead of yielding the dependency, so it was injected +unentered and never cleaned up. + +`inspect.unwrap()` resolves the decorated form, but it follows every +`__wrapped__` chain, so it also strips unrelated decorators applied with +`functools.wraps`, including `functools.lru_cache`. Applying it only in +`wired()` also splits identity: the dependency registers under the wrapped +generator function while `override_dependency()` and +`override_web_dependency()` still key on the decorator helper, so overrides +silently do nothing. + +## Decision + +One private `_normalize_factory()` recognizes the two `contextlib` +decorators by the code object their helper closures share, and returns the +generator function the helper wraps. Anything else is returned untouched. + +Normalization runs wherever a factory enters Wireme: `wired()`, +`override_dependency()` (both factories), and the bridged-adapter lookup in +`get_override_pairs()`. Declaration and override sites therefore agree on +one identity per dependency. In `get_override_pairs()` the direct FastAPI +pair keeps the callables as given, because a plain FastAPI dependency is +registered under the object passed to `Depends()`. + +`wired()` gains overloads for `AbstractContextManager[R]` and +`AbstractAsyncContextManager[R]` so a decorated factory infers `R`, matching +the generator overloads. + +## Consequences + +- Positive: context managers behave exactly like the generator functions + they wrap, including cleanup order, caching, FastAPI request lifecycle, + and both override entry points. +- Positive: unrelated `functools.wraps` decorators and caches keep working, + which unconditional unwrapping broke. +- Negative: detection depends on the closure shape of `contextlib`'s two + decorators. This is stdlib behavior stable across supported Python + versions and is regression tested; a change there degrades to injecting + the unentered manager rather than misfiring on other callables. +- Neutral: third-party context manager decorators are not recognized. Pass + the underlying generator function, or wrap it in one. diff --git a/docs/adr/README.md b/docs/adr/README.md index 25478c7..77ea105 100644 --- a/docs/adr/README.md +++ b/docs/adr/README.md @@ -18,3 +18,5 @@ defaults, member selection, and the documentation site). Decision 0018 was recorded on 2026-07-17 to establish a strict DI-only boundary. Decision 0019 was recorded on 2026-07-18 to establish one versioned history, a cohesive release tooling boundary, and immutable assets as the PyPI handoff. +Decision 0020 was recorded on 2026-08-01 to accept context manager +factories as dependencies through one normalization point. diff --git a/examples/README.md b/examples/README.md index a90deaf..d5f703d 100644 --- a/examples/README.md +++ b/examples/README.md @@ -18,6 +18,7 @@ uv run python examples/basic.py | Class, instance, and method factories | `factories.py` | | Process-wide singletons | `singletons.py` | | Generator and async resource cleanup | `resources.py` | +| Context manager factories as dependencies | `context_managers.py` | | Side-effect dependencies (`requires`) with injected context | `requires.py` | | Wiring many methods with an apply combinator | `method_wiring.py` | | Test overrides | `overrides.py` | @@ -28,6 +29,7 @@ uv run python examples/basic.py | FastAPI request-scoped resources | `fastapi_resources.py` | | FastAPI nested-safe web overrides | `fastapi_overrides.py` | | FastAPI endpoints wired directly | `fastapi_endpoints.py` | +| FastAPI context manager dependencies | `fastapi_context_managers.py` | All examples run in CI. When a public capability is added, add or extend an example and list it here. diff --git a/examples/context_managers.py b/examples/context_managers.py new file mode 100644 index 0000000..4010552 --- /dev/null +++ b/examples/context_managers.py @@ -0,0 +1,89 @@ +"""Context manager factories as dependencies. + +A factory decorated with @contextmanager or @asynccontextmanager behaves +exactly like the generator function it wraps: it is entered before the +wired call and closed afterwards, in reverse order. This is what lets an +existing context manager, such as a database session, be reused as a +dependency without rewriting it as a bare generator. +""" + +from __future__ import annotations + +import asyncio +import sqlite3 +from collections.abc import AsyncGenerator, Generator +from contextlib import asynccontextmanager, contextmanager +from typing import Annotated + +from wireme import Wired, wire, wired + +events: list[str] = [] + + +@contextmanager +def get_connection() -> Generator[sqlite3.Connection]: + events.append("open connection") + connection = sqlite3.connect(":memory:") + try: + yield connection + finally: + connection.close() + events.append("close connection") + + +type ConnectionDep = Annotated[ + sqlite3.Connection, + wired(get_connection), +] + + +@wire +def count_rows(*, connection: ConnectionDep = Wired()) -> int: + connection.execute("create table hero (name text)") + connection.executemany( + "insert into hero values (?)", + [("Deadpond",), ("Spider-Boy",)], + ) + + row = connection.execute("select count(*) from hero").fetchone() + + return int(row[0]) + + +class Client: + async def fetch(self, path: str) -> str: + return f"response from {path}" + + +@asynccontextmanager +async def get_client() -> AsyncGenerator[Client]: + events.append("open client") + try: + yield Client() + finally: + events.append("close client") + + +type ClientDep = Annotated[Client, wired(get_client)] + + +@wire +async def fetch(path: str, *, client: ClientDep = Wired()) -> str: + return await client.fetch(path) + + +async def main() -> None: + assert count_rows() == 2 + assert await fetch("/heroes") == "response from /heroes" + assert events == [ + "open connection", + "close connection", + "open client", + "close client", + ] + + print("\n".join(events)) + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/examples/fastapi_context_managers.py b/examples/fastapi_context_managers.py new file mode 100644 index 0000000..3a01a6b --- /dev/null +++ b/examples/fastapi_context_managers.py @@ -0,0 +1,77 @@ +"""Context manager factories bridged into FastAPI with FromWeb. + +The dependency is declared once with wired(...) and reused in endpoints +through FromWeb. FastAPI owns the request lifecycle: the context manager is +entered when the request needs it and exited after the response finishes. +Tests replace it with override_web_dependency like any other factory. +""" + +from __future__ import annotations + +import sqlite3 +from collections.abc import Generator +from contextlib import contextmanager +from typing import Annotated + +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from wireme import wired +from wireme.fastapi import FromWeb, override_web_dependency + +events: list[str] = [] + + +@contextmanager +def get_connection() -> Generator[sqlite3.Connection]: + events.append("open connection") + connection = sqlite3.connect(":memory:") + connection.execute("create table hero (name text)") + connection.execute("insert into hero values ('Deadpond')") + try: + yield connection + finally: + connection.close() + events.append("close connection") + + +type ConnectionDep = Annotated[ + sqlite3.Connection, + wired(get_connection), +] + + +app = FastAPI() + + +@app.get("/heroes") +def list_heroes(*, connection: FromWeb[ConnectionDep]) -> list[str]: + events.append("handle request") + return [name for (name,) in connection.execute("select name from hero")] + + +client = TestClient(app) + +assert client.get("/heroes").json() == ["Deadpond"] +assert events == ["open connection", "handle request", "close connection"] + +print("\n".join(events)) + + +@contextmanager +def get_test_connection() -> Generator[sqlite3.Connection]: + connection = sqlite3.connect(":memory:") + connection.execute("create table hero (name text)") + connection.execute("insert into hero values ('Spider-Boy')") + try: + yield connection + finally: + connection.close() + + +with override_web_dependency(app, get_connection, get_test_connection): + assert client.get("/heroes").json() == ["Spider-Boy"] + +assert client.get("/heroes").json() == ["Deadpond"] + +print("override restored") diff --git a/src/wireme/_impl.py b/src/wireme/_impl.py index 6436d7f..b735115 100644 --- a/src/wireme/_impl.py +++ b/src/wireme/_impl.py @@ -4,8 +4,10 @@ import contextlib import inspect +import types import typing from collections.abc import ( + AsyncGenerator, AsyncIterator, Awaitable, Callable, @@ -13,6 +15,7 @@ Iterator, Sequence, ) +from contextlib import AbstractAsyncContextManager, AbstractContextManager from ._core import ( _build_call_model, @@ -31,10 +34,17 @@ class _HasSignature(typing.Protocol): __signature__: inspect.Signature +class _HasWrapped(typing.Protocol): + """Represent a decorator helper exposing the callable it wraps.""" + + __wrapped__: Callable[..., object] + + __all__ = ( "Wired", "_HasSignature", "_factory_model", + "_normalize_factory", "_wire", "override_dependency", "wire", @@ -44,15 +54,50 @@ class _HasSignature(typing.Protocol): _MISSING = object() _provider = _DiProvider() +# Both decorators return a helper closure, and every helper produced by one +# decorator shares that decorator's code object. The lambdas below are only +# wrapped, never called, so their bodies are irrelevant. +_CONTEXT_MANAGER_CODES: typing.Final[frozenset[types.CodeType]] = frozenset( + { + contextlib.contextmanager( + typing.cast("Callable[[], Generator[None]]", lambda: None) + ).__code__, + contextlib.asynccontextmanager( + typing.cast("Callable[[], AsyncGenerator[None]]", lambda: None) + ).__code__, + } +) + type _DependencyFactory[R] = ( Callable[..., Awaitable[R]] | Callable[..., AsyncIterator[R]] | Callable[..., Iterator[R]] + | Callable[..., AbstractContextManager[R]] + | Callable[..., AbstractAsyncContextManager[R]] | Callable[..., R] ) +def _normalize_factory[F](factory: F, /) -> F: + """Return the generator function behind a context manager decorator. + + ``@contextmanager`` and ``@asynccontextmanager`` return a helper that + builds a context manager object rather than yielding the dependency, so + FastDepends would inject the unentered manager. Both decorators produce + helpers sharing one code object, which identifies them precisely without + unwrapping unrelated ``functools.wraps`` decorators such as caches or + instrumentation. + + Normalization is applied wherever a factory enters Wireme so declaration + and override sites agree on one identity for the same dependency. + """ + if getattr(factory, "__code__", None) in _CONTEXT_MANAGER_CODES: + return typing.cast("F", typing.cast("_HasWrapped", factory).__wrapped__) + + return factory + + def Wired() -> typing.Any: """Mark an annotated dependency as optional for static type checkers.""" return ... @@ -532,6 +577,24 @@ def wired[**P, R]( ) -> R: ... +@typing.overload +def wired[**P, R]( + factory: Callable[P, AbstractAsyncContextManager[R]], + /, + *, + use_cache: bool = True, +) -> R: ... + + +@typing.overload +def wired[**P, R]( + factory: Callable[P, AbstractContextManager[R]], + /, + *, + use_cache: bool = True, +) -> R: ... + + @typing.overload def wired[**P, R]( factory: Callable[P, R], @@ -553,9 +616,14 @@ def wired( PEP 695 aliases and postponed annotations in the factory's own parameters work at any nesting depth. + Factories decorated with ``@contextmanager`` or ``@asynccontextmanager`` + are resolved through the generator function they wrap, so they behave + exactly like the equivalent generator factory. + Raises: TypeError: If the factory uses a FastDepends CustomField marker. """ + factory = _normalize_factory(factory) _resolve_factory_signature(factory, localns=_caller_locals()) return _Depends( @@ -578,10 +646,15 @@ def override_dependency[R]( so use them for isolated tests and application setup, not concurrent request-level mutation. + Either factory may be a context manager decorated with + ``@contextmanager`` or ``@asynccontextmanager``. + Raises: TypeError: If either factory uses a FastDepends CustomField marker. """ localns = _caller_locals() + original = _normalize_factory(original) + replacement = _normalize_factory(replacement) _resolve_factory_signature(original, localns=localns) _resolve_factory_signature(replacement, localns=localns) diff --git a/src/wireme/fastapi/_dependencies.py b/src/wireme/fastapi/_dependencies.py index 13f755b..ab1d4c6 100644 --- a/src/wireme/fastapi/_dependencies.py +++ b/src/wireme/fastapi/_dependencies.py @@ -9,7 +9,7 @@ from typing import TYPE_CHECKING, Annotated, Any from wireme._core import _CallModel, _Dependant -from wireme._impl import _factory_model, _HasSignature +from wireme._impl import _factory_model, _HasSignature, _normalize_factory from ._compat import Depends @@ -132,14 +132,23 @@ def get_override_pairs( replacement: _Factory, /, ) -> tuple[tuple[_Factory, _Factory], ...]: - """Return direct and bridged FastAPI override pairs.""" + """Return direct and bridged FastAPI override pairs. + + The direct pair keeps the callables as given, because a plain FastAPI + dependency is registered under the object passed to Depends(). Bridged + adapters are looked up under the normalized factory, matching the + identity wired(...) registered for a context manager factory. + """ pairs: list[tuple[_Factory, _Factory]] = [ (original, replacement), ] - for use_cache, original_adapter in _bridges.get(original, {}).items(): + normalized_original = _normalize_factory(original) + normalized_replacement = _normalize_factory(replacement) + + for use_cache, original_adapter in _bridges.get(normalized_original, {}).items(): replacement_adapter = _bridge_factory( - replacement, + normalized_replacement, use_cache, ) pairs.append((original_adapter, replacement_adapter)) diff --git a/src/wireme/fastapi/_overrides.py b/src/wireme/fastapi/_overrides.py index bdec275..d61660c 100644 --- a/src/wireme/fastapi/_overrides.py +++ b/src/wireme/fastapi/_overrides.py @@ -5,6 +5,7 @@ import contextlib import typing from collections.abc import AsyncIterator, Awaitable, Callable, Generator, Iterator +from contextlib import AbstractAsyncContextManager, AbstractContextManager from wireme.fastapi._dependencies import get_override_pairs @@ -14,6 +15,8 @@ Callable[..., Awaitable[T]] | Callable[..., AsyncIterator[T]] | Callable[..., Iterator[T]] + | Callable[..., AbstractContextManager[T]] + | Callable[..., AbstractAsyncContextManager[T]] | Callable[..., T] ) @@ -35,8 +38,8 @@ def override_web_dependency[T]( Both direct FastAPI dependencies and adapters created by FromWeb[WiredAlias] are overridden. Nested contexts restore the previous replacement correctly, including after exceptions. Replacements may be - sync, async, generator, or async-generator factories with different - parameter lists. + sync, async, generator, async-generator, or context manager factories + with different parameter lists. Bridged adapters are discovered when the context is entered, so routes using FromWeb must be registered before entering the override context. diff --git a/tests/integration/fastapi/test_lifecycle.py b/tests/integration/fastapi/test_lifecycle.py index 86db368..39fd42b 100644 --- a/tests/integration/fastapi/test_lifecycle.py +++ b/tests/integration/fastapi/test_lifecycle.py @@ -1,5 +1,12 @@ +import contextlib import functools -from collections.abc import AsyncIterator, Callable, Iterator +from collections.abc import ( + AsyncGenerator, + AsyncIterator, + Callable, + Generator, + Iterator, +) from typing import Annotated, Any import pytest @@ -344,3 +351,91 @@ def endpoint( response = client.get("/") assert response.json() == {"first": 1, "second": 2} + + +@contextlib.contextmanager +def get_context_manager_connection() -> Generator[Connection]: + events.append("open") + try: + yield Connection("context-manager") + finally: + events.append("close") + + +type ContextManagerConnectionDep = Annotated[ + Connection, + wired(get_context_manager_connection), +] + + +@contextlib.asynccontextmanager +async def get_async_context_manager_connection() -> AsyncGenerator[Connection]: + events.append("open") + try: + yield Connection("async-context-manager") + finally: + events.append("close") + + +type AsyncContextManagerConnectionDep = Annotated[ + Connection, + wired(get_async_context_manager_connection), +] + + +async def get_context_manager_session( + *, + connection: ContextManagerConnectionDep = Wired(), +) -> AsyncIterator[Session]: + events.append("open session") + try: + yield Session(connection) + finally: + events.append("close session") + + +type ContextManagerSessionDep = Annotated[ + Session, + wired(get_context_manager_session), +] + + +def test_context_manager_factory_cleans_up_after_response() -> None: + def endpoint(connection: FromWeb[ContextManagerConnectionDep]) -> dict[str, str]: + events.append("use") + return {"name": connection.name} + + client = _create_app(endpoint) + + response = client.get("/") + + assert response.json() == {"name": "context-manager"} + assert events == ["open", "use", "close"] + + +def test_async_context_manager_factory_cleans_up_after_response() -> None: + def endpoint( + connection: FromWeb[AsyncContextManagerConnectionDep], + ) -> dict[str, str]: + events.append("use") + return {"name": connection.name} + + client = _create_app(endpoint) + + response = client.get("/") + + assert response.json() == {"name": "async-context-manager"} + assert events == ["open", "use", "close"] + + +def test_context_manager_nested_in_generator_closes_in_reverse_order() -> None: + def endpoint(session: FromWeb[ContextManagerSessionDep]) -> dict[str, str]: + events.append("use") + return {"name": session.connection.name} + + client = _create_app(endpoint) + + response = client.get("/") + + assert response.json() == {"name": "context-manager"} + assert events == ["open", "open session", "use", "close session", "close"] diff --git a/tests/integration/fastapi/test_overrides.py b/tests/integration/fastapi/test_overrides.py index 7861a1b..4aa8e8e 100644 --- a/tests/integration/fastapi/test_overrides.py +++ b/tests/integration/fastapi/test_overrides.py @@ -1,4 +1,5 @@ -from collections.abc import AsyncIterator, Iterator +import contextlib +from collections.abc import AsyncIterator, Generator, Iterator from typing import Annotated import pytest @@ -246,3 +247,46 @@ def endpoint(value: FromWeb[FreshValueDep]) -> dict[str, str]: assert client.get("/").json() == { "value": "replacement", } + + +@contextlib.contextmanager +def get_context_value() -> Generator[str]: + yield "production" + + +@contextlib.contextmanager +def get_test_context_value() -> Generator[str]: + yield "test" + + +type ContextValueDep = Annotated[ + str, + wired(get_context_value), +] + + +def _create_context_manager_app() -> FastAPI: + app = FastAPI() + + def endpoint(value: FromWeb[ContextValueDep]) -> dict[str, str]: + return {"value": value} + + app.add_api_route("/", endpoint, methods=["GET"]) + + return app + + +def test_override_web_dependency_for_context_manager_factory() -> None: + app = _create_context_manager_app() + client = TestClient(app) + + assert client.get("/").json() == {"value": "production"} + + with override_web_dependency( + app, + get_context_value, + get_test_context_value, + ): + assert client.get("/").json() == {"value": "test"} + + assert client.get("/").json() == {"value": "production"} diff --git a/tests/smoke/fastapi_integration.py b/tests/smoke/fastapi_integration.py index 45f8a51..5ef3ac0 100644 --- a/tests/smoke/fastapi_integration.py +++ b/tests/smoke/fastapi_integration.py @@ -3,7 +3,8 @@ from __future__ import annotations import typing -from collections.abc import AsyncIterator, Iterator +from collections.abc import AsyncIterator, Generator, Iterator +from contextlib import contextmanager from typing import Annotated from fastapi import FastAPI @@ -64,6 +65,21 @@ async def get_session() -> AsyncIterator[Database]: ] +@contextmanager +def get_context_connection() -> Generator[Database]: + events.append("open-context") + try: + yield Database("context-resource") + finally: + events.append("close-context") + + +type ContextConnectionDep = Annotated[ + Database, + wired(get_context_connection), +] + + def get_uncoerced_number() -> int: return typing.cast("int", "1") @@ -94,6 +110,14 @@ def async_resource_endpoint(session: FromWeb[SessionDep]) -> dict[str, str]: return {"database": session.name} +@app.get("/context-resource") +def context_resource_endpoint( + connection: FromWeb[ContextConnectionDep], +) -> dict[str, str]: + events.append("use-context") + return {"database": connection.name} + + @app.get("/uncoerced") def uncoerced_endpoint(number: FromWeb[UncoercedNumberDep]) -> dict[str, object]: return {"number": number, "type": type(number).__name__} @@ -117,6 +141,13 @@ def uncoerced_endpoint(number: FromWeb[UncoercedNumberDep]) -> dict[str, object] assert response.json() == {"database": "async-resource"} assert events == ["open-async", "use-async", "close-async"], events +events.clear() + +response = client.get("/context-resource") +assert response.status_code == 200 +assert response.json() == {"database": "context-resource"} +assert events == ["open-context", "use-context", "close-context"], events + response = client.get("/uncoerced") assert response.status_code == 200 assert response.json() == {"number": "1", "type": "str"} diff --git a/tests/typing/core.py b/tests/typing/core.py index 8aeab45..c62ff55 100644 --- a/tests/typing/core.py +++ b/tests/typing/core.py @@ -1,6 +1,7 @@ from __future__ import annotations -from collections.abc import AsyncIterator, Iterator +from collections.abc import AsyncGenerator, AsyncIterator, Generator, Iterator +from contextlib import asynccontextmanager, contextmanager from typing import Annotated, assert_type from wireme import Wired, override_dependency, wire, wired @@ -78,3 +79,23 @@ def replacement_dependency() -> str: with override_dependency(original_dependency, replacement_dependency): pass + + +@contextmanager +def context_manager_dependency() -> Generator[str]: + yield "context-manager" + + +@asynccontextmanager +async def async_context_manager_dependency() -> AsyncGenerator[str]: + yield "async-context-manager" + + +assert_type(wired(context_manager_dependency), str) +assert_type(wired(async_context_manager_dependency), str) + +with override_dependency( + context_manager_dependency, + async_context_manager_dependency, +): + pass diff --git a/tests/unit/test_core.py b/tests/unit/test_core.py index 99da2c7..83ea72a 100644 --- a/tests/unit/test_core.py +++ b/tests/unit/test_core.py @@ -1,9 +1,10 @@ from __future__ import annotations +import contextlib import functools import inspect import typing -from collections.abc import AsyncGenerator, Generator +from collections.abc import AsyncGenerator, Callable, Generator from typing import Annotated, Any import pytest @@ -582,3 +583,142 @@ def operation(value: str) -> str: assert getattr(operation, "__signature__") is custom_signature # noqa: B009 assert inspect.signature(wrapped) == custom_signature assert wrapped("value") == "value" + + +def test_context_manager_dependency() -> None: + events: list[str] = [] + + @contextlib.contextmanager + def get_session() -> Generator[str]: + events.append("open") + try: + yield "session" + finally: + events.append("close") + + @wire + def operation(value: str = wired(get_session)) -> str: + events.append("handle") + return value + + assert operation() == "session" + assert events == ["open", "handle", "close"] + + +@pytest.mark.anyio +async def test_async_context_manager_dependency() -> None: + events: list[str] = [] + + @contextlib.asynccontextmanager + async def get_session() -> AsyncGenerator[str]: + events.append("open") + try: + yield "session" + finally: + events.append("close") + + @wire + async def operation(value: str = wired(get_session)) -> str: + events.append("handle") + return value + + assert await operation() == "session" + assert events == ["open", "handle", "close"] + + +def test_context_manager_dependency_closes_after_exception() -> None: + events: list[str] = [] + + @contextlib.contextmanager + def get_session() -> Generator[str]: + try: + yield "session" + finally: + events.append("close") + + @wire + def operation(value: str = wired(get_session)) -> str: + raise RuntimeError(value) + + with pytest.raises(RuntimeError, match="session"): + operation() + + assert events == ["close"] + + +def test_context_manager_dependency_is_overridable() -> None: + @contextlib.contextmanager + def get_session() -> Generator[str]: + yield "production" + + @contextlib.contextmanager + def get_test_session() -> Generator[str]: + yield "test" + + @wire + def operation(value: str = wired(get_session)) -> str: + return value + + with override_dependency(get_session, get_test_session): + assert operation() == "test" + + assert operation() == "production" + + +def test_context_manager_dependency_accepts_plain_replacement() -> None: + @contextlib.contextmanager + def get_session() -> Generator[str]: + yield "production" + + def get_test_session() -> str: + return "test" + + @wire + def operation(value: str = wired(get_session)) -> str: + return value + + with override_dependency(get_session, get_test_session): + assert operation() == "test" + + assert operation() == "production" + + +def test_wrapped_factory_keeps_its_decorator() -> None: + events: list[str] = [] + + def audited(func: Callable[[], str]) -> Callable[[], str]: + @functools.wraps(func) + def wrapper() -> str: + events.append("audit") + return func() + + return wrapper + + @audited + def get_value() -> str: + return "value" + + @wire + def operation(value: str = wired(get_value)) -> str: + return value + + assert operation() == "value" + assert events == ["audit"] + + +def test_cached_factory_keeps_its_cache() -> None: + calls = 0 + + @functools.lru_cache + def get_value() -> str: + nonlocal calls + calls += 1 + return "value" + + @wire + def operation(value: str = wired(get_value)) -> str: + return value + + assert operation() == "value" + assert operation() == "value" + assert calls == 1 diff --git a/website/docs/guide/fastapi.md b/website/docs/guide/fastapi.md index cef6e11..ced56f3 100644 --- a/website/docs/guide/fastapi.md +++ b/website/docs/guide/fastapi.md @@ -111,6 +111,10 @@ def report(*, connection: FromWeb[ConnectionDep]) -> dict[str, str]: return {"status": connection.status()} ``` +Factories decorated with `@contextmanager` or `@asynccontextmanager` are +bridged the same way and follow the same request lifecycle. See +[Resources](resources.md#context-managers). + ## Web overrides `override_web_dependency()` temporarily replaces a dependency on one @@ -130,8 +134,10 @@ with override_web_dependency(app, get_user_service, get_test_service): client.get("/users") ``` -Replacements may be sync, async, generator, or async-generator factories -and may have different parameter lists. +Replacements may be sync, async, generator, async-generator, or context +manager factories and may have different parameter lists. The original and +the replacement may each be decorated with `@contextmanager` or +`@asynccontextmanager`. One limitation: bridged adapters are discovered when the override context is entered, so routes using `FromWeb` must be registered before entering @@ -172,6 +178,7 @@ a FastAPI error about the unresolved wired annotation. [examples/fastapi_integration.py](https://github.com/mghalix/wireme/blob/main/examples/fastapi_integration.py), [examples/fastapi_resources.py](https://github.com/mghalix/wireme/blob/main/examples/fastapi_resources.py), [examples/fastapi_overrides.py](https://github.com/mghalix/wireme/blob/main/examples/fastapi_overrides.py), -[examples/fastapi_endpoints.py](https://github.com/mghalix/wireme/blob/main/examples/fastapi_endpoints.py) +[examples/fastapi_endpoints.py](https://github.com/mghalix/wireme/blob/main/examples/fastapi_endpoints.py), +[examples/fastapi_context_managers.py](https://github.com/mghalix/wireme/blob/main/examples/fastapi_context_managers.py) Next: [Building integrations](extending.md) diff --git a/website/docs/guide/resources.md b/website/docs/guide/resources.md index bd98351..fdec5f1 100644 --- a/website/docs/guide/resources.md +++ b/website/docs/guide/resources.md @@ -34,8 +34,43 @@ Cleanup runs after the wired callable finishes. Nested resources close in reverse order. For resources that must stay open for a whole web request, see the [FastAPI integration](fastapi.md). -## Runnable example +## Context managers + +A factory decorated with `@contextmanager` or `@asynccontextmanager` works +the same way, so an existing context manager can be reused as a dependency +without rewriting it as a bare generator: + +```python +import sqlite3 +from collections.abc import Generator +from contextlib import contextmanager + + +@contextmanager +def get_connection() -> Generator[sqlite3.Connection]: + connection = sqlite3.connect("app.db") + try: + yield connection + finally: + connection.close() + + +type ConnectionDep = Annotated[sqlite3.Connection, wired(get_connection)] +``` + +`wired(...)` resolves the dependency through the generator function the +decorator wraps, so the static type stays `sqlite3.Connection` rather than +the context manager object. Declaration sites, `override_dependency`, and +`override_web_dependency` all accept the decorated factory. + +Only `contextlib`'s two decorators are treated this way. Other decorators +applied with `functools.wraps`, including caches such as +`functools.lru_cache`, stay intact and keep running. + +## Runnable examples [examples/resources.py](https://github.com/mghalix/wireme/blob/main/examples/resources.py) +[examples/context_managers.py](https://github.com/mghalix/wireme/blob/main/examples/context_managers.py) + Next: [Testing](testing.md)