From 61611dfb1bde5f0f75b1807ecb6867e83cf0dec5 Mon Sep 17 00:00:00 2001 From: Matt Van Horn <455140+mvanhorn@users.noreply.github.com> Date: Tue, 28 Jul 2026 05:13:13 -0700 Subject: [PATCH 1/2] fix: address self-review findings --- .../jetbrains/kotlinx/dataframe/api/sort.kt | 3 +- .../kotlinx/dataframe/impl/api/sort.kt | 56 ++++++++++++-- .../jetbrains/kotlinx/dataframe/api/sort.kt | 75 +++++++++++++++++++ 3 files changed, 126 insertions(+), 8 deletions(-) diff --git a/core/src/main/kotlin/org/jetbrains/kotlinx/dataframe/api/sort.kt b/core/src/main/kotlin/org/jetbrains/kotlinx/dataframe/api/sort.kt index bb7e05607c..1c24a15fe8 100644 --- a/core/src/main/kotlin/org/jetbrains/kotlinx/dataframe/api/sort.kt +++ b/core/src/main/kotlin/org/jetbrains/kotlinx/dataframe/api/sort.kt @@ -177,8 +177,7 @@ public fun GroupBy.sortByDesc(vararg cols: KProperty? sortByDesc { cols.toColumnSet() } public fun GroupBy.sortByDesc(selector: SortColumnsSelector): GroupBy { - val set = selector.toColumnSet() - return sortByImpl { set.desc() } + return sortByImpl(selector, SortFlag.Reversed) } private fun GroupBy.createColumnFromGroupExpression( diff --git a/core/src/main/kotlin/org/jetbrains/kotlinx/dataframe/impl/api/sort.kt b/core/src/main/kotlin/org/jetbrains/kotlinx/dataframe/impl/api/sort.kt index 8ee0180df5..136b78c400 100644 --- a/core/src/main/kotlin/org/jetbrains/kotlinx/dataframe/impl/api/sort.kt +++ b/core/src/main/kotlin/org/jetbrains/kotlinx/dataframe/impl/api/sort.kt @@ -5,12 +5,14 @@ import org.jetbrains.kotlinx.dataframe.DataFrame import org.jetbrains.kotlinx.dataframe.api.GroupBy import org.jetbrains.kotlinx.dataframe.api.SortColumnsSelector import org.jetbrains.kotlinx.dataframe.api.asGroupBy +import org.jetbrains.kotlinx.dataframe.api.cast import org.jetbrains.kotlinx.dataframe.api.castFrameColumn import org.jetbrains.kotlinx.dataframe.api.getFrameColumn import org.jetbrains.kotlinx.dataframe.api.toDataFrame import org.jetbrains.kotlinx.dataframe.api.update import org.jetbrains.kotlinx.dataframe.api.with import org.jetbrains.kotlinx.dataframe.columns.ColumnResolutionContext +import org.jetbrains.kotlinx.dataframe.columns.ColumnPath import org.jetbrains.kotlinx.dataframe.columns.ColumnSet import org.jetbrains.kotlinx.dataframe.columns.ColumnWithPath import org.jetbrains.kotlinx.dataframe.columns.ColumnsResolver @@ -28,21 +30,52 @@ import org.jetbrains.kotlinx.dataframe.kind import org.jetbrains.kotlinx.dataframe.nrow @Suppress("UNCHECKED_CAST", "RemoveExplicitTypeArguments") -internal fun GroupBy.sortByImpl(columns: SortColumnsSelector): GroupBy = - toDataFrame() +internal fun GroupBy.sortByImpl( + columns: SortColumnsSelector, + sortFlag: SortFlag? = null, +): GroupBy { + val groupedDf = toDataFrame() + val groupsToValidate = groups.values().toList().ifEmpty { + listOf(DataFrame.empty(groups.schema.value).cast()) + } + validateSortColumns(groupsToValidate, groupedDf, columns) + return groupedDf // sort the individual groups by the columns specified .update { groups } - .with { it.sortByImpl(UnresolvedColumnsPolicy.Skip, columns) } + .with { it.sortByImpl(UnresolvedColumnsPolicy.Skip, columns, sortFlag) } // sort the groups by the columns specified (must be either be the keys column or "groups") // will do nothing if the columns specified are not the keys column or "groups" - .sortByImpl(UnresolvedColumnsPolicy.Skip, columns as SortColumnsSelector) + .sortByImpl(UnresolvedColumnsPolicy.Skip, columns as SortColumnsSelector, sortFlag) .asGroupBy { it.getFrameColumn(groups.name()).castFrameColumn() } +} + +@Suppress("UNCHECKED_CAST") +private fun validateSortColumns( + groups: List>, + groupedDf: DataFrame, + columns: SortColumnsSelector, +) { + val missingInGroups = groups + .map { columns.missingPaths(it) } + .reduce { missingInAllGroups, missingInGroup -> missingInAllGroups.intersect(missingInGroup) } + val missingInGroupedDf = (columns as SortColumnsSelector).missingPaths(groupedDf) + val missingInBoth = missingInGroups.intersect(missingInGroupedDf) + missingInBoth.firstOrNull()?.let { missingPath -> + groups.first().getSortColumns({ missingPath }, UnresolvedColumnsPolicy.Fail) + } +} + +private fun SortColumnsSelector.missingPaths(df: DataFrame): Set = + toColumnSet().resolve(df, UnresolvedColumnsPolicy.Create) + .mapNotNull { (it.data as? MissingColumnGroup<*>)?.path } + .toSet() internal fun DataFrame.sortByImpl( unresolvedColumnsPolicy: UnresolvedColumnsPolicy = UnresolvedColumnsPolicy.Fail, columns: SortColumnsSelector, + sortFlag: SortFlag? = null, ): DataFrame { - val sortColumns = getSortColumns(columns, unresolvedColumnsPolicy) + val sortColumns = getSortColumns(columns, unresolvedColumnsPolicy, sortFlag) if (sortColumns.isEmpty()) return this val compChain = sortColumns.map { @@ -71,6 +104,7 @@ internal fun AnyCol.createComparator(nullsLast: Boolean): java.util.Comparator DataFrame.getSortColumns( columns: SortColumnsSelector, unresolvedColumnsPolicy: UnresolvedColumnsPolicy, + sortFlag: SortFlag? = null, ): List> = columns.toColumnSet().resolve(this, unresolvedColumnsPolicy) // can appear using [DataColumn?.check] with UnresolvedColumnsPolicy.Skip @@ -82,6 +116,7 @@ internal fun DataFrame.getSortColumns( else -> throw IllegalStateException("Can not use ${col.kind} as sort column") } } + .map { if (sortFlag == null) it else it.addFlag(sortFlag) } internal enum class SortFlag { Reversed, NullsLast } @@ -110,7 +145,10 @@ internal fun ColumnWithPath.addFlag(flag: SortFlag): ColumnWithPath { } internal class ColumnSetWithSortFlag(val column: ColumnsResolver, val flag: SortFlag) : ColumnSet { - override fun resolve(context: ColumnResolutionContext) = column.resolve(context).map { it.addFlag(flag) } + override fun resolve(context: ColumnResolutionContext) = + column.resolve(context).map { + if (context.allowMissingColumns && it.data is MissingColumnGroup<*>) it else it.addFlag(flag) + } } internal class SortColumnDescriptor( @@ -119,6 +157,12 @@ internal class SortColumnDescriptor( val nullsLast: Boolean = false, ) : ValueColumnInternal by column.internalValueColumn() +private fun SortColumnDescriptor<*>.addFlag(flag: SortFlag): SortColumnDescriptor<*> = + when (flag) { + SortFlag.Reversed -> SortColumnDescriptor(column, direction.reversed(), nullsLast) + SortFlag.NullsLast -> SortColumnDescriptor(column, direction, true) + } + internal enum class SortDirection { Asc, Desc } internal fun SortDirection.reversed(): SortDirection = diff --git a/core/src/test/kotlin/org/jetbrains/kotlinx/dataframe/api/sort.kt b/core/src/test/kotlin/org/jetbrains/kotlinx/dataframe/api/sort.kt index e31765d6e7..020c9e1da3 100644 --- a/core/src/test/kotlin/org/jetbrains/kotlinx/dataframe/api/sort.kt +++ b/core/src/test/kotlin/org/jetbrains/kotlinx/dataframe/api/sort.kt @@ -1,7 +1,10 @@ package org.jetbrains.kotlinx.dataframe.api +import io.kotest.assertions.throwables.shouldThrow import io.kotest.assertions.throwables.shouldThrowMessage import io.kotest.matchers.shouldBe +import io.kotest.matchers.string.shouldContain +import io.kotest.matchers.string.shouldNotContain import org.jetbrains.kotlinx.dataframe.DataColumn import org.jetbrains.kotlinx.dataframe.DataFrame import org.jetbrains.kotlinx.dataframe.io.readCsv @@ -99,4 +102,76 @@ class SortDataColumn { aggregate.sortBy { pathOf("L", "extra") } } } + + @Test + fun `group by sort validates nested paths`() { + val df = DataFrame.readCsv(testResource("ds_salaries.csv")).cast() + val grouped = df.group { salaryInUsd }.into("group").groupBy { companyLocation } + val salaryPath = pathOf("group", "group", "salary_in_usd") + + grouped.sortBy { salaryPath }.groups.values().forEach { + val salaries = it[pathOf("group", "salary_in_usd")].values().map { value -> value as Int } + salaries shouldBe salaries.sorted() + } + grouped.sortByDesc { salaryPath }.groups.values().forEach { + val salaries = it[pathOf("group", "salary_in_usd")].values().map { value -> value as Int } + salaries shouldBe salaries.sortedDescending() + } + + val invalidPath = pathOf("group", "salaryInUsd") + val ascendingError = shouldThrow { + grouped.sortBy { invalidPath } + }.message + val descendingError = shouldThrow { + grouped.sortByDesc { invalidPath } + }.message + + ascendingError shouldBe descendingError + ascendingError.orEmpty() shouldContain "group/salaryInUsd" + ascendingError.orEmpty() shouldNotContain "Can not apply sort flag to column kind" + + shouldThrowMessage(ascendingError.orEmpty()) { + grouped.sortBy { salaryPath and invalidPath } + } + } + + @Test + fun `group by sort preserves group and key sorting`() { + val df = DataFrame.readCsv(testResource("ds_salaries.csv")).cast() + + df.groupBy { companyLocation }.sortBy { salaryInUsd }.groups.values().forEach { + val salaries = it[salaryInUsd].values() + salaries shouldBe salaries.sorted() + } + + val grouped = df.group { salaryInUsd }.into("group").groupBy { companyLocation } + val sortedKeys = grouped.sortBy { companyLocation }.keys[companyLocation].values() + sortedKeys shouldBe sortedKeys.sorted() + } + + @Test + fun `group by sort accepts columns present in later heterogeneous groups`() { + val grouped = dataFrameOf("key", "groups")( + "first", + dataFrameOf("other")(2), + "second", + dataFrameOf("value")(2, 1), + ).asGroupBy("groups") + + grouped.sortBy("value").groups.values().toList()[1]["value"].values() shouldBe listOf(1, 2) + } + + @Test + fun `group by sort validates paths when no groups are present`() { + val emptySource = dataFrameOf("value") { emptyList() }.groupBy("value") + val filteredGroups = dataFrameOf("value")(1).groupBy("value").filter { false } + val invalidPath = pathOf("missing", "nested") + + shouldThrow { + emptySource.sortBy { invalidPath } + }.message.orEmpty() shouldContain "missing/nested" + shouldThrow { + filteredGroups.sortBy { invalidPath } + }.message.orEmpty() shouldContain "missing/nested" + } } From 11d663d1b3c1de65ed070220c0b669d2b111948d Mon Sep 17 00:00:00 2001 From: Matt Van Horn <455140+mvanhorn@users.noreply.github.com> Date: Tue, 28 Jul 2026 05:21:47 -0700 Subject: [PATCH 2/2] fix: address round-2 residual (surgical round) --- .../kotlinx/dataframe/impl/columns/FrameColumnImpl.kt | 9 ++++++--- .../kotlin/org/jetbrains/kotlinx/dataframe/api/sort.kt | 7 +++++++ 2 files changed, 13 insertions(+), 3 deletions(-) diff --git a/core/src/main/kotlin/org/jetbrains/kotlinx/dataframe/impl/columns/FrameColumnImpl.kt b/core/src/main/kotlin/org/jetbrains/kotlinx/dataframe/impl/columns/FrameColumnImpl.kt index 4411e05c47..9565fd7d3a 100644 --- a/core/src/main/kotlin/org/jetbrains/kotlinx/dataframe/impl/columns/FrameColumnImpl.kt +++ b/core/src/main/kotlin/org/jetbrains/kotlinx/dataframe/impl/columns/FrameColumnImpl.kt @@ -53,11 +53,14 @@ internal open class FrameColumnImpl constructor( override fun forceResolve() = ResolvingFrameColumn(this) - override fun get(indices: Iterable): FrameColumn = - DataColumn.createFrameColumn( + override fun get(indices: Iterable): FrameColumn { + val indicesList = indices.toList() + return DataColumn.createFrameColumn( name = name, - groups = indices.map { values[it] }, + groups = indicesList.map { values[it] }, + schema = schema.takeIf { indicesList.isEmpty() }, ) + } override fun get(columnName: String) = throw UnsupportedOperationException("Can not get nested column '$columnName' from FrameColumn '$name'") diff --git a/core/src/test/kotlin/org/jetbrains/kotlinx/dataframe/api/sort.kt b/core/src/test/kotlin/org/jetbrains/kotlinx/dataframe/api/sort.kt index 020c9e1da3..122d7b2b10 100644 --- a/core/src/test/kotlin/org/jetbrains/kotlinx/dataframe/api/sort.kt +++ b/core/src/test/kotlin/org/jetbrains/kotlinx/dataframe/api/sort.kt @@ -174,4 +174,11 @@ class SortDataColumn { filteredGroups.sortBy { invalidPath } }.message.orEmpty() shouldContain "missing/nested" } + + @Test + fun `group by sort preserves schema after filtering all groups`() { + val filteredGroups = dataFrameOf("key", "value")(1, 2).groupBy("key").filter { false } + + filteredGroups.sortBy("value").toDataFrame() shouldBe filteredGroups.toDataFrame() + } }