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
15 changes: 15 additions & 0 deletions docs/migration/to-4.x.md
Original file line number Diff line number Diff line change
Expand Up @@ -185,6 +185,21 @@ provider set is closed. Only code that relied on `ABCMeta` is affected:
`AbstractProvider.register(...)` raises `AttributeError`. `isinstance(x, AbstractProvider)` works
as before for every provider, and type hints that name `AbstractProvider` need no change.

### `NewType` and type alias annotations are wired

In 3.x a parameter annotated with a `NewType` or a `type X = ...` alias was treated as unannotated,
and a creator returning one got no bound type. In 4.0 both are bound types of their own: the
parameter resolves from the provider declared with `bound_type=UserId`, and `Factory(make_user_id)`
with `-> UserId` registers under `UserId`. If a group already has a provider bound to the same
`NewType` or alias, registration now raises `DuplicateProviderTypeError`; pass `bound_type=None`
to the one that should stay unregistered.

Classes whose signature comes from `__new__`, such as `NamedTuple` subclasses, now wire their
parameters from the `__new__` annotations.

A `Factory` whose creator returns a union of several types (`-> A | B`) still gets no bound type,
and now emits a `UserWarning` saying so. Pass `bound_type=` explicitly to silence it.

### The 3.x deprecations are removed

- `Container(validate=...)` raises `TypeError`, and `ValidateArgumentWarning` is gone with it. Drop
Expand Down
9 changes: 9 additions & 0 deletions docs/providers/factories.md
Original file line number Diff line number Diff line change
Expand Up @@ -136,6 +136,14 @@ Modern-DI analyzes the creator's signature to:
Explicitly sets the type for resolving by type. By default, this is automatically inferred from the creator's return type annotation.
Set to `None` to make the provider unresolvable by type.

A `NewType` or a `type X = ...` alias is a bound type of its own. A provider declared with
`bound_type=UserId` (or a creator returning `UserId`) is what a `user_id: UserId` parameter
resolves to; a provider bound to the underlying `int` is not.

A return annotation that is a union of several types (`-> A | B`) gives no bound type, and
`Factory(...)` emits a `UserWarning`. Pass `bound_type=` with the type to register under, or
`bound_type=None` if the provider is only resolved directly.

### kwargs

Manual values for creator parameters that override automatic dependency resolution.
Expand Down Expand Up @@ -203,6 +211,7 @@ The table below summarises how Modern-DI handles each parameter shape during **d
| `param: SomeClass` (plain type annotation with a registered provider) | Resolved and injected automatically. | `ArgumentResolutionError` at resolve if no provider is registered and there is no default. |
| `param: X | None` / `Optional[X]` | Provider injected if one is registered; otherwise `None`. | Never fails; see [Optional parameters](#optional-parameters). |
| `param: A | B` (union without `None`) | First registered type from the union is injected. A member that is itself a parameterized generic (e.g. `int | list[X]`) degrades to its bare origin (`list`) for matching purposes; see the note below. | `ArgumentResolutionError` at resolve if neither `A` nor `B` has a registered provider. |
| `param: UserId` (a `NewType`) or `param: Alias` (a `type Alias = ...` statement, Python 3.12+) | Resolved from the provider whose `bound_type` is that `NewType` or alias. The underlying type is not looked up. | `ArgumentResolutionError` at resolve if no provider is bound to it and there is no default. |
| `param: list[X]` / any parameterized generic, **outside a union** | **`UnsupportedCreatorParameterError` at declaration** unless the parameter has a default value or is covered by `kwargs`. | Raised at `Factory(...)` call time. |
| Positional-only param (`def f(x: T, /)`) | **`UnsupportedCreatorParameterError` at declaration** unless the parameter has a default (in which case it is silently skipped). | Raised at `Factory(...)` call time. |
| Unannotated param (`def f(x)`) | Parsed but unresolvable by type. | `ArgumentResolutionError` at resolve unless covered by `kwargs`. |
Expand Down
9 changes: 9 additions & 0 deletions modern_di/providers/factory.py
Original file line number Diff line number Diff line change
Expand Up @@ -84,6 +84,15 @@ def __init__( # noqa: PLR0913
"pass the value via the kwargs parameter or give the parameter a default"
),
)
if parsed.return_type.args and isinstance(bound_type, types.UnsetType):
members = " | ".join(getattr(t, "__name__", str(t)) for t in parsed.return_type.args)
warnings.warn(
f"The return annotation of {creator!r} is a union of {members}, so no bound_type can be "
"inferred and this provider cannot be resolved by type. Pass bound_type=OneOfThem, or "
"bound_type=None to silence this warning.",
UserWarning,
stacklevel=2,
)
self._parsed_kwargs = parsed.params
self._has_positional_only_gap = parsed.has_positional_only_gap
super().__init__(
Expand Down
37 changes: 22 additions & 15 deletions modern_di/types_parser.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import dataclasses
import inspect
import sys
import types
import typing
import warnings
Expand All @@ -8,6 +9,9 @@
from modern_di.types import UNSET


_NAMED_TYPE_FORMS = (typing.NewType,) if sys.version_info < (3, 12) else (typing.NewType, typing.TypeAliasType)


@dataclasses.dataclass(kw_only=True, slots=True, frozen=True)
class SignatureItem:
arg_type: type | None = None
Expand All @@ -24,14 +28,14 @@ def from_type(cls, type_: type, default: object = UNSET) -> "SignatureItem":
# try to resolve `NoneType` from the registry.
return cls(default=default, is_nullable=True)

# typing.Annotated
if hasattr(type_, "__metadata__"):
origin = typing.get_origin(type_)
if origin is typing.Annotated:
type_ = typing.get_args(type_)[0]
origin = typing.get_origin(type_)

result: dict[str, typing.Any] = {"default": default}

# union type
if isinstance(type_, types.UnionType) or typing.get_origin(type_) is typing.Union:
if isinstance(type_, types.UnionType) or origin is typing.Union:
# A parameterized generic member degrades to its origin (list[str] -> list); see
# test_union_member_degrades_to_bare_origin.
union_members = [typing.get_origin(x) or x for x in typing.get_args(type_)]
Expand All @@ -41,14 +45,13 @@ def from_type(cls, type_: type, default: object = UNSET) -> "SignatureItem":

if len(non_none_members) > 1:
result["args"] = non_none_members
elif non_none_members:
else:
result["arg_type"] = non_none_members[0]

# generic — parameterized generics are not resolvable by type
elif typing.get_origin(type_) is not None:
elif origin is not None:
result["raw_annotation"] = type_

elif isinstance(type_, type):
elif isinstance(type_, (type, _NAMED_TYPE_FORMS)):
result["arg_type"] = type_

return cls(**result)
Expand Down Expand Up @@ -96,6 +99,16 @@ class ParsedCreator:
has_positional_only_gap: bool


def _class_type_hints(creator: type) -> dict[str, typing.Any]:
"""Return the hints of the ``__new__`` or ``__init__`` that ``inspect.signature`` reads for a class."""
for base in creator.__mro__:
for name in ("__new__", "__init__"):
if name in base.__dict__ and inspect.isfunction(method := getattr(base, name)):
module = sys.modules.get(base.__module__)
return typing.get_type_hints(method, localns=vars(module) if module else None)
return typing.get_type_hints(creator.__init__)


def parse_creator(creator: typing.Callable[..., typing.Any]) -> ParsedCreator:
try:
sig = inspect.signature(creator)
Expand All @@ -106,10 +119,7 @@ def parse_creator(creator: typing.Callable[..., typing.Any]) -> ParsedCreator:

is_class = isinstance(creator, type)
try:
if is_class and hasattr(creator, "__init__"):
type_hints = typing.get_type_hints(creator.__init__)
else:
type_hints = typing.get_type_hints(creator)
type_hints = _class_type_hints(creator) if is_class else typing.get_type_hints(creator)
except (NameError, TypeError) as e:
warnings.warn(
f"Failed to resolve type hints for {creator}: {e}. Dependency wiring will be skipped. "
Expand All @@ -126,8 +136,6 @@ def parse_creator(creator: typing.Callable[..., typing.Any]) -> ParsedCreator:
continue
item = _parse_parameter(creator, param_name, param, type_hints)
if item is None:
# Dropped from param_hints, so a positional creator() call would bind a later
# dependency into this slot; the fast path must keep **kwargs.
has_positional_only_gap = True
continue
param_hints[param_name] = item
Expand All @@ -137,7 +145,6 @@ def parse_creator(creator: typing.Callable[..., typing.Any]) -> ParsedCreator:
elif "return" in type_hints:
return_sig = SignatureItem.from_type(type_hints["return"])
if return_sig.raw_annotation is not None:
# a parameterized generic return type degrades to its origin for bound_type
return_sig = SignatureItem(arg_type=typing.get_origin(return_sig.raw_annotation))
else:
return_sig = SignatureItem()
Expand Down
124 changes: 123 additions & 1 deletion tests/test_types_parser.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,12 @@
import dataclasses
import functools
import sys
import typing
import warnings

import pytest

from modern_di import Container, Scope, exceptions, providers, types
from modern_di import Container, Group, Scope, exceptions, providers, types
from modern_di.types_parser import SignatureItem, parse_creator


Expand Down Expand Up @@ -325,3 +326,124 @@ def test_keyword_only_signal_recorded() -> None:
params = parse_creator(_mixed_kind_creator).params
assert params["pos_or_kw"].is_keyword_only is False
assert params["kw_only"].is_keyword_only is True


class _Dep: ...


class _OtherDep: ...


class _NamedTupleCreator(typing.NamedTuple):
dep: _Dep
label: str = "x"


class _NamedTupleForwardRef(typing.NamedTuple):
dep: "_LateDep"


class _LateDep(_Dep): ...


class _NewOnlyCreator:
def __new__(cls, dep: _Dep) -> typing.Self:
instance = super().__new__(cls)
instance.dep = dep # ty: ignore[unresolved-attribute]
return instance


class _NewOnlySubclass(_NewOnlyCreator): ...


@pytest.mark.parametrize(
("creator", "dep_type"),
[
(_NamedTupleCreator, _Dep),
(_NamedTupleForwardRef, _LateDep),
(_NewOnlyCreator, _Dep),
(_NewOnlySubclass, _Dep),
],
)
def test_hints_come_from_the_callable_the_signature_reads(creator: type, dep_type: type) -> None:
assert parse_creator(creator).params["dep"] == SignatureItem(arg_type=dep_type)

class Dependencies(Group):
dep = providers.Factory(dep_type)
target = providers.Factory(creator)

container = Container(groups=[Dependencies])
assert isinstance(container.resolve(creator).dep, dep_type) # ty: ignore[unresolved-attribute]


_UserId = typing.NewType("_UserId", int)
_UserIdAlias = typing.TypeAliasType("_UserIdAlias", int) if sys.version_info >= (3, 12) else None
_NAMED_TYPE_FORMS = [
pytest.param(_UserId, id="NewType"),
pytest.param(
_UserIdAlias,
id="TypeAliasType",
marks=pytest.mark.skipif(sys.version_info < (3, 12), reason="PEP 695 type aliases need Python 3.12"),
),
]


def _greeter_of(form: object) -> typing.Callable[..., str]:
def greet(user_id) -> str: # noqa: ANN001
return f"u{user_id}"

greet.__annotations__ = {"user_id": form, "return": str}
return greet


@pytest.mark.parametrize("form", _NAMED_TYPE_FORMS)
def test_named_type_form_parameter_matches_provider_bound_to_it(form: type) -> None:
assert SignatureItem.from_type(form) == SignatureItem(arg_type=form)

class Dependencies(Group):
user_id = providers.Factory(lambda: 7, bound_type=form)
greeting = providers.Factory(_greeter_of(form))

assert Container(groups=[Dependencies]).resolve(str) == "u7"


@pytest.mark.parametrize("form", _NAMED_TYPE_FORMS)
def test_named_type_form_return_annotation_is_the_bound_type(form: type) -> None:
def make() -> int: ... # ty: ignore[empty-body]

make.__annotations__ = {"return": form}
assert providers.Factory(make).bound_type is form


@pytest.mark.parametrize("form", _NAMED_TYPE_FORMS)
def test_unregistered_named_type_form_names_the_annotation(form: type) -> None:
class Dependencies(Group):
greeting = providers.Factory(_greeter_of(form))

with pytest.raises(exceptions.ArgumentResolutionError) as exc_info:
Container(groups=[Dependencies]).resolve(str)
assert "no usable type annotation" not in str(exc_info.value)
assert "_UserId" in str(exc_info.value)


def _make_dep_or_other() -> _Dep | _OtherDep: ... # ty: ignore[empty-body]
def _make_optional_dep() -> _Dep | None: ...


def test_union_return_type_without_bound_type_warns() -> None:
with pytest.warns(UserWarning, match="bound_type") as record:
provider = providers.Factory(_make_dep_or_other)
assert provider.bound_type is None
assert "_Dep | _OtherDep" in str(record[0].message)


@pytest.mark.parametrize(
("creator", "bound_type"),
[(_make_dep_or_other, _Dep), (_make_dep_or_other, None), (_make_optional_dep, types.UNSET)],
)
def test_union_return_type_is_silent_when_bound_type_is_known(
creator: typing.Callable[..., typing.Any], bound_type: type | None
) -> None:
with warnings.catch_warnings():
warnings.simplefilter("error")
providers.Factory(creator, bound_type=bound_type)
Loading