From 44abe264d9a7aa48361168d488efee6f8197bd01 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?C=C3=A9dric=20Chantepie?= Date: Fri, 25 Sep 2026 16:16:05 +0200 Subject: [PATCH 1/2] Field TypedEncoder --- .../main/scala/frameless/RecordEncoder.scala | 13 +++++++- .../main/scala/frameless/TypedDataset.scala | 2 +- .../main/scala/frameless/TypedEncoder.scala | 9 ++++++ .../test/scala/frameless/ColumnTests.scala | 32 +++++++++++++++++++ 4 files changed, 54 insertions(+), 2 deletions(-) diff --git a/dataset/src/main/scala/frameless/RecordEncoder.scala b/dataset/src/main/scala/frameless/RecordEncoder.scala index e0ecdd47..7c4deac2 100644 --- a/dataset/src/main/scala/frameless/RecordEncoder.scala +++ b/dataset/src/main/scala/frameless/RecordEncoder.scala @@ -184,7 +184,18 @@ final class RecordFieldEncoder[T]( private[frameless] val jvmRepr: DataType, private[frameless] val fromCatalyst: Expression => Expression, private[frameless] val toCatalyst: Expression => Expression -) extends Serializable +) extends Serializable { self => + private[frameless] def toTypedEncoder = new TypedEncoder[T]()(encoder.classTag) { + def nullable: Boolean = encoder.nullable + + def jvmRepr: DataType = self.jvmRepr + def catalystRepr: DataType = encoder.catalystRepr + + def fromCatalyst(path: Expression): Expression = self.fromCatalyst(path) + + def toCatalyst(path: Expression): Expression = self.toCatalyst(path) + } +} object RecordFieldEncoder extends RecordFieldEncoderLowPriority { diff --git a/dataset/src/main/scala/frameless/TypedDataset.scala b/dataset/src/main/scala/frameless/TypedDataset.scala index 6a7780bd..b7ed61e8 100644 --- a/dataset/src/main/scala/frameless/TypedDataset.scala +++ b/dataset/src/main/scala/frameless/TypedDataset.scala @@ -956,7 +956,7 @@ class TypedDataset[T] protected[frameless] ( // now we need to unpack `Tuple1[A]` to `A` - TypedEncoder[A].catalystRepr match { + ea.catalystRepr match { case StructType(_) => // if column is struct, we use all its fields val df = diff --git a/dataset/src/main/scala/frameless/TypedEncoder.scala b/dataset/src/main/scala/frameless/TypedEncoder.scala index 8525edee..763ec80d 100644 --- a/dataset/src/main/scala/frameless/TypedEncoder.scala +++ b/dataset/src/main/scala/frameless/TypedEncoder.scala @@ -751,5 +751,14 @@ object TypedEncoder { } } + /** + * In case a type `T` encoding is supported in derivation (as a struct field), + * then this allows to resolve the corresponding `TypedEncoder[T]`, + * so the field can be handled invidiually. + */ + def usingFieldEncoder[T]( + implicit fieldEncoder: shapeless.Lazy[RecordFieldEncoder[T]] + ): TypedEncoder[T] = fieldEncoder.value.toTypedEncoder + object injections extends InjectionEnum } diff --git a/dataset/src/test/scala/frameless/ColumnTests.scala b/dataset/src/test/scala/frameless/ColumnTests.scala index baee9371..a83c0e7a 100644 --- a/dataset/src/test/scala/frameless/ColumnTests.scala +++ b/dataset/src/test/scala/frameless/ColumnTests.scala @@ -615,4 +615,36 @@ final class ColumnTests extends TypedDatasetSuite with Matchers { // we should be able to block the following as well... "ds.col(_.a.toInt)" shouldNot typeCheck } + + test("col through record encoder (for Value class)") { + import RecordEncoderTests.{Name, Person} + + val bar = new Name("bar") + val foo = new Name("foo") + + val ds: TypedDataset[Person] = + TypedDataset.create(Seq(Person(bar, 23), Person(foo, 11))) + + a[org.apache.spark.sql.AnalysisException] should be thrownBy { + // TypedEncoder[Name] is resolved using case class derivation, + // which is not compatible to the way such Value class + // is encoded as a field in another class, + // which leads to encoding/analysis error. + + ds.select(ds.col[Name](Symbol("name"))) + .collect + .run() + .toSeq shouldEqual Seq[Name](bar, foo) + } + + { + implicit def enc: TypedEncoder[Name] = + TypedEncoder.usingFieldEncoder[Name] + + ds.select(ds.col[Name](Symbol("name"))) + .collect + .run() + .toSeq shouldEqual Seq[Name](bar, foo) + } + } } From 43bc905f423ff1d3c35344861e335d7ba35a527e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?C=C3=A9dric=20Chantepie?= Date: Sun, 27 Sep 2026 18:48:11 +0200 Subject: [PATCH 2/2] WIP --- .../main/scala/frameless/TypedDataset.scala | 54 ++++++++++++++++--- .../apache/spark/sql/FramelessInternals.scala | 3 ++ .../apache/spark/sql/FramelessInternals.scala | 3 ++ .../apache/spark/sql/FramelessInternals.scala | 6 +++ dataset/src/test/resources/log4j2.properties | 6 +++ 5 files changed, 66 insertions(+), 6 deletions(-) diff --git a/dataset/src/main/scala/frameless/TypedDataset.scala b/dataset/src/main/scala/frameless/TypedDataset.scala index b7ed61e8..8f3bc047 100644 --- a/dataset/src/main/scala/frameless/TypedDataset.scala +++ b/dataset/src/main/scala/frameless/TypedDataset.scala @@ -1,14 +1,17 @@ package frameless import java.util + import frameless.functions.CatalystExplodableCollection import frameless.ops._ + import org.apache.spark.rdd.RDD import org.apache.spark.sql.{Column, DataFrame, Dataset, FramelessInternals, SparkSession} import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeReference, Literal} import org.apache.spark.sql.catalyst.plans.logical.{Join, JoinHint} import org.apache.spark.sql.catalyst.plans.Inner -import org.apache.spark.sql.types.StructType +import org.apache.spark.sql.types._ + import shapeless._ import shapeless.labelled.FieldType import shapeless.ops.hlist.{Diff, IsHCons, Mapper, Prepend, ToTraversable, Tupler} @@ -944,7 +947,7 @@ class TypedDataset[T] protected[frameless] ( /** * Type-safe projection from type T to Tuple1[A] * {{{ - * d.select( d('a), d('a)+d('b), ... ) + * d.select( d('a), d('a)+d('b), ... ) * }}} */ def select[A]( @@ -955,14 +958,13 @@ class TypedDataset[T] protected[frameless] ( val tuple1: TypedDataset[Tuple1[A]] = selectMany(ca) // now we need to unpack `Tuple1[A]` to `A` - - ea.catalystRepr match { + tuple1.dataset.schema.fields.head.dataType match { case StructType(_) => // if column is struct, we use all its fields - val df = + TypedDataset.create( tuple1.dataset.selectExpr("_1.*").as[A](TypedExpressionEncoder[A]) + ) - TypedDataset.create(df) case other => // for primitive types `Tuple1[A]` has the same schema as `A` TypedDataset.create(tuple1.dataset.as[A](TypedExpressionEncoder[A])) @@ -1189,6 +1191,30 @@ class TypedDataset[T] protected[frameless] ( } object selectMany extends ProductArgs { + private def sameShape(left: DataType, right: DataType): Boolean = + (left, right) match { + case (l: StructType, r: StructType) => + l.fields.map(_.dataType).zip(r.fields.map(_.dataType)).forall { + case (lField, rField) => sameShape(lField, rField) + } && l.fields.size == r.fields.size + + case (l: ArrayType, r: ArrayType) => + sameShape(l.elementType, r.elementType) + + case (l: MapType, r: MapType) => + sameShape(l.keyType, r.keyType) && + sameShape(l.valueType, r.valueType) + + case (_: StructType, _) | + (_, _: StructType) | + (_: ArrayType, _) | + (_, _: ArrayType) | + (_: MapType, _) | + (_, _: MapType) => + false + + case _ => true + } def applyProduct[U <: HList, Out0 <: HList, Out]( columns: U @@ -1205,6 +1231,22 @@ class TypedDataset[T] protected[frameless] ( .toList[UntypedExpression[T]] .map(c => FramelessInternals.column(c.expr)): _* ) + + val expectedSchema = TypedExpressionEncoder.targetStructType(i3) + val actualTypes = base.schema.fields.map(_.dataType) + val expectedTypes = expectedSchema.fields.map(_.dataType) + + if ( + actualTypes.size != expectedTypes.size || + !actualTypes.zip(expectedTypes).forall { + case (actual, expected) => sameShape(actual, expected) + } + ) { + throw FramelessInternals.analysisException( + s"Cannot decode selected columns as ${i3.classTag.runtimeClass.getName}: found ${base.schema}, expected $expectedSchema" + ) + } + val selected = base.as[Out](TypedExpressionEncoder[Out]) TypedDataset.create[Out](selected) diff --git a/dataset/src/main/spark-3.4+/org/apache/spark/sql/FramelessInternals.scala b/dataset/src/main/spark-3.4+/org/apache/spark/sql/FramelessInternals.scala index 79172360..fe5f91fe 100644 --- a/dataset/src/main/spark-3.4+/org/apache/spark/sql/FramelessInternals.scala +++ b/dataset/src/main/spark-3.4+/org/apache/spark/sql/FramelessInternals.scala @@ -47,6 +47,9 @@ object FramelessInternals { def getConf(ds: Dataset[_], key: String, default: String): String = ds.sqlContext.getConf(key, default) + def analysisException(message: String): AnalysisException = + new AnalysisException(message) + def joinPlan( ds: Dataset[_], plan: LogicalPlan, diff --git a/dataset/src/main/spark-3/org/apache/spark/sql/FramelessInternals.scala b/dataset/src/main/spark-3/org/apache/spark/sql/FramelessInternals.scala index 79172360..fe5f91fe 100644 --- a/dataset/src/main/spark-3/org/apache/spark/sql/FramelessInternals.scala +++ b/dataset/src/main/spark-3/org/apache/spark/sql/FramelessInternals.scala @@ -47,6 +47,9 @@ object FramelessInternals { def getConf(ds: Dataset[_], key: String, default: String): String = ds.sqlContext.getConf(key, default) + def analysisException(message: String): AnalysisException = + new AnalysisException(message) + def joinPlan( ds: Dataset[_], plan: LogicalPlan, diff --git a/dataset/src/main/spark-4/org/apache/spark/sql/FramelessInternals.scala b/dataset/src/main/spark-4/org/apache/spark/sql/FramelessInternals.scala index 850aac9e..206e8572 100644 --- a/dataset/src/main/spark-4/org/apache/spark/sql/FramelessInternals.scala +++ b/dataset/src/main/spark-4/org/apache/spark/sql/FramelessInternals.scala @@ -69,6 +69,12 @@ object FramelessInternals { def getConf(ds: Dataset[_], key: String, default: String): String = classic(ds).sparkSession.conf.get(key, default) + def analysisException(message: String): AnalysisException = + new AnalysisException( + errorClass = "UNRESOLVED_COLUMN.WITHOUT_SUGGESTION", + messageParameters = Map("objectName" -> message) + ) + def joinPlan( ds: Dataset[_], plan: LogicalPlan, diff --git a/dataset/src/test/resources/log4j2.properties b/dataset/src/test/resources/log4j2.properties index 1f672536..b626158d 100644 --- a/dataset/src/test/resources/log4j2.properties +++ b/dataset/src/test/resources/log4j2.properties @@ -23,6 +23,12 @@ logger.spark.level = warn logger.hadoop.name = org.apache.hadoop logger.hadoop.level = warn +logger.internal.name = org.apache.spark.sql.internal +logger.internal.level = error + +logger.fs.name = org.apache.hadoop.fs.FileSystem +logger.fs.level = error + # To debug expressions: #logger.codegen.name = org.apache.spark.sql.catalyst.expressions.codegen.CodeGenerator #logger.codegen.level = debug \ No newline at end of file