diff --git a/src/ytdl_sub/script/script.py b/src/ytdl_sub/script/script.py index 8523cbd8..771f51ee 100644 --- a/src/ytdl_sub/script/script.py +++ b/src/ytdl_sub/script/script.py @@ -9,6 +9,7 @@ from ytdl_sub.script.types.syntax_tree import SyntaxTree 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.name_validation import validate_variable_name # pylint: disable=missing-raises-doc @@ -97,7 +98,7 @@ class Script: 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) + validate_variable_name(name) for name in overrides.keys() if not self._is_function(name) } self._functions: Dict[str, SyntaxTree] = { diff --git a/src/ytdl_sub/script/utils/name_validation.py b/src/ytdl_sub/script/utils/name_validation.py index 462bed58..d1dd1488 100644 --- a/src/ytdl_sub/script/utils/name_validation.py +++ b/src/ytdl_sub/script/utils/name_validation.py @@ -16,7 +16,7 @@ def is_valid_name(name: str) -> bool: return re.match(_NAME_REGEX_VALIDATOR, name) is not None -def validate_variable_name(variable_name: str) -> None: +def validate_variable_name(variable_name: str) -> str: """ Raises ------ @@ -25,10 +25,18 @@ def validate_variable_name(variable_name: str) -> None: """ if not is_valid_name(variable_name): raise InvalidVariableName( - f"Variable name '{variable_name}' is invalid. " - f"Names must be lower_snake_cased and begin with a letter." + f"Variable name '{variable_name}' is invalid:" + " Names must be lower_snake_cased and begin with a letter." ) + if Functions.is_built_in(variable_name): + raise InvalidVariableName( + f"Variable name '{variable_name}' is invalid:" + " The name is used by a built-in function and cannot be overwritten." + ) + + return variable_name + def validate_custom_function_name(custom_function_name: str) -> None: """ diff --git a/tests/unit/script/types/test_variable.py b/tests/unit/script/types/test_variable.py index 071ba881..59eb6f66 100644 --- a/tests/unit/script/types/test_variable.py +++ b/tests/unit/script/types/test_variable.py @@ -52,12 +52,72 @@ class TestVariable: ): Script({"a": "a", "b": "{c}"}).resolve() - def test_invalid_variable_name_inline(self): + @pytest.mark.parametrize( + "name", + [ + "vali_LOL_INVALID", + "name!!", + "na(", + "na[", + ], + ) + def test_invalid_variable_name_inline(self, name: str): with pytest.raises( InvalidVariableName, match=re.escape( - "Variable name 'vali_LOL_INVALID' is invalid. " - "Names must be lower_snake_cased and begin with a letter." + f"Variable name '{name}' is invalid:" + " Names must be lower_snake_cased and begin with a letter." ), ): - Script({"a": "{vali_LOL_INVALID}"}).resolve() + Script({"a": f"{{{name}}}"}).resolve() + + @pytest.mark.parametrize( + "name", + ["float", "bool", "mul"], + ) + def test_invalid_variable_name_inline_is_built_in(self, name: str): + with pytest.raises( + InvalidVariableName, + match=re.escape( + f"Variable name '{name}' is invalid:" + " The name is used by a built-in function and cannot be overwritten." + ), + ): + Script({"a": f"{{{name}}}"}).resolve() + + @pytest.mark.parametrize( + "name", + [ + "vali_LOL_INVALID", + "name!!", + "na(", + "na[", + "$2232", + "1245", + "CAN_CATCH_MORE", + "{brackets_in_definition}", + ], + ) + def test_invalid_variable_name_definition(self, name: str): + with pytest.raises( + InvalidVariableName, + match=re.escape( + f"Variable name '{name}' is invalid:" + " Names must be lower_snake_cased and begin with a letter." + ), + ): + Script({f"{name}": "value"}).resolve() + + @pytest.mark.parametrize( + "name", + ["float", "bool", "mul"], + ) + def test_invalid_variable_name_definition_is_built_in(self, name: str): + with pytest.raises( + InvalidVariableName, + match=re.escape( + f"Variable name '{name}' is invalid:" + " The name is used by a built-in function and cannot be overwritten." + ), + ): + Script({f"{name}": "value"}).resolve()