diff --git a/src/ytdl_sub/script/script.py b/src/ytdl_sub/script/script.py index 3b54c9c1..acaa5d8a 100644 --- a/src/ytdl_sub/script/script.py +++ b/src/ytdl_sub/script/script.py @@ -1,4 +1,5 @@ # pylint: disable=missing-raises-doc +import copy from typing import Dict from typing import List from typing import Optional @@ -696,6 +697,11 @@ class Script: """ return set(to_function_definition_name(name) for name in self._functions.keys()) + def _to_syntax_tree(self, maybe_resolved: SyntaxTree | Resolvable) -> SyntaxTree: + if isinstance(maybe_resolved, Resolvable): + return ResolvedSyntaxTree(ast=[maybe_resolved]) + return maybe_resolved + def resolve_partial( self, unresolvable: Optional[Set[str]] = None, @@ -739,8 +745,8 @@ class Script: # which means we can iterate again partially_resolved |= definition != maybe_resolved - return Script({}).add_parsed( - self._functions - | {var.name: definition for var, definition in resolved.items()} - | unresolved + return copy.deepcopy(self).add_parsed( + {var.name: self._variables[var.name] for var in unresolvable} + | {var.name: self._to_syntax_tree(definition) for var, definition in resolved.items()} + | {var.name: self._to_syntax_tree(definition) for var, definition in unresolved.items()} ) diff --git a/src/ytdl_sub/script/types/function.py b/src/ytdl_sub/script/types/function.py index 593e6a8e..9e69c01a 100644 --- a/src/ytdl_sub/script/types/function.py +++ b/src/ytdl_sub/script/types/function.py @@ -39,25 +39,6 @@ class Function(FunctionType, VariableDependency, ABC): def _iterable_arguments(self) -> List[Argument]: return self.args - def partial_resolve( - self: TypeT, - resolved_variables: Dict[Variable, Resolvable], - custom_functions: Dict[str, "VariableDependency"], - ) -> TypeT | Resolvable: - maybe_resolvable_values, is_resolvable = VariableDependency.try_partial_resolve( - args=self.args, - resolved_variables=resolved_variables, - custom_functions=custom_functions, - ) - - if is_resolvable: - return self.resolve( - resolved_variables=resolved_variables, - custom_functions=custom_functions, - ) - - return BuiltInFunction(name=self.name, args=maybe_resolvable_values) - class CustomFunction(Function, NamedCustomFunction): def resolve( @@ -104,6 +85,28 @@ class CustomFunction(Function, NamedCustomFunction): # been checked in the parser with raise UNREACHABLE + def partial_resolve( + self, + resolved_variables: Dict[Variable, Resolvable], + custom_functions: Dict[str, "VariableDependency"], + ) -> TypeT | Resolvable: + maybe_resolvable_values, is_resolvable = VariableDependency.try_partial_resolve( + args=self.args, + resolved_variables=resolved_variables, + custom_functions=custom_functions, + ) + + if any(var not in resolved_variables for var in custom_functions[self.name].variables): + is_resolvable = False + + if is_resolvable: + return self.resolve( + resolved_variables=resolved_variables, + custom_functions=custom_functions, + ) + + return CustomFunction(name=self.name, args=maybe_resolvable_values) + class BuiltInFunction(Function, BuiltInFunctionType): def validate_args(self) -> "BuiltInFunction": @@ -332,5 +335,37 @@ class BuiltInFunction(Function, BuiltInFunctionType): f"Runtime error occurred when executing the function %{self.name}: {str(exc)}" ) from exc + def partial_resolve( + self, + resolved_variables: Dict[Variable, Resolvable], + custom_functions: Dict[str, "VariableDependency"], + ) -> TypeT | Resolvable: + conditional_return_args = self.function_spec.conditional_arg_indices( + num_input_args=len(self.args) + ) + + # If the function is conditional, only run if its entirety is resolvable + if conditional_return_args: + if self.is_subset_of(variables=resolved_variables, custom_function_definitions=custom_functions): + return self.resolve( + resolved_variables=resolved_variables, + custom_functions=custom_functions, + ) + return self + else: + maybe_resolvable_values, is_resolvable = VariableDependency.try_partial_resolve( + args=self.args, + resolved_variables=resolved_variables, + custom_functions=custom_functions, + ) + + if is_resolvable: + return self.resolve( + resolved_variables=resolved_variables, + custom_functions=custom_functions, + ) + + return BuiltInFunction(name=self.name, args=maybe_resolvable_values) + def __hash__(self): return hash((self.name, *self.args)) diff --git a/src/ytdl_sub/utils/script.py b/src/ytdl_sub/utils/script.py index 027a0bd8..5c0264f9 100644 --- a/src/ytdl_sub/utils/script.py +++ b/src/ytdl_sub/utils/script.py @@ -144,6 +144,12 @@ class ScriptUtils: """ Converts any JSON-compatible value into equivalent script syntax """ + if isinstance(value, SyntaxTree): + output = "" + for arg in value.ast: + output += cls._to_script_code(arg, top_level=True) + return output + return cls._to_script_code(cls._to_script_argument(value), top_level=True) @classmethod diff --git a/tests/unit/config/test_subscription.py b/tests/unit/config/test_subscription.py index 418e1d31..127c94b8 100644 --- a/tests/unit/config/test_subscription.py +++ b/tests/unit/config/test_subscription.py @@ -11,6 +11,7 @@ from ytdl_sub.config.config_file import ConfigFile from ytdl_sub.plugins.nfo_tags import NfoTagsOptions from ytdl_sub.subscriptions.subscription import Subscription from ytdl_sub.utils.exceptions import ValidationException +from ytdl_sub.utils.script import ScriptUtils @contextmanager @@ -616,4 +617,6 @@ def test_default_docker_config_and_subscriptions( unresolvable.add("sibling_metadata") out = default_subs[0].overrides.script.resolve_partial(unresolvable=unresolvable) + prev = ScriptUtils.to_native_script(out._variables['episode_file_name']) + print("hi")