custom function dne

This commit is contained in:
Jesse Bannon 2023-11-22 13:23:42 -08:00
parent 5f9d1f11c4
commit c90b38ae92
6 changed files with 92 additions and 28 deletions

View file

@ -2,9 +2,13 @@ from enum import Enum
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 typing import Set
from ytdl_sub.script.functions import Functions
from ytdl_sub.script.types.array import UnresolvedArray from ytdl_sub.script.types.array import UnresolvedArray
from ytdl_sub.script.types.function import Argument from ytdl_sub.script.types.function import Argument
from ytdl_sub.script.types.function import BuiltInFunction
from ytdl_sub.script.types.function import CustomFunction
from ytdl_sub.script.types.function import Function from ytdl_sub.script.types.function import Function
from ytdl_sub.script.types.map import UnresolvedMap from ytdl_sub.script.types.map import UnresolvedMap
from ytdl_sub.script.types.resolvable import Boolean from ytdl_sub.script.types.resolvable import Boolean
@ -19,6 +23,7 @@ from ytdl_sub.script.types.variable import Variable
from ytdl_sub.script.utils.exception_formatters import ParserExceptionFormatter from ytdl_sub.script.utils.exception_formatters import ParserExceptionFormatter
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 CycleDetected
from ytdl_sub.script.utils.exceptions import FunctionDoesNotExist
from ytdl_sub.script.utils.exceptions import IncompatibleFunctionArguments from ytdl_sub.script.utils.exceptions import IncompatibleFunctionArguments
from ytdl_sub.script.utils.exceptions import InvalidSyntaxException from ytdl_sub.script.utils.exceptions import InvalidSyntaxException
from ytdl_sub.script.utils.exceptions import UserException from ytdl_sub.script.utils.exceptions import UserException
@ -109,9 +114,17 @@ def _is_boolean_false(string: Optional[str]) -> bool:
class _Parser: class _Parser:
def __init__(self, text: str, custom_function_name: Optional[str]): def __init__(
self,
text: str,
name: Optional[str],
custom_function_names: Optional[Set[str]],
variable_names: Optional[Set[str]],
):
self._text = text self._text = text
self._custom_function_name = custom_function_name self._name = name
self._custom_function_names = custom_function_names
self._variable_names = variable_names
self._pos = 0 self._pos = 0
self._error_highlight_pos = 0 self._error_highlight_pos = 0
self._ast: List[Argument] = [] self._ast: List[Argument] = []
@ -195,9 +208,7 @@ class _Parser:
if not var_name: if not var_name:
raise StringFormattingException("invalid var name") raise StringFormattingException("invalid var name")
return FunctionArgument.from_idx( return FunctionArgument.from_idx(idx=int(var_name), custom_function_name=self._name)
idx=int(var_name), custom_function_name=self._custom_function_name
)
def _parse_numeric(self) -> Integer | Float: def _parse_numeric(self) -> Integer | Float:
numeric_string = "" numeric_string = ""
@ -337,17 +348,33 @@ class _Parser:
if ch == ")": if ch == ")":
# Had '(' to indicate there are args # Had '(' to indicate there are args
if function_args is not None: if function_args is not None:
if self._custom_function_name == function_name: if self._name == function_name:
self._set_highlight_position(function_start_pos) self._set_highlight_position(function_start_pos)
raise CycleDetected( raise CycleDetected(
f"The custom function %{function_name} cannot call itself." f"The custom function %{function_name} cannot call itself."
) )
try: if Functions.is_built_in(function_name):
return Function.from_name_and_args(name=function_name, args=function_args) try:
except IncompatibleFunctionArguments: return BuiltInFunction(
name=function_name, args=function_args
).validate_args()
except IncompatibleFunctionArguments:
self._set_highlight_position(function_start_pos)
raise
# Is custom function
if (
self._custom_function_names is not None
and function_name not in self._custom_function_names
):
self._set_highlight_position(function_start_pos) self._set_highlight_position(function_start_pos)
raise raise FunctionDoesNotExist(
f"Function %{function_name} does not exist as a built-in or "
"custom function."
)
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
@ -507,11 +534,21 @@ class _Parser:
return SyntaxTree(ast=self._ast) return SyntaxTree(ast=self._ast)
def parse(text: str, custom_function_name: Optional[str] = None) -> SyntaxTree: def parse(
text: str,
name: Optional[str] = None,
custom_function_names: Optional[Set[str]] = None,
variable_names: Optional[Set[str]] = None,
) -> SyntaxTree:
""" """
Entrypoint for parsing ytdl-sub code into a Syntax Tree Entrypoint for parsing ytdl-sub code into a Syntax Tree
""" """
return _Parser(text=text, custom_function_name=custom_function_name).ast return _Parser(
text=text,
name=name,
custom_function_names=custom_function_names,
variable_names=variable_names,
).ast
# pylint: enable=invalid-name # pylint: enable=invalid-name

View file

@ -1,6 +1,7 @@
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 typing import Set
from ytdl_sub.script.parser import parse from ytdl_sub.script.parser import parse
from ytdl_sub.script.types.resolvable import Resolvable from ytdl_sub.script.types.resolvable import Resolvable
@ -70,6 +71,9 @@ 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:
if dep.name not in self._functions:
continue # does not exist, will throw downstream
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])
@ -92,20 +96,35 @@ class Script:
) )
def __init__(self, overrides: Dict[str, str]): def __init__(self, overrides: Dict[str, str]):
function_names: Set[str] = {
self._function_name(name) for name in overrides.keys() if self._is_function(name)
}
variable_names: Set[str] = {
name for name in overrides.keys() if not self._is_function(name)
}
self._functions: Dict[str, SyntaxTree] = { self._functions: Dict[str, SyntaxTree] = {
# custom_function_name must be passed to properly type custom function # custom_function_name must be passed to properly type custom function
# arguments uniquely if they're nested (i.e. $0 to $custom_func___0) # arguments uniquely if they're nested (i.e. $0 to $custom_func___0)
self._function_name(function_key): parse( self._function_name(function_key): parse(
text=function_value, custom_function_name=self._function_name(function_key) text=function_value,
name=self._function_name(function_key),
custom_function_names=function_names,
variable_names=variable_names,
) )
for function_key, function_value in overrides.items() for function_key, function_value in overrides.items()
if self._is_function(function_key) if self._is_function(function_key)
} }
self._variables: Dict[str, SyntaxTree] = { self._variables: Dict[str, SyntaxTree] = {
override_name: parse(override_value) variable_key: parse(
for override_name, override_value in overrides.items() text=variable_value,
if not self._is_function(override_name) name=variable_key,
custom_function_names=function_names,
variable_names=variable_names,
)
for variable_key, variable_value in overrides.items()
if not self._is_function(variable_key)
} }
self._ensure_no_custom_function_cycles() self._ensure_no_custom_function_cycles()

View file

@ -41,12 +41,6 @@ class Function(FunctionType, VariableDependency, ABC):
def _iterable_arguments(self) -> List[Argument]: def _iterable_arguments(self) -> List[Argument]:
return self.args return self.args
@classmethod
def from_name_and_args(cls, name: str, args: List[Argument]) -> "Function":
if Functions.is_built_in(name):
return BuiltInFunction(name=name, args=args).validate_args()
return CustomFunction(name=name, args=args)
class CustomFunction(Function, NamedCustomFunction): class CustomFunction(Function, NamedCustomFunction):
def resolve( def resolve(
@ -82,9 +76,9 @@ class CustomFunction(Function, NamedCustomFunction):
custom_functions=custom_functions, custom_functions=custom_functions,
) )
raise FunctionDoesNotExist( # Implies the custom function does not exist. This should have
f"Function %{self.name} does not exist as a built-in or custom function." # been checked in the parser with
) raise UNREACHABLE
class BuiltInFunction(Function, TypeHintedFunctionType): class BuiltInFunction(Function, TypeHintedFunctionType):

View file

@ -6,6 +6,7 @@ from ytdl_sub.script.script import Script
from ytdl_sub.script.types.array import ResolvedArray from ytdl_sub.script.types.array import ResolvedArray
from ytdl_sub.script.types.resolvable import Integer from ytdl_sub.script.types.resolvable import Integer
from ytdl_sub.script.utils.exceptions import CycleDetected from ytdl_sub.script.utils.exceptions import CycleDetected
from ytdl_sub.script.utils.exceptions import FunctionDoesNotExist
class TestCustomFunction: class TestCustomFunction:
@ -66,3 +67,16 @@ class TestCustomFunction:
"output": "{%cycle_func0(1)}", "output": "{%cycle_func0(1)}",
} }
).resolve() ).resolve()
def test_custom_function_uses_non_existent_function(self):
with pytest.raises(
FunctionDoesNotExist,
match=re.escape("Function %lolnope does not exist as a built-in or custom function."),
):
Script(
{
"%func1": "{%mul(%lolnope(1), $0)}",
"%func0": "{%mul(%func1(1), $0)}",
"output": "{%func(1)}",
}
).resolve()

View file

@ -39,8 +39,8 @@ class TestFloat:
], ],
) )
def test_float(self, float_: str, expected_float: int): def test_float(self, float_: str, expected_float: int):
assert Script({"float": float_, "as_string": "{%string(float)}"}).resolve() == { assert Script({"out": float_, "as_string": "{%string(out)}"}).resolve() == {
"float": Float(expected_float), "out": Float(expected_float),
"as_string": String(str(expected_float)), "as_string": String(str(expected_float)),
} }

View file

@ -46,7 +46,7 @@ class TestString:
], ],
) )
def test_string(self, string: str, expected_string: str): def test_string(self, string: str, expected_string: str):
assert Script({"string": string}).resolve() == {"string": String(expected_string)} assert Script({"out": string}).resolve() == {"out": String(expected_string)}
@pytest.mark.parametrize( @pytest.mark.parametrize(
"string", "string",