function spec do more

This commit is contained in:
Jesse Bannon 2023-11-23 22:36:30 -08:00
parent fb02479108
commit 4278534c21
5 changed files with 57 additions and 30 deletions

View file

@ -4,6 +4,7 @@ from typing import List
from typing import Optional from typing import Optional
from typing import Set from typing import Set
from ytdl_sub.script.functions import Functions
from ytdl_sub.script.parser import parse from ytdl_sub.script.parser import parse
from ytdl_sub.script.types.resolvable import Resolvable from ytdl_sub.script.types.resolvable import Resolvable
from ytdl_sub.script.types.syntax_tree import SyntaxTree from ytdl_sub.script.types.syntax_tree import SyntaxTree
@ -144,6 +145,12 @@ class Script:
f"{nested_custom_function.num_input_args}" f"{nested_custom_function.num_input_args}"
) )
def _ensure_lambda_usage_num_input_arguments_valid(self) -> None:
for variable_name, variable_definition in self._variables.items():
for lambda_argument in variable_definition.lambda_arguments:
if Functions.is_built_in(name=lambda_argument.value):
pass
def _validate(self) -> None: def _validate(self) -> None:
self._ensure_no_custom_function_cycles() self._ensure_no_custom_function_cycles()
self._ensure_custom_function_arguments_valid() self._ensure_custom_function_arguments_valid()

View file

@ -1,9 +1,7 @@
import copy import copy
import functools import functools
import inspect
from abc import ABC from abc import ABC
from dataclasses import dataclass from dataclasses import dataclass
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
@ -29,7 +27,7 @@ from ytdl_sub.script.utils.exception_formatters import FunctionArgumentsExceptio
from ytdl_sub.script.utils.exceptions import UNREACHABLE from ytdl_sub.script.utils.exceptions import UNREACHABLE
from ytdl_sub.script.utils.exceptions import FunctionRuntimeException from ytdl_sub.script.utils.exceptions import FunctionRuntimeException
from ytdl_sub.script.utils.exceptions import UserThrownRuntimeError from ytdl_sub.script.utils.exceptions import UserThrownRuntimeError
from ytdl_sub.script.utils.type_checking import FunctionInputSpec from ytdl_sub.script.utils.type_checking import FunctionSpec
from ytdl_sub.script.utils.type_checking import is_union from ytdl_sub.script.utils.type_checking import is_union
@ -82,9 +80,9 @@ class CustomFunction(Function, NamedCustomFunction):
class BuiltInFunction(Function, TypeHintedFunctionType): class BuiltInFunction(Function, TypeHintedFunctionType):
def validate_args(self) -> "BuiltInFunction": def validate_args(self) -> "BuiltInFunction":
if not self.input_spec.is_compatible(input_args=self.args): if not self.function_spec.is_compatible(input_args=self.args):
raise FunctionArgumentsExceptionFormatter( raise FunctionArgumentsExceptionFormatter(
input_spec=self.input_spec, input_spec=self.function_spec,
function_instance=self, function_instance=self,
).highlight() ).highlight()
@ -101,21 +99,8 @@ class BuiltInFunction(Function, TypeHintedFunctionType):
raise UNREACHABLE raise UNREACHABLE
@functools.cached_property @functools.cached_property
def arg_spec(self) -> FullArgSpec: def function_spec(self) -> FunctionSpec:
return inspect.getfullargspec(self.callable) return FunctionSpec.from_callable(self.callable)
@property
def input_spec(self) -> FunctionInputSpec:
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]
)
@property
def is_lambda_function(self) -> bool:
return Lambda in (self.input_spec.args or [])
@classmethod @classmethod
def _arg_output_type(cls, arg: Argument) -> Type[Argument]: def _arg_output_type(cls, arg: Argument) -> Type[Argument]:
@ -124,19 +109,18 @@ class BuiltInFunction(Function, TypeHintedFunctionType):
return type(arg) return type(arg)
def output_type(self) -> Type[Resolvable]: def output_type(self) -> Type[Resolvable]:
output_type = self.arg_spec.annotations["return"] if is_union(self.function_spec.return_type):
if is_union(output_type):
union_types_list = [] union_types_list = []
for union_type in output_type.__args__: for union_type in self.function_spec.return_type.__args__:
if union_type in (ReturnableArgument, ReturnableArgumentA, ReturnableArgumentB): if union_type in (ReturnableArgument, ReturnableArgumentA, ReturnableArgumentB):
generic_arg_index = self.input_spec.args.index(union_type) generic_arg_index = self.function_spec.args.index(union_type)
union_types_list.append(self._arg_output_type(self.args[generic_arg_index])) union_types_list.append(self._arg_output_type(self.args[generic_arg_index]))
else: else:
union_types_list.append(union_type) union_types_list.append(union_type)
return Union[tuple(union_types_list)] return Union[tuple(union_types_list)]
return output_type return self.function_spec.return_type
def _resolve_lambda_function( def _resolve_lambda_function(
self, self,
@ -153,7 +137,7 @@ class BuiltInFunction(Function, TypeHintedFunctionType):
3. Resolve it like any other syntax 3. Resolve it like any other syntax
""" """
function_input_lambda_args = [arg for arg in resolved_arguments if isinstance(arg, Lambda)] function_input_lambda_args = [arg for arg in resolved_arguments if isinstance(arg, Lambda)]
if not self.is_lambda_function or len(function_input_lambda_args) != 1: if not self.function_spec.is_lambda_function or len(function_input_lambda_args) != 1:
raise UNREACHABLE raise UNREACHABLE
lambda_function_name = function_input_lambda_args[0].value lambda_function_name = function_input_lambda_args[0].value
@ -196,7 +180,7 @@ class BuiltInFunction(Function, TypeHintedFunctionType):
] ]
# If a lambda is in a function's arg, resolve it differently # If a lambda is in a function's arg, resolve it differently
if self.is_lambda_function: if self.function_spec.is_lambda_function:
return self._resolve_lambda_function( return self._resolve_lambda_function(
resolved_arguments=resolved_arguments, resolved_arguments=resolved_arguments,
resolved_variables=resolved_variables, resolved_variables=resolved_variables,

View file

@ -9,6 +9,7 @@ from typing import final
from ytdl_sub.script.types.resolvable import Argument from ytdl_sub.script.types.resolvable import Argument
from ytdl_sub.script.types.resolvable import FunctionType from ytdl_sub.script.types.resolvable import FunctionType
from ytdl_sub.script.types.resolvable import Lambda
from ytdl_sub.script.types.resolvable import NamedCustomFunction from ytdl_sub.script.types.resolvable import NamedCustomFunction
from ytdl_sub.script.types.resolvable import ParsedCustomFunction from ytdl_sub.script.types.resolvable import ParsedCustomFunction
from ytdl_sub.script.types.resolvable import Resolvable from ytdl_sub.script.types.resolvable import Resolvable
@ -48,6 +49,18 @@ class VariableDependency(ABC):
return output return output
@final
@property
def lambda_arguments(self) -> Set[Lambda]:
output: Set[Lambda] = set()
for arg in self._iterable_arguments:
if isinstance(arg, Lambda):
output.add(arg)
if isinstance(arg, VariableDependency):
output.update(arg.lambda_arguments)
return output
@final @final
@property @property
def custom_functions(self) -> Set[ParsedCustomFunction]: def custom_functions(self) -> Set[ParsedCustomFunction]:

View file

@ -8,7 +8,7 @@ from ytdl_sub.script.types.resolvable import NamedType
from ytdl_sub.script.types.resolvable import TypeHintedFunctionType from ytdl_sub.script.types.resolvable import TypeHintedFunctionType
from ytdl_sub.script.utils.exceptions import IncompatibleFunctionArguments from ytdl_sub.script.utils.exceptions import IncompatibleFunctionArguments
from ytdl_sub.script.utils.exceptions import UserException 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 FunctionSpec
from ytdl_sub.script.utils.type_checking import get_optional_type 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_optional
from ytdl_sub.script.utils.type_checking import is_union from ytdl_sub.script.utils.type_checking import is_union
@ -100,7 +100,7 @@ class ParserExceptionFormatter:
class FunctionArgumentsExceptionFormatter: class FunctionArgumentsExceptionFormatter:
def __init__( def __init__(
self, self,
input_spec: FunctionInputSpec, input_spec: FunctionSpec,
function_instance: TypeHintedFunctionType, function_instance: TypeHintedFunctionType,
): ):
self._args = input_spec.args self._args = input_spec.args

View file

@ -1,4 +1,7 @@
import inspect
from dataclasses import dataclass from dataclasses import dataclass
from inspect import FullArgSpec
from typing import Callable
from typing import List from typing import List
from typing import Optional from typing import Optional
from typing import Type from typing import Type
@ -7,6 +10,7 @@ from typing import get_origin
from ytdl_sub.script.types.resolvable import Argument from ytdl_sub.script.types.resolvable import Argument
from ytdl_sub.script.types.resolvable import FunctionType from ytdl_sub.script.types.resolvable import FunctionType
from ytdl_sub.script.types.resolvable import Lambda
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.resolvable import TypeHintedFunctionType from ytdl_sub.script.types.resolvable import TypeHintedFunctionType
@ -90,7 +94,8 @@ def is_type_compatible(
@dataclass(frozen=True) @dataclass(frozen=True)
class FunctionInputSpec: class FunctionSpec:
return_type: Type[Resolvable]
args: Optional[List[Type[Resolvable | Optional[Resolvable]]]] = None args: Optional[List[Type[Resolvable | Optional[Resolvable]]]] = None
varargs: Optional[Type[Resolvable]] = None varargs: Optional[Type[Resolvable]] = None
@ -131,3 +136,21 @@ class FunctionInputSpec:
return self._is_varargs_compatible(input_args=input_args) return self._is_varargs_compatible(input_args=input_args)
raise UNREACHABLE # TODO: functions with no args raise UNREACHABLE # TODO: functions with no args
@property
def is_lambda_function(self) -> bool:
return Lambda in (self.args or [])
@classmethod
def from_callable(cls, callable_ref: Callable[..., Resolvable]) -> "FunctionSpec":
arg_spec: FullArgSpec = inspect.getfullargspec(callable_ref)
if arg_spec.varargs:
return FunctionSpec(
return_type=arg_spec.annotations["return"],
varargs=arg_spec.annotations[arg_spec.varargs],
)
return FunctionSpec(
return_type=arg_spec.annotations["return"],
args=[arg_spec.annotations[arg_name] for arg_name in arg_spec.args],
)