[BACKEND] Cache cycle checks (#1420)

Small speedup to make cycle validation faster.
This commit is contained in:
Jesse Bannon 2026-01-23 10:58:14 -08:00 committed by GitHub
parent c4e112e8d5
commit 264e458c1c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194

View file

@ -1,5 +1,6 @@
# pylint: disable=missing-raises-doc # pylint: disable=missing-raises-doc
import copy import copy
from collections import defaultdict
from typing import Dict from typing import Dict
from typing import List from typing import List
from typing import Optional from typing import Optional
@ -38,59 +39,71 @@ class Script:
``{ %custom_function: syntax }`` ``{ %custom_function: syntax }``
""" """
def _ensure_no_cycle( def _throw_cycle_error(
self, name: str, dep: str, deps: List[str], definitions: Dict[str, SyntaxTree] self, name: str, dep: str, deps: List[str], definitions: Dict[str, SyntaxTree]
): ):
if dep not in definitions: type_name, pre = (
return # does not exist, will throw downstream in parser ("custom functions", "%") if definitions is self._functions else ("variables", "")
)
cycle_deps = [name] + deps + [dep]
cycle_deps_str = " -> ".join([f"{pre}{name}" for name in cycle_deps])
if name in deps + [dep]: raise CycleDetected(f"Cycle detected within these {type_name}: {cycle_deps_str}")
type_name, pre = (
("custom functions", "%") if definitions is self._functions else ("variables", "")
)
cycle_deps = [name] + deps + [dep]
cycle_deps_str = " -> ".join([f"{pre}{name}" for name in cycle_deps])
raise CycleDetected(f"Cycle detected within these {type_name}: {cycle_deps_str}")
def _traverse_variable_dependencies( def _traverse_variable_dependencies(
self, self,
variable_name: str, variable_name: str,
variable_dependency: SyntaxTree, variable_dependency: SyntaxTree,
deps: List[str], deps: List[str],
ensured: Dict[str, Set[str]],
) -> None: ) -> None:
for dep in variable_dependency.variables: for dep in variable_dependency.variables:
self._ensure_no_cycle( if variable_name == dep.name:
name=variable_name, dep=dep.name, deps=deps, definitions=self._variables self._throw_cycle_error(
) name=variable_name, dep=dep.name, deps=deps, definitions=self._variables
)
if dep.name in ensured[variable_name]:
continue
self._traverse_variable_dependencies( self._traverse_variable_dependencies(
variable_name=variable_name, variable_name=variable_name,
variable_dependency=self._variables[dep.name], variable_dependency=self._variables[dep.name],
deps=deps + [dep.name], deps=deps + [dep.name],
ensured=ensured,
) )
ensured[variable_name].add(dep.name)
for custom_func in variable_dependency.custom_function_dependencies( for custom_func in variable_dependency.custom_function_dependencies(
custom_function_definitions=self._functions custom_function_definitions=self._functions
): ):
for dep in self._functions[custom_func.name].variables: for dep in self._functions[custom_func.name].variables:
self._ensure_no_cycle( if variable_name == dep.name:
name=variable_name, self._throw_cycle_error(
dep=dep.name, name=variable_name,
deps=deps + [custom_func.definition_name()], dep=dep.name,
definitions=self._variables, deps=deps + [custom_func.definition_name()],
) definitions=self._variables,
)
if dep.name in ensured[variable_name]:
continue
self._traverse_variable_dependencies( self._traverse_variable_dependencies(
variable_name=variable_name, variable_name=variable_name,
variable_dependency=self._variables[dep.name], variable_dependency=self._variables[dep.name],
deps=deps + [custom_func.definition_name(), dep.name], deps=deps + [custom_func.definition_name(), dep.name],
ensured=ensured,
) )
ensured[variable_name].add(dep.name)
def _ensure_no_variable_cycles(self, variables: Dict[str, SyntaxTree]): def _ensure_no_variable_cycles(self, variables: Dict[str, SyntaxTree]):
ensured: Dict[str, Set[str]] = defaultdict(set)
for variable_name, variable_definition in variables.items(): for variable_name, variable_definition in variables.items():
self._traverse_variable_dependencies( self._traverse_variable_dependencies(
variable_name=variable_name, variable_name=variable_name,
variable_dependency=variable_definition, variable_dependency=variable_definition,
deps=[], deps=[],
ensured=ensured,
) )
def _traverse_custom_function_dependencies( def _traverse_custom_function_dependencies(
@ -100,9 +113,10 @@ class Script:
deps: List[str], deps: List[str],
) -> None: ) -> None:
for dep in custom_function_dependency.custom_functions: for dep in custom_function_dependency.custom_functions:
self._ensure_no_cycle( if custom_function_name == dep.name:
name=custom_function_name, dep=dep.name, deps=deps, definitions=self._functions self._throw_cycle_error(
) name=custom_function_name, dep=dep.name, deps=deps, definitions=self._functions
)
self._traverse_custom_function_dependencies( self._traverse_custom_function_dependencies(
custom_function_name=custom_function_name, custom_function_name=custom_function_name,
custom_function_dependency=self._functions[dep.name], custom_function_dependency=self._functions[dep.name],