Describe the bug
With spark.comet.exec.pyarrowUDF.enabled=true (experimental, off by default), Comet reads the batches a mapInArrow function returns at the declared output types, without checking that the Arrow types match. A mismatch that Spark reports as an error comes back from Comet as wrong values, or as reads past the end of a buffer.
PySpark does not enforce the declared types on the Python side. In PySpark 4.1.1, wrap_arrow_batch_iter_udf pairs each result batch with the declared Arrow type. ArrowStreamUDFSerializer.dump_stream then drops it (for batch, _ in iterator) and wraps the batch with the batch's own schema. Vanilla Spark catches the mismatch on the JVM side, because ArrowColumnVector chooses its accessor from the physical vector. Reading an int column that arrived as int64 throws, and a decimal is rescaled to the declared scale. On 4.1, spark.sql.execution.arrow.pyspark.validateSchema.enabled also makes Spark reject the mismatch up front.
Comet has no equivalent check. CometArrowPythonRunnerBase (around line 290) wraps each vector as it arrives with CometVector.getVector(vector, null), and CometMapInBatchExec flattens the struct without a check. CometPlainVector's getters read raw memory at the declared width, and CometVector.getDecimal applies the declared scale to the raw unscaled value. So:
int64 data under a declared int: getInt(i) reads 4 bytes at address + 4 * i, so every other row is the high half of a value, and [0, 1, 2, 3] reads as [0, 0, 1, 0]. Plain Python ints become int64 in PyArrow, so returning pa.array(values) for an int column is enough to hit this.
int32 data under a declared bigint: getLong(i) reads 8 bytes at address + 8 * i, which runs past the end of the value buffer for the second half of the batch.
decimal(4,3) data under a declared decimal(10,2): 1.234 reads as 12.34, where Spark returns 1.23.
mapInPandas is not affected, because the pandas-to-Arrow conversion uses the declared types.
Steps to reproduce
This was found by reading the Comet, PySpark 4.1.1 and Spark sources. I could not run PySpark 4.x here (no Python 3.10+ with PyArrow), so the query below has not been run. It is the shape expected to show the problem:
spark.conf.set("spark.comet.exec.pyarrowUDF.enabled", "true")
df = spark.read.parquet(path) # id: int, read by a Comet scan
def f(it):
for b in it:
# Declared "id int", but hand back int64.
yield pa.RecordBatch.from_arrays([pa.array(b.column(0).to_pylist(), pa.int64())], ["id"])
df.mapInArrow(f, "id int").collect()
Spark raises. The accelerated path is expected to return wrong integers.
Expected behavior
Match Spark: raise an error (or rescale, where Spark does) rather than reading the buffers at the wrong width.
Additional context
A cheap fix is to check the output stream's schema once, when the runner first reads it, against the declared output:
- map the schema with
Utils.fromArrowField;
- compare ignoring compatible nullability;
- treat
LargeUtf8 and LargeBinary as equal to their 32-bit forms, since CometPlainVector reads both;
- on a mismatch, raise Spark's error or cast.
A related but smaller gap: dictionary-encoded output (pa.array(...).dictionary_encode()) reaches CometVector.getVector with a null provider and fails with a NullPointerException. Vanilla Spark fails on it too, with a clearer error. Passing the runner's ArrowReader, which is a DictionaryProvider, would fix the NPE.
Describe the bug
With
spark.comet.exec.pyarrowUDF.enabled=true(experimental, off by default), Comet reads the batches amapInArrowfunction returns at the declared output types, without checking that the Arrow types match. A mismatch that Spark reports as an error comes back from Comet as wrong values, or as reads past the end of a buffer.PySpark does not enforce the declared types on the Python side. In PySpark 4.1.1,
wrap_arrow_batch_iter_udfpairs each result batch with the declared Arrow type.ArrowStreamUDFSerializer.dump_streamthen drops it (for batch, _ in iterator) and wraps the batch with the batch's own schema. Vanilla Spark catches the mismatch on the JVM side, becauseArrowColumnVectorchooses its accessor from the physical vector. Reading anintcolumn that arrived asint64throws, and a decimal is rescaled to the declared scale. On 4.1,spark.sql.execution.arrow.pyspark.validateSchema.enabledalso makes Spark reject the mismatch up front.Comet has no equivalent check.
CometArrowPythonRunnerBase(around line 290) wraps each vector as it arrives withCometVector.getVector(vector, null), andCometMapInBatchExecflattens the struct without a check.CometPlainVector's getters read raw memory at the declared width, andCometVector.getDecimalapplies the declared scale to the raw unscaled value. So:int64data under a declaredint:getInt(i)reads 4 bytes ataddress + 4 * i, so every other row is the high half of a value, and[0, 1, 2, 3]reads as[0, 0, 1, 0]. Plain Python ints becomeint64in PyArrow, so returningpa.array(values)for anintcolumn is enough to hit this.int32data under a declaredbigint:getLong(i)reads 8 bytes ataddress + 8 * i, which runs past the end of the value buffer for the second half of the batch.decimal(4,3)data under a declareddecimal(10,2): 1.234 reads as 12.34, where Spark returns 1.23.mapInPandasis not affected, because the pandas-to-Arrow conversion uses the declared types.Steps to reproduce
This was found by reading the Comet, PySpark 4.1.1 and Spark sources. I could not run PySpark 4.x here (no Python 3.10+ with PyArrow), so the query below has not been run. It is the shape expected to show the problem:
Spark raises. The accelerated path is expected to return wrong integers.
Expected behavior
Match Spark: raise an error (or rescale, where Spark does) rather than reading the buffers at the wrong width.
Additional context
A cheap fix is to check the output stream's schema once, when the runner first reads it, against the declared output:
Utils.fromArrowField;LargeUtf8andLargeBinaryas equal to their 32-bit forms, sinceCometPlainVectorreads both;A related but smaller gap: dictionary-encoded output (
pa.array(...).dictionary_encode()) reachesCometVector.getVectorwith anullprovider and fails with aNullPointerException. Vanilla Spark fails on it too, with a clearer error. Passing the runner'sArrowReader, which is aDictionaryProvider, would fix the NPE.