From 8f72030e47563ec82ae1e5b3dec8fdd605ab48c6 Mon Sep 17 00:00:00 2001 From: Artur Shiriev Date: Tue, 15 Sep 2026 20:04:07 +0300 Subject: [PATCH] fix: raise a clear RuntimeError when inject runs without the DI interceptor --- README.md | 2 +- modern_di_grpc/main.py | 19 +++++++++++++++---- tests/test_inject.py | 16 +++++++++++++++- tests/test_sync.py | 15 +++++++++++++++ 4 files changed, 46 insertions(+), 6 deletions(-) diff --git a/README.md b/README.md index 4be9718..4befeff 100644 --- a/README.md +++ b/README.md @@ -96,7 +96,7 @@ For an async server, pass `DIAioInterceptor(container)` to `grpc.aio.server(...) | `DIAioInterceptor(container)` | `grpc.aio.ServerInterceptor` for the async server. Same, with `close_async` | | `FromDI(dependency)` | Inert marker for `Annotated[T, FromDI(...)]` in servicer-method signatures; accepts a provider instance or a type | | `inject(method)` | Decorates a servicer method to resolve its `FromDI` parameters from the current RPC's child container; adapts to sync / async / async-generator methods | -| `fetch_di_container()` | Returns the current RPC's child container (raises `LookupError` outside an RPC) | +| `fetch_di_container()` | Returns the current RPC's child container (raises `RuntimeError` outside an RPC) | | `grpc_context_provider` | `ContextProvider` exposing `grpc.ServicerContext` at `Scope.REQUEST`; auto-registered by the interceptor | ## 📦 [PyPI](https://pypi.org/project/modern-di-grpc) diff --git a/modern_di_grpc/main.py b/modern_di_grpc/main.py index 137ffa1..01d191e 100644 --- a/modern_di_grpc/main.py +++ b/modern_di_grpc/main.py @@ -23,9 +23,21 @@ def _build_child(container: Container, context: ServicerContext) -> Container: return child +def _current_container() -> Container: + try: + return _request_container.get() + except LookupError: + msg = ( + "No modern-di container found for this RPC. " + "Add DIInterceptor (sync server) or DIAioInterceptor (aio server) to the server's interceptors " + "so RPCs pass through it before using @inject or fetch_di_container." + ) + raise RuntimeError(msg) from None + + def fetch_di_container() -> Container: - """Return the current RPC's child container. Raises ``LookupError`` outside an intercepted RPC.""" - return _request_container.get() + """Return the current RPC's child container. Raises ``RuntimeError`` outside an intercepted RPC.""" + return _current_container() def _ensure_context_provider(container: Container) -> None: @@ -38,8 +50,7 @@ def _ensure_context_provider(container: Container) -> None: def _resolve(di_params: dict[str, integrations.Marker[typing.Any]]) -> dict[str, typing.Any]: - container = _request_container.get() - return integrations.resolve_markers(container, di_params) + return integrations.resolve_markers(_current_container(), di_params) def inject(func: typing.Callable[..., typing.Any]) -> typing.Callable[..., typing.Any]: diff --git a/tests/test_inject.py b/tests/test_inject.py index 44408e0..975ce60 100644 --- a/tests/test_inject.py +++ b/tests/test_inject.py @@ -40,12 +40,26 @@ def test_fetch_di_container_returns_child() -> None: def test_fetch_di_container_raises_outside_rpc() -> None: def _call() -> None: - with pytest.raises(LookupError): + with pytest.raises(RuntimeError, match="DIInterceptor"): fetch_di_container() contextvars.copy_context().run(_call) # guaranteed-unset ContextVar +def test_inject_raises_without_interceptor() -> None: + @inject + def method( + _self: object, _request: str, _context: object, _app_res: typing.Annotated[AppResource, FromDI(AppResource)] + ) -> None: + pass # pragma: no cover + + def _call() -> None: + with pytest.raises(RuntimeError, match="DIInterceptor"): + method(object(), "req", object()) + + contextvars.copy_context().run(_call) # guaranteed-unset ContextVar + + def test_inject_sync_resolves() -> None: @inject def method( diff --git a/tests/test_sync.py b/tests/test_sync.py index b752680..a1039d0 100644 --- a/tests/test_sync.py +++ b/tests/test_sync.py @@ -1,5 +1,6 @@ import typing from collections.abc import Iterator +from concurrent import futures import grpc import pytest @@ -124,6 +125,20 @@ def test_unknown_method_returns_unimplemented(sync_channel: grpc.Channel) -> Non assert excinfo.value.code() == grpc.StatusCode.UNIMPLEMENTED # ty: ignore[unresolved-attribute] +def test_inject_without_interceptor_reports_missing_interceptor() -> None: + server = grpc.server(futures.ThreadPoolExecutor(max_workers=1)) + greeter_pb2_grpc.add_GreeterServicer_to_server(Servicer(), server) + port = server.add_insecure_port("127.0.0.1:0") + server.start() + try: + with grpc.insecure_channel(f"127.0.0.1:{port}") as channel, pytest.raises(grpc.RpcError) as excinfo: + greeter_pb2_grpc.GreeterStub(channel).SayHello(HelloRequest(name="a")) + finally: + server.stop(0) + assert excinfo.value.code() == grpc.StatusCode.UNKNOWN # ty: ignore[unresolved-attribute] + assert "DIInterceptor" in excinfo.value.details() # ty: ignore[unresolved-attribute] + + async def test_app_finalizer_runs_on_root_close() -> None: app_teardowns.clear() container = Container(groups=[Dependencies])