diff --git a/modern_di/providers/factory.py b/modern_di/providers/factory.py index 7887d28c..8468e462 100644 --- a/modern_di/providers/factory.py +++ b/modern_di/providers/factory.py @@ -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 @@ -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__ = ( @@ -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( @@ -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( @@ -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 diff --git a/modern_di/types_parser.py b/modern_di/types_parser.py index ac0a4792..f271b10f 100644 --- a/modern_di/types_parser.py +++ b/modern_di/types_parser.py @@ -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: @@ -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) diff --git a/tests/providers/test_factory.py b/tests/providers/test_factory.py index 1e96b42d..5062fa58 100644 --- a/tests/providers/test_factory.py +++ b/tests/providers/test_factory.py @@ -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 diff --git a/tests/test_types_parser.py b/tests/test_types_parser.py index b1713909..d6f80bd0 100644 --- a/tests/test_types_parser.py +++ b/tests/test_types_parser.py @@ -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) @@ -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: @@ -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: @@ -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