begin iff, need to handle generics
This commit is contained in:
parent
182bda9bf4
commit
fa0429bd04
4 changed files with 58 additions and 5 deletions
|
|
@ -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.numeric_functions import NumericFunctions
|
||||||
from ytdl_sub.script.functions.string_functions import StringFunctions
|
from ytdl_sub.script.functions.string_functions import StringFunctions
|
||||||
|
|
||||||
|
|
||||||
class Functions(StringFunctions, NumericFunctions):
|
class Functions(StringFunctions, NumericFunctions, ConditionalFunctions):
|
||||||
pass
|
pass
|
||||||
|
|
|
||||||
17
src/ytdl_sub/script/functions/conditional_functions.py
Normal file
17
src/ytdl_sub/script/functions/conditional_functions.py
Normal file
|
|
@ -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
|
||||||
|
|
@ -10,6 +10,7 @@ from typing import List
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
from typing import Set
|
from typing import Set
|
||||||
from typing import Type
|
from typing import Type
|
||||||
|
from typing import TypeVar
|
||||||
from typing import Union
|
from typing import Union
|
||||||
from typing import final
|
from typing import final
|
||||||
from typing import get_origin
|
from typing import get_origin
|
||||||
|
|
@ -47,6 +48,14 @@ class VariableDependency(ABC):
|
||||||
return self.variables.issubset(set(resolved_variables.keys()))
|
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)
|
@dataclass(frozen=True)
|
||||||
class FunctionInputSpec:
|
class FunctionInputSpec:
|
||||||
args: Optional[List[Type[Resolvable | Optional[Resolvable]]]] = None
|
args: Optional[List[Type[Resolvable | Optional[Resolvable]]]] = None
|
||||||
|
|
@ -63,10 +72,20 @@ class FunctionInputSpec:
|
||||||
) -> bool:
|
) -> bool:
|
||||||
input_arg_type = input_arg.__class__
|
input_arg_type = input_arg.__class__
|
||||||
|
|
||||||
if get_origin(expected_arg_type) is Union:
|
if is_union(expected_arg_type):
|
||||||
if input_arg_type not in expected_arg_type.__args__:
|
# 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
|
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 False
|
||||||
|
|
||||||
return True
|
return True
|
||||||
|
|
@ -158,7 +177,9 @@ class Function(VariableDependency):
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def output_type(self) -> Type[Resolvable]:
|
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
|
@property
|
||||||
def variables(self) -> Set[Variable]:
|
def variables(self) -> Set[Variable]:
|
||||||
|
|
|
||||||
|
|
@ -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):
|
def test_single_function_one_vararg(self):
|
||||||
parsed = parse("hello {%concat('hi mom')}")
|
parsed = parse("hello {%concat('hi mom')}")
|
||||||
assert parsed == SyntaxTree(
|
assert parsed == SyntaxTree(
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue