From fa0429bd04d9a2ca1ccb5332fa2b48e027381113 Mon Sep 17 00:00:00 2001 From: Jesse Bannon Date: Wed, 20 Sep 2023 17:03:13 -0700 Subject: [PATCH] begin iff, need to handle generics --- src/ytdl_sub/script/functions/__init__.py | 5 +++- .../script/functions/conditional_functions.py | 17 +++++++++++ src/ytdl_sub/script/types/function.py | 29 ++++++++++++++++--- tests/unit/script/test_parser.py | 12 ++++++++ 4 files changed, 58 insertions(+), 5 deletions(-) create mode 100644 src/ytdl_sub/script/functions/conditional_functions.py diff --git a/src/ytdl_sub/script/functions/__init__.py b/src/ytdl_sub/script/functions/__init__.py index b16fa070..67c83209 100644 --- a/src/ytdl_sub/script/functions/__init__.py +++ b/src/ytdl_sub/script/functions/__init__.py @@ -1,6 +1,9 @@ +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): +class Functions(StringFunctions, NumericFunctions, ConditionalFunctions): pass diff --git a/src/ytdl_sub/script/functions/conditional_functions.py b/src/ytdl_sub/script/functions/conditional_functions.py new file mode 100644 index 00000000..4247ffc4 --- /dev/null +++ b/src/ytdl_sub/script/functions/conditional_functions.py @@ -0,0 +1,17 @@ +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/types/function.py b/src/ytdl_sub/script/types/function.py index baf1b3d2..149478eb 100644 --- a/src/ytdl_sub/script/types/function.py +++ b/src/ytdl_sub/script/types/function.py @@ -10,6 +10,7 @@ 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 @@ -47,6 +48,14 @@ 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 + + @dataclass(frozen=True) class FunctionInputSpec: args: Optional[List[Type[Resolvable | Optional[Resolvable]]]] = None @@ -63,10 +72,20 @@ class FunctionInputSpec: ) -> bool: input_arg_type = input_arg.__class__ - if get_origin(expected_arg_type) is Union: - if input_arg_type not in expected_arg_type.__args__: + if is_union(expected_arg_type): + # See if the arg is a valid against the union + valid_type = False + for union_type in expected_arg_type.__args__: + if issubclass(input_arg_type, union_type): + valid_type = True + break + + if not valid_type: return False - elif input_arg_type != expected_arg_type: + 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 @@ -158,7 +177,9 @@ class Function(VariableDependency): @property def output_type(self) -> Type[Resolvable]: - return self.arg_spec.annotations["return"] + output_type = self.arg_spec.annotations["return"] + # TODO: Handle generics here + return output_type @property def variables(self) -> Set[Variable]: diff --git a/tests/unit/script/test_parser.py b/tests/unit/script/test_parser.py index 682cf1b3..f6854a8d 100644 --- a/tests/unit/script/test_parser.py +++ b/tests/unit/script/test_parser.py @@ -26,6 +26,18 @@ class TestParser: ] ) + def test_conditional(self): + parsed = parse("hello {%iff(True, 'hi', 3.4)}") + assert parsed == SyntaxTree( + [ + String("hello "), + Function( + name="iff", args=[Boolean(value=True), String(value="hi"), Float(value=3.4)] + ), + ] + ) + assert parsed.ast[1].output_type + def test_single_function_one_vararg(self): parsed = parse("hello {%concat('hi mom')}") assert parsed == SyntaxTree(