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
37 changes: 19 additions & 18 deletions modern_di/providers/factory.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@

from modern_di import exceptions, suggester, types
from modern_di.providers.abstract import AbstractProvider
from modern_di.types_parser import SignatureItem, parse_creator
from modern_di.types_parser import ParsedCreator, SignatureItem, parse_creator
from modern_di.wiring import WiringPlan


Expand All @@ -24,6 +24,13 @@ class CacheSettings(typing.Generic[types.T_co]):
def __post_init__(self) -> None:
self.is_async_finalizer = bool(self.finalizer) and inspect.iscoroutinefunction(self.finalizer)

@staticmethod
def coerce(cache: "bool | CacheSettings[types.T] | None") -> "CacheSettings[types.T] | None":
"""Read a ``Factory``'s ``cache`` argument: ``True`` is the defaults, ``False`` and ``None`` are off."""
if cache is True:
return CacheSettings()
return cache or None


class Factory(AbstractProvider[types.T_co]):
__slots__ = (
Expand All @@ -45,12 +52,6 @@ def __init__( # noqa: PLR0913
cache: bool | CacheSettings[types.T_co] | None = None,
skip_creator_parsing: bool = False,
) -> None:
if cache is True:
resolved_cache: CacheSettings[types.T_co] | None = CacheSettings()
elif cache: # a CacheSettings instance
resolved_cache = cache
else: # None or False
resolved_cache = None
if skip_creator_parsing:
if bound_type is types.UNSET:
warnings.warn(
Expand All @@ -59,15 +60,12 @@ def __init__( # noqa: PLR0913
UserWarning,
stacklevel=2,
)
parsed_type: type | None = None
parsed_kwargs: dict[str, SignatureItem] = {}
has_positional_only_gap = False
parsed = ParsedCreator(return_type=SignatureItem(), params={}, has_positional_only_gap=False)
else:
return_sig, parsed_kwargs, has_positional_only_gap = parse_creator(creator)
parsed_type = return_sig.arg_type
parsed = parse_creator(creator)
if kwargs:
self._validate_kwargs_against_signature(creator, kwargs, parsed_kwargs)
for param_name, item in parsed_kwargs.items():
self._validate_kwargs_against_signature(creator, kwargs, parsed.params)
for param_name, item in parsed.params.items():
if item.raw_annotation is None or item.default is not types.UNSET or (kwargs and param_name in kwargs):
continue
raise exceptions.UnsupportedCreatorParameterError(
Expand All @@ -79,11 +77,14 @@ def __init__( # noqa: PLR0913
f"or use skip_creator_parsing=True"
),
)
self._parsed_kwargs = parsed_kwargs
self._has_positional_only_gap = has_positional_only_gap
super().__init__(scope=scope, bound_type=parsed_type if isinstance(bound_type, types.UnsetType) else bound_type)
self._parsed_kwargs = parsed.params
self._has_positional_only_gap = parsed.has_positional_only_gap
super().__init__(
scope=scope,
bound_type=parsed.return_type.arg_type if isinstance(bound_type, types.UnsetType) else bound_type,
)
self._creator = creator
self.cache_settings = resolved_cache
self.cache_settings = CacheSettings.coerce(cache)
self._kwargs = kwargs
self._cached_definition_site: str | types.UnsetType | None = types.UNSET

Expand Down
22 changes: 15 additions & 7 deletions modern_di/types_parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -86,18 +86,26 @@ def _parse_parameter(
return item


def parse_creator(
creator: typing.Callable[..., typing.Any],
) -> tuple[SignatureItem, dict[str, SignatureItem], bool]:
"""Return (return-type item, name→param item, has_positional_only_gap).
@dataclasses.dataclass(kw_only=True, slots=True, frozen=True)
class ParsedCreator:
"""A creator's signature as the wiring reads it.

``has_positional_only_gap`` is True when a positional-only-with-default parameter was dropped
from the param map, so the map is no longer a faithful positional prefix of the signature.
from ``params``, so the map is no longer a faithful positional prefix of the signature.
"""

return_type: SignatureItem
params: dict[str, SignatureItem]
has_positional_only_gap: bool


def parse_creator(creator: typing.Callable[..., typing.Any]) -> ParsedCreator:
try:
sig = inspect.signature(creator)
except (ValueError, TypeError):
return SignatureItem.from_type(typing.cast(type, creator)), {}, False
return ParsedCreator(
return_type=SignatureItem.from_type(typing.cast(type, creator)), params={}, has_positional_only_gap=False
)

is_class = isinstance(creator, type)
try:
Expand Down Expand Up @@ -137,4 +145,4 @@ def parse_creator(
else:
return_sig = SignatureItem()

return return_sig, param_hints, has_positional_only_gap
return ParsedCreator(return_type=return_sig, params=param_hints, has_positional_only_gap=has_positional_only_gap)
13 changes: 13 additions & 0 deletions tests/providers/test_factory.py
Original file line number Diff line number Diff line change
Expand Up @@ -1207,3 +1207,16 @@ def _body_raises() -> None:
creator=_body_raises, exc=exc, resolution_step=lambda: step
)
assert result is None


@pytest.mark.parametrize(
("cache", "expected"),
[(True, providers.CacheSettings()), (False, None), (None, None)],
)
def test_cache_settings_coerce(cache: bool | None, expected: providers.CacheSettings[object] | None) -> None:
assert providers.CacheSettings.coerce(cache) == expected


def test_cache_settings_coerce_returns_an_instance_unchanged() -> None:
settings: providers.CacheSettings[object] = providers.CacheSettings(clear_cache=False)
assert providers.CacheSettings.coerce(settings) is settings
14 changes: 7 additions & 7 deletions tests/test_types_parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -71,7 +71,7 @@ def nonetype_func(hook: None = None) -> None: ...


def test_nonetype_params_keep_defaults_through_parse_creator() -> None:
_ret, params, _gap = parse_creator(nonetype_func)
params = parse_creator(nonetype_func).params
assert params["hook"] == SignatureItem(default=None, is_nullable=True)


Expand Down Expand Up @@ -210,8 +210,8 @@ def __init__(self, arg1: "WrongType", arg2: "int") -> None: ... # ty: ignore[un
],
)
def test_parse_creator(creator: type, result: tuple[SignatureItem | None, dict[str, SignatureItem]]) -> None:
return_sig, params, _gap = parse_creator(creator)
assert (return_sig, params) == result
parsed = parse_creator(creator)
assert (parsed.return_type, parsed.params) == result


def test_parse_creator_str_generic_annotation_without_default_raises() -> None:
Expand Down Expand Up @@ -303,13 +303,13 @@ def _pos_only_with_default(x: int = 0, /, y: int = 1) -> int:

def test_positional_only_param_with_default_is_skipped() -> None:
assert _pos_only_with_default(2) == _pos_only_with_default(1, 2)
return_sig, params, has_positional_only_gap = parse_creator(_pos_only_with_default)
assert (return_sig, params) == (
parsed = parse_creator(_pos_only_with_default)
assert (parsed.return_type, parsed.params) == (
SignatureItem(arg_type=int),
{"y": SignatureItem(arg_type=int, default=1)},
)
# the dropped positional-only `x` is recorded so the compiled fast path keeps **kwargs
assert has_positional_only_gap is True
assert parsed.has_positional_only_gap is True


def _mixed_kind_creator(pos_or_kw: int, *, kw_only: int) -> int:
Expand All @@ -320,6 +320,6 @@ def test_keyword_only_signal_recorded() -> None:
# A keyword-only parameter records is_keyword_only=True; a positional-or-keyword one records
# False. This is the only param-kind signal the compiled positional fast path consults.
assert _mixed_kind_creator(1, kw_only=2) == 1 + 2 # exercise the creator body for coverage
_ret, params, _gap = parse_creator(_mixed_kind_creator)
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
Loading