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
2 changes: 1 addition & 1 deletion AGENTS.md
Original file line number Diff line number Diff line change
Expand Up @@ -46,7 +46,7 @@ Every module under `modern_di/` is named for what it does; read it. What a singl
class here, never at the raise site. Submodules split by family; `__init__` re-exports every public name.
- `registries/` — `providers_registry` (type → provider, plus the shared plan/resolver memos) and
`overrides_registry` are shared tree-wide; `cache_registry` and `context_registry` are per-container.
- `dependency_graph.py` walks `WiringPlan.edges`, so what `validate()` traverses is exactly what
- `dependency_graph.py` walks `WiringPlan.provider_kwargs`, so what `validate()` traverses is exactly what
`resolve()` follows. Explicit-stack, never recursive: a caller runs it inside a `RecursionError`
handler near CPython's stack limit.
- `types.py` — `UNSET` is load-bearing on the resolve path: the miss marker for both the override
Expand Down
25 changes: 8 additions & 17 deletions modern_di/container.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,9 +3,8 @@
import threading
import typing

from modern_di import exceptions, types
from modern_di import dependency_graph, exceptions, types
from modern_di._scope_algebra import next_deeper
from modern_di.dependency_graph import DependencyGraph, build_cycle_error, collect_errors, redirect_hops
from modern_di.group import Group
from modern_di.providers.abstract import AbstractProvider
from modern_di.providers.container_provider import container_provider
Expand All @@ -20,17 +19,13 @@
def _handle_recursion_error(
provider: AbstractProvider[typing.Any], container: "Container", registry: ProvidersRegistry, exc: RecursionError
) -> typing.NoReturn:
"""Convert an escaped `RecursionError` to `CircularDependencyError`, or re-raise it unchanged.

A separate call, not inlined into `resolve_provider`: the coverage tracer re-arms on the
fresh call boundary before this raises.
"""
"""Convert an escaped `RecursionError` to `CircularDependencyError`, or re-raise it unchanged."""
if registry.is_validated():
raise exc # validated => acyclic static graph => genuine self-recursion
cycle = DependencyGraph().find_cycle_from(provider, container)
cycle = dependency_graph.find_cycle_from(provider, container)
if cycle is None:
raise exc
raise build_cycle_error(cycle, container) from exc
raise dependency_graph.build_cycle_error(cycle, container) from exc


class Container:
Expand Down Expand Up @@ -82,9 +77,7 @@ def __init__(
self._closed = False
self.scope = scope
self.parent_container = parent_container
# Ancestors only, never self: a `scope: self` entry is a reference cycle, so no container
# would ever be freed by refcounting.
# SLF001 exempts `self`/`cls` only, so it flags this same-class read; no boundary is crossed.
# Ancestors only: a `scope: self` entry is a reference cycle that refcounting never frees.
self._scope_map: dict[enum.IntEnum, typing.Self] = (
{**parent_container._scope_map, parent_container.scope: parent_container} # noqa: SLF001
if parent_container
Expand All @@ -93,9 +86,7 @@ def __init__(
self._cache_registry = CacheRegistry()
self._context_registry = ContextRegistry(copy.copy(context) if context is not None else {})
self._providers_registry: ProvidersRegistry
# Inlined rather than a helper: this runs per child build (benchmark `test_g6_build_child_container`).
if parent_container:
# SLF001 exempts `self`/`cls` only, so it flags this same-class read; no boundary is crossed.
self._lock = parent_container._lock # noqa: SLF001
self._providers_registry = parent_container._providers_registry # noqa: SLF001
else:
Expand Down Expand Up @@ -168,7 +159,7 @@ def resolve(self, dependency_type: type[types.T]) -> types.T:
except STEP_ERRORS as exc:
provider = registry.find_provider(dependency_type)
if provider is not None:
exc.prepend_step(*redirect_hops(provider, self))
exc.prepend_step(*dependency_graph.redirect_hops(provider, self))
raise

def resolve_dependency(self, dependency: "AbstractProvider[types.T] | type[types.T]") -> types.T:
Expand All @@ -194,7 +185,7 @@ def resolve_provider(self, provider: "AbstractProvider[types.T]") -> types.T:
except RecursionError as exc:
_handle_recursion_error(provider, self, registry, exc)
except STEP_ERRORS as exc:
exc.prepend_step(*redirect_hops(provider, self))
exc.prepend_step(*dependency_graph.redirect_hops(provider, self))
raise

def validate(self) -> None:
Expand All @@ -209,7 +200,7 @@ def validate(self) -> None:
if reg.is_validated():
return

if errors := collect_errors(self, reg):
if errors := dependency_graph.collect_errors(self, reg):
raise exceptions.ValidationFailedError(errors=errors)
reg.mark_validated()

Expand Down
150 changes: 73 additions & 77 deletions modern_di/dependency_graph.py
Original file line number Diff line number Diff line change
Expand Up @@ -101,89 +101,85 @@ def build_cycle_error(
)


class DependencyGraph:
"""Stateless walker over the static provider graph rooted at a container's registry."""

def walk(
self,
roots: "typing.Iterable[AbstractProvider[typing.Any]]",
container: "Container",
) -> "typing.Iterator[Event]":
"""Pre-order DFS from each root; bookkeeping is shared across roots, keyed on ``provider_id``."""
visiting: set[int] = set()
visited: set[int] = set()
for root in roots:
yield from self._walk_from(root, container, visiting, visited)

def find_cycle_from(
self,
start: "AbstractProvider[typing.Any]",
container: "Container",
) -> "list[AbstractProvider[typing.Any]] | None":
"""Return the first cycle reachable from ``start``, or None when that subgraph is acyclic."""
for event in self.walk([start], container):
if isinstance(event, Cycle):
return event.providers
return None

def _walk_from(
self,
start: "AbstractProvider[typing.Any]",
container: "Container",
visiting: set[int],
visited: set[int],
) -> "typing.Iterator[Event]":
"""Explicit-stack DFS from ``start``; skip immediately if already seen."""
if start.provider_id in visited or start.provider_id in visiting:
return

path: list[AbstractProvider[typing.Any]] = []
stack: list[typing.Iterator[tuple[str, AbstractProvider[typing.Any]]]] = []
yield from self._enter(start, container, visiting, path, stack)

while stack:
try:
name, dep = next(stack[-1])
except StopIteration:
finished = path.pop()
stack.pop()
visiting.discard(finished.provider_id)
visited.add(finished.provider_id)
continue

yield Edge(path[-1], name, dep)
if dep.provider_id in visiting:
cycle_start = next(i for i, p in enumerate(path) if p.provider_id == dep.provider_id)
yield Cycle([*path[cycle_start:], path[cycle_start]])
continue
if dep.provider_id in visited:
continue
yield from self._enter(dep, container, visiting, path, stack)

def _enter(
self,
provider: "AbstractProvider[typing.Any]",
container: "Container",
visiting: set[int],
path: "list[AbstractProvider[typing.Any]]",
stack: "list[typing.Iterator[tuple[str, AbstractProvider[typing.Any]]]]",
) -> "typing.Iterator[Event]":
"""Push ``provider`` onto the active path; a ``ResolutionError`` from it becomes ``DependenciesError``."""
visiting.add(provider.provider_id)
path.append(provider)
yield NodeEntered(provider)
def walk(
roots: "typing.Iterable[AbstractProvider[typing.Any]]",
container: "Container",
) -> "typing.Iterator[Event]":
"""Pre-order DFS from each root; bookkeeping is shared across roots, keyed on ``provider_id``."""
visiting: set[int] = set()
visited: set[int] = set()
for root in roots:
yield from _walk_from(root, container, visiting, visited)


def find_cycle_from(
start: "AbstractProvider[typing.Any]",
container: "Container",
) -> "list[AbstractProvider[typing.Any]] | None":
"""Return the first cycle reachable from ``start``, or None when that subgraph is acyclic."""
for event in walk([start], container):
if isinstance(event, Cycle):
return event.providers
return None


def _walk_from(
start: "AbstractProvider[typing.Any]",
container: "Container",
visiting: set[int],
visited: set[int],
) -> "typing.Iterator[Event]":
"""Explicit-stack DFS from ``start``; skip immediately if already seen."""
if start.provider_id in visited or start.provider_id in visiting:
return

path: list[AbstractProvider[typing.Any]] = []
stack: list[typing.Iterator[tuple[str, AbstractProvider[typing.Any]]]] = []
yield from _enter(start, container, visiting, path, stack)

while stack:
try:
dependencies = provider._get_dependencies(container) # noqa: SLF001
except exceptions.ResolutionError as exc:
yield DependenciesError(provider, exc)
dependencies = {}
stack.append(iter(dependencies.items()))
name, dep = next(stack[-1])
except StopIteration:
finished = path.pop()
stack.pop()
visiting.discard(finished.provider_id)
visited.add(finished.provider_id)
continue

yield Edge(path[-1], name, dep)
if dep.provider_id in visiting:
cycle_start = next(i for i, p in enumerate(path) if p.provider_id == dep.provider_id)
yield Cycle([*path[cycle_start:], path[cycle_start]])
continue
if dep.provider_id in visited:
continue
yield from _enter(dep, container, visiting, path, stack)


def _enter(
provider: "AbstractProvider[typing.Any]",
container: "Container",
visiting: set[int],
path: "list[AbstractProvider[typing.Any]]",
stack: "list[typing.Iterator[tuple[str, AbstractProvider[typing.Any]]]]",
) -> "typing.Iterator[Event]":
"""Push ``provider`` onto the active path; a ``ResolutionError`` from it becomes ``DependenciesError``."""
visiting.add(provider.provider_id)
path.append(provider)
yield NodeEntered(provider)
try:
dependencies = provider._get_dependencies(container) # noqa: SLF001
except exceptions.ResolutionError as exc:
yield DependenciesError(provider, exc)
dependencies = {}
stack.append(iter(dependencies.items()))


def collect_errors(container: "Container", registry: "ProvidersRegistry") -> list[Exception]:
"""Walk the graph rooted at ``registry``'s providers once; return every wiring error in walk order."""
errors: list[Exception] = []
for event in DependencyGraph().walk(registry, container):
for event in walk(registry, container):
match event:
case NodeEntered(provider):
errors.extend(provider._iter_validation_issues(container)) # noqa: SLF001
Expand Down
4 changes: 2 additions & 2 deletions modern_di/group.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,12 +6,12 @@


class Group:
def __new__(cls, *_: typing.Any, **__: typing.Any) -> typing.Self: # noqa: ANN401
def __new__(cls, *_: object, **__: object) -> typing.Self:
raise exceptions.GroupInstantiationError(group_name=cls.__name__)

_default_scope: typing.ClassVar["enum.IntEnum | None"] = None

def __init_subclass__(cls, scope: "enum.IntEnum | None" = None, **kwargs: typing.Any) -> None: # noqa: ANN401
def __init_subclass__(cls, scope: "enum.IntEnum | None" = None, **kwargs: object) -> None:
"""Record a group-default scope and stamp it onto scope-defaulted providers in this class body."""
super().__init_subclass__(**kwargs)
if scope is not None:
Expand Down
14 changes: 4 additions & 10 deletions modern_di/providers/alias.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,17 +27,11 @@ 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]":
source = container.find_provider(self._source_type)
def _get_dependencies(self, container: "Container") -> dict[str, "AbstractProvider[typing.Any]"]:
source = self._redirect_target(container)
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)}
return {"source": source}

def _redirect_target(self, container: "Container") -> "AbstractProvider[typing.Any] | None":
try:
return self._find_source(container)
except exceptions.AliasSourceNotRegisteredError:
return None
return container.find_provider(self._source_type)
Loading
Loading