lambda function helper

This commit is contained in:
Jesse Bannon 2023-11-21 16:40:01 -08:00
parent 8f220b961d
commit aafda39d80

View file

@ -144,9 +144,9 @@ class BuiltInFunction(Function, TypeHintedFunctionType):
) )
@property @property
def lambda_function(self) -> Optional[str]: def lambda_argument(self) -> Optional[Lambda]:
if Lambda in (self.input_spec.args or []): if Lambda in (self.input_spec.args or []):
return [lam for lam in self.args if isinstance(lam, Lambda)][0].function_name return [lam for lam in self.args if isinstance(lam, Lambda)][0]
return None return None
@classmethod @classmethod
@ -170,13 +170,45 @@ class BuiltInFunction(Function, TypeHintedFunctionType):
return output_type return output_type
def _resolve_lambda_function(
self,
resolved_arguments: List[Resolvable | Lambda],
resolved_variables: Dict[Variable, Resolvable],
custom_functions: Dict[str, "VariableDependency"],
) -> Resolvable:
"""
Resolve the lambda function by
1. Calling the actual built-in function, which actually forms the input args to the
lambda. NOTE: the lambda argument MUST BE the last argument in the input spec!
2. Preemptively creating the lambda's unresolved output array using output args from (1)
3. Resolve it like any other syntax
"""
assert self.lambda_argument is not None
lambda_function_name = self.lambda_argument.function_name
lambda_args = self.callable(*resolved_arguments)
assert isinstance(lambda_args, ResolvedArray)
return self._resolve_argument_type(
arg=UnresolvedArray(
[
BuiltInFunction(name=lambda_function_name, args=lambda_arg.value)
if Functions.is_built_in(lambda_function_name)
else CustomFunction(name=lambda_function_name, args=lambda_arg.value)
for lambda_arg in lambda_args.value
]
),
resolved_variables=resolved_variables,
custom_functions=custom_functions,
)
def resolve( def resolve(
self, self,
resolved_variables: Dict[Variable, Resolvable], resolved_variables: Dict[Variable, Resolvable],
custom_functions: Dict[str, "VariableDependency"], custom_functions: Dict[str, "VariableDependency"],
) -> Resolvable: ) -> Resolvable:
if lambda_function := self.lambda_function: # Resolve all non-lambda arguments
resolved_args: List[Resolvable] = [ resolved_arguments: List[Resolvable | Lambda] = [
self._resolve_argument_type( self._resolve_argument_type(
arg=arg, arg=arg,
resolved_variables=resolved_variables, resolved_variables=resolved_variables,
@ -185,33 +217,17 @@ class BuiltInFunction(Function, TypeHintedFunctionType):
for arg in self.args for arg in self.args
if not isinstance(arg, Lambda) if not isinstance(arg, Lambda)
] ]
lambda_arg = [arg for arg in self.args if isinstance(arg, Lambda)]
lambda_args = self.callable(*(resolved_args + lambda_arg)) # If a lambda is in a function's arg, resolve it differently
assert isinstance(lambda_args, ResolvedArray) if lambda_argument := self.lambda_argument:
return self._resolve_lambda_function(
return self._resolve_argument_type( resolved_arguments=resolved_arguments + [lambda_argument],
arg=UnresolvedArray(
[
BuiltInFunction(name=lambda_function, args=lambda_arg.value)
if Functions.is_built_in(lambda_function)
else CustomFunction(name=lambda_function, args=lambda_arg.value)
for lambda_arg in lambda_args.value
]
),
resolved_variables=resolved_variables, resolved_variables=resolved_variables,
custom_functions=custom_functions, custom_functions=custom_functions,
) )
resolved_args: List[Resolvable] = [
self._resolve_argument_type(
arg=arg, resolved_variables=resolved_variables, custom_functions=custom_functions
)
for arg in self.args
]
try: try:
return self.callable(*resolved_args) return self.callable(*resolved_arguments)
except UserThrownRuntimeError: except UserThrownRuntimeError:
raise raise
except Exception as exc: except Exception as exc: