From 0fffc79cfbc89515a689c3f54a88da3510ec1336 Mon Sep 17 00:00:00 2001 From: Jesse Bannon Date: Tue, 21 Nov 2023 18:40:49 -0800 Subject: [PATCH] simplify variable deps --- src/ytdl_sub/script/types/array.py | 18 ++-------- src/ytdl_sub/script/types/function.py | 28 ++------------- src/ytdl_sub/script/types/map.py | 34 +++++-------------- src/ytdl_sub/script/types/syntax_tree.py | 28 ++------------- .../script/types/variable_dependency.py | 22 ++++++++++++ 5 files changed, 36 insertions(+), 94 deletions(-) diff --git a/src/ytdl_sub/script/types/array.py b/src/ytdl_sub/script/types/array.py index 792af89a..28c86c8b 100644 --- a/src/ytdl_sub/script/types/array.py +++ b/src/ytdl_sub/script/types/array.py @@ -28,25 +28,11 @@ class UnresolvedArray(Array, VariableDependency, FutureResolvable): @property def variables(self) -> Set[Variable]: - variables: Set[Variable] = set() - for arg in self.value: - if isinstance(arg, Variable): - variables.add(arg) - elif isinstance(arg, VariableDependency): - variables.update(arg.variables) - - return variables + return self._variables(*self.value) @property def function_arguments(self) -> Set[FunctionArgument]: - variables: Set[FunctionArgument] = set() - for arg in self.value: - if isinstance(arg, FunctionArgument): - variables.add(arg) - elif isinstance(arg, VariableDependency): - variables.update(arg.function_arguments) - - return variables + return self._function_arguments(*self.value) def resolve( self, diff --git a/src/ytdl_sub/script/types/function.py b/src/ytdl_sub/script/types/function.py index 488e82e3..70852745 100644 --- a/src/ytdl_sub/script/types/function.py +++ b/src/ytdl_sub/script/types/function.py @@ -41,35 +41,11 @@ from ytdl_sub.utils.exceptions import StringFormattingException class Function(FunctionType, VariableDependency, ABC): @property def variables(self) -> Set[Variable]: - """ - Returns - ------- - All variables used within the function - """ - variables: Set[Variable] = set() - for arg in self.args: - if isinstance(arg, Variable): - variables.add(arg) - elif isinstance(arg, VariableDependency): - variables.update(arg.variables) - - return variables + return self._variables(*self.args) @property def function_arguments(self) -> Set[FunctionArgument]: - """ - Returns - ------- - All function arguments used within the function - """ - function_arguments: Set[FunctionArgument] = set() - for arg in self.args: - if isinstance(arg, FunctionArgument): - function_arguments.add(arg) - elif isinstance(arg, VariableDependency): - function_arguments.update(arg.function_arguments) - - return function_arguments + return self._function_arguments(*self.args) @classmethod def from_name_and_args(cls, name: str, args: List[ArgumentType]) -> "Function": diff --git a/src/ytdl_sub/script/types/map.py b/src/ytdl_sub/script/types/map.py index 2b95c274..b9ce5d9b 100644 --- a/src/ytdl_sub/script/types/map.py +++ b/src/ytdl_sub/script/types/map.py @@ -1,5 +1,7 @@ +import itertools from dataclasses import dataclass from typing import Dict +from typing import List from typing import Set from ytdl_sub.script.types.resolvable import ArgumentType @@ -27,37 +29,17 @@ class Map(NonHashable): class UnresolvedMap(Map, VariableDependency, FutureResolvable): value: Dict[ArgumentType, ArgumentType] + @property + def _as_flat_list(self) -> List[ArgumentType]: + return list(itertools.chain(*self.value.items())) + @property def variables(self) -> Set[Variable]: - output: Set[Variable] = set() - for key, value in self.value.items(): - if isinstance(key, Variable): - output.add(key) - elif isinstance(key, VariableDependency): - output.update(key.variables) - - if isinstance(value, Variable): - output.add(key) - elif isinstance(value, VariableDependency): - output.update(value.variables) - - return output + return self._variables(*self._as_flat_list) @property def function_arguments(self) -> Set[FunctionArgument]: - output: Set[FunctionArgument] = set() - for key, value in self.value.items(): - if isinstance(key, FunctionArgument): - output.add(key) - elif isinstance(key, VariableDependency): - output.update(key.function_arguments) - - if isinstance(value, FunctionArgument): - output.add(key) - elif isinstance(value, VariableDependency): - output.update(value.function_arguments) - - return output + return self._function_arguments(*self._as_flat_list) def resolve( self, diff --git a/src/ytdl_sub/script/types/syntax_tree.py b/src/ytdl_sub/script/types/syntax_tree.py index 7d0f931e..845514a6 100644 --- a/src/ytdl_sub/script/types/syntax_tree.py +++ b/src/ytdl_sub/script/types/syntax_tree.py @@ -19,35 +19,11 @@ class SyntaxTree(VariableDependency): @property def variables(self) -> Set[Variable]: - """ - Returns - ------- - All variables used within the SyntaxTree - """ - variables: Set[Variable] = set() - for token in self.ast: - if isinstance(token, Variable): - variables.add(token) - elif isinstance(token, VariableDependency): - variables.update(token.variables) - - return variables + return self._variables(*self.ast) @property def function_arguments(self) -> Set[FunctionArgument]: - """ - Returns - ------- - All function arguments used within the SyntaxTree - """ - function_arguments: Set[FunctionArgument] = set() - for token in self.ast: - if isinstance(token, FunctionArgument): - function_arguments.add(token) - elif isinstance(token, VariableDependency): - function_arguments.update(token.function_arguments) - - return function_arguments + return self._function_arguments(*self.ast) def resolve( self, diff --git a/src/ytdl_sub/script/types/variable_dependency.py b/src/ytdl_sub/script/types/variable_dependency.py index b48a2343..84b32a3e 100644 --- a/src/ytdl_sub/script/types/variable_dependency.py +++ b/src/ytdl_sub/script/types/variable_dependency.py @@ -15,6 +15,28 @@ from ytdl_sub.utils.exceptions import StringFormattingException @dataclass(frozen=True) class VariableDependency(ABC): + @classmethod + def _variables(cls, *args: ArgumentType) -> Set[Variable]: + output: Set[Variable] = set() + for arg in args: + if isinstance(arg, Variable): + output.add(arg) + elif isinstance(arg, VariableDependency): + output.update(arg.variables) + + return output + + @classmethod + def _function_arguments(cls, *args: ArgumentType) -> Set[FunctionArgument]: + output: Set[FunctionArgument] = set() + for arg in args: + if isinstance(arg, FunctionArgument): + output.add(arg) + elif isinstance(arg, VariableDependency): + output.update(arg.function_arguments) + + return output + @property @abstractmethod def variables(self) -> Set[Variable]: