From d5e647554e4f617786a14b75b750a56a9875591a Mon Sep 17 00:00:00 2001 From: Jesse Bannon Date: Sat, 30 Mar 2024 00:53:05 -0700 Subject: [PATCH] [DEV] Fix resolve_once script bug with custom functions (#952) * [DEV] Fix resolve_once script bug with custom functions * function name * fix adding custom functions --- .../config/validators/variable_validation.py | 4 ++ src/ytdl_sub/script/script.py | 50 +++++++++++++++---- src/ytdl_sub/utils/script.py | 6 ++- tests/unit/script/test_script.py | 24 +++++++++ 4 files changed, 72 insertions(+), 12 deletions(-) diff --git a/src/ytdl_sub/config/validators/variable_validation.py b/src/ytdl_sub/config/validators/variable_validation.py index f5de3ed5..bdeb9fac 100644 --- a/src/ytdl_sub/config/validators/variable_validation.py +++ b/src/ytdl_sub/config/validators/variable_validation.py @@ -15,6 +15,7 @@ from ytdl_sub.config.validators.options import OptionsValidator from ytdl_sub.downloaders.url.validators import MultiUrlValidator from ytdl_sub.entries.variables.override_variables import SubscriptionVariables from ytdl_sub.script.script import Script +from ytdl_sub.script.script import _is_function from ytdl_sub.utils.scriptable import BASE_SCRIPT from ytdl_sub.validators.string_formatter_validators import to_variable_dependency_format_string from ytdl_sub.validators.string_formatter_validators import validate_formatters @@ -33,6 +34,9 @@ def _add_dummy_overrides(overrides: Overrides) -> Dict[str, str]: # Have the dummy override variable contain all variable deps that it uses in the string dummy_overrides: Dict[str, str] = {} for override_name in _override_variables(overrides): + if _is_function(override_name): + continue + # pylint: disable=protected-access dummy_overrides[override_name] = to_variable_dependency_format_string( script=overrides.script, parsed_format_string=overrides.script._variables[override_name] diff --git a/src/ytdl_sub/script/script.py b/src/ytdl_sub/script/script.py index 3f3b42ff..f285eae6 100644 --- a/src/ytdl_sub/script/script.py +++ b/src/ytdl_sub/script/script.py @@ -10,6 +10,7 @@ from ytdl_sub.script.script_output import ScriptOutput from ytdl_sub.script.types.resolvable import Lambda from ytdl_sub.script.types.resolvable import Resolvable from ytdl_sub.script.types.syntax_tree import SyntaxTree +from ytdl_sub.script.types.variable import FunctionArgument from ytdl_sub.script.types.variable import Variable from ytdl_sub.script.utils.exceptions import UNREACHABLE from ytdl_sub.script.utils.exceptions import CycleDetected @@ -257,12 +258,23 @@ class Script: f"Output filter variable contains the variable {var_dep} " f"which is set as unresolvable" ) + + # Do not recurse custom function arguments since they have no deps + if isinstance(var_dep, FunctionArgument): + continue + subset_to_resolve.add(var_dep.name) subset_to_resolve |= self._recursive_get_unresolved_output_filter_variables( current_var=self._variables[var_dep.name], subset_to_resolve=subset_to_resolve, unresolvable=unresolvable, ) + for custom_func_dep in current_var.custom_functions: + subset_to_resolve |= self._recursive_get_unresolved_output_filter_variables( + current_var=self._functions[custom_func_dep.name], + subset_to_resolve=subset_to_resolve, + unresolvable=unresolvable, + ) return subset_to_resolve @@ -440,18 +452,34 @@ class Script: self """ added_variables_to_validate: Set[str] = set() - for variable_name, variable_definition in variables.items(): - self._variables[variable_name] = parse( - text=variable_definition, - name=variable_name, - custom_function_names=set(self._functions.keys()), - variable_names=set(self._variables.keys()) - .union(variables.keys()) - .union(unresolvable or set()), - ) - if self._variables[variable_name].maybe_resolvable is None: - added_variables_to_validate.add(variable_name) + functions_to_add = { + _function_name(name): definition + for name, definition in variables.items() + if _is_function(name) + } + variables_to_add = { + name: definition for name, definition in variables.items() if not _is_function(name) + } + + for definitions in [functions_to_add, variables_to_add]: + for name, definition in definitions.items(): + parsed = parse( + text=definition, + name=name, + custom_function_names=set(self._functions.keys()), + variable_names=set(self._variables.keys()) + .union(variables.keys()) + .union(unresolvable or set()), + ) + + if parsed.maybe_resolvable is None: + added_variables_to_validate.add(name) + + if name in functions_to_add: + self._functions[name] = parsed + else: + self._variables[name] = parsed if added_variables_to_validate: self._validate(added_variables=added_variables_to_validate) diff --git a/src/ytdl_sub/utils/script.py b/src/ytdl_sub/utils/script.py index 5653bc74..e2fc6d4c 100644 --- a/src/ytdl_sub/utils/script.py +++ b/src/ytdl_sub/utils/script.py @@ -3,6 +3,8 @@ import re from typing import Any from typing import Dict +from ytdl_sub.script.script import _is_function + class ScriptUtils: @classmethod @@ -11,7 +13,9 @@ class ScriptUtils: Helper to add sanitized variables to a Script """ sanitized_variables = { - f"{name}_sanitized": f"{{%sanitize({name})}}" for name in variables.keys() + f"{name}_sanitized": f"{{%sanitize({name})}}" + for name in variables.keys() + if not _is_function(name) } return dict(variables, **sanitized_variables) diff --git a/tests/unit/script/test_script.py b/tests/unit/script/test_script.py index 0b86e2be..dfa7994e 100644 --- a/tests/unit/script/test_script.py +++ b/tests/unit/script/test_script.py @@ -61,3 +61,27 @@ class TestScript: assert script.get("new_variable_upper") == String("HI MOM THE TITLE") assert script.get("new_variable_titlecase") == String("Hi Mom The Title") assert script.get("entry") == entry_map + + def test_resolve_once_with_custom_functions(self): + script = Script( + { + "%is_bilateral_url": "{ %not(%contains( $0, 'youtube.com/playlist' )) }", + "%bilateral_url": """{ + %if( + %and( + enable_bilateral_scraping, + %is_bilateral_url($0) + ), + $0, + "" + ) + }""", + "enable_bilateral_scraping": "True", + } + ) + + script.add({"%bilateral_url_wrap": "{ %bilateral_url($0) }"}) + + assert ( + script.resolve_once({"url": "{ %bilateral_url_wrap('nope') }"})["url"].native == "nope" + )