diff --git a/src/ytdl_sub/script/parser.py b/src/ytdl_sub/script/parser.py index 37d1c7ae..3284774e 100644 --- a/src/ytdl_sub/script/parser.py +++ b/src/ytdl_sub/script/parser.py @@ -108,8 +108,9 @@ def _is_boolean_false(string: Optional[str]) -> bool: class _Parser: - def __init__(self, text: str): + def __init__(self, text: str, custom_function_name: Optional[str]): self._text = text + self._custom_function_name = custom_function_name self._pos = 0 self._error_highlight_pos = 0 self._ast: List[ArgumentType] = [] @@ -193,7 +194,9 @@ class _Parser: if not var_name: raise StringFormattingException("invalid var name") - return FunctionArgument(name=f"${var_name}") + return FunctionArgument.from_idx( + idx=int(var_name), custom_function_name=self._custom_function_name + ) def _parse_numeric(self) -> Integer | Float: numeric_string = "" @@ -499,11 +502,11 @@ class _Parser: return SyntaxTree(ast=self._ast) -def parse(text: str) -> SyntaxTree: +def parse(text: str, custom_function_name: Optional[str] = None) -> SyntaxTree: """ Entrypoint for parsing ytdl-sub code into a Syntax Tree """ - return _Parser(text).ast + return _Parser(text=text, custom_function_name=custom_function_name).ast # pylint: enable=invalid-name diff --git a/src/ytdl_sub/script/script.py b/src/ytdl_sub/script/script.py index 47f9f78d..157a1128 100644 --- a/src/ytdl_sub/script/script.py +++ b/src/ytdl_sub/script/script.py @@ -21,11 +21,18 @@ class Script: @classmethod def _function_name(cls, function_key: str) -> str: + """ + Drop the % in %custom_function + """ return function_key[1:] def __init__(self, overrides: Dict[str, str]): self._functions: Dict[str, SyntaxTree] = { - self._function_name(function_key): parse(function_value) + # custom_function_name must be passed to properly type custom function + # arguments uniquely if they're nested (i.e. $0 to $custom_func___0) + self._function_name(function_key): parse( + text=function_value, custom_function_name=self._function_name(function_key) + ) for function_key, function_value in overrides.items() if self._is_function(function_key) } diff --git a/src/ytdl_sub/script/types/function.py b/src/ytdl_sub/script/types/function.py index ccbee2e3..488e82e3 100644 --- a/src/ytdl_sub/script/types/function.py +++ b/src/ytdl_sub/script/types/function.py @@ -97,7 +97,9 @@ class CustomFunction(Function): resolved_variables_with_args = copy.deepcopy(resolved_variables) for i, arg in enumerate(resolved_args): - function_arg = FunctionArgument(name=f"${i}") # Function args are 0-based + function_arg = FunctionArgument.from_idx( + idx=i, custom_function_name=self.name + ) # Function args are 0-based if function_arg in resolved_variables_with_args: raise StringFormattingException("nested custom functions???") resolved_variables_with_args[function_arg] = arg diff --git a/src/ytdl_sub/script/types/variable.py b/src/ytdl_sub/script/types/variable.py index 9b4d7948..98626d84 100644 --- a/src/ytdl_sub/script/types/variable.py +++ b/src/ytdl_sub/script/types/variable.py @@ -1,4 +1,5 @@ from dataclasses import dataclass +from typing import Optional from ytdl_sub.script.types.resolvable import ArgumentType @@ -12,4 +13,8 @@ class Variable(ArgumentType): class FunctionArgument(Variable): """Arguments for custom functions, i.e. $0, $1, etc""" - pass + @classmethod + def from_idx(cls, idx: int, custom_function_name: Optional[str]) -> "FunctionArgument": + if custom_function_name: + return FunctionArgument(name=f"${custom_function_name}___{idx}") + return FunctionArgument(name=f"${idx}") diff --git a/src/ytdl_sub/script/utils/type_checking.py b/src/ytdl_sub/script/utils/type_checking.py index 452353dc..138fb96a 100644 --- a/src/ytdl_sub/script/utils/type_checking.py +++ b/src/ytdl_sub/script/utils/type_checking.py @@ -9,6 +9,7 @@ from ytdl_sub.script.types.resolvable import ArgumentType from ytdl_sub.script.types.resolvable import FunctionType from ytdl_sub.script.types.resolvable import NamedType from ytdl_sub.script.types.resolvable import Resolvable +from ytdl_sub.script.types.resolvable import TypeHintedFunctionType from ytdl_sub.script.types.variable import Variable from ytdl_sub.script.utils.exceptions import UNREACHABLE @@ -52,8 +53,10 @@ def is_type_compatible( True if arg is compatible with expected_arg_type. False otherwise. """ arg_type: Type[NamedType] = arg.__class__ - if isinstance(arg, FunctionType): - arg_type = arg.output_type() + if isinstance(arg, FunctionType) and isinstance(arg, TypeHintedFunctionType): + arg_type = arg.output_type() # built-in function + elif isinstance(arg, FunctionType): + return True # custom-function, can be anything, so pass for now elif isinstance(arg, Variable): return True # unresolved variables can be anything, so pass for now diff --git a/tests/unit/script/types/test_function.py b/tests/unit/script/types/test_function.py index 241053a9..54a8edce 100644 --- a/tests/unit/script/types/test_function.py +++ b/tests/unit/script/types/test_function.py @@ -117,12 +117,12 @@ class TestFunction: ): Script({"dne": "{%throw}"}).resolve() - def test_lambda_function(self): + def test_lambda_with_custom_function(self): assert Script( {"%times_two": "{%mul($0, 2)}", "wip": "{%array_apply([1, 2, 3], %times_two)}"} ).resolve() == {"wip": ResolvedArray([Integer(2), Integer(4), Integer(6)])} - def test_conditional_lambda_function(self): + def test_conditional_lambda_with_custom_functions(self): assert Script( { "%times_three": "{%mul($0, 3)}", @@ -130,3 +130,22 @@ class TestFunction: "wip": "{%array_apply([1, 2, 3], %if(False, %times_two, %times_three))}", } ).resolve() == {"wip": ResolvedArray([Integer(3), Integer(6), Integer(9)])} + + def test_nested_custom_functions(self): + assert Script( + { + "%times_three": "{%mul($0, 3)}", + "%times_two": "{%mul($0, 2)}", + "identity": "{%times_three(%times_two(1))}", + } + ).resolve() == {"identity": Integer(6)} + + def test_nested_custom_functions_within_custom_functions(self): + assert Script( + { + "%power_2": "{%mul($0, 2)}", + "%power_3": "{%mul(%power_2($0), 2)}", + "%power_4": "{%mul(%power_3($0), 2)}", + "power_of_4": "{%power_4(2)}", + } + ).resolve() == {"power_of_4": Integer(16)}