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
2 changes: 1 addition & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
19 changes: 15 additions & 4 deletions modern_di_grpc/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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]:
Expand Down
16 changes: 15 additions & 1 deletion tests/test_inject.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
15 changes: 15 additions & 0 deletions tests/test_sync.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import typing
from collections.abc import Iterator
from concurrent import futures

import grpc
import pytest
Expand Down Expand Up @@ -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])
Expand Down
Loading