diff --git a/connector/protobuf/src/main/scala/org/apache/spark/sql/protobuf/utils/ProtobufOptions.scala b/connector/protobuf/src/main/scala/org/apache/spark/sql/protobuf/utils/ProtobufOptions.scala index 9a54b06a9d7f5..8732c6489d086 100644 --- a/connector/protobuf/src/main/scala/org/apache/spark/sql/protobuf/utils/ProtobufOptions.scala +++ b/connector/protobuf/src/main/scala/org/apache/spark/sql/protobuf/utils/ProtobufOptions.scala @@ -240,6 +240,15 @@ private[sql] class ProtobufOptions( val unwrapWellKnownTypes: Boolean = getBoolean("unwrap.primitive.wrapper.types", defaultValue = false) + // Whether google.protobuf.Timestamp and google.protobuf.Duration are converted to Spark's + // native TimestampType / DayTimeIntervalType. This is the default. Setting this option to + // false leaves them as ordinary messages, so they deserialize as struct (the raw proto shape) via the generic message path. Useful when a downstream + // schema was defined against the struct representation and must stay stable, or when a + // consumer needs the raw seconds/nanos rather than a micro-truncated interval/timestamp. + val convertTimestampDurationToNative: Boolean = + getBoolean("convert.timestamp.duration.to.native", defaultValue = true) + // Since Spark doesn't allow writing empty StructType, empty proto message type will be // dropped by default. Setting this option to true will insert a dummy column to empty proto // message so that the empty message will be retained. diff --git a/connector/protobuf/src/main/scala/org/apache/spark/sql/protobuf/utils/SchemaConverters.scala b/connector/protobuf/src/main/scala/org/apache/spark/sql/protobuf/utils/SchemaConverters.scala index 5a73d9aaa72be..4b0a69ab979e5 100644 --- a/connector/protobuf/src/main/scala/org/apache/spark/sql/protobuf/utils/SchemaConverters.scala +++ b/connector/protobuf/src/main/scala/org/apache/spark/sql/protobuf/utils/SchemaConverters.scala @@ -100,13 +100,15 @@ object SchemaConverters extends Logging { case BYTE_STRING => Some(BinaryType) case ENUM => if (protobufOptions.enumsAsInts) Some(IntegerType) else Some(StringType) case MESSAGE - if (fd.getMessageType.getName == "Duration" && + if (protobufOptions.convertTimestampDurationToNative && + fd.getMessageType.getName == "Duration" && fd.getMessageType.getFields.size() == 2 && fd.getMessageType.getFields.get(0).getName.equals("seconds") && fd.getMessageType.getFields.get(1).getName.equals("nanos")) => Some(DayTimeIntervalType.defaultConcreteType) case MESSAGE - if (fd.getMessageType.getName == "Timestamp" && + if (protobufOptions.convertTimestampDurationToNative && + fd.getMessageType.getName == "Timestamp" && fd.getMessageType.getFields.size() == 2 && fd.getMessageType.getFields.get(0).getName.equals("seconds") && fd.getMessageType.getFields.get(1).getName.equals("nanos")) => diff --git a/connector/protobuf/src/test/scala/org/apache/spark/sql/protobuf/ProtobufFunctionsSuite.scala b/connector/protobuf/src/test/scala/org/apache/spark/sql/protobuf/ProtobufFunctionsSuite.scala index 5c875989e7144..49b564a30f06b 100644 --- a/connector/protobuf/src/test/scala/org/apache/spark/sql/protobuf/ProtobufFunctionsSuite.scala +++ b/connector/protobuf/src/test/scala/org/apache/spark/sql/protobuf/ProtobufFunctionsSuite.scala @@ -724,6 +724,75 @@ class ProtobufFunctionsSuite extends SharedSparkSession with ProtobufTestBase } } + test("convert.timestamp.duration.to.native=false keeps Timestamp/Duration as struct") { + // Build a proto payload the normal way (native Timestamp/Duration on the write side), + // then read it back with the conversion disabled and assert both fields deserialize as + // struct -- the raw proto shape -- instead of TimestampType / + // DayTimeIntervalType. + val tsSchema = StructType( + StructField("timeStampMsg", + StructType( + StructField("key", StringType, nullable = true) :: + StructField("stmp", TimestampType, nullable = true) :: Nil + ), nullable = true) :: Nil) + val ts = Timestamp.valueOf("2016-05-09 10:12:43.999") + val tsInput = spark.createDataFrame( + spark.sparkContext.parallelize(Seq( + Row(Row("key1", ts)))), + tsSchema) + + val expectedStruct = StructType( + StructField("seconds", LongType, nullable = true) :: + StructField("nanos", IntegerType, nullable = true) :: Nil) + + checkWithFileAndClassName("timeStampMsg") { + case (name, descFilePathOpt) => + val toProtoDf = tsInput + .select(to_protobuf_wrapper($"timeStampMsg", name, descFilePathOpt) as Symbol("to_proto")) + val fromProtoDf = toProtoDf + .select(from_protobuf_wrapper($"to_proto", name, descFilePathOpt, + Map("convert.timestamp.duration.to.native" -> "false")) as Symbol("timeStampMsg")) + + val stmpField = fromProtoDf.schema("timeStampMsg").dataType + .asInstanceOf[StructType]("stmp") + assert(stmpField.dataType === expectedStruct) + // Derive the expected epoch from the input Timestamp so the assertion is independent + // of the JVM default time zone (Timestamp.valueOf parses in the local zone). + val row = fromProtoDf.select("timeStampMsg.stmp.seconds", "timeStampMsg.stmp.nanos").first() + assert(row.getLong(0) === ts.getTime / 1000) + assert(row.getInt(1) === 999000000) + } + + val durSchema = StructType( + StructField("durationMsg", + StructType( + StructField("key", StringType, nullable = true) :: + StructField("duration", DayTimeIntervalType.defaultConcreteType, nullable = true) :: + Nil + ), nullable = true) :: Nil) + val durInput = spark.createDataFrame( + spark.sparkContext.parallelize(Seq( + Row(Row("key1", Duration.ofSeconds(4).plusNanos(500000000))))), + durSchema) + + checkWithFileAndClassName("durationMsg") { + case (name, descFilePathOpt) => + val toProtoDf = durInput + .select(to_protobuf_wrapper($"durationMsg", name, descFilePathOpt) as Symbol("to_proto")) + val fromProtoDf = toProtoDf + .select(from_protobuf_wrapper($"to_proto", name, descFilePathOpt, + Map("convert.timestamp.duration.to.native" -> "false")) as Symbol("durationMsg")) + + val durField = fromProtoDf.schema("durationMsg").dataType + .asInstanceOf[StructType]("duration") + assert(durField.dataType === expectedStruct) + val row = fromProtoDf.select("durationMsg.duration.seconds", "durationMsg.duration.nanos") + .first() + assert(row.getLong(0) === 4L) + assert(row.getInt(1) === 500000000) + } + } + test("SPARK-57573: Handle TimeType between to_protobuf and from_protobuf") { val schema = StructType( StructField("timeMsg",