diff --git a/src/main/java/org/rumbledb/context/VariableValues.java b/src/main/java/org/rumbledb/context/VariableValues.java index 6b8e63002e..2d68b279b0 100644 --- a/src/main/java/org/rumbledb/context/VariableValues.java +++ b/src/main/java/org/rumbledb/context/VariableValues.java @@ -21,13 +21,11 @@ package org.rumbledb.context; import org.apache.spark.api.java.JavaRDD; -import org.apache.spark.sql.Row; import org.rumbledb.api.Item; import org.rumbledb.config.RumbleRuntimeConfiguration; import org.rumbledb.errorcodes.ErrorCode; import org.rumbledb.exceptions.*; import org.rumbledb.items.ItemFactory; -import org.rumbledb.items.parsing.RowToItemMapper; import org.rumbledb.items.structured.JSoundDataFrame; import org.rumbledb.runtime.HybridRuntimeIterator; @@ -223,7 +221,7 @@ public List getLocalVariableValue(Name varName, ExceptionMetadata metadata } JSoundDataFrame df = this.getDataFrameVariableValue(varName, metadata); return HybridRuntimeIterator.collectRDDwithLimit( - HybridRuntimeIterator.dataFrameToRDDOfItems(df, metadata), + df.toRDD(metadata), this.configuration, metadata ); @@ -271,8 +269,7 @@ public JavaRDD getRDDVariableValue(Name varName, ExceptionMetadata metadat throw new JobWithinAJobException(metadata); } JSoundDataFrame df = this.dataFrameVariableValues.get(varName); - JavaRDD rowRDD = df.javaRDD(); - return rowRDD.map(new RowToItemMapper(metadata, df.getItemType())); + return df.toRDD(metadata); } if (this.parent != null) { @@ -466,4 +463,3 @@ public void changeVariableValue(Name varName, JSoundDataFrame value) { nodeWithVariableDecl.dataFrameVariableValues.put(varName, value); } } - diff --git a/src/main/java/org/rumbledb/exceptions/ExitStatementException.java b/src/main/java/org/rumbledb/exceptions/ExitStatementException.java index 28ddc7508b..dfdd68f7e8 100644 --- a/src/main/java/org/rumbledb/exceptions/ExitStatementException.java +++ b/src/main/java/org/rumbledb/exceptions/ExitStatementException.java @@ -8,8 +8,6 @@ import java.io.Serial; import java.util.List; -import static org.rumbledb.runtime.HybridRuntimeIterator.dataFrameToRDDOfItems; - public class ExitStatementException extends RuntimeException { @Serial private static final long serialVersionUID = 1L; @@ -43,7 +41,7 @@ public List getLocalResult() { } else if (hasRDDResult()) { return this.rddResult.collect(); } else if (hasDataFrameResult()) { - return dataFrameToRDDOfItems(this.dataFrameResult, this.exceptionMetadata).collect(); + return this.dataFrameResult.toRDD(this.exceptionMetadata).collect(); } throw new OurBadException("Expected local result but there was nothing to return from the exit statement!"); } diff --git a/src/main/java/org/rumbledb/items/structured/JSoundDataFrame.java b/src/main/java/org/rumbledb/items/structured/JSoundDataFrame.java index 4f542af224..2ccffce6b2 100644 --- a/src/main/java/org/rumbledb/items/structured/JSoundDataFrame.java +++ b/src/main/java/org/rumbledb/items/structured/JSoundDataFrame.java @@ -18,6 +18,8 @@ import org.rumbledb.exceptions.ExceptionMetadata; import org.rumbledb.exceptions.OurBadException; import org.rumbledb.items.parsing.ItemParser; +import org.rumbledb.items.parsing.RowToItemMapper; +import org.rumbledb.runtime.dataframe.RuntimeDataFrame; import org.rumbledb.runtime.flwor.FlworDataFrameColumn; import org.rumbledb.types.BuiltinTypesCatalogue; import org.rumbledb.types.ItemType; @@ -25,7 +27,7 @@ import sparksoniq.spark.SparkSessionManager; -public class JSoundDataFrame implements Serializable { +public class JSoundDataFrame implements RuntimeDataFrame, Serializable { @Serial private static final long serialVersionUID = 1L; @@ -126,6 +128,7 @@ public static JSoundDataFrame emptyDataFrame() { ); } + @Override public Dataset getDataFrame() { return this.dataFrame; } @@ -135,8 +138,15 @@ public void show() { this.dataFrame.show(); } - public JavaRDD javaRDD() { - return this.dataFrame.javaRDD(); + /** + * Converts this JSONiq DataFrame to its logical item representation. + * + * @param metadata query metadata used if a row cannot be decoded + * @return an RDD containing the represented items + */ + @Override + public JavaRDD toRDD(ExceptionMetadata metadata) { + return this.dataFrame.javaRDD().map(new RowToItemMapper(metadata, this.itemType)); } public long count() { diff --git a/src/main/java/org/rumbledb/runtime/HybridRuntimeIterator.java b/src/main/java/org/rumbledb/runtime/HybridRuntimeIterator.java index 31595bfc60..2374ac46b6 100644 --- a/src/main/java/org/rumbledb/runtime/HybridRuntimeIterator.java +++ b/src/main/java/org/rumbledb/runtime/HybridRuntimeIterator.java @@ -21,7 +21,6 @@ package org.rumbledb.runtime; import org.apache.spark.api.java.JavaRDD; -import org.apache.spark.sql.Row; import org.rumbledb.api.Item; import org.rumbledb.config.RumbleRuntimeConfiguration; import org.rumbledb.context.DynamicContext; @@ -32,9 +31,7 @@ import org.rumbledb.exceptions.MoreThanOneItemException; import org.rumbledb.exceptions.NoItemException; import org.rumbledb.expressions.ExecutionMode; -import org.rumbledb.items.parsing.RowToItemMapper; -import org.rumbledb.items.structured.JSoundDataFrame; - +import org.rumbledb.runtime.dataframe.RuntimeDataFrame; import sparksoniq.spark.SparkSessionManager; import java.io.Serial; @@ -100,10 +97,7 @@ public boolean hasNext() { this.currentResultIndex = 0; JavaRDD rdd = null; if (!isRDD() && implementsDataFrames()) { - rdd = dataFrameToRDDOfItems( - this.getDataFrame(this.currentDynamicContextForLocalExecution), - this.getMetadata() - ); + rdd = this.getDataFrame(this.currentDynamicContextForLocalExecution).toRDD(this.getMetadata()); } else { rdd = this.getRDDAux(this.currentDynamicContextForLocalExecution); } @@ -141,8 +135,8 @@ public Item next() { @Override public JavaRDD getRDD(DynamicContext context) { if ((isDataFrame() && implementsDataFrames()) || (isRDD() && implementsDataFrames() && !implementsRDD())) { - JSoundDataFrame df = this.getDataFrame(context); - return dataFrameToRDDOfItems(df, getMetadata()); + RuntimeDataFrame df = this.getDataFrame(context); + return df.toRDD(getMetadata()); } if (isRDDOrDataFrame()) { return getRDDAux(context); @@ -151,11 +145,6 @@ public JavaRDD getRDD(DynamicContext context) { return SparkSessionManager.getInstance().getJavaSparkContext().parallelize(contents); } - public static JavaRDD dataFrameToRDDOfItems(JSoundDataFrame df, ExceptionMetadata metadata) { - JavaRDD rowRDD = df.javaRDD(); - return rowRDD.map(new RowToItemMapper(metadata, df.getItemType())); - } - public static List collectRDDwithLimit( JavaRDD rdd, RumbleRuntimeConfiguration configuration, diff --git a/src/main/java/org/rumbledb/runtime/RuntimeIterator.java b/src/main/java/org/rumbledb/runtime/RuntimeIterator.java index 1cf58b27f3..6013e0f101 100644 --- a/src/main/java/org/rumbledb/runtime/RuntimeIterator.java +++ b/src/main/java/org/rumbledb/runtime/RuntimeIterator.java @@ -42,16 +42,13 @@ import org.rumbledb.expressions.ExecutionMode; import org.rumbledb.expressions.comparison.ComparisonExpression.ComparisonOperator; import org.rumbledb.items.structured.JSoundDataFrame; +import org.rumbledb.runtime.dataframe.ItemRuntimeDataFrameFactory; import org.rumbledb.runtime.flwor.NativeClauseContext; import org.rumbledb.runtime.misc.ComparisonIterator; -import org.rumbledb.runtime.typing.TypeInferrenceUtils; -import org.rumbledb.runtime.typing.ValidateTypeIterator; import org.rumbledb.runtime.update.PendingUpdateList; import org.rumbledb.types.BuiltinTypesCatalogue; -import org.rumbledb.types.ItemType; import org.rumbledb.types.SequenceType; - public abstract class RuntimeIterator implements RuntimeIteratorInterface { protected static final String FLOW_EXCEPTION_MESSAGE = "Invalid next() call; "; @@ -311,57 +308,9 @@ public final JSoundDataFrame getOrCreateDataFrame(DynamicContext context) { return this.getDataFrame(context); } if (isRDD()) { - if (this.getStaticType().getItemType().isCompatibleWithDataFrames(this.getConfiguration())) { - return ValidateTypeIterator.convertRDDToValidDataFrame( - this.getRDD(context), - this.getStaticType().getItemType(), - context, - true, - this.staticContext - ); - } else { - JavaRDD rdd = this.getRDD(context); - ItemType type = TypeInferrenceUtils.inferItemTypeOfRDDItems( - rdd, - getMetadata(), - TypeInferrenceUtils.TypeMergeMode.LAX - ); - return ValidateTypeIterator.convertRDDToValidDataFrame( - rdd, - type, - context, - true, - this.staticContext - ); - } - } - List items = new ArrayList<>(); - materialize(context, items); - if (this.getStaticType().getItemType().isCompatibleWithDataFrames(this.getConfiguration())) { - return ValidateTypeIterator.convertLocalItemsToDataFrame( - items, - this.getStaticType().getItemType(), - context, - true, - this.staticContext - ); - } else { - ItemType type = TypeInferrenceUtils.inferItemTypeOfLocalItems( - items, - getMetadata(), - TypeInferrenceUtils.TypeMergeMode.LAX - ); - if (this.getConfiguration().printInferredTypes()) { - System.err.println("Inferred DataFrame type:\n" + this.getStaticType().getItemType()); - } - return ValidateTypeIterator.convertLocalItemsToDataFrame( - items, - type, - context, - true, - this.staticContext - ); + return ItemRuntimeDataFrameFactory.INSTANCE.fromRDD(this.getRDD(context), context, this.staticContext); } + return ItemRuntimeDataFrameFactory.INSTANCE.fromLocal(this.materialize(context), context, this.staticContext); } public boolean isUpdating() { diff --git a/src/main/java/org/rumbledb/runtime/dataframe/ItemRuntimeDataFrameFactory.java b/src/main/java/org/rumbledb/runtime/dataframe/ItemRuntimeDataFrameFactory.java new file mode 100644 index 0000000000..25e39091bc --- /dev/null +++ b/src/main/java/org/rumbledb/runtime/dataframe/ItemRuntimeDataFrameFactory.java @@ -0,0 +1,71 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0. + */ + +package org.rumbledb.runtime.dataframe; + +import java.io.Serial; +import java.util.List; + +import org.apache.spark.api.java.JavaRDD; +import org.rumbledb.api.Item; +import org.rumbledb.context.DynamicContext; +import org.rumbledb.context.RuntimeStaticContext; +import org.rumbledb.items.structured.JSoundDataFrame; +import org.rumbledb.runtime.typing.TypeInferrenceUtils; +import org.rumbledb.runtime.typing.ValidateTypeIterator; +import org.rumbledb.types.ItemType; + +/** + * Encodes item RDDs as {@link JSoundDataFrame}s. + */ +public final class ItemRuntimeDataFrameFactory implements RuntimeDataFrameFactory { + + @Serial + private static final long serialVersionUID = 1L; + + public static final ItemRuntimeDataFrameFactory INSTANCE = new ItemRuntimeDataFrameFactory(); + + private ItemRuntimeDataFrameFactory() { + } + + @Override + public JSoundDataFrame fromLocal( + List items, + DynamicContext context, + RuntimeStaticContext staticContext + ) { + ItemType itemType = staticContext.getStaticType().getItemType(); + if (!itemType.isCompatibleWithDataFrames(staticContext.getConfiguration())) { + itemType = TypeInferrenceUtils.inferItemTypeOfLocalItems( + items, + staticContext.getMetadata(), + TypeInferrenceUtils.TypeMergeMode.LAX + ); + if (staticContext.getConfiguration().printInferredTypes()) { + System.err.println("Inferred DataFrame type:\n" + itemType); + } + } + return ValidateTypeIterator.convertLocalItemsToDataFrame(items, itemType, context, true, staticContext); + } + + @Override + public JSoundDataFrame fromRDD( + JavaRDD rdd, + DynamicContext context, + RuntimeStaticContext staticContext + ) { + ItemType itemType = staticContext.getStaticType().getItemType(); + if (!itemType.isCompatibleWithDataFrames(staticContext.getConfiguration())) { + itemType = TypeInferrenceUtils.inferItemTypeOfRDDItems( + rdd, + staticContext.getMetadata(), + TypeInferrenceUtils.TypeMergeMode.LAX + ); + } + return ValidateTypeIterator.convertRDDToValidDataFrame(rdd, itemType, context, true, staticContext); + } +} diff --git a/src/main/java/org/rumbledb/runtime/dataframe/RuntimeDataFrame.java b/src/main/java/org/rumbledb/runtime/dataframe/RuntimeDataFrame.java new file mode 100644 index 0000000000..69be8dc7cd --- /dev/null +++ b/src/main/java/org/rumbledb/runtime/dataframe/RuntimeDataFrame.java @@ -0,0 +1,41 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0. + */ + +package org.rumbledb.runtime.dataframe; + +import java.io.Serializable; + +import org.apache.spark.api.java.JavaRDD; +import org.apache.spark.sql.Dataset; +import org.apache.spark.sql.Row; +import org.rumbledb.exceptions.ExceptionMetadata; + +/** + * A Spark DataFrame whose rows represent runtime values of type {@code T}. + * + *

+ * Spark stores the physical representation as {@link Row}; implementations define how rows are mapped back to the + * logical runtime type. + *

+ * + * @param the logical runtime value represented by each row + */ +public interface RuntimeDataFrame extends Serializable { + + /** + * Returns the underlying physical Spark DataFrame. + */ + Dataset getDataFrame(); + + /** + * Converts this DataFrame to its logical runtime representation. + * + * @param metadata query metadata used if a row cannot be decoded + * @return an RDD of logical runtime values + */ + JavaRDD toRDD(ExceptionMetadata metadata); +} diff --git a/src/main/java/org/rumbledb/runtime/dataframe/RuntimeDataFrameFactory.java b/src/main/java/org/rumbledb/runtime/dataframe/RuntimeDataFrameFactory.java new file mode 100644 index 0000000000..e55a0eaaef --- /dev/null +++ b/src/main/java/org/rumbledb/runtime/dataframe/RuntimeDataFrameFactory.java @@ -0,0 +1,35 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0. + */ + +package org.rumbledb.runtime.dataframe; + +import java.io.Serializable; +import java.util.List; + +import org.apache.spark.api.java.JavaRDD; +import org.rumbledb.context.DynamicContext; +import org.rumbledb.context.RuntimeStaticContext; + +/** + * Creates a typed runtime DataFrame from the logical values stored in an RDD. + * + * @param the logical runtime value represented by each DataFrame row + */ +public interface RuntimeDataFrameFactory extends Serializable { + + RuntimeDataFrame fromLocal( + List values, + DynamicContext context, + RuntimeStaticContext staticContext + ); + + RuntimeDataFrame fromRDD( + JavaRDD rdd, + DynamicContext context, + RuntimeStaticContext staticContext + ); +} diff --git a/src/main/java/org/rumbledb/runtime/flwor/FlworDataFrame.java b/src/main/java/org/rumbledb/runtime/flwor/FlworDataFrame.java index 8c155cb33b..f96a3dc87f 100644 --- a/src/main/java/org/rumbledb/runtime/flwor/FlworDataFrame.java +++ b/src/main/java/org/rumbledb/runtime/flwor/FlworDataFrame.java @@ -7,16 +7,21 @@ import java.util.List; import java.util.Map; +import org.apache.spark.api.java.JavaRDD; import org.apache.spark.sql.Dataset; import org.apache.spark.sql.Row; import org.apache.spark.sql.UDFRegistration; import org.apache.spark.sql.types.StructType; import org.rumbledb.context.DynamicContext; import org.rumbledb.context.Name; +import org.rumbledb.exceptions.ExceptionMetadata; import org.rumbledb.exceptions.OurBadException; +import org.rumbledb.runtime.dataframe.RuntimeDataFrame; import org.rumbledb.types.SequenceType; -public class FlworDataFrame implements Serializable { +import sparksoniq.jsoniq.tuple.FlworTuple; + +public class FlworDataFrame implements RuntimeDataFrame, Serializable { @Serial private static final long serialVersionUID = 1L; @@ -37,10 +42,19 @@ public FlworDataFrame(Dataset dataFrame) { } } + @Override public Dataset getDataFrame() { return this.dataFrame; } + @Override + public JavaRDD toRDD(ExceptionMetadata metadata) { + throw new OurBadException( + "Converting a FLWOR DataFrame to an RDD of tuples is not implemented.", + metadata + ); + } + public List getColumns() { return this.columns; } diff --git a/src/main/java/org/rumbledb/runtime/functions/input/ParallelizeFunctionIterator.java b/src/main/java/org/rumbledb/runtime/functions/input/ParallelizeFunctionIterator.java index 10d03acefc..23dda26e9b 100644 --- a/src/main/java/org/rumbledb/runtime/functions/input/ParallelizeFunctionIterator.java +++ b/src/main/java/org/rumbledb/runtime/functions/input/ParallelizeFunctionIterator.java @@ -58,7 +58,7 @@ public JavaRDD getRDDAux(DynamicContext context) { List contents = new ArrayList<>(); if (this.sequenceIterator.isDataFrame()) { JSoundDataFrame dataFrame = this.sequenceIterator.getDataFrame(context); - rdd = dataFrameToRDDOfItems(dataFrame, this.getMetadata()); + rdd = dataFrame.toRDD(this.getMetadata()); if (this.getChildren().size() == 1) { return rdd; } else { diff --git a/src/main/java/org/rumbledb/runtime/navigation/SequenceLookupIterator.java b/src/main/java/org/rumbledb/runtime/navigation/SequenceLookupIterator.java index d5580a6ba6..bd3b7a8b1f 100644 --- a/src/main/java/org/rumbledb/runtime/navigation/SequenceLookupIterator.java +++ b/src/main/java/org/rumbledb/runtime/navigation/SequenceLookupIterator.java @@ -43,8 +43,6 @@ import java.util.Arrays; import java.util.List; -import static org.rumbledb.runtime.HybridRuntimeIterator.dataFrameToRDDOfItems; - public class SequenceLookupIterator extends AtMostOneItemLocalRuntimeIterator { @Serial @@ -123,10 +121,7 @@ public Item lookupDF(DynamicContext dynamicContext) { ), df.getItemType() ); - JavaRDD rdd = dataFrameToRDDOfItems( - df, - this.getMetadata() - ); + JavaRDD rdd = df.toRDD(this.getMetadata()); List results = rdd.take(1); if (results.isEmpty()) { diff --git a/src/main/java/org/rumbledb/runtime/typing/ValidateTypeIterator.java b/src/main/java/org/rumbledb/runtime/typing/ValidateTypeIterator.java index 66f8f62d50..fe2b7ad123 100644 --- a/src/main/java/org/rumbledb/runtime/typing/ValidateTypeIterator.java +++ b/src/main/java/org/rumbledb/runtime/typing/ValidateTypeIterator.java @@ -83,7 +83,7 @@ public JSoundDataFrame getDataFrame(DynamicContext context) { if (actualType.isSubtypeOf(this.itemType)) { return inputDataAsDataFrame; } - JavaRDD inputDataAsRDDOfItems = dataFrameToRDDOfItems(inputDataAsDataFrame, getMetadata()); + JavaRDD inputDataAsRDDOfItems = inputDataAsDataFrame.toRDD(getMetadata()); return convertRDDToValidDataFrame( inputDataAsRDDOfItems, this.itemType, diff --git a/src/main/java/sparksoniq/jsoniq/tuple/FlworTuple.java b/src/main/java/sparksoniq/jsoniq/tuple/FlworTuple.java index 97eb91f382..640c7c4cd1 100644 --- a/src/main/java/sparksoniq/jsoniq/tuple/FlworTuple.java +++ b/src/main/java/sparksoniq/jsoniq/tuple/FlworTuple.java @@ -21,13 +21,11 @@ package sparksoniq.jsoniq.tuple; import org.apache.spark.api.java.JavaRDD; -import org.apache.spark.sql.Row; import org.rumbledb.api.Item; import org.rumbledb.config.RumbleRuntimeConfiguration; import org.rumbledb.context.Name; import org.rumbledb.exceptions.ExceptionMetadata; import org.rumbledb.exceptions.OurBadException; -import org.rumbledb.items.parsing.RowToItemMapper; import org.rumbledb.items.structured.JSoundDataFrame; import org.rumbledb.runtime.HybridRuntimeIterator; @@ -137,8 +135,7 @@ public JavaRDD getRDDValue(Name key, ExceptionMetadata metadata) { } if (this.dataFrameVariables.containsKey(key)) { JSoundDataFrame df = this.dataFrameVariables.get(key); - JavaRDD rowRDD = df.javaRDD(); - return rowRDD.map(new RowToItemMapper(metadata, df.getItemType())); + return df.toRDD(metadata); } throw new OurBadException("Undeclared FLOWR variable", metadata); } diff --git a/src/main/java/sparksoniq/spark/ml/AnnotateFunctionIterator.java b/src/main/java/sparksoniq/spark/ml/AnnotateFunctionIterator.java index f771510e68..06ca8e3c92 100644 --- a/src/main/java/sparksoniq/spark/ml/AnnotateFunctionIterator.java +++ b/src/main/java/sparksoniq/spark/ml/AnnotateFunctionIterator.java @@ -44,7 +44,7 @@ public JSoundDataFrame getDataFrame(DynamicContext context) { if (actualSchemaType.isSubtypeOf(schemaType)) { return inputDataAsDataFrame; } - JavaRDD inputDataAsRDDOfItems = dataFrameToRDDOfItems(inputDataAsDataFrame, getMetadata()); + JavaRDD inputDataAsRDDOfItems = inputDataAsDataFrame.toRDD(getMetadata()); return ValidateTypeIterator.convertRDDToValidDataFrame( inputDataAsRDDOfItems, schemaType, diff --git a/src/test/java/iq/JavaAPITest.java b/src/test/java/iq/JavaAPITest.java index b7144d7bb3..f55df0568c 100644 --- a/src/test/java/iq/JavaAPITest.java +++ b/src/test/java/iq/JavaAPITest.java @@ -170,6 +170,19 @@ public void testGetAsDataFramePreservesArrayMembers() throws Throwable { Assertions.assertEquals("s", items.get(1).getItemByKey("arr").getItemAt(0).getItemByKey("x").getStringValue()); } + @Test + @Timeout(1000) + public void testGetAsDataFrameFromRDD() throws Throwable { + Rumble rumble = new Rumble(RumbleRuntimeConfiguration.getDefaultConfiguration()); + SequenceOfItems iterator = rumble.runQuery("parallelize(({ \"x\" : 1 }, { \"x\" : 2 }))"); + + List rows = iterator.getAsDataFrame().collectAsList(); + + Assertions.assertEquals(2, rows.size()); + Assertions.assertEquals(1, ((Number) rows.get(0).get(0)).intValue()); + Assertions.assertEquals(2, ((Number) rows.get(1).get(0)).intValue()); + } + @Test @Timeout(1000) public void testHtmlSerializationRejectsEmptyMap() {