Skip to content

Commit 2a7c085

Browse files
committed
cache nodes, limit recursion depth
1 parent 48152f1 commit 2a7c085

3 files changed

Lines changed: 55 additions & 35 deletions

File tree

src/dishka/container.py

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -51,7 +51,7 @@ def __init__(
5151
):
5252
self.registry = registry
5353
self.child_registries = child_registries
54-
self._context = {DependencyKey(type(self), DEFAULT_COMPONENT): self}
54+
self._context = {CONTAINER_KEY: self}
5555
if context:
5656
for key, value in context.items():
5757
if not isinstance(key, DependencyKey):
@@ -252,3 +252,6 @@ def make_container(
252252
close_parent=True,
253253
)
254254
return container
255+
256+
257+
CONTAINER_KEY = DependencyKey(Container, DEFAULT_COMPONENT)

src/dishka/graph_compiler.py

Lines changed: 27 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,9 @@
1111
from .text_rendering import get_name
1212

1313

14+
MAX_DEPTH = 5 # max code depth, otherwise we get too big file
15+
16+
1417
class Node(FactoryData):
1518
__slots__ = (
1619
"dependencies",
@@ -106,7 +109,6 @@ def make_args(args: list[str], kwargs: dict[str, str]) -> str:
106109
FactoryType.VALUE: VALUE,
107110
FactoryType.CONTEXT: CONTEXT,
108111
FactoryType.ALIAS: ALIAS,
109-
None: GO_PARENT,
110112
}
111113
FUNC_TEMPLATE = """
112114
{async_}def {func_name}(getter, exits, context):
@@ -116,19 +118,19 @@ def make_args(args: list[str], kwargs: dict[str, str]) -> str:
116118
"""
117119

118120
IF_TEMPLATE = """
119-
if {var} := cache_getter({key}, None):
120-
pass # cache found
121-
else:
121+
if ({var} := cache_getter({key}, ...)) is ...:
122122
{deps}
123123
{body}
124124
{cache}
125125
"""
126126
CACHE = "context[{key}] = {var}"
127127

128-
128+
builtins = {getattr(__builtins__, name): name for name in dir(__builtins__)}
129129
def make_name(obj: Any, ns: dict[Any, str]) -> str:
130+
if obj in builtins:
131+
return builtins[obj]
130132
if isinstance(obj, DependencyKey):
131-
key = get_name(obj.type_hint, include_module=False) + obj.component
133+
key = get_name(obj.type_hint, include_module=False) +"_"+ obj.component
132134
else:
133135
key = get_name(obj, include_module=False)
134136
key = re.sub(r"\W", "_", key)
@@ -153,24 +155,33 @@ def make_var(node: Node, ns: dict[Any, str]):
153155

154156

155157
def make_if(
156-
node: Node, node_var: str, ns: dict[Any, str], is_async: bool,
158+
node: Node, node_var: str, ns: dict[Any, str],
159+
is_async: bool,
160+
depth: int,
157161
) -> str:
158162
node_key = ns[node.provides]
159163
node_source = ns[node.source]
164+
if depth > MAX_DEPTH or node.type is None:
165+
if is_async:
166+
return GO_PARENT.format(
167+
var=node_var,
168+
key=node_real_key,
169+
)
170+
else:
171+
return GO_PARENT.format(
172+
var=node_var,
173+
key=node_key,
174+
)
160175

161176
deps = "".join(
162-
make_if(dep, make_var(dep, ns), ns, is_async)
177+
make_if(dep, make_var(dep, ns), ns, is_async, depth+1)
163178
for dep in node.dependencies
164179
)
165180
deps += "".join(
166-
make_if(dep, make_var(dep, ns), ns, is_async)
181+
make_if(dep, make_var(dep, ns), ns, is_async, depth+1)
167182
for dep in node.kw_dependencies.values()
168183
)
169184
deps = indent(deps, " ")
170-
if node.cache:
171-
cache = CACHE.format(var=node_var, key=node_key)
172-
else:
173-
cache = "# no cache"
174185

175186
args = [make_var(dep, ns) for dep in node.dependencies]
176187
kwargs = {
@@ -192,6 +203,7 @@ def make_if(
192203
)
193204

194205
if node.cache:
206+
cache = CACHE.format(var=node_var, key=node_key)
195207
body_str = indent(body_str, " ")
196208
return IF_TEMPLATE.format(
197209
var=node_var,
@@ -201,14 +213,14 @@ def make_if(
201213
cache=cache,
202214
)
203215
else:
204-
return "\n".join([deps, body_str, cache])
216+
return "\n".join([deps, body_str])
205217

206218

207219
def make_func(
208220
node: Node, ns: dict[Any, str], func_name: str, is_async: bool,
209221
) -> str:
210222
node_var = make_var(node, ns)
211-
body = make_if(node, node_var, ns, is_async)
223+
body = make_if(node, node_var, ns, is_async, 0)
212224
body = indent(body, " ")
213225
return FUNC_TEMPLATE.format(
214226
async_="async " if is_async else "",

src/dishka/registry.py

Lines changed: 24 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,8 @@
1+
import time
12
from collections.abc import Callable
3+
from linecache import cache
24
from typing import Any, TypeVar, get_args, get_origin
35

4-
from pydantic.v1 import compiled
56

67
from ._adaptix.type_tools.fundamentals import get_type_vars
78
from .container_objects import CompiledFactory
@@ -12,7 +13,6 @@
1213
from .entities.factory_type import FactoryType
1314
from .entities.key import DependencyKey
1415
from .entities.scope import BaseScope
15-
from .factory_compiler import compile_factory
1616
from .graph_compiler import Node, compile_graph
1717

1818

@@ -153,10 +153,12 @@ def _specialize_generic(
153153
)
154154

155155

156-
def make_node(registry: Registry, key: DependencyKey) -> Node:
156+
def make_node(registry: Registry, key: DependencyKey, cache: dict| None = None) -> Node:
157+
if cache is None:
158+
cache = {}
157159
factory = registry.get_factory(key)
158160
if not factory:
159-
return Node(
161+
node = Node(
160162
provides=key,
161163
scope=registry.scope,
162164
type_=None,
@@ -165,18 +167,21 @@ def make_node(registry: Registry, key: DependencyKey) -> Node:
165167
cache=False,
166168
source=None,
167169
)
168-
return Node(
169-
provides=factory.provides,
170-
scope=factory.scope,
171-
source=factory.source,
172-
type_=factory.type,
173-
cache=factory.cache,
174-
dependencies=[
175-
make_node(registry, dep)
176-
for dep in factory.dependencies
177-
],
178-
kw_dependencies={
179-
key: make_node(registry, dep)
180-
for key, dep in factory.kw_dependencies.items()
181-
},
182-
)
170+
else:
171+
node = Node(
172+
provides=factory.provides,
173+
scope=factory.scope,
174+
source=factory.source,
175+
type_=factory.type,
176+
cache=factory.cache,
177+
dependencies=[
178+
make_node(registry, dep, cache)
179+
for dep in factory.dependencies
180+
],
181+
kw_dependencies={
182+
key: make_node(registry, dep, cache)
183+
for key, dep in factory.kw_dependencies.items()
184+
},
185+
)
186+
cache[key] = node
187+
return node

0 commit comments

Comments
 (0)