From 4f3f2f36fd8894cb3e476c254f42ee506234176a Mon Sep 17 00:00:00 2001 From: Jesse Bannon Date: Wed, 20 Sep 2023 23:58:16 -0700 Subject: [PATCH] syntax tree testing begins --- .../script/functions/special_functions.py | 14 ----- src/ytdl_sub/script/overrides_resolver.py | 7 +-- src/ytdl_sub/script/syntax_tree.py | 59 +++++-------------- src/ytdl_sub/script/types/function.py | 2 +- src/ytdl_sub/script/types/resolvable.py | 10 +--- tests/unit/script/test_syntax_tree.py | 40 +++++++++++++ 6 files changed, 61 insertions(+), 71 deletions(-) create mode 100644 tests/unit/script/test_syntax_tree.py diff --git a/src/ytdl_sub/script/functions/special_functions.py b/src/ytdl_sub/script/functions/special_functions.py index 4eeabc8f..0d390834 100644 --- a/src/ytdl_sub/script/functions/special_functions.py +++ b/src/ytdl_sub/script/functions/special_functions.py @@ -1,7 +1,5 @@ -from ytdl_sub.entries.entry import Entry from ytdl_sub.script.types.resolvable import Boolean from ytdl_sub.script.types.resolvable import Resolvable -from ytdl_sub.script.types.resolvable import String class SpecialFunctions: @@ -10,15 +8,3 @@ class SpecialFunctions: if condition.value: return true return false - - @staticmethod - def entry_contains(entry: Entry, key: String) -> Boolean: - return Boolean(entry.kwargs_contains(key=key.value)) - - @staticmethod - def entry(entry: Entry, key: String) -> Resolvable: - return entry.kwargs(key=key.value) - - @staticmethod - def entry_get(entry: Entry, key: String, default: Resolvable) -> Resolvable: - return entry.kwargs_get(key=key.value, default=default.value) diff --git a/src/ytdl_sub/script/overrides_resolver.py b/src/ytdl_sub/script/overrides_resolver.py index 39955f78..d78ff7b5 100644 --- a/src/ytdl_sub/script/overrides_resolver.py +++ b/src/ytdl_sub/script/overrides_resolver.py @@ -37,7 +37,7 @@ class OverridesResolver: for variable in variable_dependencies.keys(): _traverse(variable) - def resolve_overrides(self) -> Dict[str, str]: + def resolve_overrides(self) -> Dict[str, Resolvable]: self._ensure_no_cycles() unresolved_variables: List[Variable] = list(self.overrides.keys()) @@ -59,7 +59,4 @@ class OverridesResolver: len(unresolved_variables) != unresolved_count ), "did not resolve any variables, cycle detected" - return { - variable.name: resolvable.resolve() - for variable, resolvable in resolved_variables.items() - } + return {variable.name: resolvable for variable, resolvable in resolved_variables.items()} diff --git a/src/ytdl_sub/script/syntax_tree.py b/src/ytdl_sub/script/syntax_tree.py index 3df074c8..f1ef3298 100644 --- a/src/ytdl_sub/script/syntax_tree.py +++ b/src/ytdl_sub/script/syntax_tree.py @@ -33,51 +33,26 @@ class SyntaxTree(VariableDependency): return variables def resolve(self, resolved_variables: Dict[Variable, Resolvable]) -> Resolvable: - output: str = "" + resolved: List[Resolvable] = [] for token in self.ast: - if isinstance(token, String): - output += token.resolve() + if isinstance(token, Resolvable): + resolved.append(token) elif isinstance(token, Variable): - output += resolved_variables[token].resolve() + resolved.append(resolved_variables[token]) elif isinstance(token, Function): - output += token.resolve(resolved_variables=resolved_variables) + resolved.append(token.resolve(resolved_variables=resolved_variables)) else: assert False, "should never reach" - return String(output) + # If only one resolvable resides in the AST, return as that + if len(resolved) == 1: + return resolved[0] + + # Otherwise, to concat multiple resolved outputs, we must concat as strings + return String("".join([str(res) for res in resolved])) @classmethod - def detect_cycles(cls, parsed_overrides: Dict[str, "SyntaxTree"]) -> None: - """ - Parameters - ---------- - parsed_overrides - ``overrides`` in a subscription, parsed into a SyntaxTree - """ - variable_dependencies: Dict[Variable, Set[Variable]] = { - Variable(name): ast.variables for name, ast in parsed_overrides.items() - } - - def _traverse( - to_variable: Variable, visited_variables: Optional[List[Variable]] = None - ) -> None: - if visited_variables is None: - visited_variables = [] - - if to_variable in visited_variables: - raise StringFormattingException("Detected cycle in variables") - visited_variables.append(to_variable) - - for dep in variable_dependencies[to_variable]: - _traverse(to_variable=dep, visited_variables=visited_variables) - - for variable in variable_dependencies.keys(): - _traverse(variable) - - @classmethod - def resolve_overrides(cls, parsed_overrides: Dict[str, "SyntaxTree"]) -> Dict[str, str]: - cls.detect_cycles(parsed_overrides=parsed_overrides) - + def resolve_overrides(cls, parsed_overrides: Dict[str, "SyntaxTree"]) -> Dict[str, Resolvable]: overrides: Dict[Variable, "SyntaxTree"] = { Variable(name): ast for name, ast in parsed_overrides.items() } @@ -97,11 +72,7 @@ class SyntaxTree(VariableDependency): ) unresolved_variables.remove(variable) - assert ( - len(unresolved_variables) != unresolved_count - ), "did not resolve any variables, cycle detected" + if len(unresolved_variables) == unresolved_count: + raise StringFormattingException("did not resolve any variables, cycle detected") - return { - variable.name: resolvable.resolve() - for variable, resolvable in resolved_variables.items() - } + return {variable.name: resolvable for variable, resolvable in resolved_variables.items()} diff --git a/src/ytdl_sub/script/types/function.py b/src/ytdl_sub/script/types/function.py index 04a15b9b..36e29965 100644 --- a/src/ytdl_sub/script/types/function.py +++ b/src/ytdl_sub/script/types/function.py @@ -45,7 +45,7 @@ class VariableDependency(ABC): ------- True if variable dependency. False otherwise. """ - return self.variables.issubset(set(resolved_variables.keys())) + return not self.variables.issubset(set(resolved_variables.keys())) def is_union(arg_type: Type) -> bool: diff --git a/src/ytdl_sub/script/types/resolvable.py b/src/ytdl_sub/script/types/resolvable.py index 63030339..9b135506 100644 --- a/src/ytdl_sub/script/types/resolvable.py +++ b/src/ytdl_sub/script/types/resolvable.py @@ -14,18 +14,14 @@ NumericT = TypeVar("NumericT", bound=int | float) class Resolvable(ABC): value: Any - @abstractmethod - def resolve(self) -> str: - ... + def __str__(self) -> str: + return str(self.value) @dataclass(frozen=True) class ResolvableT(Resolvable, ABC, Generic[T]): value: T - def resolve(self) -> str: - return str(self.value) - @dataclass(frozen=True) class Numeric(ResolvableT[NumericT], ABC, Generic[NumericT]): @@ -56,7 +52,7 @@ class String(ResolvableT[str]): class _List(Resolvable, Generic[T], ABC): value: List[T] - def resolve(self) -> str: + def __str__(self) -> str: return f"[{', '.join([str(val) for val in self.value])}]" diff --git a/tests/unit/script/test_syntax_tree.py b/tests/unit/script/test_syntax_tree.py new file mode 100644 index 00000000..9c3a0b0f --- /dev/null +++ b/tests/unit/script/test_syntax_tree.py @@ -0,0 +1,40 @@ +from typing import Dict +from typing import Union + +import pytest + +from ytdl_sub.script.parser import parse +from ytdl_sub.script.syntax_tree import SyntaxTree +from ytdl_sub.script.types.function import Function +from ytdl_sub.script.types.function import IfFunction +from ytdl_sub.script.types.resolvable import Boolean +from ytdl_sub.script.types.resolvable import Float +from ytdl_sub.script.types.resolvable import Integer +from ytdl_sub.script.types.resolvable import String +from ytdl_sub.script.types.variable import Variable +from ytdl_sub.utils.exceptions import StringFormattingException + + +class TestSyntaxTree: + def test_simple(self): + overrides: Dict[str, SyntaxTree] = { + "a": SyntaxTree(ast=[String("a")]), + "b": SyntaxTree(ast=[Variable("b_")]), + "b_": SyntaxTree(ast=[String("b")]), + } + + resolved = SyntaxTree.resolve_overrides(parsed_overrides=overrides) + assert resolved == { + "a": String(value="a"), + "b": String(value="b"), + "b_": String(value="b"), + } + + def test_simple_cycle(self): + overrides: Dict[str, SyntaxTree] = { + "a": SyntaxTree(ast=[Variable("b")]), + "b": SyntaxTree(ast=[Variable("a")]), + } + + with pytest.raises(StringFormattingException): + _ = SyntaxTree.resolve_overrides(parsed_overrides=overrides)