From 3a59da2f1834dcf6b20e4f866f57ef160cb538f6 Mon Sep 17 00:00:00 2001 From: Artur Shiriev Date: Sun, 4 Oct 2026 22:34:04 +0300 Subject: [PATCH] fix!: wire NamedTuple, __new__-only classes, NewType and type aliases (#571) --- docs/migration/to-4.x.md | 15 ++++ docs/providers/factories.md | 9 +++ modern_di/providers/factory.py | 9 +++ modern_di/types_parser.py | 37 ++++++---- tests/test_types_parser.py | 124 ++++++++++++++++++++++++++++++++- 5 files changed, 178 insertions(+), 16 deletions(-) diff --git a/docs/migration/to-4.x.md b/docs/migration/to-4.x.md index 0c42218f..fabe68be 100644 --- a/docs/migration/to-4.x.md +++ b/docs/migration/to-4.x.md @@ -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 diff --git a/docs/providers/factories.md b/docs/providers/factories.md index ca3f0ea8..798fdc4c 100644 --- a/docs/providers/factories.md +++ b/docs/providers/factories.md @@ -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. @@ -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`. | diff --git a/modern_di/providers/factory.py b/modern_di/providers/factory.py index beba982b..0461e6cc 100644 --- a/modern_di/providers/factory.py +++ b/modern_di/providers/factory.py @@ -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__( diff --git a/modern_di/types_parser.py b/modern_di/types_parser.py index 39caa375..8e0f7593 100644 --- a/modern_di/types_parser.py +++ b/modern_di/types_parser.py @@ -1,5 +1,6 @@ import dataclasses import inspect +import sys import types import typing import warnings @@ -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 @@ -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_)] @@ -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) @@ -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) @@ -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. " @@ -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 @@ -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() diff --git a/tests/test_types_parser.py b/tests/test_types_parser.py index 5b9fe28a..b2ca6993 100644 --- a/tests/test_types_parser.py +++ b/tests/test_types_parser.py @@ -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 @@ -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)