num args into input functions validated
This commit is contained in:
parent
38c85562a7
commit
dc8b600d9b
5 changed files with 117 additions and 6 deletions
|
|
@ -387,7 +387,10 @@ class _Parser:
|
||||||
"custom function."
|
"custom function."
|
||||||
)
|
)
|
||||||
|
|
||||||
return CustomFunction(name=function_name, args=function_args)
|
return CustomFunction(
|
||||||
|
name=function_name,
|
||||||
|
args=function_args,
|
||||||
|
)
|
||||||
|
|
||||||
# Go back one so the parent function can close using the ')'
|
# Go back one so the parent function can close using the ')'
|
||||||
self._pos -= 1
|
self._pos -= 1
|
||||||
|
|
|
||||||
|
|
@ -4,6 +4,7 @@ from typing import Optional
|
||||||
from typing import Set
|
from typing import Set
|
||||||
|
|
||||||
from ytdl_sub.script.parser import parse
|
from ytdl_sub.script.parser import parse
|
||||||
|
from ytdl_sub.script.types.function import CustomFunction
|
||||||
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
|
||||||
|
|
@ -109,6 +110,39 @@ class Script:
|
||||||
f"do not increment from $0 to ${len(indices) - 1}."
|
f"do not increment from $0 to ${len(indices) - 1}."
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _ensure_custom_function_usage_num_input_arguments_valid(self):
|
||||||
|
for variable_name, variable_definition in self._variables.items():
|
||||||
|
for nested_custom_function in variable_definition.custom_functions:
|
||||||
|
if nested_custom_function.num_input_args != (
|
||||||
|
expected_num_args := len(
|
||||||
|
self._functions[nested_custom_function.name].function_arguments
|
||||||
|
)
|
||||||
|
):
|
||||||
|
raise InvalidCustomFunctionArguments(
|
||||||
|
f"Variable {variable_name} has invalid usage of the custom "
|
||||||
|
f"function %{nested_custom_function.name}: Expects {expected_num_args} "
|
||||||
|
f"argument{'s' if expected_num_args > 1 else ''} but received "
|
||||||
|
f"{nested_custom_function.num_input_args}"
|
||||||
|
)
|
||||||
|
|
||||||
|
for function_name, function_definition in self._functions.items():
|
||||||
|
for nested_custom_function in function_definition.custom_functions:
|
||||||
|
if nested_custom_function.name == function_name:
|
||||||
|
# Do not need to validate a cycle that should not exist
|
||||||
|
continue
|
||||||
|
|
||||||
|
if nested_custom_function.num_input_args != (
|
||||||
|
expected_num_args := len(
|
||||||
|
self._functions[nested_custom_function.name].function_arguments
|
||||||
|
)
|
||||||
|
):
|
||||||
|
raise InvalidCustomFunctionArguments(
|
||||||
|
f"Custom function %{function_name} has invalid usage of the custom "
|
||||||
|
f"function %{nested_custom_function.name}: Expects {expected_num_args} "
|
||||||
|
f"argument{'s' if expected_num_args > 1 else ''} but received "
|
||||||
|
f"{nested_custom_function.num_input_args}"
|
||||||
|
)
|
||||||
|
|
||||||
def __init__(self, overrides: Dict[str, str]):
|
def __init__(self, overrides: Dict[str, str]):
|
||||||
function_names: Set[str] = {
|
function_names: Set[str] = {
|
||||||
self._function_name(name) for name in overrides.keys() if self._is_function(name)
|
self._function_name(name) for name in overrides.keys() if self._is_function(name)
|
||||||
|
|
@ -144,6 +178,7 @@ class Script:
|
||||||
self._ensure_no_custom_function_cycles()
|
self._ensure_no_custom_function_cycles()
|
||||||
self._ensure_custom_function_arguments_valid()
|
self._ensure_custom_function_arguments_valid()
|
||||||
self._ensure_no_variable_cycles()
|
self._ensure_no_variable_cycles()
|
||||||
|
self._ensure_custom_function_usage_num_input_arguments_valid()
|
||||||
|
|
||||||
def resolve(
|
def resolve(
|
||||||
self, pre_resolved_variables: Optional[Dict[Variable, Resolvable]] = None
|
self, pre_resolved_variables: Optional[Dict[Variable, Resolvable]] = None
|
||||||
|
|
|
||||||
|
|
@ -134,8 +134,13 @@ class String(ResolvableT[str], Hashable, Argument):
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
class NamedCustomFunction(Argument, ABC):
|
class NamedCustomFunction(NamedArgument, ABC):
|
||||||
name: str
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class ParsedCustomFunction(NamedCustomFunction):
|
||||||
|
num_input_args: int
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
|
|
|
||||||
|
|
@ -7,7 +7,9 @@ from typing import Set
|
||||||
from typing import final
|
from typing import final
|
||||||
|
|
||||||
from ytdl_sub.script.types.resolvable import Argument
|
from ytdl_sub.script.types.resolvable import Argument
|
||||||
|
from ytdl_sub.script.types.resolvable import FunctionType
|
||||||
from ytdl_sub.script.types.resolvable import NamedCustomFunction
|
from ytdl_sub.script.types.resolvable import NamedCustomFunction
|
||||||
|
from ytdl_sub.script.types.resolvable import ParsedCustomFunction
|
||||||
from ytdl_sub.script.types.resolvable import Resolvable
|
from ytdl_sub.script.types.resolvable import Resolvable
|
||||||
from ytdl_sub.script.types.variable import FunctionArgument
|
from ytdl_sub.script.types.variable import FunctionArgument
|
||||||
from ytdl_sub.script.types.variable import Variable
|
from ytdl_sub.script.types.variable import Variable
|
||||||
|
|
@ -47,12 +49,16 @@ class VariableDependency(ABC):
|
||||||
|
|
||||||
@final
|
@final
|
||||||
@property
|
@property
|
||||||
def custom_functions(self) -> Set[NamedCustomFunction]:
|
def custom_functions(self) -> Set[ParsedCustomFunction]:
|
||||||
output: Set[NamedCustomFunction] = set()
|
output: Set[ParsedCustomFunction] = set()
|
||||||
for arg in self._iterable_arguments:
|
for arg in self._iterable_arguments:
|
||||||
if isinstance(arg, NamedCustomFunction):
|
if isinstance(arg, NamedCustomFunction):
|
||||||
|
if not isinstance(arg, FunctionType):
|
||||||
|
# A NamedCustomFunction should also always be a FunctionType
|
||||||
|
raise UNREACHABLE
|
||||||
|
|
||||||
# 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(ParsedCustomFunction(name=arg.name, num_input_args=len(arg.args)))
|
||||||
if isinstance(arg, VariableDependency):
|
if isinstance(arg, VariableDependency):
|
||||||
output.update(arg.custom_functions)
|
output.update(arg.custom_functions)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -156,3 +156,65 @@ class TestCustomFunction:
|
||||||
"%func1": f"{{[{arguments}]}}",
|
"%func1": f"{{[{arguments}]}}",
|
||||||
}
|
}
|
||||||
).resolve()
|
).resolve()
|
||||||
|
|
||||||
|
def test_custom_function_uses_custom_function_wrong_number_of_arguments(self):
|
||||||
|
with pytest.raises(
|
||||||
|
InvalidCustomFunctionArguments,
|
||||||
|
match=re.escape(
|
||||||
|
"Custom function %func0 has invalid usage of the custom function %func1: "
|
||||||
|
"Expects 1 argument but received 2"
|
||||||
|
),
|
||||||
|
):
|
||||||
|
Script(
|
||||||
|
{
|
||||||
|
"%func1": "{%mul(1, $0)}",
|
||||||
|
"%func0": "{%mul(%func1(1, 2), $0)}",
|
||||||
|
"output": "{%func0(1)}",
|
||||||
|
}
|
||||||
|
).resolve()
|
||||||
|
|
||||||
|
def test_custom_function_uses_custom_function_wrong_number_of_arguments_plural(self):
|
||||||
|
with pytest.raises(
|
||||||
|
InvalidCustomFunctionArguments,
|
||||||
|
match=re.escape(
|
||||||
|
"Custom function %func0 has invalid usage of the custom function %func1: "
|
||||||
|
"Expects 2 arguments but received 1"
|
||||||
|
),
|
||||||
|
):
|
||||||
|
Script(
|
||||||
|
{
|
||||||
|
"%func1": "{%mul($1, $0)}",
|
||||||
|
"%func0": "{%mul(%func1(1), $0)}",
|
||||||
|
"output": "{%func0(1)}",
|
||||||
|
}
|
||||||
|
).resolve()
|
||||||
|
|
||||||
|
def test_variable_uses_custom_function_wrong_number_of_arguments(self):
|
||||||
|
with pytest.raises(
|
||||||
|
InvalidCustomFunctionArguments,
|
||||||
|
match=re.escape(
|
||||||
|
"Variable output has invalid usage of the custom function %func1: "
|
||||||
|
"Expects 1 argument but received 2"
|
||||||
|
),
|
||||||
|
):
|
||||||
|
Script(
|
||||||
|
{
|
||||||
|
"%func1": "{%mul(1, $0)}",
|
||||||
|
"output": "{%mul(%func1(1, 2), 1)}",
|
||||||
|
}
|
||||||
|
).resolve()
|
||||||
|
|
||||||
|
def test_variable_uses_custom_function_wrong_number_of_arguments_plural(self):
|
||||||
|
with pytest.raises(
|
||||||
|
InvalidCustomFunctionArguments,
|
||||||
|
match=re.escape(
|
||||||
|
"Variable output has invalid usage of the custom function %func1: "
|
||||||
|
"Expects 2 arguments but received 1"
|
||||||
|
),
|
||||||
|
):
|
||||||
|
Script(
|
||||||
|
{
|
||||||
|
"%func1": "{%mul($1, $0)}",
|
||||||
|
"output": "{%mul(%func1(1), 1)}",
|
||||||
|
}
|
||||||
|
).resolve()
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue