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
24 changes: 24 additions & 0 deletions docs/migration/to-4.x.md
Original file line number Diff line number Diff line change
Expand Up @@ -157,6 +157,30 @@ except* exceptions.AsyncFinalizerInSyncCloseError:
...
```

### `Factory(cache=)` takes only a bool or a `CacheSettings`

`cache=None` raises `TypeError`, and so does any other value that is not `True`, `False` or a
`CacheSettings`. Replace `cache=None` with `cache=False`, or drop the argument: an uncached
factory is still the default.

### Provider internals are no longer public

These methods were never documented and no integration calls them. In 4.0 they are private, so
code that calls them raises `AttributeError`:

- `AbstractProvider.get_dependencies()`, `AbstractProvider.redirect_target()` and
`AbstractProvider.iter_validation_issues()`. Call `container.validate()` to check the graph.
- `Factory.wiring_plan()`, `Factory.can_call_positionally()` and `Factory.resolution_step()`.
- `Alias.find_source()`. Resolve the alias instead.
- `CacheSettings.coerce()`. Pass `True`, `False` or a `CacheSettings` to `Factory(cache=)`.

### `AbstractProvider` is a plain class

`AbstractProvider` no longer derives from `abc.ABC`. It never declared an abstract method, and the
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.

### The 3.x deprecations are removed

- `Container(validate=...)` raises `TypeError`, and `ValidateArgumentWarning` is gone with it. Drop
Expand Down
4 changes: 2 additions & 2 deletions docs/providers/factories.md
Original file line number Diff line number Diff line change
Expand Up @@ -112,7 +112,7 @@ class Dependencies(Group):

## Parameters

`Factory(creator, *, scope=Scope.APP, bound_type=UNSET, kwargs=None, cache=None, skip_creator_parsing=False)`.
`Factory(creator, *, scope=Scope.APP, bound_type=UNSET, kwargs=None, cache=False, skip_creator_parsing=False)`.
The `creator` may also be passed as a keyword (`creator=`).

When creating a Factory provider, you can configure several parameters:
Expand Down Expand Up @@ -143,7 +143,7 @@ Use this to provide specific values for parameters or override automatically res

### cache

Enables caching for the provider. Pass `cache=True` to cache with default settings (no finalizer, cache cleared on close), or `cache=providers.CacheSettings(...)` to tune the finalizer and/or `clear_cache` behavior. Absent, `None`, or `False` means a fresh instance is created on every resolve. See [Lifecycle](lifecycle.md) for how caching, finalizers, and `close_async()` fit together.
Enables caching for the provider. Pass `cache=True` to cache with default settings (no finalizer, cache cleared on close), or `cache=providers.CacheSettings(...)` to tune the finalizer and/or `clear_cache` behavior. With `cache=False`, the default, a fresh instance is created on every resolve. Any other value, `None` included, raises `TypeError`. See [Lifecycle](lifecycle.md) for how caching, finalizers, and `close_async()` fit together.

### skip_creator_parsing

Expand Down
2 changes: 1 addition & 1 deletion modern_di/__init__.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
from modern_di import exceptions, integrations
from modern_di import exceptions, integrations, providers
from modern_di.container import Container
from modern_di.group import Group
from modern_di.scope import Scope
Expand Down
8 changes: 4 additions & 4 deletions modern_di/dependency_graph.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,14 +51,14 @@ class DependenciesError(NamedTuple):
def terminal_chain(
provider: "AbstractProvider[typing.Any]", container: "Container"
) -> "list[AbstractProvider[typing.Any]]":
"""Follow ``redirect_target`` hops from ``provider``, ``provider`` first.
"""Follow ``_redirect_target`` hops from ``provider``, ``provider`` first.

A redirect cycle collapses the chain to the single provider the repeat was detected at, so
``effective_scope`` reports that provider's own scope; ``walk()`` reports the cycle itself.
"""
chain = [provider]
seen: set[int] = set()
while (nxt := provider.redirect_target(container)) is not None:
while (nxt := provider._redirect_target(container)) is not None: # noqa: SLF001
if provider.provider_id in seen:
return [provider]
seen.add(provider.provider_id)
Expand Down Expand Up @@ -174,7 +174,7 @@ def _enter(
path.append(provider)
yield NodeEntered(provider)
try:
dependencies = provider.get_dependencies(container)
dependencies = provider._get_dependencies(container) # noqa: SLF001
except exceptions.ResolutionError as exc:
yield DependenciesError(provider, exc)
dependencies = {}
Expand All @@ -187,7 +187,7 @@ def collect_errors(container: "Container") -> list[Exception]:
for event in DependencyGraph().walk(container.providers_registry, container):
match event:
case NodeEntered(provider):
errors.extend(provider.iter_validation_issues(container))
errors.extend(provider._iter_validation_issues(container)) # noqa: SLF001
case DependenciesError(_, error):
errors.append(error)
case Edge(parent, name, dep):
Expand Down
9 changes: 4 additions & 5 deletions modern_di/providers/abstract.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,3 @@
import abc
import enum
import itertools
import typing
Expand All @@ -13,7 +12,7 @@
_provider_id_counter = itertools.count()


class AbstractProvider(abc.ABC, typing.Generic[types.T_co]):
class AbstractProvider(typing.Generic[types.T_co]):
__slots__ = ("_explicit_scope", "_group_claim", "_registered", "bound_type", "provider_id")

_takes_group_scope: typing.ClassVar[bool] = True
Expand Down Expand Up @@ -81,13 +80,13 @@ def definition_site(self) -> str | None:
"""``module:line`` of the provider's declaration when known; None by default (no creator)."""
return None

def get_dependencies(self, container: "Container") -> dict[str, "AbstractProvider[typing.Any]"]: # noqa: ARG002
def _get_dependencies(self, container: "Container") -> dict[str, "AbstractProvider[typing.Any]"]: # noqa: ARG002
return {}

def redirect_target(self, container: "Container") -> "AbstractProvider[typing.Any] | None": # noqa: ARG002
def _redirect_target(self, container: "Container") -> "AbstractProvider[typing.Any] | None": # noqa: ARG002
"""Return the provider this transparently forwards to, or None if resolution terminates here."""
return None

def iter_validation_issues(self, container: "Container") -> typing.Iterable[Exception]: # noqa: ARG002
def _iter_validation_issues(self, container: "Container") -> typing.Iterable[Exception]: # noqa: ARG002
"""Yield validation-time issues for this provider. Default: no issues."""
return iter(())
10 changes: 5 additions & 5 deletions modern_di/providers/alias.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,17 +27,17 @@ def __init__(
def __repr__(self) -> str:
return f"Alias(source_type={self._source_type!r}, bound_type={self.bound_type!r}, scope={self.scope!r})"

def find_source(self, container: "Container") -> "AbstractProvider[types.T_co]":
def _find_source(self, container: "Container") -> "AbstractProvider[types.T_co]":
source = container.providers_registry.find_provider(self._source_type)
if source is None:
raise exceptions.AliasSourceNotRegisteredError(source_type=self._source_type)
return source

def get_dependencies(self, container: "Container") -> dict[str, "AbstractProvider[typing.Any]"]:
return {"source": self.find_source(container)}
def _get_dependencies(self, container: "Container") -> dict[str, "AbstractProvider[typing.Any]"]:
return {"source": self._find_source(container)}

def redirect_target(self, container: "Container") -> "AbstractProvider[typing.Any] | None":
def _redirect_target(self, container: "Container") -> "AbstractProvider[typing.Any] | None":
try:
return self.find_source(container)
return self._find_source(container)
except exceptions.AliasSourceNotRegisteredError:
return None
32 changes: 20 additions & 12 deletions modern_di/providers/factory.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,11 +25,19 @@ 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."""
def _coerce(cache: "bool | CacheSettings[types.T]") -> "CacheSettings[types.T] | None":
"""Read a ``Factory``'s ``cache`` argument: ``True`` is the defaults, ``False`` is off."""
if isinstance(cache, CacheSettings):
return cache
if cache is True:
return CacheSettings()
return cache or None
if cache is False:
return None
msg = (
f"Factory cache= takes a bool or a CacheSettings; got {cache!r}. "
"Pass cache=False, or leave it out, for an uncached factory."
)
raise TypeError(msg)


class Factory(AbstractProvider[types.T_co]):
Expand All @@ -49,7 +57,7 @@ def __init__( # noqa: PLR0913
scope: enum.IntEnum | types.UnsetType = types.UNSET,
bound_type: type | types.UnsetType | None = types.UNSET,
kwargs: dict[str, typing.Any] | None = None,
cache: bool | CacheSettings[types.T_co] | None = None,
cache: bool | CacheSettings[types.T_co] = False,
skip_creator_parsing: bool = False,
) -> None:
if skip_creator_parsing:
Expand Down Expand Up @@ -83,7 +91,7 @@ def __init__( # noqa: PLR0913
bound_type=parsed.return_type.arg_type if isinstance(bound_type, types.UnsetType) else bound_type,
)
self._creator = creator
self.cache_settings = CacheSettings.coerce(cache)
self.cache_settings = CacheSettings._coerce(cache) # noqa: SLF001
self._kwargs = kwargs
self._cached_definition_site: str | types.UnsetType | None = types.UNSET

Expand Down Expand Up @@ -145,7 +153,7 @@ def _compute_definition_site(self) -> str | None:
return None
return f"{module}:{lineno}"

def resolution_step(self) -> exceptions.ResolutionStep:
def _resolution_step(self) -> exceptions.ResolutionStep:
return exceptions.ResolutionStep(scope=self.scope, name=self.display_name, location=self.definition_site)

def _argument_resolution_error(
Expand All @@ -161,11 +169,11 @@ def _argument_resolution_error(
member_types=item.args,
)

def wiring_plan(self, registry: "ProvidersRegistry") -> WiringPlan:
def _wiring_plan(self, registry: "ProvidersRegistry") -> WiringPlan:
"""Return this factory's wiring plan, memoized on the tree-wide providers registry."""
return registry.plan_for(self, self._parsed_kwargs, self._kwargs)

def can_call_positionally(self, plan: WiringPlan) -> bool:
def _can_call_positionally(self, plan: WiringPlan) -> bool:
"""Whether this creator can be called positionally under `plan`.

True when every parsed parameter is a positional-or-keyword provider dependency, in signature
Expand All @@ -180,12 +188,12 @@ def can_call_positionally(self, plan: WiringPlan) -> bool:
return False
return not (names and self._has_positional_only_gap)

def get_dependencies(self, container: "Container") -> dict[str, "AbstractProvider[typing.Any]"]:
def _get_dependencies(self, container: "Container") -> dict[str, "AbstractProvider[typing.Any]"]:
"""Return parameter name → dependency provider: a pure registry lookup, no scope or cache touched."""
return self.wiring_plan(container.providers_registry).edges
return self._wiring_plan(container.providers_registry).edges

def iter_validation_issues(self, container: "Container") -> typing.Iterable[Exception]:
def _iter_validation_issues(self, container: "Container") -> typing.Iterable[Exception]:
"""Yield ArgumentResolutionError for parameters with no provider, no default, no static kwarg."""
plan = self.wiring_plan(container.providers_registry)
plan = self._wiring_plan(container.providers_registry)
for name, item in plan.unwireable:
yield self._argument_resolution_error(arg_name=name, item=item, registry=container.providers_registry)
8 changes: 4 additions & 4 deletions modern_di/resolver_compiler.py
Original file line number Diff line number Diff line change
Expand Up @@ -184,10 +184,10 @@ def _code(arity: int, names: tuple[str, ...] | None, static: bool, cached: bool)


def _compile_factory(f: "Factory[typing.Any]", registry: "ProvidersRegistry") -> "Resolver":
plan = f.wiring_plan(registry)
plan = f._wiring_plan(registry)
if plan.unwireable:
return _compile_unwireable_factory(f, plan)
positional = f.can_call_positionally(plan)
positional = f._can_call_positionally(plan)
code, arg_lines = _code(
len(plan.provider_kwargs),
None if positional else tuple(plan.provider_kwargs),
Expand All @@ -199,7 +199,7 @@ def _compile_factory(f: "Factory[typing.Any]", registry: "ProvidersRegistry") ->
"pid": f.provider_id,
"scope": f.scope,
"creator": f._creator,
"resolution_step": f.resolution_step,
"resolution_step": f._resolution_step,
"edges": plan.provider_kwargs,
"arg_lines": arg_lines,
"static": plan.static_kwargs,
Expand Down Expand Up @@ -229,7 +229,7 @@ def resolve(_: "Container") -> typing.Any:
def _compile_unwireable_factory(f: "Factory[typing.Any]", plan: "WiringPlan") -> "Resolver":
"""Compile a resolver that always raises for the factory's first unwireable parameter, freshly built per call."""
scope = f.scope
resolution_step = f.resolution_step
resolution_step = f._resolution_step
build_error = f._argument_resolution_error
arg_name, item = plan.unwireable[0]

Expand Down
6 changes: 3 additions & 3 deletions tests/providers/test_alias.py
Original file line number Diff line number Diff line change
Expand Up @@ -424,14 +424,14 @@ def test_redirect_target_default_none() -> None:
class X: ...

factory = providers.Factory(scope=Scope.APP, creator=X)
assert factory.redirect_target(None) is None # ty: ignore[invalid-argument-type]
assert factory._redirect_target(None) is None # ty: ignore[invalid-argument-type]


def test_alias_redirect_target_returns_source() -> None:
container = Container(groups=[MyGroup])
source = container.providers_registry.find_provider(PostgresRepository)
assert source is not None
target = MyGroup.abstract_repo.redirect_target(container)
target = MyGroup.abstract_repo._redirect_target(container)
assert target is not None
assert target.provider_id == source.provider_id

Expand All @@ -441,7 +441,7 @@ class G(Group):
abstract = providers.Alias(source_type=PostgresRepository, bound_type=AbstractRepository)

container = Container(groups=[G])
assert G.abstract.redirect_target(container) is None
assert G.abstract._redirect_target(container) is None


# A genuine scope inversion through an alias must still raise, even measured against a custom
Expand Down
20 changes: 13 additions & 7 deletions tests/providers/test_factory.py
Original file line number Diff line number Diff line change
Expand Up @@ -770,10 +770,9 @@ class G(Group):
assert G.f.cache_settings is None


@pytest.mark.parametrize("cache_value", [False, None])
def test_cache_falsy_disables_caching(cache_value: bool | None) -> None:
def test_cache_false_disables_caching() -> None:
class G(Group):
f = providers.Factory(creator=SimpleCreator, kwargs={"dep1": "x"}, cache=cache_value)
f = providers.Factory(creator=SimpleCreator, kwargs={"dep1": "x"}, cache=False)

container = Container(groups=[G])
container.open()
Expand Down Expand Up @@ -1226,12 +1225,19 @@ def _body_raises() -> None:

@pytest.mark.parametrize(
("cache", "expected"),
[(True, providers.CacheSettings()), (False, None), (None, None)],
[(True, providers.CacheSettings()), (False, 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(cache: bool, 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
assert providers.CacheSettings._coerce(settings) is settings


@pytest.mark.parametrize("cache", [None, 1, "yes"])
def test_factory_cache_rejects_anything_but_bool_or_cache_settings(cache: object) -> None:
with pytest.raises(TypeError, match=r"cache= takes a bool or a CacheSettings") as exc_info:
providers.Factory(SimpleCreator, cache=cache) # ty: ignore[invalid-argument-type]
assert repr(cache) in str(exc_info.value)
2 changes: 1 addition & 1 deletion tests/registries/test_providers_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -292,7 +292,7 @@ def hold_open(
return plan

monkeypatch.setattr(pr_mod.WiringPlan, "build", staticmethod(hold_open))
worker = threading.Thread(target=lambda: svc.wiring_plan(registry))
worker = threading.Thread(target=lambda: svc._wiring_plan(registry))
worker.start()
try:
assert built.wait(5), "plan build never reached the publication window"
Expand Down
9 changes: 7 additions & 2 deletions tests/test_container.py
Original file line number Diff line number Diff line change
Expand Up @@ -201,10 +201,10 @@ class Top:
class _CountingFactory(providers.Factory[Bottom]):
__slots__ = ()

def get_dependencies(self, container: Container) -> dict[str, AbstractProvider[typing.Any]]:
def _get_dependencies(self, container: Container) -> dict[str, AbstractProvider[typing.Any]]:
nonlocal call_count
call_count += 1
return super().get_dependencies(container)
return super()._get_dependencies(container)

bottom_provider = _CountingFactory(creator=Bottom)

Expand Down Expand Up @@ -971,6 +971,11 @@ class _UnknownProvider(AbstractProvider[object]):
container.resolve_provider(provider)


def test_abstract_provider_is_a_plain_class() -> None:
"""`resolve_dependency` runs `isinstance(x, AbstractProvider)` per call; ABCMeta would make it Python-level."""
assert type(AbstractProvider) is type


# --- validate() is the only trigger: construction, open(), resolve() and add_providers() never ----
# --- walk the graph on their own. -------------------------------------------------------------

Expand Down
9 changes: 9 additions & 0 deletions tests/test_packaging.py
Original file line number Diff line number Diff line change
Expand Up @@ -53,3 +53,12 @@ def test_modern_di_imports_without_typing_extensions() -> None:
)
assert result.returncode == 0, result.stderr
assert "OK REQUEST" in result.stdout


def test_package_init_imports_every_name_it_exports() -> None:
"""Every name in `modern_di.__all__` is bound by an import in `__init__` itself."""
tree = ast.parse((_PKG_ROOT / "__init__.py").read_text(encoding="utf-8"))
imported = {
alias.asname or alias.name for node in tree.body if isinstance(node, ast.ImportFrom) for alias in node.names
}
assert set(modern_di.__all__) <= imported
Loading
Loading