function input spec
This commit is contained in:
parent
22188d84f8
commit
ca1fc6a227
4 changed files with 158 additions and 26 deletions
|
|
@ -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))
|
||||||
|
|
|
||||||
|
|
@ -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]:
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue