better variable cycle detection

This commit is contained in:
Jesse Bannon 2023-11-22 12:37:27 -08:00
parent ef58092ff3
commit f74d43a67c
6 changed files with 54 additions and 14 deletions

View file

@ -47,6 +47,8 @@ class ArrayFunctions:
""" """
return Array(list(reversed(array.value))) return Array(list(reversed(array.value)))
# pylint: disable=unused-argument
@staticmethod @staticmethod
def array_apply(array: Array, lambda_function: Lambda) -> Array: def array_apply(array: Array, lambda_function: Lambda) -> Array:
""" """
@ -63,3 +65,5 @@ class ArrayFunctions:
return ResolvedArray( return ResolvedArray(
[ResolvedArray([Integer(idx), val]) for idx, val in enumerate(array.value)] [ResolvedArray([Integer(idx), val]) for idx, val in enumerate(array.value)]
) )
# pylint: enable=unused-argument

View file

@ -6,8 +6,11 @@ from ytdl_sub.script.parser import parse
from ytdl_sub.script.types.resolvable import Resolvable from ytdl_sub.script.types.resolvable import Resolvable
from ytdl_sub.script.types.syntax_tree import SyntaxTree from ytdl_sub.script.types.syntax_tree import SyntaxTree
from ytdl_sub.script.types.variable import Variable from ytdl_sub.script.types.variable import Variable
from ytdl_sub.script.utils.exceptions import UNREACHABLE
from ytdl_sub.script.utils.exceptions import CycleDetected from ytdl_sub.script.utils.exceptions import CycleDetected
# pylint: disable=missing-raises-doc
class Script: class Script:
""" """
@ -28,6 +31,32 @@ class Script:
""" """
return function_key[1:] return function_key[1:]
def _traverse_variable_dependencies(
self,
variable_name: str,
variable_dependency: SyntaxTree,
deps: List[str],
) -> None:
for dep in variable_dependency.variables:
if variable_name in deps + [dep.name]:
cycle_deps = [variable_name] + deps + [dep.name]
cycle_deps_str = " -> ".join(cycle_deps)
raise CycleDetected(f"Cycle detected within these variables: {cycle_deps_str}")
self._traverse_variable_dependencies(
variable_name=variable_name,
variable_dependency=self._variables[dep.name],
deps=deps + [dep.name],
)
def _ensure_no_variable_cycles(self):
for variable_name, variable_definition in self._variables.items():
self._traverse_variable_dependencies(
variable_name=variable_name,
variable_dependency=variable_definition,
deps=[],
)
def _traverse_custom_function_dependencies( def _traverse_custom_function_dependencies(
self, self,
custom_function_name: str, custom_function_name: str,
@ -38,7 +67,9 @@ class Script:
if custom_function_name in deps + [dep.name]: if custom_function_name in deps + [dep.name]:
cycle_deps = [custom_function_name] + deps + [dep.name] cycle_deps = [custom_function_name] + deps + [dep.name]
cycle_deps_str = " -> ".join([f"%{name}" for name in cycle_deps]) cycle_deps_str = " -> ".join([f"%{name}" for name in cycle_deps])
raise CycleDetected(f"Custom functions contain a cycle: {cycle_deps_str}") raise CycleDetected(
f"Cycle detected within these custom functions: {cycle_deps_str}"
)
self._traverse_custom_function_dependencies( self._traverse_custom_function_dependencies(
custom_function_name=custom_function_name, custom_function_name=custom_function_name,
@ -72,6 +103,7 @@ class Script:
} }
self._ensure_no_custom_function_cycles() self._ensure_no_custom_function_cycles()
self._ensure_no_variable_cycles()
def resolve( def resolve(
self, pre_resolved_variables: Optional[Dict[Variable, Resolvable]] = None self, pre_resolved_variables: Optional[Dict[Variable, Resolvable]] = None
@ -109,9 +141,8 @@ class Script:
unresolved_variables.remove(variable) unresolved_variables.remove(variable)
if len(unresolved_variables) == unresolved_count: if len(unresolved_variables) == unresolved_count:
raise CycleDetected( # Implies a cycle within the variables. Should never reach
f"Cycle detected within these variables: " # since cycles are detected in __init__
f"{', '.join(sorted([var.name for var in unresolved_variables]))}" raise UNREACHABLE
)
return {variable.name: resolvable for variable, resolvable in resolved_variables.items()} return {variable.name: resolvable for variable, resolvable in resolved_variables.items()}

View file

@ -7,13 +7,10 @@ from inspect import FullArgSpec
from typing import Callable from typing import Callable
from typing import Dict from typing import Dict
from typing import List from typing import List
from typing import Optional
from typing import Set
from typing import Type from typing import Type
from typing import Union from typing import Union
from ytdl_sub.script.functions import Functions from ytdl_sub.script.functions import Functions
from ytdl_sub.script.types.array import Array
from ytdl_sub.script.types.array import ResolvedArray from ytdl_sub.script.types.array import ResolvedArray
from ytdl_sub.script.types.array import UnresolvedArray from ytdl_sub.script.types.array import UnresolvedArray
from ytdl_sub.script.types.resolvable import Argument from ytdl_sub.script.types.resolvable import Argument
@ -30,7 +27,6 @@ from ytdl_sub.script.types.variable import Variable
from ytdl_sub.script.types.variable_dependency import VariableDependency from ytdl_sub.script.types.variable_dependency import VariableDependency
from ytdl_sub.script.utils.exception_formatters import FunctionArgumentsExceptionFormatter from ytdl_sub.script.utils.exception_formatters import FunctionArgumentsExceptionFormatter
from ytdl_sub.script.utils.exceptions import UNREACHABLE from ytdl_sub.script.utils.exceptions import UNREACHABLE
from ytdl_sub.script.utils.exceptions import CycleDetected
from ytdl_sub.script.utils.exceptions import FunctionDoesNotExist from ytdl_sub.script.utils.exceptions import FunctionDoesNotExist
from ytdl_sub.script.utils.exceptions import FunctionRuntimeException from ytdl_sub.script.utils.exceptions import FunctionRuntimeException
from ytdl_sub.script.utils.exceptions import UserThrownRuntimeError from ytdl_sub.script.utils.exceptions import UserThrownRuntimeError

View file

@ -2,7 +2,6 @@ import itertools
from dataclasses import dataclass from dataclasses import dataclass
from typing import Dict from typing import Dict
from typing import List from typing import List
from typing import Set
from ytdl_sub.script.types.resolvable import AnyArgument from ytdl_sub.script.types.resolvable import AnyArgument
from ytdl_sub.script.types.resolvable import Argument from ytdl_sub.script.types.resolvable import Argument

View file

@ -1,7 +1,10 @@
import re
import pytest import pytest
from ytdl_sub.script.script import Script from ytdl_sub.script.script import Script
from ytdl_sub.script.types.resolvable import String from ytdl_sub.script.types.resolvable import String
from ytdl_sub.script.utils.exceptions import CycleDetected
from ytdl_sub.utils.exceptions import StringFormattingException from ytdl_sub.utils.exceptions import StringFormattingException
@ -38,9 +41,15 @@ class TestSyntaxTree:
} }
def test_simple_cycle(self): def test_simple_cycle(self):
with pytest.raises(StringFormattingException): with pytest.raises(
CycleDetected,
match=re.escape("Cycle detected within these variables: " "a -> b -> a"),
):
Script({"a": "{b}", "b": "{a}"}).resolve() Script({"a": "{b}", "b": "{a}"}).resolve()
def test_simple_cycle_with_function(self): def test_simple_cycle_with_function(self):
with pytest.raises(StringFormattingException): with pytest.raises(
CycleDetected,
match=re.escape("Cycle detected within these variables: " "b -> b_ -> b"),
):
Script({"b": "{%capitalize(b_)}", "b_": "{b}"}).resolve() Script({"b": "{%capitalize(b_)}", "b_": "{b}"}).resolve()

View file

@ -29,7 +29,8 @@ class TestCustomFunction:
with pytest.raises( with pytest.raises(
CycleDetected, CycleDetected,
match=re.escape( match=re.escape(
"Custom functions contain a cycle: %cycle_func1 -> %cycle_func0 -> %cycle_func1" "Cycle detected within these custom functions: "
"%cycle_func1 -> %cycle_func0 -> %cycle_func1"
), ),
): ):
Script( Script(
@ -44,7 +45,7 @@ class TestCustomFunction:
with pytest.raises( with pytest.raises(
CycleDetected, CycleDetected,
match=re.escape( match=re.escape(
"Custom functions contain a cycle: " "Cycle detected within these custom functions: "
"%cycle_func4 -> " "%cycle_func4 -> "
"%cycle_func0 -> " "%cycle_func0 -> "
"%cycle_func1 -> " "%cycle_func1 -> "