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())
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(
self,
unresolvable: Optional[Set[str]] = None,
@ -759,6 +753,6 @@ class Script:
return copy.deepcopy(self).add_parsed(
{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: self._to_syntax_tree(definition) for var, definition in unresolved.items()}
| {var.name: SyntaxTree(ast=[definition]) for var, definition in resolved.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,
)
out = UnresolvedArray(value=maybe_resolvable_values)
if is_resolvable:
return self.resolve(
resolved_variables=resolved_variables,
custom_functions=custom_functions,
return out.resolve(
resolved_variables=resolved_variables, custom_functions=custom_functions
)
return UnresolvedArray(value=maybe_resolvable_values)
return out
def future_resolvable_type(self) -> Type[Resolvable]:
return Array

View file

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

View file

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

View file

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