function input spec

This commit is contained in:
Jesse Bannon 2023-09-20 15:51:20 -07:00
parent 22188d84f8
commit ca1fc6a227
4 changed files with 158 additions and 26 deletions

View file

@ -1,3 +1,6 @@
from typing import Optional
from ytdl_sub.script.types.resolvable import Integer
from ytdl_sub.script.types.resolvable import String from ytdl_sub.script.types.resolvable import String
@ -30,5 +33,14 @@ class StringFunctions:
return String(string.value.capitalize()) return String(string.value.capitalize())
@staticmethod @staticmethod
def concat(l_string: String, r_string: String) -> String: def replace(
return String(f"{l_string}{r_string}") string: String, old: String, new: String, count: Optional[Integer] = None
) -> String:
if count:
return String(string.value.replace(old.value, new.value, count.value))
return String(string.value.replace(old.value, new.value))
@staticmethod
def concat(*args: String) -> String:
return String("".join(*args))

View file

@ -7,10 +7,12 @@ from inspect import FullArgSpec
from typing import Callable from typing import Callable
from typing import Dict from typing import Dict
from typing import List from typing import List
from typing import Optional
from typing import Set from typing import Set
from typing import Type from typing import Type
from typing import Union from typing import Union
from typing import final from typing import final
from typing import get_origin
from ytdl_sub.script.functions import Functions from ytdl_sub.script.functions import Functions
from ytdl_sub.script.types.resolvable import Boolean from ytdl_sub.script.types.resolvable import Boolean
@ -45,42 +47,99 @@ class VariableDependency(ABC):
return self.variables.issubset(set(resolved_variables.keys())) return self.variables.issubset(set(resolved_variables.keys()))
@dataclass(frozen=True)
class FunctionInputSpec:
args: Optional[List[Type[Resolvable | Optional[Resolvable]]]] = None
varargs: Optional[Type[Resolvable]] = None
def __post_init__(self):
assert (self.args is None) ^ (self.varargs is None)
@classmethod
def _is_type_compatible(
cls,
input_arg: Optional[Resolvable],
expected_arg_type: Type[Resolvable | Optional[Resolvable]],
) -> 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__:
return False
elif input_arg_type != expected_arg_type:
return False
return True
def _is_args_compatible(self, input_args: List[Resolvable | Optional[Resolvable]]) -> bool:
assert self.args is not None
if len(input_args) > len(self.args):
return False
for idx in range(len(self.args)):
input_arg = input_args[idx] if idx < len(input_args) else None
if not self._is_type_compatible(input_arg=input_arg, expected_arg_type=self.args[idx]):
return False
return True
def _is_varargs_compatible(self, input_args: List[Resolvable | Optional[Resolvable]]) -> bool:
assert self.varargs is not None
for input_arg in input_args:
if not self._is_type_compatible(input_arg=input_arg, expected_arg_type=self.varargs):
return False
return True
def is_compatible(self, input_args: List[Resolvable | Optional[Resolvable]]) -> bool:
if self.args is not None:
return self._is_args_compatible(input_args=input_args)
elif self.varargs is not None:
return self._is_varargs_compatible(input_args=input_args)
else:
assert False, "should never reach here"
def expected_args_str(self) -> str:
if self.args is not None:
return f"({', '.join([type_.__name__ for type_ in self.args])})"
elif self.varargs is not None:
return f"({self.varargs.__name__}, ...)"
@classmethod
def from_function(cls, func: "Function") -> "FunctionInputSpec":
if func.arg_spec.varargs:
return FunctionInputSpec(varargs=func.arg_spec.annotations[func.arg_spec.varargs])
return FunctionInputSpec(
args=[func.arg_spec.annotations[arg_name] for arg_name in func.arg_spec.args]
)
@dataclass(frozen=True) @dataclass(frozen=True)
class Function(VariableDependency): class Function(VariableDependency):
name: str name: str
args: List[ArgumentType] args: List[ArgumentType]
def __post_init__(self): def __post_init__(self):
# TODO: Figure out resolution via introspecting args and outputs of function if not self.input_spec.is_compatible(input_args=self.args):
if len(self.args) != len(self.input_types):
raise StringFormattingException(
f"Unequal amount of arguments passed to function {self.name}.\n"
f"{self._expected_received_error_msg()}"
)
for input_arg, input_arg_type in zip(self.args, self.input_types):
if isinstance(input_arg, Function):
input_arg = input_arg.output_type
elif isinstance(input_arg, Variable):
pass # cannot evaluate the variable yet, so pass
if not issubclass(input_arg.__class__, input_arg_type):
raise StringFormattingException( raise StringFormattingException(
f"Invalid arguments passed to function {self.name}.\n" f"Invalid arguments passed to function {self.name}.\n"
f"{self._expected_received_error_msg()}" f"{self._expected_received_error_msg()}"
) )
def _expected_received_error_msg(self) -> str: def _expected_received_error_msg(self) -> str:
output_type_names: List[str] = [] received_type_names: List[str] = []
for arg in self.args: for arg in self.args:
if isinstance(arg, Function): if isinstance(arg, Function):
output_type_names.append(f"{arg.name}(...)->{arg.output_type.__name__}") received_type_names.append(f"{arg.name}(...)->{arg.output_type.__name__}")
else: else:
output_type_names.append(arg.__class__.__name__) received_type_names.append(arg.__class__.__name__)
return ( received_args_str = f"({', '.join([name for name in received_type_names])})"
f"Expected ({', '.join([type_.__name__ for type_ in self.input_types])}).\n"
f"Received ({', '.join([output_type_name for output_type_name in output_type_names])})" return f"Expected {self.input_spec.expected_args_str()}.\nReceived ({received_args_str})"
)
@property @property
def callable(self) -> Callable[..., Resolvable]: def callable(self) -> Callable[..., Resolvable]:
@ -94,8 +153,8 @@ class Function(VariableDependency):
return inspect.getfullargspec(self.callable) return inspect.getfullargspec(self.callable)
@property @property
def input_types(self) -> List[Type[Resolvable]]: def input_spec(self) -> FunctionInputSpec:
return [self.arg_spec.annotations[arg_name] for arg_name in self.arg_spec.args] return FunctionInputSpec.from_function(self)
@property @property
def output_type(self) -> Type[Resolvable]: def output_type(self) -> Type[Resolvable]:

View file

@ -2,6 +2,7 @@ from abc import ABC
from abc import abstractmethod from abc import abstractmethod
from dataclasses import dataclass from dataclasses import dataclass
from typing import Generic from typing import Generic
from typing import List
from typing import TypeVar from typing import TypeVar
T = TypeVar("T") T = TypeVar("T")
@ -16,7 +17,7 @@ class Resolvable(ABC):
@dataclass(frozen=True) @dataclass(frozen=True)
class ResolvableT(Resolvable, Generic[T]): class ResolvableT(Resolvable, ABC, Generic[T]):
value: T value: T
def resolve(self) -> str: def resolve(self) -> str:
@ -46,3 +47,16 @@ class Boolean(ResolvableT[bool]):
@dataclass(frozen=True) @dataclass(frozen=True)
class String(ResolvableT[str]): class String(ResolvableT[str]):
pass pass
@dataclass(frozen=True)
class _List(Resolvable, Generic[T], ABC):
value: List[T]
def resolve(self) -> str:
return f"[{', '.join([str(val) for val in self.value])}]"
@dataclass(frozen=True)
class StringList(_List[String]):
pass

View file

@ -26,6 +26,53 @@ class TestParser:
] ]
) )
def test_single_function_one_vararg(self):
parsed = parse("hello {%concat('hi mom')}")
assert parsed == SyntaxTree(
[
String("hello "),
Function(name="concat", args=[String(value="hi mom")]),
]
)
def test_single_function_many_vararg(self):
parsed = parse("hello {%concat('hi', 'mom')}")
assert parsed == SyntaxTree(
[
String("hello "),
Function(name="concat", args=[String(value="hi"), String(value="mom")]),
]
)
def test_single_function_many_args_with_optional_none(self):
parsed = parse("hello {%replace('hi mom', 'hi', '')}")
assert parsed == SyntaxTree(
[
String("hello "),
Function(
name="replace",
args=[String(value="hi mom"), String(value="hi"), String(value="")],
),
]
)
def test_single_function_many_args_with_optional_provided(self):
parsed = parse("hello {%replace('hi mom', 'hi', '', 1)}")
assert parsed == SyntaxTree(
[
String("hello "),
Function(
name="replace",
args=[
String(value="hi mom"),
String(value="hi"),
String(value=""),
Integer(value=1),
],
),
]
)
@pytest.mark.parametrize("whitespace", ["", " ", " ", "\n", " \n "]) @pytest.mark.parametrize("whitespace", ["", " ", " ", "\n", " \n "])
def test_single_function_multiple_args(self, whitespace: str): def test_single_function_multiple_args(self, whitespace: str):
s = whitespace s = whitespace