Skip to content

Commit e434792

Browse files
authored
Merge pull request #650 from reagento/fix/decorate_override
Check override works after decoration
2 parents 530ced8 + 946156b commit e434792

7 files changed

Lines changed: 204 additions & 70 deletions

File tree

src/dishka/code_tools/factory_compiler.py

Lines changed: 45 additions & 38 deletions
Original file line numberDiff line numberDiff line change
@@ -170,52 +170,53 @@ def _context_factory_body(
170170
),
171171
)
172172

173-
174173
def _selector_factory_body(
175174
builder: FactoryBuilder, source_call: str, factory: Factory,
176175
) -> None:
176+
error_call = builder.call(
177+
builder.global_(NoActiveFactoryError),
178+
builder.global_(factory.provides),
179+
builder.global_(factory.when_dependencies, "when_dependencies"),
180+
)
181+
builder.raise_(error_call)
182+
183+
ASYNC_TYPES = (FactoryType.ASYNC_FACTORY, FactoryType.ASYNC_GENERATOR)
184+
BODY_GENERATORS = {
185+
FactoryType.FACTORY: _sync_factory_body,
186+
FactoryType.ASYNC_FACTORY: _async_factory_body,
187+
FactoryType.GENERATOR: _generator_body,
188+
FactoryType.ASYNC_GENERATOR: _async_generator_body,
189+
FactoryType.CONTEXT: _context_factory_body,
190+
FactoryType.VALUE: _value_factory_body,
191+
FactoryType.ALIAS: _alias_factory_body,
192+
FactoryType.SELECTOR: _selector_factory_body,
193+
}
194+
195+
196+
def _select_when_dependency(
197+
builder: FactoryBuilder,
198+
factory: Factory,
199+
) -> bool:
200+
"""return True if there is assignment in any case"""
177201
first = True
178202
for variant in factory.when_dependencies:
179203
condition = builder.when(variant.when_override, factory.when_component)
180204
solved_value = builder.getter(variant.provides)
181205
if first and not condition:
182206
builder.assign_solved(solved_value)
207+
return True
183208
elif first:
184209
with builder.if_(condition):
185210
builder.assign_solved(solved_value)
186-
first = False
187211
elif not condition:
188212
with builder.else_():
189213
builder.assign_solved(solved_value)
190-
first = True
214+
return True
191215
else:
192216
with builder.elif_(condition):
193217
builder.assign_solved(solved_value)
194-
# if-chain not closed with else or not generated at all
195-
if not first or not factory.when_dependencies:
196-
error_call = builder.call(
197-
builder.global_(NoActiveFactoryError),
198-
builder.global_(factory.provides),
199-
builder.global_(factory.when_dependencies, "when_dependencies"),
200-
)
201-
if first: # no options at all
202-
builder.raise_(error_call)
203-
else:
204-
with builder.else_():
205-
builder.raise_(error_call)
206-
207-
208-
ASYNC_TYPES = (FactoryType.ASYNC_FACTORY, FactoryType.ASYNC_GENERATOR)
209-
BODY_GENERATORS = {
210-
FactoryType.FACTORY: _sync_factory_body,
211-
FactoryType.ASYNC_FACTORY: _async_factory_body,
212-
FactoryType.GENERATOR: _generator_body,
213-
FactoryType.ASYNC_GENERATOR: _async_generator_body,
214-
FactoryType.CONTEXT: _context_factory_body,
215-
FactoryType.VALUE: _value_factory_body,
216-
FactoryType.ALIAS: _alias_factory_body,
217-
FactoryType.SELECTOR: _selector_factory_body,
218-
}
218+
first = False
219+
return False
219220

220221

221222
def compile_factory(*, factory: Factory, is_async: bool) -> CompiledFactory:
@@ -228,16 +229,22 @@ def compile_factory(*, factory: Factory, is_async: bool) -> CompiledFactory:
228229
builder.register_provides(factory.provides)
229230

230231
with builder.make_getter():
231-
source_call = builder.call(
232-
builder.global_(factory.source),
233-
*(builder.getter(dep) for dep in factory.dependencies),
234-
**{
235-
name: builder.getter(dep)
236-
for name, dep in factory.kw_dependencies.items()
237-
},
238-
)
239-
body_generator = BODY_GENERATORS[factory.type]
240-
body_generator(builder, source_call, factory)
232+
if not _select_when_dependency(builder, factory):
233+
source_call = builder.call(
234+
builder.global_(factory.source),
235+
*(builder.getter(dep) for dep in factory.dependencies),
236+
**{
237+
name: builder.getter(dep)
238+
for name, dep in factory.kw_dependencies.items()
239+
},
240+
)
241+
body_generator = BODY_GENERATORS[factory.type]
242+
if factory.when_dependencies: # conditions generated
243+
with builder.else_():
244+
body_generator(builder, source_call, factory)
245+
else: # no options at all
246+
body_generator(builder, source_call, factory)
247+
241248
if factory.cache:
242249
builder.cache()
243250
builder.return_("solved")

src/dishka/dependency_source/decorator.py

Lines changed: 2 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -42,8 +42,6 @@ def as_factory(
4242
new_dependency: DependencyKey,
4343
cache: bool,
4444
component: Component,
45-
when_override: BaseMarker | None,
46-
when_active: BaseMarker | None,
4745
) -> Factory:
4846
typevar_replacement = get_typevar_replacement(
4947
self.provides.type_hint,
@@ -71,8 +69,8 @@ def as_factory(
7169
},
7270
type_=self.factory.type,
7371
cache=cache,
74-
when_override=combine_when(self.when, when_override),
75-
when_active=combine_when(self.when, when_active),
72+
when_override=self.when,
73+
when_active=self.when,
7674
when_component=self.factory.when_component or component,
7775
when_dependencies=[],
7876
)

src/dishka/dependency_source/factory.py

Lines changed: 23 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,8 @@
44
Mapping,
55
Sequence,
66
)
7-
from typing import Any
7+
from enum import Enum
8+
from typing import Any, TypeAlias, TypeVar
89

910
from dishka.entities.component import Component
1011
from dishka.entities.factory_type import FactoryData, FactoryType
@@ -13,6 +14,19 @@
1314
from dishka.entities.scope import BaseScope
1415

1516

17+
class Special(Enum):
18+
OMITTED = "omitted"
19+
20+
21+
T = TypeVar("T")
22+
MayBe: TypeAlias = Special | T
23+
24+
def coalesce(a: MayBe[T], b: T) -> T:
25+
if a is Special.OMITTED:
26+
return b
27+
return a
28+
29+
1630
class Factory(FactoryData):
1731
__slots__ = (
1832
"cache",
@@ -140,19 +154,22 @@ def with_scope(self, scope: BaseScope) -> Factory:
140154

141155
def replace(
142156
self,
143-
provides: DependencyKey | None = None,
157+
provides: MayBe[DependencyKey] = Special.OMITTED,
158+
when_active: MayBe[BaseMarker|None] = Special.OMITTED,
159+
when_override: MayBe[BaseMarker|None] = Special.OMITTED,
160+
when_component: MayBe[Component] = Special.OMITTED,
144161
) -> Factory:
145162
return Factory(
146163
dependencies=list(self.dependencies),
147164
kw_dependencies=dict(self.kw_dependencies),
148165
source=self.source,
149-
provides=provides or self.provides,
166+
provides=coalesce(provides, self.provides),
150167
scope=self.scope,
151168
is_to_bind=self.is_to_bind,
152169
cache=self.cache,
153170
type_=self.type,
154-
when_override=self.when_override,
155-
when_active=self.when_active,
156-
when_component=self.when_component,
171+
when_override=coalesce(when_override, self.when_override),
172+
when_active=coalesce(when_active, self.when_active),
173+
when_component=coalesce(when_component, self.when_component),
157174
when_dependencies=self.when_dependencies,
158175
)

src/dishka/registry_builder.py

Lines changed: 28 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -229,7 +229,6 @@ def _ensure_override_flags(
229229
) -> None:
230230
if self.skip_validation:
231231
return
232-
233232
if (
234233
not prev_factory and
235234
self.validation_settings.nothing_overridden and
@@ -253,8 +252,10 @@ def _unite_factories_group(
253252
group: list[Factory],
254253
) -> dict[DependencyKey, list[Factory]]:
255254
if len(group) == 1:
256-
self._ensure_override_flags(group[0], None)
257-
return {}
255+
factory = group[0]
256+
self._ensure_override_flags(factory, None)
257+
if factory.when_override in (None, BoolMarker(True)):
258+
return {}
258259

259260
when_dependencies: list[Factory] = []
260261
moved_factories: dict[DependencyKey, list[Factory]] = {}
@@ -279,10 +280,12 @@ def _unite_factories_group(
279280
new_factory = factory.replace(provides=new_provides)
280281
moved_factories[new_provides] = [new_factory]
281282
when_dependencies.append(new_factory)
282-
if len(moved_factories) == 1:
283-
self.processed_factories[provides] = [
284-
cast(Factory, prev_factory), # at least one factory found
285-
]
283+
if (
284+
len(moved_factories) == 1 and
285+
prev_factory and # at least one factory found
286+
prev_factory.when_override in (None, BoolMarker(True))
287+
):
288+
self.processed_factories[provides] = [prev_factory]
286289
return {}
287290

288291
scope = max(
@@ -435,8 +438,6 @@ def _process_generic_decorator(
435438
new_dependency=provides,
436439
cache=False,
437440
component=provider.component,
438-
when_override=None,
439-
when_active=None,
440441
)],
441442
)
442443

@@ -456,8 +457,6 @@ def _process_normal_decorator(
456457
new_dependency=provides,
457458
cache=False,
458459
component=provider.component,
459-
when_active=None,
460-
when_override=None,
461460
)],
462461
)
463462
self._decorate_factory(
@@ -497,8 +496,6 @@ def _decorate_factory(
497496
f"no factory for {provides}",
498497
)
499498

500-
if decorator.when not in (None, BoolMarker(True)):
501-
group_replacement.extend(old_group)
502499
for old_factory in old_group:
503500

504501
depth = self.decorator_depth[provides]
@@ -515,18 +512,32 @@ def _decorate_factory(
515512
):
516513
return
517514

518-
new_factory = old_factory.replace(provides=decorated_provides)
515+
new_factory = old_factory.replace(
516+
provides=decorated_provides,
517+
when_active=None,
518+
when_override=None,
519+
)
519520
decorated_groups[decorated_provides] = [new_factory]
520521
decorated_factory = decorator.as_factory(
521522
scope=cast(BaseScope, old_factory.scope),
522523
new_dependency=decorated_provides,
523524
cache=old_factory.cache,
524525
component=provides.component,
525-
when_override=old_factory.when_override,
526+
).replace(
527+
provides=provides,
526528
when_active=old_factory.when_active,
527-
).replace(provides=provides)
529+
when_override=old_factory.when_override,
530+
when_component=cast(Component, old_factory.when_component),
531+
)
532+
if decorator.when is not None:
533+
conditional_factory = new_factory.replace(
534+
when_override=~decorator.when,
535+
when_active=~decorator.when,
536+
when_component=provides.component,
537+
)
538+
decorated_factory.when_dependencies=[conditional_factory]
539+
self._register_when(conditional_factory)
528540
group_replacement.append(decorated_factory)
529-
self._register_when(decorated_factory)
530541

531542
self.processed_factories[provides] = group_replacement
532543
self.processed_factories.update(decorated_groups)

tests/unit/container/test_decorator.py

Lines changed: 48 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44

55
from dishka import (
66
DEFAULT_COMPONENT,
7+
STRICT_VALIDATION,
78
DependencyKey,
89
Has,
910
Provider,
@@ -13,7 +14,11 @@
1314
make_container,
1415
provide,
1516
)
16-
from dishka.exceptions import CycleDependenciesError, NoFactoryError
17+
from dishka.exceptions import (
18+
CycleDependenciesError,
19+
ImplicitOverrideDetectedError,
20+
NoFactoryError,
21+
)
1722

1823

1924
class A:
@@ -398,3 +403,45 @@ def dec(self, a: A) -> A:
398403

399404
with pytest.raises(NoFactoryError):
400405
make_container(MyProvider())
406+
407+
408+
def test_decorate_override():
409+
class MyProvider(Provider):
410+
@provide
411+
def make_str(self) -> str:
412+
return "a"
413+
414+
@provide(override=True)
415+
def make_str2(self) -> str:
416+
return "b"
417+
418+
@decorate
419+
def decorate_str(self, old_value: str) -> str:
420+
return old_value + "d"
421+
422+
c = make_container(
423+
MyProvider(scope=Scope.APP),
424+
validation_settings=STRICT_VALIDATION,
425+
)
426+
assert c.get(str) == "bd"
427+
428+
429+
def test_decorate_override_implicit():
430+
class MyProvider(Provider):
431+
@provide
432+
def make_str(self) -> str:
433+
return "a"
434+
435+
@provide
436+
def make_str2(self) -> str:
437+
return "b"
438+
439+
@decorate
440+
def decorate_str(self, old_value: str) -> str:
441+
return old_value + "d"
442+
443+
with pytest.raises(ImplicitOverrideDetectedError):
444+
make_container(
445+
MyProvider(scope=Scope.APP),
446+
validation_settings=STRICT_VALIDATION,
447+
)

tests/unit/container/test_exceptions.py

Lines changed: 13 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -335,11 +335,20 @@ def test_no_active_factory_smoke(path_len: int, variants_count: int) -> None:
335335

336336

337337
@pytest.mark.parametrize("is_async", [True, False])
338+
@pytest.mark.parametrize(("has_variants", "variants_count"), [
339+
([float], 1),
340+
([float, complex], 2),
341+
])
338342
@pytest.mark.asyncio
339-
async def test_no_active_factory(*, is_async: bool) -> None:
343+
async def test_no_active_factory(
344+
*,
345+
is_async: bool,
346+
has_variants: list[object],
347+
variants_count: int,
348+
) -> None:
340349
provider = Provider(scope=Scope.APP)
341-
provider.provide(int, when=Has(float))
342-
provider.provide(int, when=Has(complex))
350+
for variant in has_variants:
351+
provider.provide(int, when=Has(variant))
343352

344353
@provider.provide
345354
def get_str(value: int) -> str:
@@ -356,5 +365,5 @@ def get_str(value: int) -> str:
356365

357366

358367
assert str(e.value)
359-
assert len(e.value.variants) == 2
368+
assert len(e.value.variants) == variants_count
360369
assert len(e.value.path) == 2

0 commit comments

Comments
 (0)