Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -105,6 +105,12 @@ public String toJson(ZoneId zoneId) {
return new Variant(value, metadata).toJson(zoneId);
}

// Rejects a value nested more deeply than `maxNestingDepth` when that argument is positive.
// A non-positive value imposes no limit and preserves the previous behavior.
public String toJson(ZoneId zoneId, int maxNestingDepth) {
return new Variant(value, metadata).toJson(zoneId, maxNestingDepth);
}

/**
* @return A human-readable representation of the Variant value. It is always a JSON string at
* this moment.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -237,8 +237,15 @@ public Variant getElementAtIndex(int index) {
// Stringify the variant in JSON format.
// Throw `MALFORMED_VARIANT` if the variant is malformed.
public String toJson(ZoneId zoneId) {
return toJson(zoneId, -1);
}

// Stringify the variant in JSON format, rejecting a value nested more deeply than
// `maxNestingDepth` when that argument is positive. A non-positive value imposes no limit and
// preserves the previous behavior.
public String toJson(ZoneId zoneId, int maxNestingDepth) {
StringBuilder sb = new StringBuilder();
toJsonImpl(value, metadata, pos, sb, zoneId);
toJsonImpl(value, metadata, pos, sb, zoneId, 1, maxNestingDepth);
return sb.toString();
}

Expand Down Expand Up @@ -280,6 +287,15 @@ private static Instant microsToInstant(long timestamp) {
}

static void toJsonImpl(byte[] value, byte[] metadata, int pos, StringBuilder sb, ZoneId zoneId) {
toJsonImpl(value, metadata, pos, sb, zoneId, 1, -1);
}

static void toJsonImpl(byte[] value, byte[] metadata, int pos, StringBuilder sb, ZoneId zoneId,
int depth, int maxNestingDepth) {
if (maxNestingDepth > 0 && depth > maxNestingDepth) {
throw new IllegalStateException(
"Variant value nesting depth exceeds the configured maximum of " + maxNestingDepth + ".");
}
switch (VariantUtil.getType(value, pos)) {
case OBJECT:
handleObject(value, pos, (size, idSize, offsetSize, idStart, offsetStart, dataStart) -> {
Expand All @@ -291,7 +307,7 @@ static void toJsonImpl(byte[] value, byte[] metadata, int pos, StringBuilder sb,
if (i != 0) sb.append(',');
sb.append(escapeJson(getMetadataKey(metadata, id)));
sb.append(':');
toJsonImpl(value, metadata, elementPos, sb, zoneId);
toJsonImpl(value, metadata, elementPos, sb, zoneId, depth + 1, maxNestingDepth);
}
sb.append('}');
return null;
Expand All @@ -304,7 +320,7 @@ static void toJsonImpl(byte[] value, byte[] metadata, int pos, StringBuilder sb,
int offset = readUnsigned(value, offsetStart + offsetSize * i, offsetSize);
int elementPos = dataStart + offset;
if (i != 0) sb.append(',');
toJsonImpl(value, metadata, elementPos, sb, zoneId);
toJsonImpl(value, metadata, elementPos, sb, zoneId, depth + 1, maxNestingDepth);
}
sb.append(']');
return null;
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -810,7 +810,8 @@ case class Cast(
private lazy val castArgs = variant.VariantCastArgs(
evalMode != EvalMode.TRY,
timeZoneId,
zoneId)
zoneId,
SQLConf.get.getConf(SQLConf.VARIANT_MAX_NESTING_DEPTH))

def needsTimeZone: Boolean = Cast.needsTimeZone(child.dataType, dataType)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -449,7 +449,8 @@ case class VariantGet(
private lazy val castArgs = VariantCastArgs(
failOnError,
timeZoneId,
zoneId)
zoneId,
SQLConf.get.getConf(SQLConf.VARIANT_MAX_NESTING_DEPTH))

override def eval(input: InternalRow): Any = {
val _ = parsedPath
Expand Down Expand Up @@ -512,7 +513,8 @@ case class VariantGet(
case class VariantCastArgs(
failOnError: Boolean,
zoneStr: Option[String],
zoneId: ZoneId)
zoneId: ZoneId,
maxNestingDepth: Int = -1)

case object VariantGet {
/**
Expand Down Expand Up @@ -594,7 +596,8 @@ case object VariantGet {
def cast(v: Variant, dataType: DataType, castArgs: VariantCastArgs): Any = {
def invalidCast(): Any = {
if (castArgs.failOnError) {
throw QueryExecutionErrors.invalidVariantCast(v.toJson(castArgs.zoneId), dataType)
throw QueryExecutionErrors.invalidVariantCast(
v.toJson(castArgs.zoneId, castArgs.maxNestingDepth), dataType)
} else {
null
}
Expand Down Expand Up @@ -622,7 +625,7 @@ case object VariantGet {
val input = variantType match {
case Type.OBJECT | Type.ARRAY =>
return if (dataType.isInstanceOf[StringType]) {
UTF8String.fromString(v.toJson(castArgs.zoneId))
UTF8String.fromString(v.toJson(castArgs.zoneId, castArgs.maxNestingDepth))
} else {
invalidCast()
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@ import org.apache.spark.sql.catalyst.expressions.SpecializedGetters
import org.apache.spark.sql.catalyst.util._
import org.apache.spark.sql.catalyst.util.LegacyDateFormats.FAST_DATE_FORMAT
import org.apache.spark.sql.errors.QueryExecutionErrors
import org.apache.spark.sql.internal.SQLConf
import org.apache.spark.sql.types._
import org.apache.spark.unsafe.types.VariantVal
import org.apache.spark.util.ArrayImplicits._
Expand Down Expand Up @@ -111,6 +112,9 @@ class JacksonGenerator(
case customPattern => TimeFormatter(customPattern, isParsing = false)
}

// Read once at construction (per task) to avoid any per-row lookup overhead.
private val variantMaxNestingDepth = SQLConf.get.getConf(SQLConf.VARIANT_MAX_NESTING_DEPTH)

private def makeWriter(dataType: DataType): ValueWriter = dataType match {
case NullType =>
(row: SpecializedGetters, ordinal: Int) =>
Expand Down Expand Up @@ -351,7 +355,7 @@ class JacksonGenerator(
}

def write(v: VariantVal): Unit = {
gen.writeRawValue(v.toJson(options.zoneId))
gen.writeRawValue(v.toJson(options.zoneId, variantMaxNestingDepth))
}

def writeLineEnding(): Unit = {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -7074,6 +7074,18 @@ object SQLConf {
.booleanConf
.createWithDefault(true)

val VARIANT_MAX_NESTING_DEPTH =
buildConf("spark.sql.variant.maxNestingDepth")
.internal()
.doc("The maximum nesting depth allowed when converting a variant value to its JSON " +
"string form. When set to a positive value, converting a variant nested more deeply " +
"than this limit fails instead of recursing. A non-positive value (the default) " +
"imposes no limit and preserves the previous behavior.")
.version("4.3.0")
.withBindingPolicy(ConfigBindingPolicy.NOT_APPLICABLE)
.intConf
.createWithDefault(-1)

val PUSH_VARIANT_INTO_SCAN =
buildConf("spark.sql.variant.pushVariantIntoScan")
.internal()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -2058,6 +2058,25 @@ class VariantExpressionSuite extends SparkFunSuite with ExpressionEvalHelper {
name => Map("sizeLimit" -> "16.0 MiB", "functionName" -> s"`$name`"))
}

test("Variant.toJson enforces the configured maximum nesting depth") {
// A value nested `depth` arrays deep: [[[ ... 1 ... ]]].
val depth = 50
val json = ("[" * depth) + "1" + ("]" * depth)
val variant = VariantBuilder.parseJson(json, false)

// A non-positive limit imposes no bound and reproduces the previous behavior.
assert(variant.toJson(ZoneOffset.UTC) == json)
assert(variant.toJson(ZoneOffset.UTC, -1) == json)
// A limit at least as large as the actual nesting still renders the whole value.
assert(variant.toJson(ZoneOffset.UTC, depth * 2) == json)

// A limit smaller than the actual nesting is rejected instead of recursing all the way down.
val e = intercept[IllegalStateException] {
variant.toJson(ZoneOffset.UTC, 10)
}
assert(e.getMessage.contains("nesting depth"))
}

test("variant_strip_nulls") {
// Strip `input`, render the result back to JSON, and compare. `includeArrays` defaults to true.
def check(input: String, expected: String, includeArrays: Boolean = true): Unit = {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@ import org.apache.spark.sql.errors.{QueryCompilationErrors, QueryExecutionErrors
import org.apache.spark.sql.execution.RowToColumnConverter
import org.apache.spark.sql.execution.datasources.VariantMetadata
import org.apache.spark.sql.execution.vectorized.WritableColumnVector
import org.apache.spark.sql.internal.SQLConf
import org.apache.spark.sql.types._
import org.apache.spark.types.variant._
import org.apache.spark.types.variant.VariantUtil.Type
Expand Down Expand Up @@ -211,7 +212,8 @@ class ParquetVariantReader(
protected final def invalidCast(row: InternalRow, topLevelMetadata: Array[Byte]): Any = {
if (castArgs.failOnError) {
throw QueryExecutionErrors.invalidVariantCast(
rebuildVariant(row, topLevelMetadata).toJson(castArgs.zoneId), targetType)
rebuildVariant(row, topLevelMetadata).toJson(
castArgs.zoneId, castArgs.maxNestingDepth), targetType)
} else {
null
}
Expand Down Expand Up @@ -444,7 +446,8 @@ private[this] final class ScalarReader(
override def readFromTyped(row: InternalRow, topLevelMetadata: Array[Byte]): Any = {
if (castProject == null) {
return if (targetType.isInstanceOf[StringType]) {
UTF8String.fromString(rebuildVariant(row, topLevelMetadata).toJson(castArgs.zoneId))
UTF8String.fromString(
rebuildVariant(row, topLevelMetadata).toJson(castArgs.zoneId, castArgs.maxNestingDepth))
} else {
invalidCast(row, topLevelMetadata)
}
Expand Down Expand Up @@ -769,7 +772,8 @@ case object SparkShreddingUtils {
val reader = ParquetVariantReader(schema, f.dataType, VariantCastArgs(
metadata.failOnError,
Some(metadata.timeZoneId),
DateTimeUtils.getZoneId(metadata.timeZoneId)),
DateTimeUtils.getZoneId(metadata.timeZoneId),
SQLConf.get.getConf(SQLConf.VARIANT_MAX_NESTING_DEPTH)),
isTopLevelUnshredded = schemaPath.isEmpty && inputSchema.isUnshredded)
val castErrorOrdinal = companionIdxByDataName.getOrElse(f.name, -1)
if (castErrorOrdinal >= 0) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -523,4 +523,36 @@ class VariantEndToEndSuite extends SharedSparkSession {
}
}
}

test("SPARK-59325: spark.sql.variant.maxNestingDepth bounds rendering to JSON") {
def nestedObject(depth: Int): String =
if (depth == 0) "1" else s"""{"a":${nestedObject(depth - 1)}}"""
def nestedArray(depth: Int): String =
if (depth == 0) "1" else s"[${nestedArray(depth - 1)}]"
def causeMessages(t: Throwable): String =
if (t == null) "" else t.getMessage + " " + causeMessages(t.getCause)

Seq(nestedObject(20), nestedArray(20)).foreach { json =>
val df = Seq(json).toDF("v")

// Default (-1) and a generous limit render fully (to_json and CAST AS STRING).
checkAnswer(df.select(to_json(parse_json(col("v")))), Seq(Row(json)))
checkAnswer(df.selectExpr("CAST(parse_json(v) AS STRING)"), Seq(Row(json)))
withSQLConf(SQLConf.VARIANT_MAX_NESTING_DEPTH.key -> "100") {
checkAnswer(df.select(to_json(parse_json(col("v")))), Seq(Row(json)))
checkAnswer(df.selectExpr("CAST(parse_json(v) AS STRING)"), Seq(Row(json)))
}

// A limit below the nesting depth rejects on both render paths.
withSQLConf(SQLConf.VARIANT_MAX_NESTING_DEPTH.key -> "5") {
val renders = Seq(
() => df.select(to_json(parse_json(col("v")))).collect(),
() => df.selectExpr("CAST(parse_json(v) AS STRING)").collect())
renders.foreach { render =>
val e = intercept[Exception](render())
assert(causeMessages(e).contains("nesting depth exceeds the configured maximum"))
}
}
}
}
}