improve cycle WIP
This commit is contained in:
parent
d3e127783e
commit
333cd329a1
4 changed files with 87 additions and 7 deletions
|
|
@ -22,6 +22,7 @@ from ytdl_sub.script.types.resolvable import AnyTypeReturnableB
|
||||||
from ytdl_sub.script.types.resolvable import ArgumentType
|
from ytdl_sub.script.types.resolvable import ArgumentType
|
||||||
from ytdl_sub.script.types.resolvable import FunctionType
|
from ytdl_sub.script.types.resolvable import FunctionType
|
||||||
from ytdl_sub.script.types.resolvable import Lambda
|
from ytdl_sub.script.types.resolvable import Lambda
|
||||||
|
from ytdl_sub.script.types.resolvable import NamedCustomFunction
|
||||||
from ytdl_sub.script.types.resolvable import Resolvable
|
from ytdl_sub.script.types.resolvable import Resolvable
|
||||||
from ytdl_sub.script.types.resolvable import TypeHintedFunctionType
|
from ytdl_sub.script.types.resolvable import TypeHintedFunctionType
|
||||||
from ytdl_sub.script.types.variable import FunctionArgument
|
from ytdl_sub.script.types.variable import FunctionArgument
|
||||||
|
|
@ -29,6 +30,7 @@ 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
|
||||||
|
|
@ -50,12 +52,15 @@ class Function(FunctionType, VariableDependency, ABC):
|
||||||
return CustomFunction(name=name, args=args)
|
return CustomFunction(name=name, args=args)
|
||||||
|
|
||||||
|
|
||||||
class CustomFunction(Function):
|
class CustomFunction(Function, NamedCustomFunction):
|
||||||
def resolve(
|
def resolve(
|
||||||
self,
|
self,
|
||||||
resolved_variables: Dict[Variable, Resolvable],
|
resolved_variables: Dict[Variable, Resolvable],
|
||||||
custom_functions: Dict[str, "VariableDependency"],
|
custom_functions: Dict[str, "VariableDependency"],
|
||||||
) -> Resolvable:
|
) -> Resolvable:
|
||||||
|
if NamedCustomFunction(name=self.name) in self.custom_functions:
|
||||||
|
raise CycleDetected("ackkk!")
|
||||||
|
|
||||||
resolved_args: List[Resolvable] = [
|
resolved_args: List[Resolvable] = [
|
||||||
self._resolve_argument_type(
|
self._resolve_argument_type(
|
||||||
arg=arg, resolved_variables=resolved_variables, custom_functions=custom_functions
|
arg=arg, resolved_variables=resolved_variables, custom_functions=custom_functions
|
||||||
|
|
@ -72,8 +77,8 @@ class CustomFunction(Function):
|
||||||
function_arg = FunctionArgument.from_idx(
|
function_arg = FunctionArgument.from_idx(
|
||||||
idx=i, custom_function_name=self.name
|
idx=i, custom_function_name=self.name
|
||||||
) # Function args are 0-based
|
) # Function args are 0-based
|
||||||
if function_arg in resolved_variables_with_args:
|
# if function_arg in resolved_variables_with_args:
|
||||||
raise StringFormattingException("nested custom functions???")
|
# raise StringFormattingException("nested custom functions???")
|
||||||
resolved_variables_with_args[function_arg] = arg
|
resolved_variables_with_args[function_arg] = arg
|
||||||
|
|
||||||
return custom_functions[self.name].resolve(
|
return custom_functions[self.name].resolve(
|
||||||
|
|
|
||||||
|
|
@ -1,13 +1,16 @@
|
||||||
|
import copy
|
||||||
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 Optional
|
from typing import Optional
|
||||||
|
|
||||||
from ytdl_sub.script.types.resolvable import ArgumentType
|
from ytdl_sub.script.types.resolvable import ArgumentType
|
||||||
|
from ytdl_sub.script.types.resolvable import NamedCustomFunction
|
||||||
from ytdl_sub.script.types.resolvable import Resolvable
|
from ytdl_sub.script.types.resolvable import Resolvable
|
||||||
from ytdl_sub.script.types.resolvable import String
|
from ytdl_sub.script.types.resolvable import String
|
||||||
from ytdl_sub.script.types.variable import Variable
|
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.exceptions import CycleDetected
|
||||||
from ytdl_sub.utils.exceptions import StringFormattingException
|
from ytdl_sub.utils.exceptions import StringFormattingException
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -41,6 +44,47 @@ class SyntaxTree(VariableDependency):
|
||||||
# Otherwise, to concat multiple resolved outputs, we must concat as strings
|
# Otherwise, to concat multiple resolved outputs, we must concat as strings
|
||||||
return String("".join([str(res) for res in resolved]))
|
return String("".join([str(res) for res in resolved]))
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _get_custom_function_dependencies(
|
||||||
|
cls,
|
||||||
|
custom_function_name: str,
|
||||||
|
custom_function_dependency: "SyntaxTree",
|
||||||
|
custom_functions: Dict[str, "SyntaxTree"],
|
||||||
|
deps: List[str],
|
||||||
|
) -> List[str]:
|
||||||
|
deps = copy.deepcopy(deps) # do not work with references
|
||||||
|
|
||||||
|
for dep in custom_function_dependency.custom_functions:
|
||||||
|
# Skip leaf functions since they will never cause a cycle
|
||||||
|
if not custom_functions[dep.name].custom_functions:
|
||||||
|
continue
|
||||||
|
|
||||||
|
deps.append(dep.name)
|
||||||
|
|
||||||
|
if custom_function_name in deps:
|
||||||
|
cycle_deps = [custom_function_name] + deps[0 : deps.index(custom_function_name) + 1]
|
||||||
|
cycle_deps_str = " -> ".join([f"%{name}" for name in cycle_deps])
|
||||||
|
raise CycleDetected(f"Custom functions contain a cycle: {cycle_deps_str}")
|
||||||
|
|
||||||
|
deps += cls._get_custom_function_dependencies(
|
||||||
|
custom_function_name=custom_function_name,
|
||||||
|
custom_function_dependency=custom_functions[dep.name],
|
||||||
|
custom_functions=custom_functions,
|
||||||
|
deps=deps,
|
||||||
|
)
|
||||||
|
|
||||||
|
return deps
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _ensure_no_custom_function_cycles(cls, custom_functions: Dict[str, "SyntaxTree"]):
|
||||||
|
for custom_function_name, custom_function in custom_functions.items():
|
||||||
|
_ = cls._get_custom_function_dependencies(
|
||||||
|
custom_function_name=custom_function_name,
|
||||||
|
custom_function_dependency=custom_function,
|
||||||
|
custom_functions=custom_functions,
|
||||||
|
deps=[],
|
||||||
|
)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def resolve_overrides(
|
def resolve_overrides(
|
||||||
cls,
|
cls,
|
||||||
|
|
@ -57,6 +101,8 @@ class SyntaxTree(VariableDependency):
|
||||||
pre_resolved_variables if pre_resolved_variables else {}
|
pre_resolved_variables if pre_resolved_variables else {}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
cls._ensure_no_custom_function_cycles(custom_functions=custom_functions)
|
||||||
|
|
||||||
while unresolved_variables:
|
while unresolved_variables:
|
||||||
unresolved_count: int = len(unresolved_variables)
|
unresolved_count: int = len(unresolved_variables)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -29,7 +29,7 @@ class VariableDependency(ABC):
|
||||||
for arg in self._iterable_arguments:
|
for arg in self._iterable_arguments:
|
||||||
if isinstance(arg, Variable):
|
if isinstance(arg, Variable):
|
||||||
output.add(arg)
|
output.add(arg)
|
||||||
elif isinstance(arg, VariableDependency):
|
if isinstance(arg, VariableDependency):
|
||||||
output.update(arg.variables)
|
output.update(arg.variables)
|
||||||
|
|
||||||
return output
|
return output
|
||||||
|
|
@ -41,7 +41,7 @@ class VariableDependency(ABC):
|
||||||
for arg in self._iterable_arguments:
|
for arg in self._iterable_arguments:
|
||||||
if isinstance(arg, FunctionArgument):
|
if isinstance(arg, FunctionArgument):
|
||||||
output.add(arg)
|
output.add(arg)
|
||||||
elif isinstance(arg, VariableDependency):
|
if isinstance(arg, VariableDependency):
|
||||||
output.update(arg.function_arguments)
|
output.update(arg.function_arguments)
|
||||||
|
|
||||||
return output
|
return output
|
||||||
|
|
@ -54,7 +54,7 @@ class VariableDependency(ABC):
|
||||||
if isinstance(arg, NamedCustomFunction):
|
if isinstance(arg, NamedCustomFunction):
|
||||||
# Custom funcs aren't hashable, so recreate just the base-class portion
|
# Custom funcs aren't hashable, so recreate just the base-class portion
|
||||||
output.add(NamedCustomFunction(name=arg.name))
|
output.add(NamedCustomFunction(name=arg.name))
|
||||||
elif isinstance(arg, VariableDependency):
|
if isinstance(arg, VariableDependency):
|
||||||
output.update(arg.custom_functions)
|
output.update(arg.custom_functions)
|
||||||
|
|
||||||
return output
|
return output
|
||||||
|
|
|
||||||
|
|
@ -128,7 +128,10 @@ class TestFunction:
|
||||||
|
|
||||||
def test_custom_function_chained_cycle(self):
|
def test_custom_function_chained_cycle(self):
|
||||||
with pytest.raises(
|
with pytest.raises(
|
||||||
CycleDetected, match=re.escape("The custom function %cycle_func cannot call itself.")
|
CycleDetected,
|
||||||
|
match=re.escape(
|
||||||
|
"Custom functions contain a cycle: %cycle_func1 -> %cycle_func0 -> %cycle_func1"
|
||||||
|
),
|
||||||
):
|
):
|
||||||
Script(
|
Script(
|
||||||
{
|
{
|
||||||
|
|
@ -138,6 +141,32 @@ class TestFunction:
|
||||||
}
|
}
|
||||||
).resolve()
|
).resolve()
|
||||||
|
|
||||||
|
def test_custom_function_deep_chained_cycle(self):
|
||||||
|
with pytest.raises(
|
||||||
|
CycleDetected,
|
||||||
|
match=re.escape(
|
||||||
|
"Custom functions contain a cycle: "
|
||||||
|
"%cycle_func4 -> "
|
||||||
|
"%cycle_func0 -> "
|
||||||
|
"%cycle_func1 -> "
|
||||||
|
"%cycle_func2 -> "
|
||||||
|
"%cycle_func3 -> "
|
||||||
|
"%cycle_func4"
|
||||||
|
),
|
||||||
|
):
|
||||||
|
Script(
|
||||||
|
{
|
||||||
|
"%nested_safe_func": "{%mul($0, 1)}",
|
||||||
|
"%safe_func": "{%nested_safe_func($0, 1)}",
|
||||||
|
"%cycle_func4": "{%mul(%cycle_func0(1), %safe_func($0))}",
|
||||||
|
"%cycle_func3": "{%mul(%cycle_func4(1), %safe_func($0))}",
|
||||||
|
"%cycle_func2": "{%mul(%cycle_func3(1), %safe_func($0))}",
|
||||||
|
"%cycle_func1": "{%mul(%cycle_func2(1), %safe_func($0))}",
|
||||||
|
"%cycle_func0": "{%mul(%cycle_func1(1), %safe_func($0))}",
|
||||||
|
"output": "{%cycle_func0(1)}",
|
||||||
|
}
|
||||||
|
).resolve()
|
||||||
|
|
||||||
def test_lambda_with_custom_function(self):
|
def test_lambda_with_custom_function(self):
|
||||||
assert Script(
|
assert Script(
|
||||||
{"%times_two": "{%mul($0, 2)}", "wip": "{%array_apply([1, 2, 3], %times_two)}"}
|
{"%times_two": "{%mul($0, 2)}", "wip": "{%array_apply([1, 2, 3], %times_two)}"}
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue