Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view

This file was deleted.

Original file line number Diff line number Diff line change
Expand Up @@ -267,7 +267,9 @@ public StaticContext visitFunctionCall(FunctionCallExpression expression, Static
expression.getMetadata()
);
}
if (BuiltinFunctionCatalogue.exists(expression.getFunctionIdentifier())) {
if (expression.isPartialApplication()) {
expression.setHighestExecutionMode(ExecutionMode.LOCAL);
} else if (BuiltinFunctionCatalogue.exists(expression.getFunctionIdentifier())) {
BuiltinFunction builtinFunction = BuiltinFunctionCatalogue.getBuiltinFunction(
expression.getFunctionIdentifier()
);
Expand Down
8 changes: 0 additions & 8 deletions src/main/java/org/rumbledb/compiler/InferTypeVisitor.java
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,6 @@
import org.rumbledb.exceptions.OurBadException;
import org.rumbledb.exceptions.UnexpectedStaticTypeException;
import org.rumbledb.exceptions.UnknownFunctionCallException;
import org.rumbledb.exceptions.UnsupportedFeatureException;
import org.rumbledb.expressions.AbstractNodeVisitor;
import org.rumbledb.expressions.CommaExpression;
import org.rumbledb.expressions.Expression;
Expand Down Expand Up @@ -761,13 +760,6 @@ public StaticContext visitFunctionCall(FunctionCallExpression expression, Static
visitDescendants(expression, argument);

if (BuiltinFunctionCatalogue.exists(expression.getFunctionIdentifier())) {
if (expression.isPartialApplication()) {
/// This should never be reached because partial application on built-in functions should have been rewritten before
throw new UnsupportedFeatureException(
"Partial application on built-in functions are not supported.",
expression.getMetadata()
);
}
BuiltinFunction builtinFunction = BuiltinFunctionCatalogue.getBuiltinFunction(
expression.getFunctionIdentifier()
);
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1330,7 +1330,18 @@ public RuntimeIterator visitFunctionCall(FunctionCallExpression expression, Runt
FunctionIdentifier identifier = new FunctionIdentifier(fnName, arity);

RuntimeIterator runtimeIterator = null;
if (BuiltinFunctionCatalogue.exists(identifier)) {
if (expression.isPartialApplication()) {
runtimeIterator = new DynamicFunctionCallIterator(
new NamedFunctionRefRuntimeIterator(
identifier,
expression.getStaticContextForRuntime(this.config, this.visitorConfig)
),
arguments,
expression.getStaticContextForRuntime(this.config, this.visitorConfig)
);
}

else if (BuiltinFunctionCatalogue.exists(identifier)) {
runtimeIterator = NamedFunctions.getBuiltInFunctionIterator(
identifier,
arguments,
Expand Down
9 changes: 0 additions & 9 deletions src/main/java/org/rumbledb/compiler/VisitorHelpers.java
Original file line number Diff line number Diff line change
Expand Up @@ -65,15 +65,6 @@ private static void inferTypes(Module module, RumbleRuntimeConfiguration conf) {

private static MainModule applyTypeIndependentOptimizations(MainModule module, RumbleRuntimeConfiguration conf) {
MainModule result = module;
if (conf.debug()) {
System.err.println("***************************************");
System.err.println("Builtin Partial Application Rewrite Visitor");
System.err.println("***************************************");
}
result = (MainModule) new BuiltinPartialApplicationRewriteVisitor().visit(result, null);
if (conf.debug()) {
printTree(result, conf);
}
// Annotate recursive functions as such
if (conf.debug()) {
System.err.println("***************************************");
Expand Down
157 changes: 149 additions & 8 deletions src/main/java/org/rumbledb/context/NamedFunctions.java
Original file line number Diff line number Diff line change
Expand Up @@ -29,13 +29,21 @@
import org.rumbledb.exceptions.DuplicateFunctionIdentifierException;
import org.rumbledb.exceptions.ExceptionMetadata;
import org.rumbledb.exceptions.OurBadException;
import org.rumbledb.exceptions.UnsupportedFeatureException;
import org.rumbledb.exceptions.UnknownFunctionCallException;
import org.rumbledb.expressions.ExecutionMode;
import org.rumbledb.items.FunctionItem;
import org.rumbledb.items.PartiallyAppliedFunctionItem;
import org.rumbledb.items.PartiallyAppliedFunctionItem.ArgumentBinding;
import org.rumbledb.items.PartiallyAppliedFunctionItem.DataFrameBinding;
import org.rumbledb.items.PartiallyAppliedFunctionItem.LocalBinding;
import org.rumbledb.items.PartiallyAppliedFunctionItem.PlaceholderBinding;
import org.rumbledb.items.PartiallyAppliedFunctionItem.RddBinding;
import org.rumbledb.runtime.RuntimeIterator;
import org.rumbledb.runtime.functions.BuiltinFunctionItemCallIterator;
import org.rumbledb.runtime.functions.CapturedFunctionArgumentIterator;
import org.rumbledb.runtime.functions.FunctionCallArgumentCoercion;
import org.rumbledb.runtime.functions.FunctionItemCallIterator;
import org.rumbledb.runtime.functions.PartialFunctionCallIterator;
import org.rumbledb.runtime.functions.sequences.general.DataFunctionIterator;
import org.rumbledb.runtime.typing.AtMostOneItemTypePromotionIterator;
import org.rumbledb.runtime.typing.TypePromotionIterator;
Expand All @@ -44,6 +52,7 @@

import java.io.Serializable;
import java.lang.reflect.Constructor;
import java.util.ArrayList;
import java.util.Collections;
import java.util.HashMap;
import java.util.List;
Expand All @@ -52,6 +61,12 @@ public class NamedFunctions implements Serializable, KryoSerializable {

private static final long serialVersionUID = 1L;

public record ResolvedFunctionCall(
Item functionItem,
List<RuntimeIterator> arguments,
ExecutionMode executionMode) {
}

// two maps for User defined function are needed as execution mode is known at
// static analysis phase
// but functions items are fully known at runtimeIterator generation
Expand Down Expand Up @@ -104,7 +119,56 @@ public static RuntimeIterator buildFunctionItemCallIterator(
List<RuntimeIterator> arguments,
boolean isTailOptimization
) {
ExceptionMetadata metadata = callerRuntimeContext.getMetadata();
if (functionItem instanceof PartiallyAppliedFunctionItem) {
return buildResolvedFunctionItemCallIterator(
resolveFunctionItemCall(functionItem, arguments, callerRuntimeContext),
callerRuntimeContext
);
}
return buildDirectFunctionItemCallIterator(
functionItem,
callerRuntimeContext,
executionModeForFunctionCall,
arguments,
isTailOptimization
);
}

public static RuntimeIterator buildResolvedFunctionItemCallIterator(
ResolvedFunctionCall resolvedCall,
RuntimeStaticContext callerRuntimeContext
) {
return buildDirectFunctionItemCallIterator(
resolvedCall.functionItem(),
callerRuntimeContext,
resolvedCall.executionMode(),
new ArrayList<>(resolvedCall.arguments()),
false
);
}

private static RuntimeIterator buildDirectFunctionItemCallIterator(
Item functionItem,
RuntimeStaticContext callerRuntimeContext,
ExecutionMode executionModeForFunctionCall,
List<RuntimeIterator> arguments,
boolean isTailOptimization
) {
if (isTailOptimization) {
return new PartialFunctionCallIterator(
functionItem,
arguments,
callerRuntimeContext.withExecutionMode(ExecutionMode.LOCAL),
Name.TAIL_CALL_OPTIMIZATION
);
}
if (arguments.stream().anyMatch(a -> a == null)) {
return new PartialFunctionCallIterator(
functionItem,
arguments,
callerRuntimeContext.withExecutionMode(ExecutionMode.LOCAL)
);
}
SequenceType sequenceType = functionItem.getSignature().getReturnType();
SequenceType innerSequenceType = functionItem.getBodyIterator().getStaticType();
RuntimeStaticContext outerStaticContext = callerRuntimeContext.withStaticType(
Expand All @@ -118,12 +182,6 @@ public static RuntimeIterator buildFunctionItemCallIterator(
).withExecutionMode(executionModeForFunctionCall);
RuntimeIterator functionCallIterator;
if (functionItem.isBuiltinFunction()) {
if (arguments.stream().anyMatch(a -> a == null)) {
throw new UnsupportedFeatureException(
"Partial application of builtin named function references is not supported yet.",
metadata
);
}
functionCallIterator = new BuiltinFunctionItemCallIterator(
functionItem,
arguments,
Expand Down Expand Up @@ -169,6 +227,89 @@ public static RuntimeIterator buildFunctionItemCallIterator(
}
}

public static ResolvedFunctionCall resolveFunctionItemCall(
Item functionItem,
List<RuntimeIterator> arguments,
RuntimeStaticContext callerRuntimeContext
) {
Item resolvedFunction = functionItem;
List<RuntimeIterator> resolvedArguments = new ArrayList<>(arguments);
while (resolvedFunction instanceof PartiallyAppliedFunctionItem partiallyAppliedFunction) {
FunctionCallArgumentCoercion.validateArity(
resolvedFunction,
resolvedArguments,
callerRuntimeContext.getMetadata()
);
FunctionCallArgumentCoercion.wrapAccordingToSignature(
resolvedFunction,
resolvedArguments,
callerRuntimeContext
);
resolvedArguments = expandPartialArguments(
partiallyAppliedFunction,
resolvedArguments,
callerRuntimeContext
);
resolvedFunction = partiallyAppliedFunction.getTargetFunction();
}
ExecutionMode executionMode = resolvedArguments.stream().anyMatch(argument -> argument == null)
? ExecutionMode.LOCAL
: resolveDirectFunctionItemExecutionMode(resolvedFunction, resolvedArguments, callerRuntimeContext);
return new ResolvedFunctionCall(resolvedFunction, resolvedArguments, executionMode);
}

private static ExecutionMode resolveDirectFunctionItemExecutionMode(
Item functionItem,
List<RuntimeIterator> arguments,
RuntimeStaticContext callerRuntimeContext
) {
if (functionItem.isBuiltinFunction()) {
BuiltinFunction builtin = BuiltinFunctionCatalogue.getBuiltinFunction(functionItem.getIdentifier());
ExecutionMode firstArgumentMode = arguments.isEmpty() || arguments.get(0) == null
? ExecutionMode.LOCAL
: arguments.get(0).getHighestExecutionMode();
return BuiltinFunctionExecutionModes.resolve(
builtin,
firstArgumentMode,
callerRuntimeContext.getConfiguration()
);
}
return functionItem.getBodyIterator().getHighestExecutionMode();
}

private static List<RuntimeIterator> expandPartialArguments(
PartiallyAppliedFunctionItem functionItem,
List<RuntimeIterator> suppliedArguments,
RuntimeStaticContext callerRuntimeContext
) {
List<RuntimeIterator> result = new ArrayList<>();
int suppliedIndex = 0;
for (ArgumentBinding binding : functionItem.getArgumentBindings()) {
if (binding instanceof PlaceholderBinding) {
result.add(suppliedArguments.get(suppliedIndex++));
continue;
}
ExecutionMode executionMode;
if (binding instanceof DataFrameBinding) {
executionMode = ExecutionMode.DATAFRAME;
} else if (binding instanceof RddBinding) {
executionMode = ExecutionMode.RDD;
} else if (binding instanceof LocalBinding) {
executionMode = ExecutionMode.LOCAL;
} else {
throw new OurBadException("Unsupported partial-function argument binding.");
}
result.add(
CapturedFunctionArgumentIterator.create(
binding,
callerRuntimeContext.withStaticType(binding.sequenceType())
.withExecutionMode(executionMode)
)
);
}
return result;
}

public void addUserDefinedFunction(Item function, ExceptionMetadata meta) {
if (!function.isFunction()) {
throw new OurBadException("Only a function item can be added as a user-defined function.");
Expand Down
Loading
Loading