clean up function file
This commit is contained in:
parent
f0cbc1bbd3
commit
f4cd0fc781
7 changed files with 218 additions and 189 deletions
|
|
@ -15,11 +15,11 @@ from ytdl_sub.script.types.resolvable import String
|
||||||
from ytdl_sub.script.types.syntax_tree import SyntaxTree
|
from ytdl_sub.script.types.syntax_tree import SyntaxTree
|
||||||
from ytdl_sub.script.types.variable import FunctionArgument
|
from ytdl_sub.script.types.variable import FunctionArgument
|
||||||
from ytdl_sub.script.types.variable import Variable
|
from ytdl_sub.script.types.variable import Variable
|
||||||
|
from ytdl_sub.script.utils.exception_formatters import ParserExceptionFormatter
|
||||||
from ytdl_sub.script.utils.exceptions import UNREACHABLE
|
from ytdl_sub.script.utils.exceptions import UNREACHABLE
|
||||||
from ytdl_sub.script.utils.exceptions import IncompatibleFunctionArguments
|
from ytdl_sub.script.utils.exceptions import IncompatibleFunctionArguments
|
||||||
from ytdl_sub.script.utils.exceptions import InvalidSyntaxException
|
from ytdl_sub.script.utils.exceptions import InvalidSyntaxException
|
||||||
from ytdl_sub.script.utils.exceptions import UserException
|
from ytdl_sub.script.utils.exceptions import UserException
|
||||||
from ytdl_sub.script.utils.parser_exception_formatter import ParserExceptionFormatter
|
|
||||||
from ytdl_sub.utils.exceptions import StringFormattingException
|
from ytdl_sub.utils.exceptions import StringFormattingException
|
||||||
from ytdl_sub.validators.string_formatter_validators import is_valid_source_variable_name
|
from ytdl_sub.validators.string_formatter_validators import is_valid_source_variable_name
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -7,7 +7,6 @@ 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
|
||||||
|
|
@ -17,94 +16,17 @@ from ytdl_sub.script.types.resolvable import AnyTypeReturnable
|
||||||
from ytdl_sub.script.types.resolvable import AnyTypeReturnableA
|
from ytdl_sub.script.types.resolvable import AnyTypeReturnableA
|
||||||
from ytdl_sub.script.types.resolvable import AnyTypeReturnableB
|
from ytdl_sub.script.types.resolvable import AnyTypeReturnableB
|
||||||
from ytdl_sub.script.types.resolvable import ArgumentType
|
from ytdl_sub.script.types.resolvable import ArgumentType
|
||||||
from ytdl_sub.script.types.resolvable import NamedType
|
from ytdl_sub.script.types.resolvable import FunctionLike
|
||||||
from ytdl_sub.script.types.resolvable import Resolvable
|
from ytdl_sub.script.types.resolvable import Resolvable
|
||||||
from ytdl_sub.script.types.variable import FunctionArgument
|
from ytdl_sub.script.types.variable import FunctionArgument
|
||||||
from ytdl_sub.script.types.variable import Variable
|
from ytdl_sub.script.types.variable import Variable
|
||||||
from ytdl_sub.script.types.variable_dependency import VariableDependency
|
from ytdl_sub.script.types.variable_dependency import VariableDependency
|
||||||
from ytdl_sub.script.utils.exceptions import UNREACHABLE
|
from ytdl_sub.script.utils.exception_formatters import FunctionArgumentsExceptionFormatter
|
||||||
from ytdl_sub.script.utils.exceptions import IncompatibleFunctionArguments
|
from ytdl_sub.script.utils.type_checking import FunctionInputSpec
|
||||||
from ytdl_sub.script.utils.type_checking import get_optional_type
|
|
||||||
from ytdl_sub.script.utils.type_checking import is_optional
|
|
||||||
from ytdl_sub.script.utils.type_checking import is_type_compatible
|
|
||||||
from ytdl_sub.script.utils.type_checking import is_union
|
from ytdl_sub.script.utils.type_checking import is_union
|
||||||
from ytdl_sub.utils.exceptions import StringFormattingException
|
from ytdl_sub.utils.exceptions import StringFormattingException
|
||||||
|
|
||||||
|
|
||||||
@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_arg_compatible(
|
|
||||||
cls, arg: NamedType, expected_arg_type: Type[Resolvable | Optional[Resolvable]]
|
|
||||||
):
|
|
||||||
input_arg_type = arg.__class__
|
|
||||||
if isinstance(arg, BuiltInFunction):
|
|
||||||
input_arg_type = arg.output_type
|
|
||||||
elif isinstance(arg, Variable):
|
|
||||||
return True # unresolved variables can be anything, so pass for now
|
|
||||||
|
|
||||||
return is_type_compatible(arg_type=input_arg_type, expected_arg_type=expected_arg_type)
|
|
||||||
|
|
||||||
def _is_args_compatible(self, input_args: List[ArgumentType]) -> bool:
|
|
||||||
assert self.args is not None
|
|
||||||
|
|
||||||
if len(input_args) > len(self.args):
|
|
||||||
return False
|
|
||||||
|
|
||||||
for idx, arg in enumerate(self.args):
|
|
||||||
input_arg = input_args[idx] if idx < len(input_args) else None
|
|
||||||
if not self._is_arg_compatible(arg=input_arg, expected_arg_type=arg):
|
|
||||||
return False
|
|
||||||
|
|
||||||
return True
|
|
||||||
|
|
||||||
def _is_varargs_compatible(self, input_args: List[ArgumentType]) -> bool:
|
|
||||||
assert self.varargs is not None
|
|
||||||
|
|
||||||
for input_arg in input_args:
|
|
||||||
if not self._is_arg_compatible(arg=input_arg, expected_arg_type=self.varargs):
|
|
||||||
return False
|
|
||||||
|
|
||||||
return True
|
|
||||||
|
|
||||||
def is_compatible(self, input_args: List[ArgumentType]) -> bool:
|
|
||||||
if self.args is not None:
|
|
||||||
return self._is_args_compatible(input_args=input_args)
|
|
||||||
if self.varargs is not None:
|
|
||||||
return self._is_varargs_compatible(input_args=input_args)
|
|
||||||
|
|
||||||
raise UNREACHABLE # TODO: functions with no args
|
|
||||||
|
|
||||||
def expected_args_str(self) -> str:
|
|
||||||
def to_human_readable_name(python_type: Type[NamedType] | Type[Union[NamedType]]) -> str:
|
|
||||||
if is_optional(python_type):
|
|
||||||
return f"Optional[{to_human_readable_name(get_optional_type(python_type))}]"
|
|
||||||
if is_union(python_type):
|
|
||||||
return ", ".join(to_human_readable_name(arg) for arg in python_type.__args__)
|
|
||||||
return python_type.type_name()
|
|
||||||
|
|
||||||
if self.args is not None:
|
|
||||||
return f"({', '.join([to_human_readable_name(type_) for type_ in self.args])})"
|
|
||||||
if self.varargs is not None:
|
|
||||||
return f"({to_human_readable_name(self.varargs.__name__)}, ...)"
|
|
||||||
return "()"
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def from_function(cls, func: "BuiltInFunction") -> "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, ArgumentType, ABC):
|
class Function(VariableDependency, ArgumentType, ABC):
|
||||||
name: str
|
name: str
|
||||||
|
|
@ -182,33 +104,14 @@ class CustomFunction(Function):
|
||||||
raise StringFormattingException(f"Custom function {self.name} does not exist")
|
raise StringFormattingException(f"Custom function {self.name} does not exist")
|
||||||
|
|
||||||
|
|
||||||
class BuiltInFunction(Function):
|
class BuiltInFunction(Function, FunctionLike):
|
||||||
def _expected_received_error_msg(self) -> str:
|
|
||||||
received_type_names: List[str] = []
|
|
||||||
for arg in self.args:
|
|
||||||
if isinstance(arg, BuiltInFunction):
|
|
||||||
if is_union(arg.output_type):
|
|
||||||
# TODO: Move naming to separate function, deal with Union input naming
|
|
||||||
received_type_names.append(
|
|
||||||
f"%{arg.name}(...)->Union["
|
|
||||||
f"{', '.join(type_.type_name() for type_ in arg.output_type.__args__)}"
|
|
||||||
f"]"
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
received_type_names.append(f"%{arg.name}(...)->{arg.output_type.type_name()}")
|
|
||||||
else:
|
|
||||||
received_type_names.append(arg.type_name())
|
|
||||||
|
|
||||||
received_args_str = f"({', '.join(name for name in received_type_names)})"
|
|
||||||
|
|
||||||
return f"Expected {self.input_spec.expected_args_str()}\nReceived {received_args_str}"
|
|
||||||
|
|
||||||
def validate_args(self) -> "BuiltInFunction":
|
def validate_args(self) -> "BuiltInFunction":
|
||||||
if not self.input_spec.is_compatible(input_args=self.args):
|
if not self.input_spec.is_compatible(input_args=self.args):
|
||||||
raise IncompatibleFunctionArguments(
|
raise FunctionArgumentsExceptionFormatter(
|
||||||
f"Incompatible arguments passed to function {self.name}.\n"
|
input_spec=self.input_spec,
|
||||||
f"{self._expected_received_error_msg()}"
|
function_instance=self,
|
||||||
)
|
).highlight()
|
||||||
|
|
||||||
return self
|
return self
|
||||||
|
|
||||||
@property
|
@property
|
||||||
|
|
@ -226,15 +129,19 @@ class BuiltInFunction(Function):
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def input_spec(self) -> FunctionInputSpec:
|
def input_spec(self) -> FunctionInputSpec:
|
||||||
return FunctionInputSpec.from_function(self)
|
if self.arg_spec.varargs:
|
||||||
|
return FunctionInputSpec(varargs=self.arg_spec.annotations[self.arg_spec.varargs])
|
||||||
|
|
||||||
|
return FunctionInputSpec(
|
||||||
|
args=[self.arg_spec.annotations[arg_name] for arg_name in self.arg_spec.args]
|
||||||
|
)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def _arg_output_type(cls, arg: ArgumentType) -> Type[ArgumentType]:
|
def _arg_output_type(cls, arg: ArgumentType) -> Type[ArgumentType]:
|
||||||
if isinstance(arg, BuiltInFunction):
|
if isinstance(arg, BuiltInFunction):
|
||||||
return arg.output_type
|
return arg.output_type()
|
||||||
return type(arg)
|
return type(arg)
|
||||||
|
|
||||||
@property
|
|
||||||
def output_type(self) -> Type[Resolvable]:
|
def output_type(self) -> Type[Resolvable]:
|
||||||
output_type = self.arg_spec.annotations["return"]
|
output_type = self.arg_spec.annotations["return"]
|
||||||
if is_union(output_type):
|
if is_union(output_type):
|
||||||
|
|
|
||||||
|
|
@ -1,8 +1,11 @@
|
||||||
import json
|
import json
|
||||||
from abc import ABC
|
from abc import ABC
|
||||||
|
from abc import abstractmethod
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Any
|
from typing import Any
|
||||||
from typing import Generic
|
from typing import Generic
|
||||||
|
from typing import List
|
||||||
|
from typing import Type
|
||||||
from typing import TypeVar
|
from typing import TypeVar
|
||||||
|
|
||||||
T = TypeVar("T")
|
T = TypeVar("T")
|
||||||
|
|
@ -111,3 +114,12 @@ class Boolean(ResolvableT[bool], Hashable, ArgumentType):
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
class String(ResolvableT[str], Hashable, ArgumentType):
|
class String(ResolvableT[str], Hashable, ArgumentType):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class FunctionLike(NamedType):
|
||||||
|
name: str
|
||||||
|
args: List[ArgumentType]
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def output_type(self) -> Type[Resolvable]:
|
||||||
|
pass
|
||||||
|
|
|
||||||
136
src/ytdl_sub/script/utils/exception_formatters.py
Normal file
136
src/ytdl_sub/script/utils/exception_formatters.py
Normal file
|
|
@ -0,0 +1,136 @@
|
||||||
|
import sys
|
||||||
|
from typing import List
|
||||||
|
from typing import Type
|
||||||
|
from typing import TypeVar
|
||||||
|
from typing import Union
|
||||||
|
|
||||||
|
from ytdl_sub.script.types.resolvable import FunctionLike
|
||||||
|
from ytdl_sub.script.types.resolvable import NamedType
|
||||||
|
from ytdl_sub.script.utils.exceptions import IncompatibleFunctionArguments
|
||||||
|
from ytdl_sub.script.utils.exceptions import UserException
|
||||||
|
from ytdl_sub.script.utils.type_checking import FunctionInputSpec
|
||||||
|
from ytdl_sub.script.utils.type_checking import get_optional_type
|
||||||
|
from ytdl_sub.script.utils.type_checking import is_optional
|
||||||
|
from ytdl_sub.script.utils.type_checking import is_union
|
||||||
|
|
||||||
|
TUserException = TypeVar("TUserException", bound=UserException)
|
||||||
|
|
||||||
|
|
||||||
|
class ParserExceptionFormatter:
|
||||||
|
def __init__(self, text: str, start: int, end: int, exception: TUserException):
|
||||||
|
self._text = text
|
||||||
|
self._start = start
|
||||||
|
self._end = end
|
||||||
|
self._exception = exception
|
||||||
|
|
||||||
|
def exception_text(self, border: int):
|
||||||
|
text_left = max(0, self._start - border)
|
||||||
|
text_right = min(len(self._text), self._start + border)
|
||||||
|
relative_start = self._start - text_left
|
||||||
|
|
||||||
|
exception_text: str = ""
|
||||||
|
if text_left > 3:
|
||||||
|
exception_text = "… "
|
||||||
|
relative_start += len(exception_text)
|
||||||
|
|
||||||
|
exception_text += self._text[text_left:text_right]
|
||||||
|
|
||||||
|
if text_right < len(self._text) - 3:
|
||||||
|
exception_text += " …"
|
||||||
|
|
||||||
|
exception_text += "\n"
|
||||||
|
exception_text += f"{' ' * relative_start}^"
|
||||||
|
|
||||||
|
return "\n" + exception_text
|
||||||
|
|
||||||
|
@property
|
||||||
|
def is_multi_line(self) -> bool:
|
||||||
|
return "\n" in self._text
|
||||||
|
|
||||||
|
def exception_text_lines(self, border_lines: int = 0) -> str:
|
||||||
|
split_text = self._text.split("\n")
|
||||||
|
|
||||||
|
start_line: int = sys.maxsize
|
||||||
|
end_line: int = -1
|
||||||
|
pos: int = 0
|
||||||
|
for idx, line in enumerate(split_text):
|
||||||
|
if self._start <= pos < self._end:
|
||||||
|
start_line = min(start_line, idx)
|
||||||
|
end_line = max(end_line, idx + 1)
|
||||||
|
pos += len(line)
|
||||||
|
|
||||||
|
true_start_line = start_line
|
||||||
|
start_line = max(0, start_line - border_lines)
|
||||||
|
end_line = min(len(split_text), end_line + border_lines)
|
||||||
|
|
||||||
|
# Get min leading spaces between all lines to return
|
||||||
|
min_leading_spaces = sys.maxsize
|
||||||
|
for line in split_text[start_line:end_line]:
|
||||||
|
min_leading_spaces = min(min_leading_spaces, len(line) - len(line.lstrip()))
|
||||||
|
|
||||||
|
to_return: List[str] = []
|
||||||
|
for idx in range(start_line, end_line):
|
||||||
|
if idx == true_start_line:
|
||||||
|
to_return.append(f">>> {split_text[idx][min_leading_spaces:]}")
|
||||||
|
else:
|
||||||
|
to_return.append(f" {split_text[idx][min_leading_spaces:]}")
|
||||||
|
|
||||||
|
return "\n" + "\n".join(to_return)
|
||||||
|
|
||||||
|
def highlight(self) -> TUserException:
|
||||||
|
if self.is_multi_line:
|
||||||
|
invalid_syntax = self.exception_text_lines(border_lines=3)
|
||||||
|
else:
|
||||||
|
invalid_syntax = self.exception_text(border=20)
|
||||||
|
|
||||||
|
return self._exception.__class__(f"{invalid_syntax}\n{str(self._exception)}")
|
||||||
|
|
||||||
|
|
||||||
|
class FunctionArgumentsExceptionFormatter:
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
input_spec: FunctionInputSpec,
|
||||||
|
function_instance: FunctionLike,
|
||||||
|
):
|
||||||
|
self._args = input_spec.args
|
||||||
|
self._varargs = input_spec.varargs
|
||||||
|
self._name = function_instance.name
|
||||||
|
self._input_args = function_instance.args
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _to_human_readable_name(cls, python_type: Type[NamedType] | Type[Union[NamedType]]) -> str:
|
||||||
|
if is_optional(python_type):
|
||||||
|
return f"Optional[{cls._to_human_readable_name(get_optional_type(python_type))}]"
|
||||||
|
if is_union(python_type):
|
||||||
|
return ", ".join(cls._to_human_readable_name(arg) for arg in python_type.__args__)
|
||||||
|
return python_type.type_name()
|
||||||
|
|
||||||
|
def _expected_args_str(self) -> str:
|
||||||
|
if self._args is not None:
|
||||||
|
return f"({', '.join([self._to_human_readable_name(type_) for type_ in self._args])})"
|
||||||
|
if self._varargs is not None:
|
||||||
|
return f"({self._to_human_readable_name(self._varargs)}, ...)"
|
||||||
|
return "()"
|
||||||
|
|
||||||
|
def highlight(self) -> IncompatibleFunctionArguments:
|
||||||
|
received_type_names: List[str] = []
|
||||||
|
for arg in self._input_args:
|
||||||
|
if isinstance(arg, FunctionLike):
|
||||||
|
if is_union(arg.output_type()):
|
||||||
|
# TODO: Move naming to separate function, deal with Union input naming
|
||||||
|
received_type_names.append(
|
||||||
|
f"%{arg.name}(...)->Union["
|
||||||
|
f"{', '.join(type_.type_name() for type_ in arg.output_type().__args__)}"
|
||||||
|
f"]"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
received_type_names.append(f"%{arg.name}(...)->{arg.output_type().type_name()}")
|
||||||
|
else:
|
||||||
|
received_type_names.append(arg.type_name())
|
||||||
|
|
||||||
|
received_args_str = f"({', '.join(name for name in received_type_names)})"
|
||||||
|
|
||||||
|
return IncompatibleFunctionArguments(
|
||||||
|
f"Incompatible arguments passed to function {self._name}.\n"
|
||||||
|
f"Expected {self._expected_args_str()}\nReceived {received_args_str}"
|
||||||
|
)
|
||||||
|
|
@ -1,77 +0,0 @@
|
||||||
import sys
|
|
||||||
from typing import List
|
|
||||||
from typing import TypeVar
|
|
||||||
|
|
||||||
from ytdl_sub.script.utils.exceptions import UserException
|
|
||||||
|
|
||||||
TUserException = TypeVar("TUserException", bound=UserException)
|
|
||||||
|
|
||||||
|
|
||||||
class ParserExceptionFormatter:
|
|
||||||
def __init__(self, text: str, start: int, end: int, exception: TUserException):
|
|
||||||
self._text = text
|
|
||||||
self._start = start
|
|
||||||
self._end = end
|
|
||||||
self._exception = exception
|
|
||||||
|
|
||||||
def exception_text(self, border: int):
|
|
||||||
text_left = max(0, self._start - border)
|
|
||||||
text_right = min(len(self._text), self._start + border)
|
|
||||||
relative_start = self._start - text_left
|
|
||||||
|
|
||||||
exception_text: str = ""
|
|
||||||
if text_left > 3:
|
|
||||||
exception_text = "… "
|
|
||||||
relative_start += len(exception_text)
|
|
||||||
|
|
||||||
exception_text += self._text[text_left:text_right]
|
|
||||||
|
|
||||||
if text_right < len(self._text) - 3:
|
|
||||||
exception_text += " …"
|
|
||||||
|
|
||||||
exception_text += "\n"
|
|
||||||
exception_text += f"{' ' * relative_start}^"
|
|
||||||
|
|
||||||
return "\n" + exception_text
|
|
||||||
|
|
||||||
@property
|
|
||||||
def is_multi_line(self) -> bool:
|
|
||||||
return "\n" in self._text
|
|
||||||
|
|
||||||
def exception_text_lines(self, border_lines: int = 0) -> str:
|
|
||||||
split_text = self._text.split("\n")
|
|
||||||
|
|
||||||
start_line: int = sys.maxsize
|
|
||||||
end_line: int = -1
|
|
||||||
pos: int = 0
|
|
||||||
for idx, line in enumerate(split_text):
|
|
||||||
if self._start <= pos < self._end:
|
|
||||||
start_line = min(start_line, idx)
|
|
||||||
end_line = max(end_line, idx + 1)
|
|
||||||
pos += len(line)
|
|
||||||
|
|
||||||
true_start_line = start_line
|
|
||||||
start_line = max(0, start_line - border_lines)
|
|
||||||
end_line = min(len(split_text), end_line + border_lines)
|
|
||||||
|
|
||||||
# Get min leading spaces between all lines to return
|
|
||||||
min_leading_spaces = sys.maxsize
|
|
||||||
for line in split_text[start_line:end_line]:
|
|
||||||
min_leading_spaces = min(min_leading_spaces, len(line) - len(line.lstrip()))
|
|
||||||
|
|
||||||
to_return: List[str] = []
|
|
||||||
for idx in range(start_line, end_line):
|
|
||||||
if idx == true_start_line:
|
|
||||||
to_return.append(f">>> {split_text[idx][min_leading_spaces:]}")
|
|
||||||
else:
|
|
||||||
to_return.append(f" {split_text[idx][min_leading_spaces:]}")
|
|
||||||
|
|
||||||
return "\n" + "\n".join(to_return)
|
|
||||||
|
|
||||||
def highlight(self) -> TUserException:
|
|
||||||
if self.is_multi_line:
|
|
||||||
invalid_syntax = self.exception_text_lines(border_lines=3)
|
|
||||||
else:
|
|
||||||
invalid_syntax = self.exception_text(border=20)
|
|
||||||
|
|
||||||
return self._exception.__class__(f"{invalid_syntax}\n{str(self._exception)}")
|
|
||||||
|
|
@ -1,10 +1,16 @@
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import List
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
from typing import Type
|
from typing import Type
|
||||||
from typing import Union
|
from typing import Union
|
||||||
from typing import get_origin
|
from typing import get_origin
|
||||||
|
|
||||||
|
from ytdl_sub.script.types.resolvable import ArgumentType
|
||||||
|
from ytdl_sub.script.types.resolvable import FunctionLike
|
||||||
from ytdl_sub.script.types.resolvable import NamedType
|
from ytdl_sub.script.types.resolvable import NamedType
|
||||||
from ytdl_sub.script.types.resolvable import Resolvable
|
from ytdl_sub.script.types.resolvable import Resolvable
|
||||||
|
from ytdl_sub.script.types.variable import Variable
|
||||||
|
from ytdl_sub.script.utils.exceptions import UNREACHABLE
|
||||||
|
|
||||||
|
|
||||||
def is_union(arg_type: Type) -> bool:
|
def is_union(arg_type: Type) -> bool:
|
||||||
|
|
@ -20,9 +26,15 @@ def get_optional_type(optional_type: Type) -> Type[NamedType]:
|
||||||
|
|
||||||
|
|
||||||
def is_type_compatible(
|
def is_type_compatible(
|
||||||
arg_type: Type[NamedType],
|
arg: NamedType,
|
||||||
expected_arg_type: Type[Resolvable | Optional[Resolvable]],
|
expected_arg_type: Type[Resolvable | Optional[Resolvable]],
|
||||||
) -> bool:
|
) -> bool:
|
||||||
|
arg_type: Type[NamedType] = arg.__class__
|
||||||
|
if isinstance(arg, FunctionLike):
|
||||||
|
arg_type = arg.output_type()
|
||||||
|
elif isinstance(arg, Variable):
|
||||||
|
return True # unresolved variables can be anything, so pass for now
|
||||||
|
|
||||||
if is_union(expected_arg_type):
|
if is_union(expected_arg_type):
|
||||||
# See if the arg is a valid against the union
|
# See if the arg is a valid against the union
|
||||||
valid_type = False
|
valid_type = False
|
||||||
|
|
@ -50,3 +62,42 @@ def is_type_compatible(
|
||||||
return False
|
return False
|
||||||
|
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
@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)
|
||||||
|
|
||||||
|
def _is_args_compatible(self, input_args: List[ArgumentType]) -> bool:
|
||||||
|
assert self.args is not None
|
||||||
|
|
||||||
|
if len(input_args) > len(self.args):
|
||||||
|
return False
|
||||||
|
|
||||||
|
for idx, arg in enumerate(self.args):
|
||||||
|
input_arg = input_args[idx] if idx < len(input_args) else None
|
||||||
|
if not is_type_compatible(arg=input_arg, expected_arg_type=arg):
|
||||||
|
return False
|
||||||
|
|
||||||
|
return True
|
||||||
|
|
||||||
|
def _is_varargs_compatible(self, input_args: List[ArgumentType]) -> bool:
|
||||||
|
assert self.varargs is not None
|
||||||
|
|
||||||
|
for input_arg in input_args:
|
||||||
|
if not is_type_compatible(arg=input_arg, expected_arg_type=self.varargs):
|
||||||
|
return False
|
||||||
|
|
||||||
|
return True
|
||||||
|
|
||||||
|
def is_compatible(self, input_args: List[ArgumentType]) -> bool:
|
||||||
|
if self.args is not None:
|
||||||
|
return self._is_args_compatible(input_args=input_args)
|
||||||
|
if self.varargs is not None:
|
||||||
|
return self._is_varargs_compatible(input_args=input_args)
|
||||||
|
|
||||||
|
raise UNREACHABLE # TODO: functions with no args
|
||||||
|
|
|
||||||
|
|
@ -44,7 +44,7 @@ class TestParser:
|
||||||
),
|
),
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
assert parsed.ast[1].output_type == Union[String, Float]
|
assert parsed.ast[1].output_type() == Union[String, Float]
|
||||||
|
|
||||||
def test_conditional_as_input_same_outputs(self):
|
def test_conditional_as_input_same_outputs(self):
|
||||||
parsed = parse("hello {%concat(%if(True, 'hi', 'mom'), 'and dad')}")
|
parsed = parse("hello {%concat(%if(True, 'hi', 'mom'), 'and dad')}")
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue