less duplication

This commit is contained in:
Jesse Bannon 2026-01-12 16:52:59 -08:00
parent 9b79a78dec
commit 26c2045a74
5 changed files with 18 additions and 19 deletions

View file

@ -699,12 +699,6 @@ class Script:
""" """
return set(to_function_definition_name(name) for name in self._functions.keys()) return set(to_function_definition_name(name) for name in self._functions.keys())
def _to_syntax_tree(self, maybe_resolved: SyntaxTree | Argument) -> SyntaxTree:
if isinstance(maybe_resolved, SyntaxTree):
return maybe_resolved
return SyntaxTree(ast=[maybe_resolved])
def resolve_partial( def resolve_partial(
self, self,
unresolvable: Optional[Set[str]] = None, unresolvable: Optional[Set[str]] = None,
@ -759,6 +753,6 @@ class Script:
return copy.deepcopy(self).add_parsed( return copy.deepcopy(self).add_parsed(
{var_name: self._variables[var_name] for var_name in unresolvable} {var_name: self._variables[var_name] for var_name in unresolvable}
| {var.name: self._to_syntax_tree(definition) for var, definition in resolved.items()} | {var.name: SyntaxTree(ast=[definition]) for var, definition in resolved.items()}
| {var.name: self._to_syntax_tree(definition) for var, definition in unresolved.items()} | {var.name: SyntaxTree(ast=[definition]) for var, definition in unresolved.items()}
) )

View file

@ -60,13 +60,13 @@ class UnresolvedArray(_Array, VariableDependency, FutureResolvable):
custom_functions=custom_functions, custom_functions=custom_functions,
) )
out = UnresolvedArray(value=maybe_resolvable_values)
if is_resolvable: if is_resolvable:
return self.resolve( return out.resolve(
resolved_variables=resolved_variables, resolved_variables=resolved_variables, custom_functions=custom_functions
custom_functions=custom_functions,
) )
return UnresolvedArray(value=maybe_resolvable_values) return out
def future_resolvable_type(self) -> Type[Resolvable]: def future_resolvable_type(self) -> Type[Resolvable]:
return Array return Array

View file

@ -101,11 +101,12 @@ class CustomFunction(Function, NamedCustomFunction):
for i in range(len(self.args)): for i in range(len(self.args)):
function_arg = FunctionArgument.from_idx(idx=i, custom_function_name=self.name) function_arg = FunctionArgument.from_idx(idx=i, custom_function_name=self.name)
function_value = maybe_resolvable_args[i]
if isinstance(maybe_resolvable_args[i], Resolvable): if isinstance(function_value, Resolvable):
resolved_variables[function_arg] = maybe_resolvable_args[i] resolved_variables[function_arg] = function_value
else: else:
unresolved_variables[function_arg] = maybe_resolvable_args[i] unresolved_variables[function_arg] = function_value
assert len(custom_functions[self.name].iterable_arguments) == 1 assert len(custom_functions[self.name].iterable_arguments) == 1
custom_function_definition = custom_functions[self.name].iterable_arguments[0] custom_function_definition = custom_functions[self.name].iterable_arguments[0]
@ -472,13 +473,14 @@ class BuiltInFunction(Function, BuiltInFunctionType):
custom_functions=custom_functions, custom_functions=custom_functions,
) )
out = BuiltInFunction(name=self.name, args=maybe_resolvable_values)
if is_resolvable: if is_resolvable:
return BuiltInFunction(name=self.name, args=maybe_resolvable_values).resolve( return out.resolve(
resolved_variables=resolved_variables, resolved_variables=resolved_variables,
custom_functions=custom_functions, custom_functions=custom_functions,
) )
return BuiltInFunction(name=self.name, args=maybe_resolvable_values) return out
def __hash__(self): def __hash__(self):
return hash((self.name, *self.args)) return hash((self.name, *self.args))

View file

@ -75,13 +75,14 @@ class UnresolvedMap(_Map, VariableDependency, FutureResolvable):
custom_functions=custom_functions, custom_functions=custom_functions,
) )
out = UnresolvedMap(value=dict(zip(maybe_resolvable_keys, maybe_resolvable_values)))
if is_keys_resolvable and is_values_resolvable: if is_keys_resolvable and is_values_resolvable:
return self.resolve( return out.resolve(
resolved_variables=resolved_variables, resolved_variables=resolved_variables,
custom_functions=custom_functions, custom_functions=custom_functions,
) )
return UnresolvedMap(value=dict(zip(maybe_resolvable_keys, maybe_resolvable_values))) return out
def future_resolvable_type(self) -> Type[Resolvable]: def future_resolvable_type(self) -> Type[Resolvable]:
return Map return Map

View file

@ -47,6 +47,7 @@ class SyntaxTree(VariableDependency):
unresolved_variables: Dict[Variable, Argument], unresolved_variables: Dict[Variable, Argument],
custom_functions: Dict[str, VariableDependency], custom_functions: Dict[str, VariableDependency],
) -> Argument | Resolvable: ) -> Argument | Resolvable:
# Ensure this does not get returned as a SyntaxTree since nesting them is not supported.
maybe_resolvable_values, _ = VariableDependency.try_partial_resolve( maybe_resolvable_values, _ = VariableDependency.try_partial_resolve(
args=self.ast, args=self.ast,
resolved_variables=resolved_variables, resolved_variables=resolved_variables,
@ -54,6 +55,7 @@ class SyntaxTree(VariableDependency):
custom_functions=custom_functions, custom_functions=custom_functions,
) )
# Mimic the above resolve behavior
if len(maybe_resolvable_values) > 1: if len(maybe_resolvable_values) > 1:
return BuiltInFunction(name="concat", args=maybe_resolvable_values) return BuiltInFunction(name="concat", args=maybe_resolvable_values)