Skip to content
Open
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
84 changes: 50 additions & 34 deletions test/unit/test_inject_from_container_optional_types.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,9 @@
import warnings
from dataclasses import dataclass
from typing import Annotated, Optional, Union
from typing import Annotated, Optional

import pytest
import wireup
from wireup._annotations import Inject, Injected
from wireup.errors import DuplicateServiceRegistrationError


class MaybeThing: ...
Expand Down Expand Up @@ -50,61 +49,78 @@ def main(
main()


def test_getting_optional_service_via_plain_type_emits_deprecation_warning() -> None:
@wireup.injectable
def test_getting_optional_service_via_plain_type_resolves_silently() -> None:
class Foo:
pass

@wireup.injectable
def make_foo() -> Foo | None:
return Foo()

container = wireup.create_sync_container(injectables=[make_foo])
container = wireup.create_sync_container(injectables=[wireup.injectable(make_foo)])

with pytest.warns(DeprecationWarning) as record:
with warnings.catch_warnings():
warnings.simplefilter("error")
inst = container.get(Foo)

assert len(record) == 1
assert "registered as optional" in str(record[0].message)

assert isinstance(inst, Foo)
assert inst is container.get(optional_hint(Foo))
assert inst is container.get(Foo | None)


def test_optional_factory_with_qualifier() -> None:
async def test_getting_optional_service_via_plain_type_resolves_silently_async() -> None:
class Foo:
pass

def make_foo() -> Foo | None:
return Foo()

container = wireup.create_async_container(injectables=[wireup.injectable(make_foo)])

with warnings.catch_warnings():
warnings.simplefilter("error")
inst = await container.get(Foo)

assert isinstance(inst, Foo)
assert inst is await container.get(Foo | None)


def test_getting_qualified_optional_service_via_plain_type_keeps_qualifier() -> None:
# https://github.com/maldoinc/wireup/issues/138
# A factory returning Optional[T] with a qualifier must preserve the qualifier
# on the backwards-compatibility alias factory created for the raw type T.
@wireup.injectable
class AuthContext:
# A qualified Optional[T] factory must keep its qualifier so container.get(T, qualifier=...)
# resolves it instead of failing with a spurious self-dependency error.
class Foo:
pass

@wireup.injectable
def require_authentication() -> AuthContext:
return AuthContext()
@wireup.injectable(qualifier="primary")
def make_foo() -> Foo | None:
return Foo()

@wireup.injectable(qualifier="optional")
def maybe_get_authentication() -> AuthContext | None:
return None
container = wireup.create_sync_container(injectables=[make_foo])

container = wireup.create_sync_container(
injectables=[require_authentication, maybe_get_authentication],
)
with warnings.catch_warnings():
warnings.simplefilter("error")
inst = container.get(Foo, qualifier="primary")

assert isinstance(container.get(AuthContext), AuthContext)
assert container.get(AuthContext | None, qualifier="optional") is None
assert isinstance(inst, Foo)
assert inst is container.get(Foo | None, qualifier="primary")


def test_registering_optional_and_plain_type_raises_duplicate() -> None:
@wireup.injectable
def test_registering_optional_and_plain_type_are_distinct() -> None:
class Foo:
pass

@wireup.injectable
def make_optional() -> Foo | None:
return None

# Registering both an Optional[T] factory and a T factory together raises since
# wireup will add a backwards-compatible factory for T when registering it as optional.
with pytest.raises(DuplicateServiceRegistrationError):
wireup.create_sync_container(injectables=[make_optional, Foo])
def make_plain() -> Foo:
return Foo()

# The Optional[T] factory and a plain T factory are distinct registration keys:
# T | None and T no longer collide, so retrieving each returns its own instance
# and container.get(T) resolves the plain factory directly (not the optional fallback).
container = wireup.create_sync_container(
injectables=[wireup.injectable(make_optional), wireup.injectable(make_plain)]
)

assert container.get(Foo | None) is None
assert isinstance(container.get(Foo), Foo)
3 changes: 3 additions & 0 deletions wireup/ioc/container/async_container.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,9 @@ async def get(

return await res if compiled_factory.is_async else res # type:ignore[no-any-return]

if (optional_klass := self._optional_compat_klass(klass, qualifier)) is not None:
return await self.get(optional_klass, qualifier)

raise UnknownServiceRequestedError(klass, qualifier)

async def close(self) -> None:
Expand Down
19 changes: 19 additions & 0 deletions wireup/ioc/container/base_container.py
Original file line number Diff line number Diff line change
Expand Up @@ -87,6 +87,22 @@ def override(self) -> OverrideManager:
"""Override registered container injectables with new values."""
return self._override_mgr

def _optional_compat_klass(self, klass: Callable[..., T] | None, qualifier: Qualifier | None) -> Any:
"""Backwards-compat fallback for ``container.get``.

A factory registered as ``Optional[T]`` is also retrievable via ``container.get(T)``.
Returns ``T | None`` when that registration exists (so the caller can resolve against it),
otherwise ``None``. Only applies to ``container.get``; injected dependencies resolve by
their actual annotated type.
"""
try:
optional_klass = klass | None # type: ignore[operator]
except TypeError:
return None

obj_id = optional_klass if qualifier is None else (optional_klass, qualifier)
return optional_klass if obj_id in self._factories else None

@overload
def _synchronous_get(self, klass: type[T], qualifier: Qualifier | None = None) -> T: ...
@overload
Expand Down Expand Up @@ -126,6 +142,9 @@ def _synchronous_get(

return compiled_factory.factory(self) # type:ignore[no-any-return]

if (optional_klass := self._optional_compat_klass(klass, qualifier)) is not None:
return self._synchronous_get(optional_klass, qualifier) # type: ignore[no-any-return]

raise UnknownServiceRequestedError(klass, qualifier)

def _recompile(self) -> None:
Expand Down
39 changes: 1 addition & 38 deletions wireup/ioc/registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,6 @@
get_container_object_id,
)
from wireup.ioc.util import ensure_is_type, get_callable_type, get_globals
from wireup.util import stringify_type

if TYPE_CHECKING:
from wireup._annotations import AbstractDeclaration, InjectableDeclaration
Expand Down Expand Up @@ -248,7 +247,7 @@ def _has_user_defined_mapping_key(self, collection_key: Any) -> bool:
def _register_sequence_collections(self) -> None:
for klass, qualifiers in dict(self.impls).items():
# Only real registration keys should contribute to collection members.
# Synthetic aliases like raw optional-compat or Sequence[T] keys must not participate here.
# Synthetic keys such as Sequence[T] must not participate here.
real_qualifiers = [
qualifier
for qualifier in qualifiers
Expand Down Expand Up @@ -376,42 +375,6 @@ def discover_interfaces(bases: tuple[type, ...]) -> None:
)
self.impls[klass].append(qualifier)

if type_analysis.is_optional:
# Backwards compatibility: In earlier versions when a factory returned T | None
# you could do container.get(T). Alias that type to the normalized T | None factory.
# Create a fake factory that warns and returns the original instance.
# https://github.com/maldoinc/wireup/commit/00590dc741035a4c7042c5b6fc434ed08e27f5c0
def compat_fn(raw_type_instance: Any) -> Any:
type_name = type_analysis.raw_type.__name__
deprecated_msg = (
f"Deprecated: {stringify_type(type_analysis.raw_type)} was registered as optional "
f"and retrieving it via container.get({type_name}) is deprecated. "
f"Please use container.get({type_name} | None) or container.get(Optional[{type_name}]) instead."
)

warnings.warn(deprecated_msg, DeprecationWarning, stacklevel=4)

return raw_type_instance

compat_fn.__signature__ = inspect.Signature( # type: ignore[attr-defined]
parameters=[
inspect.Parameter(
"raw_type_instance",
kind=inspect.Parameter.POSITIONAL_OR_KEYWORD,
annotation=Annotated[klass, InjectableQualifier(qualifier)],
)
],
)

self._register(
type_analysis.raw_type,
factory_fn=compat_fn,
lifetime=lifetime,
qualifier=qualifier,
auto_discover_interfaces=True,
is_synthetic_factory=True,
)

def _register_abstract(self, klass: type) -> None:
self.interfaces[klass] = {}

Expand Down
Loading