diff --git a/src/main/java/org/rumbledb/compiler/RuntimeIteratorVisitor.java b/src/main/java/org/rumbledb/compiler/RuntimeIteratorVisitor.java index 0009d88e7a..a8554a1fe6 100644 --- a/src/main/java/org/rumbledb/compiler/RuntimeIteratorVisitor.java +++ b/src/main/java/org/rumbledb/compiler/RuntimeIteratorVisitor.java @@ -398,7 +398,9 @@ private RuntimeTupleIterator visitFlowrClause( groupByExpressionIterator, variableName, clause.getMetadata(), - var.getCollationURI(), + var.getCollationURI() == null + ? clause.getStaticContext().getDefaultCollation() + : var.getCollationURI(), var.getActualSequenceType() ) ); diff --git a/src/main/java/org/rumbledb/runtime/flwor/udfs/GroupClauseCreateColumnsUDF.java b/src/main/java/org/rumbledb/runtime/flwor/udfs/GroupClauseCreateColumnsUDF.java index cc71aad1b0..5f13fc4bac 100644 --- a/src/main/java/org/rumbledb/runtime/flwor/udfs/GroupClauseCreateColumnsUDF.java +++ b/src/main/java/org/rumbledb/runtime/flwor/udfs/GroupClauseCreateColumnsUDF.java @@ -32,6 +32,7 @@ import org.rumbledb.exceptions.UnexpectedTypeException; import org.rumbledb.runtime.flwor.FlworDataFrameColumn; import org.rumbledb.runtime.flwor.expression.GroupByClauseSparkIteratorExpression; +import org.rumbledb.runtime.misc.CollationSupport; import org.rumbledb.runtime.typing.InstanceOfIterator; import org.rumbledb.types.SequenceType; @@ -45,6 +46,7 @@ public class GroupClauseCreateColumnsUDF implements UDF1 { private static final long serialVersionUID = 1L; private final DataFrameContext dataFrameContext; private final List groupingVariableNames; + private final List groupingCollationURIs; private final List groupingSequenceTypes; private final List results; @@ -71,9 +73,11 @@ public GroupClauseCreateColumnsUDF( ) { this.dataFrameContext = new DataFrameContext(context, columns); this.groupingVariableNames = new ArrayList<>(); + this.groupingCollationURIs = new ArrayList<>(); this.groupingSequenceTypes = new ArrayList<>(); for (GroupByClauseSparkIteratorExpression expression : groupingExpressions) { this.groupingVariableNames.add(expression.getVariableName()); + this.groupingCollationURIs.add(expression.getCollationURI()); this.groupingSequenceTypes.add(expression.getSequenceType()); } this.results = new ArrayList<>(); @@ -88,6 +92,7 @@ public Row call(Row row) { for (int i = 0; i < this.groupingVariableNames.size(); i++) { Name groupingVariableName = this.groupingVariableNames.get(i); + String collationURI = this.groupingCollationURIs.get(i); SequenceType declaredType = this.groupingSequenceTypes.get(i); List items = this.dataFrameContext.getContext() .getVariableValues() @@ -108,7 +113,11 @@ public Row call(Row row) { continue; } - Item nextItem = atomizedGroupingKey.get(0); + Item nextItem = CollationSupport.normalizeItemForCollation( + atomizedGroupingKey.get(0), + collationURI, + this.metadata + ); this.createColumnsForItem(nextItem); }