From 2c2148495891ab822f24fc997ca7a530297b47da Mon Sep 17 00:00:00 2001 From: Jesse Bannon Date: Wed, 20 Sep 2023 22:51:24 -0700 Subject: [PATCH] If working --- src/ytdl_sub/script/functions/__init__.py | 3 +- .../script/functions/conditional_functions.py | 17 -------- .../script/functions/special_functions.py | 24 +++++++++++ src/ytdl_sub/script/parser.py | 5 +++ src/ytdl_sub/script/types/function.py | 43 ++++++++++++------- src/ytdl_sub/script/types/resolvable.py | 3 ++ tests/unit/script/test_parser.py | 26 +++++++++-- 7 files changed, 84 insertions(+), 37 deletions(-) delete mode 100644 src/ytdl_sub/script/functions/conditional_functions.py create mode 100644 src/ytdl_sub/script/functions/special_functions.py diff --git a/src/ytdl_sub/script/functions/__init__.py b/src/ytdl_sub/script/functions/__init__.py index 67c83209..6a1dfc95 100644 --- a/src/ytdl_sub/script/functions/__init__.py +++ b/src/ytdl_sub/script/functions/__init__.py @@ -1,9 +1,8 @@ from typing import Optional -from ytdl_sub.script.functions.conditional_functions import ConditionalFunctions from ytdl_sub.script.functions.numeric_functions import NumericFunctions from ytdl_sub.script.functions.string_functions import StringFunctions -class Functions(StringFunctions, NumericFunctions, ConditionalFunctions): +class Functions(StringFunctions, NumericFunctions): pass diff --git a/src/ytdl_sub/script/functions/conditional_functions.py b/src/ytdl_sub/script/functions/conditional_functions.py deleted file mode 100644 index 4247ffc4..00000000 --- a/src/ytdl_sub/script/functions/conditional_functions.py +++ /dev/null @@ -1,17 +0,0 @@ -from typing import TypeVar - -from ytdl_sub.script.types.resolvable import Boolean -from ytdl_sub.script.types.resolvable import Resolvable - -ResolvableTrue = TypeVar("ResolvableTrue", bound=Resolvable) -ResolvableFalse = TypeVar("ResolvableFalse", bound=Resolvable) - - -class ConditionalFunctions: - @staticmethod - def iff( - condition: Boolean, true: ResolvableTrue, false: ResolvableFalse - ) -> ResolvableTrue | ResolvableFalse: - if condition.value: - return true - return false diff --git a/src/ytdl_sub/script/functions/special_functions.py b/src/ytdl_sub/script/functions/special_functions.py new file mode 100644 index 00000000..4eeabc8f --- /dev/null +++ b/src/ytdl_sub/script/functions/special_functions.py @@ -0,0 +1,24 @@ +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: + @staticmethod + def if_(condition: Boolean, true: Resolvable, false: Resolvable) -> Resolvable: + 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/parser.py b/src/ytdl_sub/script/parser.py index 5bc59e62..2bd30b2e 100644 --- a/src/ytdl_sub/script/parser.py +++ b/src/ytdl_sub/script/parser.py @@ -4,6 +4,7 @@ from typing import Optional from ytdl_sub.script.syntax_tree import SyntaxTree from ytdl_sub.script.types.function import ArgumentType 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 @@ -157,6 +158,10 @@ class _Parser: while ch := self._read(): if ch == ")": + # Special case for If functions since it can return a Union based on input types + if function_name == "if": + return IfFunction(name=function_name, args=function_args) + return Function(name=function_name, args=function_args) if ch != "(": diff --git a/src/ytdl_sub/script/types/function.py b/src/ytdl_sub/script/types/function.py index 149478eb..47aaa7ae 100644 --- a/src/ytdl_sub/script/types/function.py +++ b/src/ytdl_sub/script/types/function.py @@ -10,12 +10,12 @@ from typing import List from typing import Optional from typing import Set from typing import Type -from typing import TypeVar from typing import Union from typing import final from typing import get_origin from ytdl_sub.script.functions import Functions +from ytdl_sub.script.functions.special_functions import SpecialFunctions from ytdl_sub.script.types.resolvable import Boolean from ytdl_sub.script.types.resolvable import Float from ytdl_sub.script.types.resolvable import Integer @@ -48,12 +48,8 @@ class VariableDependency(ABC): return self.variables.issubset(set(resolved_variables.keys())) -def is_union(type: Type) -> bool: - return get_origin(type) is Union - - -def is_generic(type: Type) -> bool: - return type.__class__ is TypeVar +def is_union(arg_type: Type) -> bool: + return get_origin(arg_type) is Union @dataclass(frozen=True) @@ -67,10 +63,15 @@ class FunctionInputSpec: @classmethod def _is_type_compatible( cls, - input_arg: Optional[Resolvable], + input_arg: ArgumentType, expected_arg_type: Type[Resolvable | Optional[Resolvable]], ) -> bool: - input_arg_type = input_arg.__class__ + if isinstance(input_arg, Function): + input_arg_type = input_arg.output_type + elif isinstance(input_arg, Variable): + return True # unresolved variables can be anything, so pass for now + else: + input_arg_type = input_arg.__class__ if is_union(expected_arg_type): # See if the arg is a valid against the union @@ -82,15 +83,12 @@ class FunctionInputSpec: if not valid_type: return False - elif is_generic(expected_arg_type): - # TypeVars (generics) support any type of input - return True elif not issubclass(input_arg_type, expected_arg_type): return False return True - def _is_args_compatible(self, input_args: List[Resolvable | Optional[Resolvable]]) -> bool: + def _is_args_compatible(self, input_args: List[ArgumentType]) -> bool: assert self.args is not None if len(input_args) > len(self.args): @@ -103,7 +101,7 @@ class FunctionInputSpec: return True - def _is_varargs_compatible(self, input_args: List[Resolvable | Optional[Resolvable]]) -> bool: + def _is_varargs_compatible(self, input_args: List[ArgumentType]) -> bool: assert self.varargs is not None for input_arg in input_args: @@ -112,7 +110,7 @@ class FunctionInputSpec: return True - def is_compatible(self, input_args: List[Resolvable | Optional[Resolvable]]) -> bool: + def is_compatible(self, input_args: List[ArgumentType]) -> bool: if self.args is not None: return self._is_args_compatible(input_args=input_args) elif self.varargs is not None: @@ -199,3 +197,18 @@ class Function(VariableDependency): def resolve(self, resolved_variables: Dict[Variable, Resolvable]) -> Resolvable: raise NotImplemented() + + +@dataclass(frozen=True) +class IfFunction(Function): + def __post_init__(self): + super().__post_init__() + assert len(self.args) == 3 # bool, true, false + + @property + def callable(self) -> Callable[..., Resolvable]: + return SpecialFunctions.if_ + + @property + def output_type(self) -> Type[Resolvable]: + return Union[self.args[1].__class__, self.args[2].__class__] diff --git a/src/ytdl_sub/script/types/resolvable.py b/src/ytdl_sub/script/types/resolvable.py index 815cebb1..63030339 100644 --- a/src/ytdl_sub/script/types/resolvable.py +++ b/src/ytdl_sub/script/types/resolvable.py @@ -1,6 +1,7 @@ from abc import ABC from abc import abstractmethod from dataclasses import dataclass +from typing import Any from typing import Generic from typing import List from typing import TypeVar @@ -11,6 +12,8 @@ NumericT = TypeVar("NumericT", bound=int | float) @dataclass(frozen=True) class Resolvable(ABC): + value: Any + @abstractmethod def resolve(self) -> str: ... diff --git a/tests/unit/script/test_parser.py b/tests/unit/script/test_parser.py index f6854a8d..f57c9795 100644 --- a/tests/unit/script/test_parser.py +++ b/tests/unit/script/test_parser.py @@ -1,8 +1,11 @@ +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 @@ -27,16 +30,33 @@ class TestParser: ) def test_conditional(self): - parsed = parse("hello {%iff(True, 'hi', 3.4)}") + parsed = parse("hello {%if(True, 'hi', 3.4)}") + assert parsed == SyntaxTree( + [ + String("hello "), + IfFunction( + name="if", args=[Boolean(value=True), String(value="hi"), Float(value=3.4)] + ), + ] + ) + assert parsed.ast[1].output_type == Union[String, Float] + + def test_conditional_as_input(self): + parsed = parse("hello {%concat(%if(True, 'hi', 'mom'), 'and dad')}") assert parsed == SyntaxTree( [ String("hello "), Function( - name="iff", args=[Boolean(value=True), String(value="hi"), Float(value=3.4)] + name="concat", + args=[ + IfFunction( + name="if", args=[Boolean(value=True), String("hi"), String("mom")] + ), + String(value="and dad"), + ], ), ] ) - assert parsed.ast[1].output_type def test_single_function_one_vararg(self): parsed = parse("hello {%concat('hi mom')}")