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
40 changes: 3 additions & 37 deletions modern_di/container.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,16 +8,7 @@
from types import FrameType

from modern_di import exceptions, types
from modern_di.dependency_graph import (
Cycle,
DependenciesError,
DependencyGraph,
Edge,
NodeEntered,
build_cycle_error,
effective_scope,
terminal_chain,
)
from modern_di.dependency_graph import DependencyGraph, build_cycle_error, collect_errors
from modern_di.group import Group
from modern_di.providers.abstract import AbstractProvider
from modern_di.providers.container_provider import container_provider
Expand Down Expand Up @@ -215,30 +206,6 @@ def resolve_provider(self, provider: "AbstractProvider[types.T]") -> types.T:
except RecursionError as exc:
_handle_recursion_error(provider, self, exc)

def _walk_errors(self) -> list[Exception]:
"""Walk the graph once, returning every wiring error in walk order."""
errors: list[Exception] = []
graph = DependencyGraph()
for event in graph.walk(self.providers_registry, self):
match event:
case NodeEntered(provider):
errors.extend(provider.iter_validation_issues(self))
case DependenciesError(_, error):
errors.append(error)
case Edge(parent, name, dep):
dep_chain = terminal_chain(dep, self)
if dep_chain[-1].scope > effective_scope(parent, self):
errors.append(
exceptions.InvalidScopeDependencyError(
provider=parent,
parameter_name=name,
dep_chain=dep_chain,
)
)
case Cycle(providers):
errors.append(build_cycle_error(providers, self))
return errors

def validate(self) -> None:
"""Walk the static provider graph and raise on any wiring error.

Expand All @@ -251,9 +218,8 @@ def validate(self) -> None:
if reg.is_validated():
return

validation_errors = self._walk_errors()
if validation_errors:
raise exceptions.ValidationFailedError(errors=validation_errors)
if errors := collect_errors(self):
raise exceptions.ValidationFailedError(errors=errors)
reg.mark_validated()

def add_providers(self, *providers: AbstractProvider[typing.Any]) -> None:
Expand Down
24 changes: 24 additions & 0 deletions modern_di/dependency_graph.py
Original file line number Diff line number Diff line change
Expand Up @@ -180,3 +180,27 @@ def _enter(
yield DependenciesError(provider, exc)
dependencies = {}
stack.append(iter(dependencies.items()))


def collect_errors(container: "Container") -> list[Exception]:
"""Walk ``container``'s provider graph once and return every wiring error in walk order."""
errors: list[Exception] = []
for event in DependencyGraph().walk(container.providers_registry, container):
match event:
case NodeEntered(provider):
errors.extend(provider.iter_validation_issues(container))
case DependenciesError(_, error):
errors.append(error)
case Edge(parent, name, dep):
dep_chain = terminal_chain(dep, container)
if dep_chain[-1].scope > effective_scope(parent, container):
errors.append(
exceptions.InvalidScopeDependencyError(
provider=parent,
parameter_name=name,
dep_chain=dep_chain,
)
)
case Cycle(providers):
errors.append(build_cycle_error(providers, container))
return errors
5 changes: 3 additions & 2 deletions tests/test_container.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@

from modern_di import Container, Group, Scope, exceptions, providers, suggester
from modern_di import container as container_module
from modern_di.dependency_graph import collect_errors
from modern_di.exceptions import (
ArgumentResolutionError,
ChildContainerRegistrationError,
Expand Down Expand Up @@ -335,7 +336,7 @@ class G(Group):
assert CircularDependencyError in error_types


def test_walk_errors_returns_flat_list_in_walk_order() -> None:
def test_collect_errors_returns_flat_list_in_walk_order() -> None:
class _Missing: ...

@dataclasses.dataclass(kw_only=True, slots=True)
Expand All @@ -348,7 +349,7 @@ class G(Group):
svc = providers.Factory(creator=_NeedsMissing)

container = Container(scope=Scope.APP, groups=[G])
errors = container._walk_errors()
errors = collect_errors(container)

# Root order is registration order (a, b, svc): the cycle closes while walking from root
# `a`, so it is appended before `svc`'s missing dependency is reached.
Expand Down
Loading