syntax tree testing begins

This commit is contained in:
Jesse Bannon 2023-09-20 23:58:16 -07:00
parent 1c2ba91853
commit 4f3f2f36fd
6 changed files with 61 additions and 71 deletions

View file

@ -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)

View file

@ -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()}

View file

@ -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()}

View file

@ -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:

View file

@ -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])}]"

View file

@ -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)