From 04f78db8f19e88cf9a368c81e1a09360b607660d Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Thu, 2 Jul 2026 17:39:52 -0600 Subject: [PATCH 01/17] feat: add DataSketches HLL wrapper and Spark compatibility spike Wrap the pure-Rust datasketches crate's HLL_8 sketch/union behind SparkHllSketch/SparkHllUnion, hashing inputs via the crate's hash_value wrappers (raw_bytes, sign_extend) so MurmurHash3-x64-128 input matches DataSketches-Java. Verified cross-engine: Comet-produced sketches are byte-identical to Spark hll_sketch_agg output for HLL-array mode and mutually readable for low-cardinality List/Set mode. --- native/Cargo.lock | 7 + native/Cargo.toml | 1 + native/spark-expr/Cargo.toml | 1 + native/spark-expr/src/agg_funcs/hll_sketch.rs | 209 ++++++++++++++++++ native/spark-expr/src/agg_funcs/mod.rs | 2 + .../testdata/hll_sketch_spark_lgk12.bin | Bin 0 -> 4136 bytes 6 files changed, 220 insertions(+) create mode 100644 native/spark-expr/src/agg_funcs/hll_sketch.rs create mode 100644 native/spark-expr/src/agg_funcs/testdata/hll_sketch_spark_lgk12.bin diff --git a/native/Cargo.lock b/native/Cargo.lock index adb764fbfbe..4fc27f0299c 100644 --- a/native/Cargo.lock +++ b/native/Cargo.lock @@ -2141,6 +2141,7 @@ dependencies = [ "datafusion", "datafusion-comet-common", "datafusion-comet-jni-bridge", + "datasketches", "futures", "hex", "jni 0.22.4", @@ -2735,6 +2736,12 @@ dependencies = [ "sqlparser", ] +[[package]] +name = "datasketches" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "46c4cf71a36b46dcfc00e5014c0c20ccad2b1b6a008304d7d57d2749b2d41b3d" + [[package]] name = "debugid" version = "0.8.0" diff --git a/native/Cargo.toml b/native/Cargo.toml index 3e797eb9683..0c79dad57d3 100644 --- a/native/Cargo.toml +++ b/native/Cargo.toml @@ -49,6 +49,7 @@ datafusion-comet-proto = { path = "proto" } datafusion-comet-shuffle = { path = "shuffle" } chrono = { version = "0.4", default-features = false, features = ["clock"] } chrono-tz = { version = "0.10" } +datasketches = { version = "0.3.0", features = ["hll"] } futures = "0.3.32" num = "0.4" rand = "0.10" diff --git a/native/spark-expr/Cargo.toml b/native/spark-expr/Cargo.toml index 800fe3ecb17..d54094ff40a 100644 --- a/native/spark-expr/Cargo.toml +++ b/native/spark-expr/Cargo.toml @@ -44,6 +44,7 @@ twox-hash = "2.1.2" rand = { workspace = true } hex = "0.4.3" base64 = "0.22.1" +datasketches = { workspace = true } [dev-dependencies] arrow = {workspace = true} diff --git a/native/spark-expr/src/agg_funcs/hll_sketch.rs b/native/spark-expr/src/agg_funcs/hll_sketch.rs new file mode 100644 index 00000000000..5e40c54aaec --- /dev/null +++ b/native/spark-expr/src/agg_funcs/hll_sketch.rs @@ -0,0 +1,209 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +//! Thin wrapper over the `datasketches` crate's HLL sketch, isolating all +//! crate-specific API so Comet's aggregate/scalar code depends on a stable +//! surface. Every sketch uses `HllType::Hll8` and DataSketches' +//! `DEFAULT_UPDATE_SEED` (9001), matching Spark's `HllSketchAgg`. +//! +//! Input hashing goes through the crate's `hash_value` wrappers +//! (`raw_bytes` for strings/binary without Rust's length prefix, `sign_extend` +//! for narrow integers) so the MurmurHash3-x64-128 input bytes are identical to +//! DataSketches-Java. This makes the sketches mutually readable with Spark. +//! +//! Note: the crate serializes List/Set (low-cardinality) modes in DataSketches +//! *compact* form, whereas Spark emits the *updatable* form. The bytes are +//! therefore not byte-identical to Spark's output for small inputs, but +//! DataSketches `deserialize` reads both forms, so estimates round-trip in both +//! directions. Comet must own both Partial and Final aggregation +//! (`supportsMixedPartialFinal = false`) so this compact intermediate is only +//! ever read back by Comet. + +use datafusion::error::DataFusionError; +use datasketches::hash_value::{raw_bytes, sign_extend}; +use datasketches::hll::{HllSketch, HllType, HllUnion}; + +/// A DataSketches HLL_8 sketch configured to match Spark's `HllSketchAgg`. +#[derive(Debug)] +pub struct SparkHllSketch { + inner: HllSketch, +} + +impl SparkHllSketch { + /// Create an empty HLL_8 sketch with the given `lgConfigK`. + pub fn new(lg_config_k: u8) -> Self { + Self { + inner: HllSketch::new(lg_config_k, HllType::Hll8), + } + } + + /// Update with a 64-bit integer. Spark widens narrower integrals to `long` + /// before hashing; callers should pass the already-widened value here. + /// Rust's `Hash` for `i64` writes 8 little-endian bytes with no prefix, + /// matching DataSketches-Java `update(long)`. + pub fn update_i64(&mut self, v: i64) { + self.inner.update(v); + } + + /// Update with a narrow signed integer, sign-extending to 64 bits exactly as + /// Spark's `toLong` does before hashing. + pub fn update_i32(&mut self, v: i32) { + self.inner.update(sign_extend::from_i32(v)); + } + pub fn update_i16(&mut self, v: i16) { + self.inner.update(sign_extend::from_i16(v)); + } + pub fn update_i8(&mut self, v: i8) { + self.inner.update(sign_extend::from_i8(v)); + } + + /// Update with raw bytes (used for both StringType UTF-8 bytes and + /// BinaryType), hashing without Rust's slice length prefix. Empty inputs are + /// skipped, matching DataSketches (and Spark), which ignore empty values. + pub fn update_bytes(&mut self, v: &[u8]) { + if v.is_empty() { + return; + } + self.inner.update(raw_bytes::from_slice(v)); + } + + /// Serialize to DataSketches bytes (compact for List/Set modes, full for HLL + /// array modes). Readable by Spark's `hll_sketch_estimate` / `hll_union_agg`. + pub fn to_sketch_bytes(&self) -> Vec { + self.inner.serialize() + } + + /// Deserialize a DataSketches sketch (either compact or updatable form). + pub fn from_bytes(bytes: &[u8]) -> Result { + HllSketch::deserialize(bytes) + .map(|inner| Self { inner }) + .map_err(|e| DataFusionError::Internal(format!("invalid HLL sketch bytes: {e}"))) + } + + /// The configured `lgConfigK`. + pub fn lg_config_k(&self) -> u8 { + self.inner.lg_config_k() + } + + /// Raw cardinality estimate (caller rounds to `i64` for Spark). + pub fn estimate(&self) -> f64 { + self.inner.estimate() + } + + /// Merge another sketch into this one via a union, keeping HLL_8 output. + pub fn merge_sketch(&mut self, other: &SparkHllSketch) { + let mut u = HllUnion::new(self.lg_config_k()); + u.update(&self.inner); + u.update(&other.inner); + self.inner = u.to_sketch(HllType::Hll8); + } +} + +/// A DataSketches HLL union configured to match Spark's `HllUnionAgg`. +#[derive(Debug)] +pub struct SparkHllUnion { + inner: HllUnion, +} + +impl SparkHllUnion { + /// Create an empty union with the given `lgMaxK` (Spark fixes this at 12). + pub fn new(lg_max_k: u8) -> Self { + Self { + inner: HllUnion::new(lg_max_k), + } + } + + /// Merge a sketch into the union. + pub fn merge(&mut self, sketch: &SparkHllSketch) { + self.inner.update(&sketch.inner); + } + + /// The union result as an HLL_8 sketch's serialized bytes. + pub fn to_sketch_bytes(&self) -> Vec { + self.inner.to_sketch(HllType::Hll8).serialize() + } +} + +/// Estimate the distinct count from serialized sketch bytes, rounded to the +/// nearest `i64` (Spark's `hll_sketch_estimate` returns a `Long`). +pub fn estimate_from_bytes(bytes: &[u8]) -> Result { + let sketch = SparkHllSketch::from_bytes(bytes)?; + Ok(sketch.estimate().round() as i64) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn sketch_roundtrips_and_estimates() { + let mut s = SparkHllSketch::new(12); + for i in 0..1000i64 { + s.update_i64(i); + } + let bytes = s.to_sketch_bytes(); + let est = estimate_from_bytes(&bytes).unwrap(); + assert!((est - 1000).abs() <= 30, "estimate {est} not within 3% of 1000"); + } + + #[test] + fn union_merges_two_sketches() { + let mut a = SparkHllSketch::new(12); + for i in 0..1000i64 { + a.update_i64(i); + } + let mut b = SparkHllSketch::new(12); + for i in 500..1500i64 { + b.update_i64(i); + } + let mut u = SparkHllUnion::new(12); + u.merge(&a); + u.merge(&b); + let est = estimate_from_bytes(&u.to_sketch_bytes()).unwrap(); + assert!((est - 1500).abs() <= 45, "union estimate {est} not within 3% of 1500"); + } + + /// A sketch built from raw bytes (StringType/BinaryType path) round-trips and + /// estimates. Empty inputs are skipped, so they do not affect the estimate. + #[test] + fn byte_input_roundtrips_and_estimates() { + let mut s = SparkHllSketch::new(12); + for i in 0..1000i64 { + s.update_bytes(format!("val-{i}").as_bytes()); + } + s.update_bytes(b""); // skipped, no effect + let est = estimate_from_bytes(&s.to_sketch_bytes()).unwrap(); + assert!((est - 1000).abs() <= 30, "estimate {est} not within 3% of 1000"); + } + + /// Cross-engine regression guard: `testdata/hll_sketch_spark_lgk12.bin` was + /// produced by Spark 3.5's `hll_sketch_agg(id)` over `range(0, 1000)`. Comet + /// must read it and estimate the distinct count, proving the crate's + /// serialization stays DataSketches-Java compatible across crate upgrades. + /// (For this HLL_8 input the Comet-produced bytes are byte-identical to + /// Spark's; low-cardinality List/Set sketches differ in bytes but remain + /// mutually readable.) + #[test] + fn reads_spark_produced_sketch() { + let bytes = include_bytes!("testdata/hll_sketch_spark_lgk12.bin"); + let est = estimate_from_bytes(bytes).unwrap(); + assert!( + (est - 1000).abs() <= 30, + "estimate {est} of Spark-produced sketch not within 3% of 1000" + ); + } +} diff --git a/native/spark-expr/src/agg_funcs/mod.rs b/native/spark-expr/src/agg_funcs/mod.rs index 2a0322e46c9..57af25b35e4 100644 --- a/native/spark-expr/src/agg_funcs/mod.rs +++ b/native/spark-expr/src/agg_funcs/mod.rs @@ -19,6 +19,7 @@ mod avg; mod avg_decimal; mod correlation; mod covariance; +mod hll_sketch; mod stddev; mod sum_decimal; mod sum_int; @@ -29,6 +30,7 @@ pub use avg::Avg; pub use avg_decimal::AvgDecimal; pub use correlation::Correlation; pub use covariance::Covariance; +pub use hll_sketch::{estimate_from_bytes, SparkHllSketch, SparkHllUnion}; pub use stddev::Stddev; pub use sum_decimal::SumDecimal; pub use sum_int::SumInteger; diff --git a/native/spark-expr/src/agg_funcs/testdata/hll_sketch_spark_lgk12.bin b/native/spark-expr/src/agg_funcs/testdata/hll_sketch_spark_lgk12.bin new file mode 100644 index 0000000000000000000000000000000000000000..761477d5f69f2b44be248c34f56667948df43310 GIT binary patch literal 4136 zcmZWs+mX~j4D@UFs4KqtClewd1geOE1jv9YK1hNj>6bL}GH~{c9<7!}V|(BCw~yPl zy??(xef{zI*B8ux{`vNoca7iFVjR2gz8>2-mo>M|!R1x5W21E6jwhwEkO2T15)=HF zc)JR@wAZpL%yloL23y___#q}OPK1wz9o~%4-tj7DYn1{??fkb}d4{G2y?0w4}Yy)g%jM&6G^1U(Xa4I1dCO-8g@d1i6F9! z+J+NIBGe5R7N`2NOjWW5j{~DqmgS6IGM6LkdrK0R)W^JkwQaePuM7BK47d^?ua}V2rve0txIpt1Eb8TmBkB$+g@HH`9YES zR!d^4(>emjQe|Ryu|m4wjAAUt*kraTb^{Dns24g^!AqwE=5;41eMn?ffucY)jwU&= z*6|7;DuyrCcOR%ho*ZjLZ^tB>cubS&0r}j{aNrH1xOx7;GF-=a3T1|Rc{u>7QQ(JO zWkA)IFi>S4-KzBG08FpqkL_PoXTN)MA zOcI;SvEC8Nd~}obGLju7 zC&+28I1bgtB8Ttv^93Xj`LWrb>MN4iE*iRQlX$xrrxnLQP zNdvDqGIPG!A(h>@4PRyaQ!Mig$K05i%ca&|Ypgy&5FP<@J}6 z@Is%^qNy=qO4A2_#CO zQt7b7D2@G`GUx8Q^_#egVXx^aiOlu(Ahus4K*6}&NdAsvCWz=?Z2PI8G`@YJnw<&J c+9;Xmtur1CRPLc&ktAExT_@w26MgXa4=a5PVE_OC literal 0 HcmV?d00001 From 501e8ab7a0f3b2cbfc7d0beecc45a0c093adf3b0 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Thu, 2 Jul 2026 17:45:15 -0600 Subject: [PATCH 02/17] feat: add version-specific aggregate serde registration hook --- .../apache/comet/serde/QueryPlanSerde.scala | 45 ++++++++++--------- .../apache/comet/shims/CometExprShim.scala | 4 +- .../apache/comet/shims/CometExprShim.scala | 4 +- .../comet/shims/Spark4xCometExprShim.scala | 4 +- 4 files changed, 33 insertions(+), 24 deletions(-) diff --git a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala index b752f41d74c..379e6951925 100644 --- a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala +++ b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala @@ -384,27 +384,30 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim { /** * Mapping of Spark aggregate expression class to Comet expression handler. */ - val aggrSerdeMap: Map[Class[_], CometAggregateExpressionSerde[_]] = Map( - classOf[Average] -> CometAverage, - classOf[BitAndAgg] -> CometBitAndAgg, - classOf[BitOrAgg] -> CometBitOrAgg, - classOf[BitXorAgg] -> CometBitXOrAgg, - classOf[BloomFilterAggregate] -> CometBloomFilterAggregate, - classOf[CollectSet] -> CometCollectSet, - classOf[Corr] -> CometCorr, - classOf[Count] -> CometCount, - classOf[CovPopulation] -> CometCovPopulation, - classOf[CovSample] -> CometCovSample, - classOf[First] -> CometFirst, - classOf[Last] -> CometLast, - classOf[Max] -> CometMax, - classOf[Min] -> CometMin, - classOf[Percentile] -> CometPercentile, - classOf[StddevPop] -> CometStddevPop, - classOf[StddevSamp] -> CometStddevSamp, - classOf[Sum] -> CometSum, - classOf[VariancePop] -> CometVariancePop, - classOf[VarianceSamp] -> CometVarianceSamp) + val aggrSerdeMap: Map[Class[_], CometAggregateExpressionSerde[_]] = { + val base: Map[Class[_], CometAggregateExpressionSerde[_]] = Map( + classOf[Average] -> CometAverage, + classOf[BitAndAgg] -> CometBitAndAgg, + classOf[BitOrAgg] -> CometBitOrAgg, + classOf[BitXorAgg] -> CometBitXOrAgg, + classOf[BloomFilterAggregate] -> CometBloomFilterAggregate, + classOf[CollectSet] -> CometCollectSet, + classOf[Corr] -> CometCorr, + classOf[Count] -> CometCount, + classOf[CovPopulation] -> CometCovPopulation, + classOf[CovSample] -> CometCovSample, + classOf[First] -> CometFirst, + classOf[Last] -> CometLast, + classOf[Max] -> CometMax, + classOf[Min] -> CometMin, + classOf[Percentile] -> CometPercentile, + classOf[StddevPop] -> CometStddevPop, + classOf[StddevSamp] -> CometStddevSamp, + classOf[Sum] -> CometSum, + classOf[VariancePop] -> CometVariancePop, + classOf[VarianceSamp] -> CometVarianceSamp) + base ++ sparkVersionSpecificAggregates + } /** * Returns true if all aggregate expressions in the list have intermediate buffer formats that diff --git a/spark/src/main/spark-3.4/org/apache/comet/shims/CometExprShim.scala b/spark/src/main/spark-3.4/org/apache/comet/shims/CometExprShim.scala index 1ad9ec75bf0..43837f266f7 100644 --- a/spark/src/main/spark-3.4/org/apache/comet/shims/CometExprShim.scala +++ b/spark/src/main/spark-3.4/org/apache/comet/shims/CometExprShim.scala @@ -23,7 +23,7 @@ import org.apache.spark.sql.catalyst.expressions._ import org.apache.spark.sql.catalyst.expressions.aggregate.Sum import org.apache.comet.expressions.CometEvalMode -import org.apache.comet.serde.{CometExpressionSerde, CometStringDecode} +import org.apache.comet.serde.{CometAggregateExpressionSerde, CometExpressionSerde, CometStringDecode} import org.apache.comet.serde.ExprOuterClass.{BinaryOutputStyle, Expr} /** @@ -44,6 +44,8 @@ trait CometExprShim { Map.empty def sparkVersionSpecificMapExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = Map.empty + def sparkVersionSpecificAggregates: Map[Class[_], CometAggregateExpressionSerde[_]] = + Map.empty def sparkVersionSpecificExprToProtoInternal( expr: Expression, diff --git a/spark/src/main/spark-3.5/org/apache/comet/shims/CometExprShim.scala b/spark/src/main/spark-3.5/org/apache/comet/shims/CometExprShim.scala index 0be1185f592..42290ab5654 100644 --- a/spark/src/main/spark-3.5/org/apache/comet/shims/CometExprShim.scala +++ b/spark/src/main/spark-3.5/org/apache/comet/shims/CometExprShim.scala @@ -23,7 +23,7 @@ import org.apache.spark.sql.catalyst.expressions._ import org.apache.spark.sql.catalyst.expressions.aggregate.Sum import org.apache.comet.expressions.CometEvalMode -import org.apache.comet.serde.{CometExpressionSerde, CometStringDecode, CometToPrettyString, CometWidthBucket} +import org.apache.comet.serde.{CometAggregateExpressionSerde, CometExpressionSerde, CometStringDecode, CometToPrettyString, CometWidthBucket} import org.apache.comet.serde.ExprOuterClass.{BinaryOutputStyle, Expr} /** @@ -44,6 +44,8 @@ trait CometExprShim { Map(classOf[ToPrettyString] -> CometToPrettyString) def sparkVersionSpecificMapExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = Map.empty + def sparkVersionSpecificAggregates: Map[Class[_], CometAggregateExpressionSerde[_]] = + Map.empty def sparkVersionSpecificExprToProtoInternal( expr: Expression, diff --git a/spark/src/main/spark-4.x/org/apache/comet/shims/Spark4xCometExprShim.scala b/spark/src/main/spark-4.x/org/apache/comet/shims/Spark4xCometExprShim.scala index 7efd17f68a8..9f9290920b0 100644 --- a/spark/src/main/spark-4.x/org/apache/comet/shims/Spark4xCometExprShim.scala +++ b/spark/src/main/spark-4.x/org/apache/comet/shims/Spark4xCometExprShim.scala @@ -27,7 +27,7 @@ import org.apache.spark.sql.types.ArrayType import org.apache.comet.CometExplainInfo import org.apache.comet.expressions.CometEvalMode -import org.apache.comet.serde.{CometExpressionSerde, CometMapSort, CometToPrettyString, CometWidthBucket} +import org.apache.comet.serde.{CometAggregateExpressionSerde, CometExpressionSerde, CometMapSort, CometToPrettyString, CometWidthBucket} import org.apache.comet.serde.ExprOuterClass.Expr import org.apache.comet.serde.QueryPlanSerde.{exprToProtoInternal, optExprWithFallbackReason, scalarFunctionExprToProtoWithReturnType} @@ -49,6 +49,8 @@ trait Spark4xCometExprShim extends CometExprShim4x { Map(classOf[ToPrettyString] -> CometToPrettyString) def sparkVersionSpecificMapExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = Map(classOf[MapSort] -> CometMapSort) + def sparkVersionSpecificAggregates: Map[Class[_], CometAggregateExpressionSerde[_]] = + Map.empty def sparkVersionSpecificExprToProtoInternal( expr: Expression, From 12ccf610e65cf2c49326faad233d4a3f885b62b2 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Thu, 2 Jul 2026 17:49:23 -0600 Subject: [PATCH 03/17] feat: add HllSketchAgg and HllUnionAgg proto messages Adds the two protobuf messages needed to serialize Spark's hll_sketch_agg and hll_union_agg aggregate functions, wired into the AggExpr oneof as field numbers 19 and 20. --- native/proto/src/proto/expr.proto | 16 ++++++++++++++++ 1 file changed, 16 insertions(+) diff --git a/native/proto/src/proto/expr.proto b/native/proto/src/proto/expr.proto index 5b2a6ce9eed..978c0e0186b 100644 --- a/native/proto/src/proto/expr.proto +++ b/native/proto/src/proto/expr.proto @@ -145,6 +145,8 @@ message AggExpr { BloomFilterAgg bloomFilterAgg = 16; CollectSet collectSet = 17; Percentile percentile = 18; + HllSketchAgg hllSketchAgg = 19; + HllUnionAgg hllUnionAgg = 20; } // Optional filter expression for SQL FILTER (WHERE ...) clause. @@ -271,6 +273,20 @@ enum BloomFilterVersion { BLOOM_FILTER_VERSION_V2 = 2; } +message HllSketchAgg { + // Child value expression (integral, string, or binary). + Expr child = 1; + // DataSketches lgConfigK (log2 of the number of buckets), Spark default 12. + int32 lg_config_k = 2; +} + +message HllUnionAgg { + // Child sketch expression (Binary column of serialized HLL sketches). + Expr child = 1; + // When false, Spark errors if input sketches have differing lgConfigK. + bool allow_different_lg_config_k = 2; +} + message CollectSet { Expr child = 1; DataType datatype = 2; From 5dc29b0045f2912405183f59f9895b1f8ca03b83 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Thu, 2 Jul 2026 18:04:45 -0600 Subject: [PATCH 04/17] feat: add native hll_sketch_agg accumulator and planner arm Wire up HllSketchAgg as an AggregateUDFImpl backed by SparkHllSketch, accepting Int8/16/32/64, Utf8, and Binary inputs and returning a serialized HLL sketch as Binary. Null groups evaluate to NULL, matching Spark's HllSketchAgg. Wires the new AggExprStruct::HllSketchAgg arm into the native planner's create_agg_expr. Also add a placeholder HllUnionAgg planner arm returning a clear "not yet supported" error, since that oneof variant already exists in the proto but its native accumulator lands in a follow-on task. --- native/core/src/execution/planner.rs | 13 +- .../src/agg_funcs/hll_sketch_agg.rs | 196 ++++++++++++++++++ native/spark-expr/src/agg_funcs/mod.rs | 2 + 3 files changed, 209 insertions(+), 2 deletions(-) create mode 100644 native/spark-expr/src/agg_funcs/hll_sketch_agg.rs diff --git a/native/core/src/execution/planner.rs b/native/core/src/execution/planner.rs index 25162332fd6..c4bbdca2237 100644 --- a/native/core/src/execution/planner.rs +++ b/native/core/src/execution/planner.rs @@ -130,8 +130,8 @@ use datafusion_comet_proto::{ use datafusion_comet_spark_expr::{ jvm_udf::JvmScalarUdfExpr, ArrayInsert, Avg, AvgDecimal, Cast, CheckOverflow, Correlation, Covariance, CreateNamedStruct, DecimalRescaleCheckOverflow, GetArrayStructFields, - GetStructField, IfExpr, ListExtract, NormalizeNaNAndZero, SparkCastOptions, Stddev, SumDecimal, - ToJson, UnboundColumn, Variance, WideDecimalBinaryExpr, WideDecimalOp, + GetStructField, HllSketchAgg, IfExpr, ListExtract, NormalizeNaNAndZero, SparkCastOptions, + Stddev, SumDecimal, ToJson, UnboundColumn, Variance, WideDecimalBinaryExpr, WideDecimalOp, }; use itertools::Itertools; use jni::objects::{Global, JObject}; @@ -2648,11 +2648,20 @@ impl PhysicalPlanner { )); Self::create_aggr_func_expr("bloom_filter_agg", schema, vec![child], func) } + AggExprStruct::HllSketchAgg(expr) => { + let child = self.create_expr(expr.child.as_ref().unwrap(), Arc::clone(&schema))?; + let func = AggregateUDF::new_from_impl(HllSketchAgg::new(expr.lg_config_k)); + Self::create_aggr_func_expr("hll_sketch_agg", schema, vec![child], func) + } AggExprStruct::CollectSet(expr) => { let child = self.create_expr(expr.child.as_ref().unwrap(), Arc::clone(&schema))?; let func = AggregateUDF::new_from_impl(SparkCollectSet::new()); Self::create_aggr_func_expr("collect_set", schema, vec![child], func) } + // hll_union_agg's native accumulator + planner arm is wired up in a follow-on task. + AggExprStruct::HllUnionAgg(_) => Err(ExecutionError::GeneralError( + "hll_union_agg is not yet supported".to_string(), + )), } } diff --git a/native/spark-expr/src/agg_funcs/hll_sketch_agg.rs b/native/spark-expr/src/agg_funcs/hll_sketch_agg.rs new file mode 100644 index 00000000000..e962f984f14 --- /dev/null +++ b/native/spark-expr/src/agg_funcs/hll_sketch_agg.rs @@ -0,0 +1,196 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +use crate::agg_funcs::hll_sketch::SparkHllSketch; +use arrow::array::Array; +use arrow::array::ArrayRef; +use arrow::array::BinaryArray; +use arrow::datatypes::{DataType, Field, FieldRef}; +use datafusion::common::{downcast_value, ScalarValue}; +use datafusion::error::{DataFusionError, Result}; +use datafusion::logical_expr::function::{AccumulatorArgs, StateFieldsArgs}; +use datafusion::logical_expr::{AggregateUDFImpl, Signature, Volatility}; +use datafusion::physical_plan::Accumulator; +use std::sync::Arc; + +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +pub struct HllSketchAgg { + signature: Signature, + lg_config_k: i32, +} + +impl HllSketchAgg { + pub fn new(lg_config_k: i32) -> Self { + Self { + signature: Signature::uniform( + 1, + vec![ + DataType::Int8, + DataType::Int16, + DataType::Int32, + DataType::Int64, + DataType::Utf8, + DataType::Binary, + ], + Volatility::Immutable, + ), + lg_config_k, + } + } +} + +impl AggregateUDFImpl for HllSketchAgg { + fn name(&self) -> &str { + "hll_sketch_agg" + } + fn signature(&self) -> &Signature { + &self.signature + } + fn return_type(&self, _: &[DataType]) -> Result { + Ok(DataType::Binary) + } + fn accumulator(&self, _: AccumulatorArgs) -> Result> { + Ok(Box::new(HllSketchAccumulator::new(self.lg_config_k as u8))) + } + fn state_fields(&self, _: StateFieldsArgs) -> Result> { + Ok(vec![Arc::new(Field::new("sketch", DataType::Binary, true))]) + } + fn groups_accumulator_supported(&self, _: AccumulatorArgs) -> bool { + false + } +} + +#[derive(Debug)] +pub struct HllSketchAccumulator { + sketch: SparkHllSketch, + saw_input: bool, +} + +impl HllSketchAccumulator { + pub fn new(lg_config_k: u8) -> Self { + Self { + sketch: SparkHllSketch::new(lg_config_k), + saw_input: false, + } + } +} + +impl Accumulator for HllSketchAccumulator { + fn update_batch(&mut self, values: &[ArrayRef]) -> Result<()> { + if values.is_empty() { + return Ok(()); + } + let arr = &values[0]; + (0..arr.len()).try_for_each(|i| { + match ScalarValue::try_from_array(arr, i)? { + ScalarValue::Int8(Some(v)) => { + self.sketch.update_i64(v as i64); + self.saw_input = true; + } + ScalarValue::Int16(Some(v)) => { + self.sketch.update_i64(v as i64); + self.saw_input = true; + } + ScalarValue::Int32(Some(v)) => { + self.sketch.update_i64(v as i64); + self.saw_input = true; + } + ScalarValue::Int64(Some(v)) => { + self.sketch.update_i64(v); + self.saw_input = true; + } + ScalarValue::Utf8(Some(v)) => { + self.sketch.update_bytes(v.as_bytes()); + self.saw_input = true; + } + ScalarValue::Binary(Some(v)) => { + self.sketch.update_bytes(&v); + self.saw_input = true; + } + // Spark's HllSketchAgg ignores null inputs. + ScalarValue::Int8(None) + | ScalarValue::Int16(None) + | ScalarValue::Int32(None) + | ScalarValue::Int64(None) + | ScalarValue::Utf8(None) + | ScalarValue::Binary(None) => {} + other => { + return Err(DataFusionError::Internal(format!( + "hll_sketch_agg received an unsupported input type: {other:?}" + ))) + } + } + Ok(()) + }) + } + + fn evaluate(&mut self) -> Result { + // Spark returns a non-null sketch even for empty groups only when it saw input; + // an empty group yields NULL. + if !self.saw_input { + return Ok(ScalarValue::Binary(None)); + } + Ok(ScalarValue::Binary(Some(self.sketch.to_sketch_bytes()))) + } + + fn size(&self) -> usize { + std::mem::size_of_val(self) + } + + fn state(&mut self) -> Result> { + if !self.saw_input { + return Ok(vec![ScalarValue::Binary(None)]); + } + Ok(vec![ScalarValue::Binary(Some( + self.sketch.to_sketch_bytes(), + ))]) + } + + fn merge_batch(&mut self, states: &[ArrayRef]) -> Result<()> { + let arr = downcast_value!(states[0], BinaryArray); + for i in 0..arr.len() { + if arr.is_null(i) { + continue; + } + let peer = SparkHllSketch::from_bytes(arr.value(i))?; + // Merge peer into self by unioning; reuse SparkHllUnion via sketch merge. + self.sketch.merge_sketch(&peer); + self.saw_input = true; + } + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use arrow::array::Int64Array; + use datafusion::physical_plan::Accumulator; + use std::sync::Arc; + + #[test] + fn accumulates_and_estimates() { + let mut acc = HllSketchAccumulator::new(12); + let arr = Arc::new(Int64Array::from((0..1000i64).collect::>())); + acc.update_batch(&[arr]).unwrap(); + let ScalarValue::Binary(Some(bytes)) = acc.evaluate().unwrap() else { + panic!("expected binary") + }; + let est = crate::agg_funcs::estimate_from_bytes(&bytes).unwrap(); + assert!((est - 1000).abs() <= 30, "estimate {est}"); + } +} diff --git a/native/spark-expr/src/agg_funcs/mod.rs b/native/spark-expr/src/agg_funcs/mod.rs index 57af25b35e4..3211de7db2a 100644 --- a/native/spark-expr/src/agg_funcs/mod.rs +++ b/native/spark-expr/src/agg_funcs/mod.rs @@ -20,6 +20,7 @@ mod avg_decimal; mod correlation; mod covariance; mod hll_sketch; +mod hll_sketch_agg; mod stddev; mod sum_decimal; mod sum_int; @@ -31,6 +32,7 @@ pub use avg_decimal::AvgDecimal; pub use correlation::Correlation; pub use covariance::Covariance; pub use hll_sketch::{estimate_from_bytes, SparkHllSketch, SparkHllUnion}; +pub use hll_sketch_agg::HllSketchAgg; pub use stddev::Stddev; pub use sum_decimal::SumDecimal; pub use sum_int::SumInteger; From a52944d18e23325e742e59af6f5bf644badd2207 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Thu, 2 Jul 2026 18:17:12 -0600 Subject: [PATCH 05/17] feat: add hll_sketch_agg Scala serde for Spark 4.x --- .../comet/serde/CometHllSketchAgg.scala | 91 +++++++++++++++++++ .../comet/shims/Spark4xCometExprShim.scala | 5 +- 2 files changed, 94 insertions(+), 2 deletions(-) create mode 100644 spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchAgg.scala diff --git a/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchAgg.scala b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchAgg.scala new file mode 100644 index 00000000000..7d5c907df63 --- /dev/null +++ b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchAgg.scala @@ -0,0 +1,91 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.comet.serde + +import org.apache.spark.sql.catalyst.expressions.Attribute +import org.apache.spark.sql.catalyst.expressions.aggregate.{AggregateExpression, HllSketchAgg} +import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.types.{BinaryType, IntegerType, LongType, StringType} + +import org.apache.comet.CometSparkSessionExtensions.withFallbackReason +import org.apache.comet.serde.QueryPlanSerde.exprToProto + +// In Spark 4.0, HllSketchAgg's fields are `left` (the child value expression) and `right` (the +// lgConfigK expression), not `child`/`lgConfigKExpression`. Accepted input types are Integer, +// Long, String, and Binary (Byte/Short are not accepted by Spark's HllSketchAgg). +object CometHllSketchAgg extends CometAggregateExpressionSerde[HllSketchAgg] { + + // DataSketches valid lgConfigK range; outside this Spark itself throws, so we + // fall back rather than forward an out-of-range value to native. + private val MinLgConfigK = 4 + private val MaxLgConfigK = 21 + + private val nonLiteralLgConfigKReason = + "The lgConfigK argument must be a foldable literal." + private val inputTypeReason = + "Only int, long, string, and binary input types are supported." + + override def getUnsupportedReasons(): Seq[String] = + Seq(nonLiteralLgConfigKReason, inputTypeReason) + + override def getSupportLevel(expr: HllSketchAgg): SupportLevel = { + if (!expr.right.foldable) { + return Unsupported(Some(nonLiteralLgConfigKReason)) + } + val lgConfigK = expr.right.eval() match { + case i: Int => i + case l: Long => l.toInt + case _ => return Unsupported(Some(nonLiteralLgConfigKReason)) + } + if (lgConfigK < MinLgConfigK || lgConfigK > MaxLgConfigK) { + return Unsupported(Some(s"lgConfigK must be in [$MinLgConfigK, $MaxLgConfigK]")) + } + expr.left.dataType match { + case IntegerType | LongType | StringType | BinaryType => + Compatible(None) + case _ => Unsupported(Some(inputTypeReason)) + } + } + + override def convert( + aggExpr: AggregateExpression, + expr: HllSketchAgg, + inputs: Seq[Attribute], + binding: Boolean, + conf: SQLConf): Option[ExprOuterClass.AggExpr] = { + val childExpr = exprToProto(expr.left, inputs, binding) + val lgConfigK = expr.right.eval() match { + case i: Int => i + case l: Long => l.toInt + case other => + withFallbackReason(aggExpr, s"Unsupported lgConfigK literal: $other", expr.left) + return None + } + if (childExpr.isDefined) { + val builder = ExprOuterClass.HllSketchAgg.newBuilder() + builder.setChild(childExpr.get) + builder.setLgConfigK(lgConfigK) + Some(ExprOuterClass.AggExpr.newBuilder().setHllSketchAgg(builder).build()) + } else { + withFallbackReason(aggExpr, expr.left) + None + } + } +} diff --git a/spark/src/main/spark-4.x/org/apache/comet/shims/Spark4xCometExprShim.scala b/spark/src/main/spark-4.x/org/apache/comet/shims/Spark4xCometExprShim.scala index 9f9290920b0..1f11b04c3a7 100644 --- a/spark/src/main/spark-4.x/org/apache/comet/shims/Spark4xCometExprShim.scala +++ b/spark/src/main/spark-4.x/org/apache/comet/shims/Spark4xCometExprShim.scala @@ -20,6 +20,7 @@ package org.apache.comet.shims import org.apache.spark.sql.catalyst.expressions._ +import org.apache.spark.sql.catalyst.expressions.aggregate.HllSketchAgg import org.apache.spark.sql.catalyst.expressions.json.{JsonExpressionUtils, StructsToJsonEvaluator} import org.apache.spark.sql.catalyst.expressions.objects.{Invoke, StaticInvoke} import org.apache.spark.sql.catalyst.expressions.url.ParseUrlEvaluator @@ -27,7 +28,7 @@ import org.apache.spark.sql.types.ArrayType import org.apache.comet.CometExplainInfo import org.apache.comet.expressions.CometEvalMode -import org.apache.comet.serde.{CometAggregateExpressionSerde, CometExpressionSerde, CometMapSort, CometToPrettyString, CometWidthBucket} +import org.apache.comet.serde.{CometAggregateExpressionSerde, CometExpressionSerde, CometHllSketchAgg, CometMapSort, CometToPrettyString, CometWidthBucket} import org.apache.comet.serde.ExprOuterClass.Expr import org.apache.comet.serde.QueryPlanSerde.{exprToProtoInternal, optExprWithFallbackReason, scalarFunctionExprToProtoWithReturnType} @@ -50,7 +51,7 @@ trait Spark4xCometExprShim extends CometExprShim4x { def sparkVersionSpecificMapExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = Map(classOf[MapSort] -> CometMapSort) def sparkVersionSpecificAggregates: Map[Class[_], CometAggregateExpressionSerde[_]] = - Map.empty + Map(classOf[HllSketchAgg] -> CometHllSketchAgg) def sparkVersionSpecificExprToProtoInternal( expr: Expression, From f31cbf82c594ffa128548b20c3037f19c8bc8ba3 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Thu, 2 Jul 2026 18:25:19 -0600 Subject: [PATCH 06/17] feat: add hll_sketch_estimate scalar function --- native/spark-expr/src/comet_scalar_funcs.rs | 5 ++ native/spark-expr/src/hll_scalar.rs | 60 +++++++++++++++++++ native/spark-expr/src/lib.rs | 2 + .../comet/serde/CometHllSketchEstimate.scala | 40 +++++++++++++ .../comet/shims/Spark4xCometExprShim.scala | 6 +- 5 files changed, 111 insertions(+), 2 deletions(-) create mode 100644 native/spark-expr/src/hll_scalar.rs create mode 100644 spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchEstimate.scala diff --git a/native/spark-expr/src/comet_scalar_funcs.rs b/native/spark-expr/src/comet_scalar_funcs.rs index 42ee72c82af..d52212873dd 100644 --- a/native/spark-expr/src/comet_scalar_funcs.rs +++ b/native/spark-expr/src/comet_scalar_funcs.rs @@ -16,6 +16,7 @@ // under the License. use crate::hash_funcs::*; +use crate::hll_scalar::spark_hll_sketch_estimate; use crate::json_funcs::JsonArrayLength; use crate::map_funcs::spark_map_sort; use crate::math_funcs::abs::abs; @@ -221,6 +222,10 @@ pub fn create_comet_physical_fun_with_eval_mode( let func = Arc::new(spark_map_sort); make_comet_scalar_udf!("spark_map_sort", func, without data_type) } + "hll_sketch_estimate" => { + let func = Arc::new(|args: &[ColumnarValue]| spark_hll_sketch_estimate(args)); + make_comet_scalar_udf!("hll_sketch_estimate", func, without data_type) + } "to_time" => { make_comet_scalar_udf!("to_time", spark_to_time, without data_type, fail_on_error) } diff --git a/native/spark-expr/src/hll_scalar.rs b/native/spark-expr/src/hll_scalar.rs new file mode 100644 index 00000000000..755b5df5fa8 --- /dev/null +++ b/native/spark-expr/src/hll_scalar.rs @@ -0,0 +1,60 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +use crate::agg_funcs::estimate_from_bytes; +use arrow::array::{Array, BinaryArray, Int64Array}; +use datafusion::common::Result; +use datafusion::physical_plan::ColumnarValue; +use std::sync::Arc; + +/// Spark hll_sketch_estimate: Binary sketch -> Long distinct-count estimate. +pub fn spark_hll_sketch_estimate(args: &[ColumnarValue]) -> Result { + let arrays = ColumnarValue::values_to_arrays(args)?; + let input = arrays[0].as_any().downcast_ref::().unwrap(); + let mut out = Int64Array::builder(input.len()); + for i in 0..input.len() { + if input.is_null(i) { + out.append_null(); + } else { + out.append_value(estimate_from_bytes(input.value(i))?); + } + } + Ok(ColumnarValue::Array(Arc::new(out.finish()))) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::agg_funcs::SparkHllSketch; + + #[test] + fn estimates_from_sketch_column() { + let mut s = SparkHllSketch::new(12); + for i in 0..1000i64 { + s.update_i64(i); + } + let arr = Arc::new(BinaryArray::from(vec![Some( + s.to_sketch_bytes().as_slice(), + )])); + let out = spark_hll_sketch_estimate(&[ColumnarValue::Array(arr)]).unwrap(); + let ColumnarValue::Array(a) = out else { + panic!() + }; + let est = a.as_any().downcast_ref::().unwrap().value(0); + assert!((est - 1000).abs() <= 30, "estimate {est}"); + } +} diff --git a/native/spark-expr/src/lib.rs b/native/spark-expr/src/lib.rs index 174a4ada9a0..baca331dc9d 100644 --- a/native/spark-expr/src/lib.rs +++ b/native/spark-expr/src/lib.rs @@ -59,6 +59,8 @@ pub mod jvm_udf; mod conditional_funcs; mod conversion_funcs; +mod hll_scalar; +pub use hll_scalar::spark_hll_sketch_estimate; mod map_funcs; pub use map_funcs::spark_map_sort; mod math_funcs; diff --git a/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchEstimate.scala b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchEstimate.scala new file mode 100644 index 00000000000..1f032abdc38 --- /dev/null +++ b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchEstimate.scala @@ -0,0 +1,40 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.comet.serde + +import org.apache.spark.sql.catalyst.expressions.{Attribute, HllSketchEstimate} +import org.apache.spark.sql.types.LongType + +import org.apache.comet.serde.QueryPlanSerde.{exprToProtoInternal, optExprWithFallbackReason, scalarFunctionExprToProtoWithReturnType} + +object CometHllSketchEstimate extends CometExpressionSerde[HllSketchEstimate] { + override def convert( + expr: HllSketchEstimate, + inputs: Seq[Attribute], + binding: Boolean): Option[ExprOuterClass.Expr] = { + val childExpr = exprToProtoInternal(expr.child, inputs, binding) + val estimateExpr = scalarFunctionExprToProtoWithReturnType( + "hll_sketch_estimate", + LongType, + failOnError = false, + childExpr) + optExprWithFallbackReason(estimateExpr, expr, expr.child) + } +} diff --git a/spark/src/main/spark-4.x/org/apache/comet/shims/Spark4xCometExprShim.scala b/spark/src/main/spark-4.x/org/apache/comet/shims/Spark4xCometExprShim.scala index 1f11b04c3a7..bf8c86b1939 100644 --- a/spark/src/main/spark-4.x/org/apache/comet/shims/Spark4xCometExprShim.scala +++ b/spark/src/main/spark-4.x/org/apache/comet/shims/Spark4xCometExprShim.scala @@ -28,7 +28,7 @@ import org.apache.spark.sql.types.ArrayType import org.apache.comet.CometExplainInfo import org.apache.comet.expressions.CometEvalMode -import org.apache.comet.serde.{CometAggregateExpressionSerde, CometExpressionSerde, CometHllSketchAgg, CometMapSort, CometToPrettyString, CometWidthBucket} +import org.apache.comet.serde.{CometAggregateExpressionSerde, CometExpressionSerde, CometHllSketchAgg, CometHllSketchEstimate, CometMapSort, CometToPrettyString, CometWidthBucket} import org.apache.comet.serde.ExprOuterClass.Expr import org.apache.comet.serde.QueryPlanSerde.{exprToProtoInternal, optExprWithFallbackReason, scalarFunctionExprToProtoWithReturnType} @@ -47,7 +47,9 @@ trait Spark4xCometExprShim extends CometExprShim4x { def sparkVersionSpecificMathExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = Map(classOf[WidthBucket] -> CometWidthBucket) def sparkVersionSpecificMiscExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = - Map(classOf[ToPrettyString] -> CometToPrettyString) + Map( + classOf[ToPrettyString] -> CometToPrettyString, + classOf[HllSketchEstimate] -> CometHllSketchEstimate) def sparkVersionSpecificMapExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = Map(classOf[MapSort] -> CometMapSort) def sparkVersionSpecificAggregates: Map[Class[_], CometAggregateExpressionSerde[_]] = From 4f0f289939f49c1cecbda25de8b89807f595d81b Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Thu, 2 Jul 2026 18:56:03 -0600 Subject: [PATCH 07/17] feat: mark HLL expressions incompatible and add opt-in end-to-end test [skip ci] --- .../comet/serde/CometHllSketchAgg.scala | 7 +++++- .../comet/serde/CometHllSketchEstimate.scala | 8 ++++++ .../apache/comet/CometExpressionSuite.scala | 25 +++++++++++++++++++ 3 files changed, 39 insertions(+), 1 deletion(-) diff --git a/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchAgg.scala b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchAgg.scala index 7d5c907df63..5141e0c9526 100644 --- a/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchAgg.scala +++ b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchAgg.scala @@ -41,10 +41,15 @@ object CometHllSketchAgg extends CometAggregateExpressionSerde[HllSketchAgg] { "The lgConfigK argument must be a foldable literal." private val inputTypeReason = "Only int, long, string, and binary input types are supported." + private val incompatReason = + "Comet uses a Rust DataSketches port; HLL sketch bytes and estimates may differ " + + "slightly from Spark." override def getUnsupportedReasons(): Seq[String] = Seq(nonLiteralLgConfigKReason, inputTypeReason) + override def getIncompatibleReasons(): Seq[String] = Seq(incompatReason) + override def getSupportLevel(expr: HllSketchAgg): SupportLevel = { if (!expr.right.foldable) { return Unsupported(Some(nonLiteralLgConfigKReason)) @@ -59,7 +64,7 @@ object CometHllSketchAgg extends CometAggregateExpressionSerde[HllSketchAgg] { } expr.left.dataType match { case IntegerType | LongType | StringType | BinaryType => - Compatible(None) + Incompatible(Some(incompatReason)) case _ => Unsupported(Some(inputTypeReason)) } } diff --git a/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchEstimate.scala b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchEstimate.scala index 1f032abdc38..454b53eed80 100644 --- a/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchEstimate.scala +++ b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchEstimate.scala @@ -25,6 +25,14 @@ import org.apache.spark.sql.types.LongType import org.apache.comet.serde.QueryPlanSerde.{exprToProtoInternal, optExprWithFallbackReason, scalarFunctionExprToProtoWithReturnType} object CometHllSketchEstimate extends CometExpressionSerde[HllSketchEstimate] { + private val incompatReason = + "Comet uses a Rust DataSketches port; HLL estimates may differ slightly from Spark." + + override def getIncompatibleReasons(): Seq[String] = Seq(incompatReason) + + override def getSupportLevel(expr: HllSketchEstimate): SupportLevel = + Incompatible(Some(incompatReason)) + override def convert( expr: HllSketchEstimate, inputs: Seq[Attribute], diff --git a/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala b/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala index 390a7c4908a..8faecb24a52 100644 --- a/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala @@ -3380,4 +3380,29 @@ class CometExpressionSuite extends CometTestBase with AdaptiveSparkPlanHelper { } } + test("hll_sketch_agg and hll_sketch_estimate (incompatible, opt-in)") { + assume(isSpark40Plus) + // HLL is approximate: Comet's Rust DataSketches estimator differs slightly from + // Spark's after a merge, so these functions are Incompatible. Opt in, assert the + // query runs natively (no fallback), and that the estimate is within HLL error of + // the TRUE distinct count (700). Do NOT compare bit-exactly to Spark. + withSQLConf( + "spark.comet.expression.HllSketchAgg.allowIncompatible" -> "true", + "spark.comet.expression.HllSketchEstimate.allowIncompatible" -> "true") { + withParquetTable((0 until 1000).map(i => (i % 700, i)), "tbl") { + def checkEstimate(query: String): Unit = { + val df = sql(query) + checkCometOperators(stripAQEPlan(df.queryExecution.executedPlan)) + val est = df.collect().head.getLong(0) + assert( + math.abs(est - 700).toDouble / 700 <= 0.05, + s"estimate $est not within 5% of the true distinct count 700 for: $query") + } + checkEstimate("SELECT hll_sketch_estimate(hll_sketch_agg(_1)) FROM tbl") + checkEstimate("SELECT hll_sketch_estimate(hll_sketch_agg(_1, 14)) FROM tbl") + checkEstimate("SELECT hll_sketch_estimate(hll_sketch_agg(cast(_1 as string))) FROM tbl") + } + } + } + } From 342bd6e18def1c80c40bef5f84ab2528fa560549 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Thu, 2 Jul 2026 19:05:00 -0600 Subject: [PATCH 08/17] feat: add native and serde for hll_union_agg [skip ci] --- native/core/src/execution/planner.rs | 15 +- .../spark-expr/src/agg_funcs/hll_union_agg.rs | 177 ++++++++++++++++++ native/spark-expr/src/agg_funcs/mod.rs | 2 + .../apache/comet/serde/CometHllUnionAgg.scala | 74 ++++++++ .../comet/shims/Spark4xCometExprShim.scala | 6 +- 5 files changed, 265 insertions(+), 9 deletions(-) create mode 100644 native/spark-expr/src/agg_funcs/hll_union_agg.rs create mode 100644 spark/src/main/spark-4.x/org/apache/comet/serde/CometHllUnionAgg.scala diff --git a/native/core/src/execution/planner.rs b/native/core/src/execution/planner.rs index c4bbdca2237..e63e09c8eea 100644 --- a/native/core/src/execution/planner.rs +++ b/native/core/src/execution/planner.rs @@ -130,8 +130,9 @@ use datafusion_comet_proto::{ use datafusion_comet_spark_expr::{ jvm_udf::JvmScalarUdfExpr, ArrayInsert, Avg, AvgDecimal, Cast, CheckOverflow, Correlation, Covariance, CreateNamedStruct, DecimalRescaleCheckOverflow, GetArrayStructFields, - GetStructField, HllSketchAgg, IfExpr, ListExtract, NormalizeNaNAndZero, SparkCastOptions, - Stddev, SumDecimal, ToJson, UnboundColumn, Variance, WideDecimalBinaryExpr, WideDecimalOp, + GetStructField, HllSketchAgg, HllUnionAgg, IfExpr, ListExtract, NormalizeNaNAndZero, + SparkCastOptions, Stddev, SumDecimal, ToJson, UnboundColumn, Variance, WideDecimalBinaryExpr, + WideDecimalOp, }; use itertools::Itertools; use jni::objects::{Global, JObject}; @@ -2658,10 +2659,12 @@ impl PhysicalPlanner { let func = AggregateUDF::new_from_impl(SparkCollectSet::new()); Self::create_aggr_func_expr("collect_set", schema, vec![child], func) } - // hll_union_agg's native accumulator + planner arm is wired up in a follow-on task. - AggExprStruct::HllUnionAgg(_) => Err(ExecutionError::GeneralError( - "hll_union_agg is not yet supported".to_string(), - )), + AggExprStruct::HllUnionAgg(expr) => { + let child = self.create_expr(expr.child.as_ref().unwrap(), Arc::clone(&schema))?; + let func = + AggregateUDF::new_from_impl(HllUnionAgg::new(expr.allow_different_lg_config_k)); + Self::create_aggr_func_expr("hll_union_agg", schema, vec![child], func) + } } } diff --git a/native/spark-expr/src/agg_funcs/hll_union_agg.rs b/native/spark-expr/src/agg_funcs/hll_union_agg.rs new file mode 100644 index 00000000000..9d3658cef78 --- /dev/null +++ b/native/spark-expr/src/agg_funcs/hll_union_agg.rs @@ -0,0 +1,177 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +use crate::agg_funcs::hll_sketch::{SparkHllSketch, SparkHllUnion}; +use arrow::array::{Array, ArrayRef, BinaryArray}; +use arrow::datatypes::{DataType, Field, FieldRef}; +use datafusion::common::{downcast_value, ScalarValue}; +use datafusion::error::{DataFusionError, Result}; +use datafusion::logical_expr::function::{AccumulatorArgs, StateFieldsArgs}; +use datafusion::logical_expr::{AggregateUDFImpl, Signature, Volatility}; +use datafusion::physical_plan::Accumulator; +use std::sync::Arc; + +// NOTE: matches bloom_filter_agg.rs for DataFusion 54.0.0 - no `as_any` method on +// AggregateUDFImpl, and PartialEq/Eq/Hash are required (DynEq/DynHash). +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +pub struct HllUnionAgg { + signature: Signature, + allow_different_lg_config_k: bool, +} + +impl HllUnionAgg { + pub fn new(allow_different_lg_config_k: bool) -> Self { + Self { + signature: Signature::uniform(1, vec![DataType::Binary], Volatility::Immutable), + allow_different_lg_config_k, + } + } +} + +impl AggregateUDFImpl for HllUnionAgg { + fn name(&self) -> &str { + "hll_union_agg" + } + fn signature(&self) -> &Signature { + &self.signature + } + fn return_type(&self, _: &[DataType]) -> Result { + Ok(DataType::Binary) + } + fn accumulator(&self, _: AccumulatorArgs) -> Result> { + Ok(Box::new(HllUnionAccumulator::new( + self.allow_different_lg_config_k, + ))) + } + fn state_fields(&self, _: StateFieldsArgs) -> Result> { + Ok(vec![Arc::new(Field::new("sketch", DataType::Binary, true))]) + } + fn groups_accumulator_supported(&self, _: AccumulatorArgs) -> bool { + false + } +} + +#[derive(Debug)] +pub struct HllUnionAccumulator { + // Spark's HllUnionAgg defers creating the Union until the first sketch is seen, + // then builds `new Union(sketch.getLgConfigK)` - so lgMaxK is NOT a fixed 12. + union: Option, + allow_different_lg_config_k: bool, + seen_lg_config_k: Option, +} + +impl HllUnionAccumulator { + pub fn new(allow_different_lg_config_k: bool) -> Self { + Self { + union: None, + allow_different_lg_config_k, + seen_lg_config_k: None, + } + } + + fn absorb(&mut self, bytes: &[u8]) -> Result<()> { + let sketch = SparkHllSketch::from_bytes(bytes)?; + let k = sketch.lg_config_k(); + match self.seen_lg_config_k { + None => { + // Lazily instantiate the union from the first sketch's lgConfigK. + self.seen_lg_config_k = Some(k); + self.union = Some(SparkHllUnion::new(k)); + } + Some(prev) if prev != k && !self.allow_different_lg_config_k => { + return Err(DataFusionError::Execution(format!( + "Sketches have different lgConfigK values: {prev} and {k}. \ + Set allowDifferentLgConfigK to true to enable unions of different lgConfigK." + ))); + } + _ => {} + } + self.union.as_mut().unwrap().merge(&sketch); + Ok(()) + } +} + +impl Accumulator for HllUnionAccumulator { + fn update_batch(&mut self, values: &[ArrayRef]) -> Result<()> { + if values.is_empty() { + return Ok(()); + } + let arr = downcast_value!(values[0], BinaryArray); + for i in 0..arr.len() { + if !arr.is_null(i) { + self.absorb(arr.value(i))?; + } + } + Ok(()) + } + fn evaluate(&mut self) -> Result { + match &self.union { + Some(u) => Ok(ScalarValue::Binary(Some(u.to_sketch_bytes()))), + None => Ok(ScalarValue::Binary(None)), + } + } + fn size(&self) -> usize { + std::mem::size_of_val(self) + } + fn state(&mut self) -> Result> { + match &self.union { + Some(u) => Ok(vec![ScalarValue::Binary(Some(u.to_sketch_bytes()))]), + None => Ok(vec![ScalarValue::Binary(None)]), + } + } + fn merge_batch(&mut self, states: &[ArrayRef]) -> Result<()> { + let arr = downcast_value!(states[0], BinaryArray); + for i in 0..arr.len() { + if !arr.is_null(i) { + self.absorb(arr.value(i))?; + } + } + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::agg_funcs::hll_sketch::SparkHllSketch; + use arrow::array::BinaryArray; + use datafusion::physical_plan::Accumulator; + use std::sync::Arc; + + #[test] + fn unions_sketch_column() { + let mut a = SparkHllSketch::new(12); + for i in 0..1000i64 { + a.update_i64(i); + } + let mut b = SparkHllSketch::new(12); + for i in 500..1500i64 { + b.update_i64(i); + } + let arr = Arc::new(BinaryArray::from(vec![ + Some(a.to_sketch_bytes().as_slice()), + Some(b.to_sketch_bytes().as_slice()), + ])); + let mut acc = HllUnionAccumulator::new(false); + acc.update_batch(&[arr]).unwrap(); + let ScalarValue::Binary(Some(bytes)) = acc.evaluate().unwrap() else { + panic!() + }; + let est = crate::agg_funcs::estimate_from_bytes(&bytes).unwrap(); + assert!((est - 1500).abs() <= 45, "union estimate {est}"); + } +} diff --git a/native/spark-expr/src/agg_funcs/mod.rs b/native/spark-expr/src/agg_funcs/mod.rs index 3211de7db2a..f74e88b8486 100644 --- a/native/spark-expr/src/agg_funcs/mod.rs +++ b/native/spark-expr/src/agg_funcs/mod.rs @@ -21,6 +21,7 @@ mod correlation; mod covariance; mod hll_sketch; mod hll_sketch_agg; +mod hll_union_agg; mod stddev; mod sum_decimal; mod sum_int; @@ -33,6 +34,7 @@ pub use correlation::Correlation; pub use covariance::Covariance; pub use hll_sketch::{estimate_from_bytes, SparkHllSketch, SparkHllUnion}; pub use hll_sketch_agg::HllSketchAgg; +pub use hll_union_agg::HllUnionAgg; pub use stddev::Stddev; pub use sum_decimal::SumDecimal; pub use sum_int::SumInteger; diff --git a/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllUnionAgg.scala b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllUnionAgg.scala new file mode 100644 index 00000000000..156f26e675d --- /dev/null +++ b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllUnionAgg.scala @@ -0,0 +1,74 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.comet.serde + +import org.apache.spark.sql.catalyst.expressions.Attribute +import org.apache.spark.sql.catalyst.expressions.aggregate.{AggregateExpression, HllUnionAgg} +import org.apache.spark.sql.internal.SQLConf + +import org.apache.comet.CometSparkSessionExtensions.withFallbackReason +import org.apache.comet.serde.QueryPlanSerde.exprToProto + +// IMPORTANT: In Spark 4.0, HllUnionAgg's fields are `left` (the child binary sketch +// expression) and `right` (the allowDifferentLgConfigK boolean expression) - NOT +// `child`/`allowDifferentLgConfigKExpression`. (Verified against Spark 4.0 source.) +object CometHllUnionAgg extends CometAggregateExpressionSerde[HllUnionAgg] { + + private val incompatReason = + "Comet uses a Rust DataSketches port; HLL sketch bytes and estimates may differ slightly " + + "from Spark." + + override def getIncompatibleReasons(): Seq[String] = Seq(incompatReason) + + override def getSupportLevel(expr: HllUnionAgg): SupportLevel = { + if (!expr.right.foldable) { + Unsupported(Some("The allowDifferentLgConfigK argument must be a foldable literal.")) + } else { + Incompatible(Some(incompatReason)) + } + } + + override def convert( + aggExpr: AggregateExpression, + expr: HllUnionAgg, + inputs: Seq[Attribute], + binding: Boolean, + conf: SQLConf): Option[ExprOuterClass.AggExpr] = { + val childExpr = exprToProto(expr.left, inputs, binding) + val allow = expr.right.eval() match { + case b: Boolean => b + case other => + withFallbackReason( + aggExpr, + s"Unsupported allowDifferentLgConfigK literal: $other", + expr.left) + return None + } + if (childExpr.isDefined) { + val builder = ExprOuterClass.HllUnionAgg.newBuilder() + builder.setChild(childExpr.get) + builder.setAllowDifferentLgConfigK(allow) + Some(ExprOuterClass.AggExpr.newBuilder().setHllUnionAgg(builder).build()) + } else { + withFallbackReason(aggExpr, expr.left) + None + } + } +} diff --git a/spark/src/main/spark-4.x/org/apache/comet/shims/Spark4xCometExprShim.scala b/spark/src/main/spark-4.x/org/apache/comet/shims/Spark4xCometExprShim.scala index bf8c86b1939..cfab80033af 100644 --- a/spark/src/main/spark-4.x/org/apache/comet/shims/Spark4xCometExprShim.scala +++ b/spark/src/main/spark-4.x/org/apache/comet/shims/Spark4xCometExprShim.scala @@ -20,7 +20,7 @@ package org.apache.comet.shims import org.apache.spark.sql.catalyst.expressions._ -import org.apache.spark.sql.catalyst.expressions.aggregate.HllSketchAgg +import org.apache.spark.sql.catalyst.expressions.aggregate.{HllSketchAgg, HllUnionAgg} import org.apache.spark.sql.catalyst.expressions.json.{JsonExpressionUtils, StructsToJsonEvaluator} import org.apache.spark.sql.catalyst.expressions.objects.{Invoke, StaticInvoke} import org.apache.spark.sql.catalyst.expressions.url.ParseUrlEvaluator @@ -28,7 +28,7 @@ import org.apache.spark.sql.types.ArrayType import org.apache.comet.CometExplainInfo import org.apache.comet.expressions.CometEvalMode -import org.apache.comet.serde.{CometAggregateExpressionSerde, CometExpressionSerde, CometHllSketchAgg, CometHllSketchEstimate, CometMapSort, CometToPrettyString, CometWidthBucket} +import org.apache.comet.serde.{CometAggregateExpressionSerde, CometExpressionSerde, CometHllSketchAgg, CometHllSketchEstimate, CometHllUnionAgg, CometMapSort, CometToPrettyString, CometWidthBucket} import org.apache.comet.serde.ExprOuterClass.Expr import org.apache.comet.serde.QueryPlanSerde.{exprToProtoInternal, optExprWithFallbackReason, scalarFunctionExprToProtoWithReturnType} @@ -53,7 +53,7 @@ trait Spark4xCometExprShim extends CometExprShim4x { def sparkVersionSpecificMapExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = Map(classOf[MapSort] -> CometMapSort) def sparkVersionSpecificAggregates: Map[Class[_], CometAggregateExpressionSerde[_]] = - Map(classOf[HllSketchAgg] -> CometHllSketchAgg) + Map(classOf[HllSketchAgg] -> CometHllSketchAgg, classOf[HllUnionAgg] -> CometHllUnionAgg) def sparkVersionSpecificExprToProtoInternal( expr: Expression, From a9f88933175ef18f3ec86a495f724fb595b15d83 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Thu, 2 Jul 2026 19:12:16 -0600 Subject: [PATCH 09/17] feat: add hll_union scalar function [skip ci] --- native/spark-expr/src/comet_scalar_funcs.rs | 6 +- native/spark-expr/src/hll_scalar.rs | 77 ++++++++++++++++++- native/spark-expr/src/lib.rs | 1 + .../apache/comet/serde/CometHllUnion.scala | 59 ++++++++++++++ .../comet/shims/Spark4xCometExprShim.scala | 5 +- 5 files changed, 144 insertions(+), 4 deletions(-) create mode 100644 spark/src/main/spark-4.x/org/apache/comet/serde/CometHllUnion.scala diff --git a/native/spark-expr/src/comet_scalar_funcs.rs b/native/spark-expr/src/comet_scalar_funcs.rs index d52212873dd..15de6bac49c 100644 --- a/native/spark-expr/src/comet_scalar_funcs.rs +++ b/native/spark-expr/src/comet_scalar_funcs.rs @@ -16,7 +16,7 @@ // under the License. use crate::hash_funcs::*; -use crate::hll_scalar::spark_hll_sketch_estimate; +use crate::hll_scalar::{spark_hll_sketch_estimate, spark_hll_union}; use crate::json_funcs::JsonArrayLength; use crate::map_funcs::spark_map_sort; use crate::math_funcs::abs::abs; @@ -226,6 +226,10 @@ pub fn create_comet_physical_fun_with_eval_mode( let func = Arc::new(|args: &[ColumnarValue]| spark_hll_sketch_estimate(args)); make_comet_scalar_udf!("hll_sketch_estimate", func, without data_type) } + "hll_union" => { + let func = Arc::new(|args: &[ColumnarValue]| spark_hll_union(args)); + make_comet_scalar_udf!("hll_union", func, without data_type) + } "to_time" => { make_comet_scalar_udf!("to_time", spark_to_time, without data_type, fail_on_error) } diff --git a/native/spark-expr/src/hll_scalar.rs b/native/spark-expr/src/hll_scalar.rs index 755b5df5fa8..f9a9648b20c 100644 --- a/native/spark-expr/src/hll_scalar.rs +++ b/native/spark-expr/src/hll_scalar.rs @@ -17,7 +17,7 @@ use crate::agg_funcs::estimate_from_bytes; use arrow::array::{Array, BinaryArray, Int64Array}; -use datafusion::common::Result; +use datafusion::common::{DataFusionError, Result}; use datafusion::physical_plan::ColumnarValue; use std::sync::Arc; @@ -36,6 +36,43 @@ pub fn spark_hll_sketch_estimate(args: &[ColumnarValue]) -> Result Result { + use crate::agg_funcs::{SparkHllSketch, SparkHllUnion}; + use arrow::array::BooleanArray; + let arrays = ColumnarValue::values_to_arrays(args)?; + let a = arrays[0].as_any().downcast_ref::().unwrap(); + let b = arrays[1].as_any().downcast_ref::().unwrap(); + let allow = arrays[2].as_any().downcast_ref::().unwrap(); + let mut out = arrow::array::BinaryBuilder::new(); + for i in 0..a.len() { + if a.is_null(i) || b.is_null(i) { + out.append_null(); + continue; + } + let sa = SparkHllSketch::from_bytes(a.value(i))?; + let sb = SparkHllSketch::from_bytes(b.value(i))?; + let allow_i = !allow.is_null(i) && allow.value(i); + if !allow_i && sa.lg_config_k() != sb.lg_config_k() { + return Err(DataFusionError::Execution(format!( + "Sketches have different lgConfigK values: {} and {}. \ + Set allowDifferentLgConfigK to true to enable unions of different lgConfigK.", + sa.lg_config_k(), + sb.lg_config_k() + ))); + } + // Spark builds `new Union(min(k1, k2))`. + let mut u = SparkHllUnion::new(sa.lg_config_k().min(sb.lg_config_k())); + u.merge(&sa); + u.merge(&sb); + out.append_value(u.to_sketch_bytes()); + } + Ok(ColumnarValue::Array(Arc::new(out.finish()))) +} + #[cfg(test)] mod tests { use super::*; @@ -58,3 +95,41 @@ mod tests { assert!((est - 1000).abs() <= 30, "estimate {est}"); } } + +#[cfg(test)] +mod union_tests { + use super::*; + use crate::agg_funcs::SparkHllSketch; + use arrow::array::BooleanArray; + + #[test] + fn unions_two_sketch_columns() { + let mut a = SparkHllSketch::new(12); + for i in 0..1000i64 { + a.update_i64(i); + } + let mut b = SparkHllSketch::new(12); + for i in 500..1500i64 { + b.update_i64(i); + } + let aa = Arc::new(BinaryArray::from(vec![Some( + a.to_sketch_bytes().as_slice(), + )])); + let bb = Arc::new(BinaryArray::from(vec![Some( + b.to_sketch_bytes().as_slice(), + )])); + let allow = Arc::new(BooleanArray::from(vec![false])); + let out = spark_hll_union(&[ + ColumnarValue::Array(aa), + ColumnarValue::Array(bb), + ColumnarValue::Array(allow), + ]) + .unwrap(); + let ColumnarValue::Array(arr) = out else { + panic!() + }; + let est = estimate_from_bytes(arr.as_any().downcast_ref::().unwrap().value(0)) + .unwrap(); + assert!((est - 1500).abs() <= 45, "estimate {est}"); + } +} diff --git a/native/spark-expr/src/lib.rs b/native/spark-expr/src/lib.rs index baca331dc9d..22609531552 100644 --- a/native/spark-expr/src/lib.rs +++ b/native/spark-expr/src/lib.rs @@ -61,6 +61,7 @@ mod conditional_funcs; mod conversion_funcs; mod hll_scalar; pub use hll_scalar::spark_hll_sketch_estimate; +pub use hll_scalar::spark_hll_union; mod map_funcs; pub use map_funcs::spark_map_sort; mod math_funcs; diff --git a/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllUnion.scala b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllUnion.scala new file mode 100644 index 00000000000..af7b3309666 --- /dev/null +++ b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllUnion.scala @@ -0,0 +1,59 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.comet.serde + +import org.apache.spark.sql.catalyst.expressions.{Attribute, HllUnion} +import org.apache.spark.sql.types.BinaryType + +import org.apache.comet.serde.QueryPlanSerde.{exprToProtoInternal, optExprWithFallbackReason, scalarFunctionExprToProtoWithReturnType} + +// Spark 4.0 HllUnion is a TernaryExpression: `first`, `second` (binary sketches), +// `third` (allowDifferentLgConfigK boolean, default Literal(false)). All three are +// passed to the native `hll_union`, which enforces the flag and unions with min lgConfigK. +object CometHllUnion extends CometExpressionSerde[HllUnion] { + private val incompatReason = + "Comet uses a Rust DataSketches port; HLL sketch bytes and estimates may differ slightly from Spark." + + override def getIncompatibleReasons(): Seq[String] = Seq(incompatReason) + + override def getSupportLevel(expr: HllUnion): SupportLevel = { + if (!expr.third.foldable) { + Unsupported(Some("The allowDifferentLgConfigK argument must be a foldable literal.")) + } else { + Incompatible(Some(incompatReason)) + } + } + override def convert( + expr: HllUnion, + inputs: Seq[Attribute], + binding: Boolean): Option[ExprOuterClass.Expr] = { + val first = exprToProtoInternal(expr.first, inputs, binding) + val second = exprToProtoInternal(expr.second, inputs, binding) + val third = exprToProtoInternal(expr.third, inputs, binding) + val unionExpr = scalarFunctionExprToProtoWithReturnType( + "hll_union", + BinaryType, + failOnError = false, + first, + second, + third) + optExprWithFallbackReason(unionExpr, expr, expr.first, expr.second, expr.third) + } +} diff --git a/spark/src/main/spark-4.x/org/apache/comet/shims/Spark4xCometExprShim.scala b/spark/src/main/spark-4.x/org/apache/comet/shims/Spark4xCometExprShim.scala index cfab80033af..93cc2ef1848 100644 --- a/spark/src/main/spark-4.x/org/apache/comet/shims/Spark4xCometExprShim.scala +++ b/spark/src/main/spark-4.x/org/apache/comet/shims/Spark4xCometExprShim.scala @@ -28,7 +28,7 @@ import org.apache.spark.sql.types.ArrayType import org.apache.comet.CometExplainInfo import org.apache.comet.expressions.CometEvalMode -import org.apache.comet.serde.{CometAggregateExpressionSerde, CometExpressionSerde, CometHllSketchAgg, CometHllSketchEstimate, CometHllUnionAgg, CometMapSort, CometToPrettyString, CometWidthBucket} +import org.apache.comet.serde.{CometAggregateExpressionSerde, CometExpressionSerde, CometHllSketchAgg, CometHllSketchEstimate, CometHllUnion, CometHllUnionAgg, CometMapSort, CometToPrettyString, CometWidthBucket} import org.apache.comet.serde.ExprOuterClass.Expr import org.apache.comet.serde.QueryPlanSerde.{exprToProtoInternal, optExprWithFallbackReason, scalarFunctionExprToProtoWithReturnType} @@ -49,7 +49,8 @@ trait Spark4xCometExprShim extends CometExprShim4x { def sparkVersionSpecificMiscExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = Map( classOf[ToPrettyString] -> CometToPrettyString, - classOf[HllSketchEstimate] -> CometHllSketchEstimate) + classOf[HllSketchEstimate] -> CometHllSketchEstimate, + classOf[HllUnion] -> CometHllUnion) def sparkVersionSpecificMapExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = Map(classOf[MapSort] -> CometMapSort) def sparkVersionSpecificAggregates: Map[Class[_], CometAggregateExpressionSerde[_]] = From 93dedb3fdf3d06c709f90d8167b99f97a659b41e Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Thu, 2 Jul 2026 19:19:42 -0600 Subject: [PATCH 10/17] test: add hll_union end-to-end and error tests, apply formatting [skip ci] --- native/spark-expr/src/agg_funcs/hll_sketch.rs | 15 ++++-- .../apache/comet/CometExpressionSuite.scala | 47 +++++++++++++++++++ 2 files changed, 59 insertions(+), 3 deletions(-) diff --git a/native/spark-expr/src/agg_funcs/hll_sketch.rs b/native/spark-expr/src/agg_funcs/hll_sketch.rs index 5e40c54aaec..68c8d039432 100644 --- a/native/spark-expr/src/agg_funcs/hll_sketch.rs +++ b/native/spark-expr/src/agg_funcs/hll_sketch.rs @@ -157,7 +157,10 @@ mod tests { } let bytes = s.to_sketch_bytes(); let est = estimate_from_bytes(&bytes).unwrap(); - assert!((est - 1000).abs() <= 30, "estimate {est} not within 3% of 1000"); + assert!( + (est - 1000).abs() <= 30, + "estimate {est} not within 3% of 1000" + ); } #[test] @@ -174,7 +177,10 @@ mod tests { u.merge(&a); u.merge(&b); let est = estimate_from_bytes(&u.to_sketch_bytes()).unwrap(); - assert!((est - 1500).abs() <= 45, "union estimate {est} not within 3% of 1500"); + assert!( + (est - 1500).abs() <= 45, + "union estimate {est} not within 3% of 1500" + ); } /// A sketch built from raw bytes (StringType/BinaryType path) round-trips and @@ -187,7 +193,10 @@ mod tests { } s.update_bytes(b""); // skipped, no effect let est = estimate_from_bytes(&s.to_sketch_bytes()).unwrap(); - assert!((est - 1000).abs() <= 30, "estimate {est} not within 3% of 1000"); + assert!( + (est - 1000).abs() <= 30, + "estimate {est} not within 3% of 1000" + ); } /// Cross-engine regression guard: `testdata/hll_sketch_spark_lgk12.bin` was diff --git a/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala b/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala index 8faecb24a52..2b58a9e7913 100644 --- a/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala @@ -3405,4 +3405,51 @@ class CometExpressionSuite extends CometTestBase with AdaptiveSparkPlanHelper { } } + test("hll_union_agg and hll_union (incompatible, opt-in)") { + assume(isSpark40Plus) + withSQLConf( + "spark.comet.expression.HllSketchAgg.allowIncompatible" -> "true", + "spark.comet.expression.HllSketchEstimate.allowIncompatible" -> "true", + "spark.comet.expression.HllUnionAgg.allowIncompatible" -> "true", + "spark.comet.expression.HllUnion.allowIncompatible" -> "true") { + withParquetTable((0 until 1000).map(i => (i % 3, i)), "tbl") { + // hll_union_agg: union the per-group sketches -> ~1000 distinct. + val aggDf = sql( + "SELECT hll_sketch_estimate(hll_union_agg(s)) FROM " + + "(SELECT _1 AS g, hll_sketch_agg(_2) AS s FROM tbl GROUP BY _1)") + checkCometOperators(stripAQEPlan(aggDf.queryExecution.executedPlan)) + val aggEst = aggDf.collect().head.getLong(0) + assert(math.abs(aggEst - 1000).toDouble / 1000 <= 0.05, s"union_agg estimate $aggEst") + + // hll_union: union two disjoint group sketches -> ~667 distinct. + val unionDf = sql( + "SELECT hll_sketch_estimate(hll_union(a.s, b.s)) FROM " + + "(SELECT hll_sketch_agg(_2) AS s FROM tbl WHERE _1 = 0) a, " + + "(SELECT hll_sketch_agg(_2) AS s FROM tbl WHERE _1 = 1) b") + checkCometOperators(stripAQEPlan(unionDf.queryExecution.executedPlan)) + val unionEst = unionDf.collect().head.getLong(0) + assert(math.abs(unionEst - 667).toDouble / 667 <= 0.05, s"union estimate $unionEst") + } + } + } + + test("hll_union_agg rejects different lgConfigK when not allowed") { + assume(isSpark40Plus) + withSQLConf( + "spark.comet.expression.HllSketchAgg.allowIncompatible" -> "true", + "spark.comet.expression.HllUnionAgg.allowIncompatible" -> "true") { + withParquetTable((0 until 100).map(i => Tuple1(i)), "tbl") { + // A lgConfigK=10 sketch unioned with a lgConfigK=12 sketch (allowDifferentLgConfigK + // defaults false) must throw in BOTH Spark and Comet. + val df = sql( + "SELECT hll_union_agg(s) FROM (" + + " SELECT hll_sketch_agg(_1, 10) AS s FROM tbl UNION ALL" + + " SELECT hll_sketch_agg(_1, 12) AS s FROM tbl)") + val (sparkErr, cometErr) = checkSparkAnswerMaybeThrows(df) + assert(sparkErr.isDefined, "expected Spark to throw on different lgConfigK") + assert(cometErr.isDefined, "expected Comet to throw on different lgConfigK") + } + } + } + } From ad39ec54d2f55260e99a010a86abf012b654f4b8 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Thu, 2 Jul 2026 19:37:58 -0600 Subject: [PATCH 11/17] fix: HLL empty-group returns empty sketch and size() accounts for sketch heap [skip ci] --- .../src/agg_funcs/hll_sketch_agg.rs | 47 ++++++++++------- .../spark-expr/src/agg_funcs/hll_union_agg.rs | 51 ++++++++++++++++++- .../apache/comet/CometExpressionSuite.scala | 20 ++++++++ 3 files changed, 98 insertions(+), 20 deletions(-) diff --git a/native/spark-expr/src/agg_funcs/hll_sketch_agg.rs b/native/spark-expr/src/agg_funcs/hll_sketch_agg.rs index e962f984f14..3453b646573 100644 --- a/native/spark-expr/src/agg_funcs/hll_sketch_agg.rs +++ b/native/spark-expr/src/agg_funcs/hll_sketch_agg.rs @@ -77,14 +77,12 @@ impl AggregateUDFImpl for HllSketchAgg { #[derive(Debug)] pub struct HllSketchAccumulator { sketch: SparkHllSketch, - saw_input: bool, } impl HllSketchAccumulator { pub fn new(lg_config_k: u8) -> Self { Self { sketch: SparkHllSketch::new(lg_config_k), - saw_input: false, } } } @@ -99,27 +97,21 @@ impl Accumulator for HllSketchAccumulator { match ScalarValue::try_from_array(arr, i)? { ScalarValue::Int8(Some(v)) => { self.sketch.update_i64(v as i64); - self.saw_input = true; } ScalarValue::Int16(Some(v)) => { self.sketch.update_i64(v as i64); - self.saw_input = true; } ScalarValue::Int32(Some(v)) => { self.sketch.update_i64(v as i64); - self.saw_input = true; } ScalarValue::Int64(Some(v)) => { self.sketch.update_i64(v); - self.saw_input = true; } ScalarValue::Utf8(Some(v)) => { self.sketch.update_bytes(v.as_bytes()); - self.saw_input = true; } ScalarValue::Binary(Some(v)) => { self.sketch.update_bytes(&v); - self.saw_input = true; } // Spark's HllSketchAgg ignores null inputs. ScalarValue::Int8(None) @@ -139,22 +131,18 @@ impl Accumulator for HllSketchAccumulator { } fn evaluate(&mut self) -> Result { - // Spark returns a non-null sketch even for empty groups only when it saw input; - // an empty group yields NULL. - if !self.saw_input { - return Ok(ScalarValue::Binary(None)); - } + // Spark's HllSketchAgg is declared non-nullable: an empty/all-null group + // still returns a serialized empty sketch (which estimates to 0), never NULL. Ok(ScalarValue::Binary(Some(self.sketch.to_sketch_bytes()))) } fn size(&self) -> usize { - std::mem::size_of_val(self) + // An HLL_8 sketch at lgConfigK=k can heap-allocate up to 1 << k bytes; + // account for that so memory reservation reflects actual usage. + std::mem::size_of_val(self) + (1usize << self.sketch.lg_config_k() as usize) } fn state(&mut self) -> Result> { - if !self.saw_input { - return Ok(vec![ScalarValue::Binary(None)]); - } Ok(vec![ScalarValue::Binary(Some( self.sketch.to_sketch_bytes(), ))]) @@ -169,7 +157,6 @@ impl Accumulator for HllSketchAccumulator { let peer = SparkHllSketch::from_bytes(arr.value(i))?; // Merge peer into self by unioning; reuse SparkHllUnion via sketch merge. self.sketch.merge_sketch(&peer); - self.saw_input = true; } Ok(()) } @@ -193,4 +180,28 @@ mod tests { let est = crate::agg_funcs::estimate_from_bytes(&bytes).unwrap(); assert!((est - 1000).abs() <= 30, "estimate {est}"); } + + /// Spark's `HllSketchAgg` is non-nullable: an empty/all-null group still + /// produces a serialized empty sketch (estimate 0), not NULL. + #[test] + fn empty_group_evaluates_to_empty_sketch_not_null() { + let mut acc = HllSketchAccumulator::new(12); + let ScalarValue::Binary(Some(bytes)) = acc.evaluate().unwrap() else { + panic!("expected Binary(Some(_)) for an empty group, got NULL") + }; + let est = crate::agg_funcs::estimate_from_bytes(&bytes).unwrap(); + assert_eq!(est, 0, "empty sketch should estimate to 0, got {est}"); + } + + #[test] + fn size_accounts_for_sketch_heap() { + let mut acc = HllSketchAccumulator::new(12); + let arr = Arc::new(Int64Array::from((0..10000i64).collect::>())); + acc.update_batch(&[arr]).unwrap(); + assert!( + acc.size() > 1000, + "size() should account for the sketch heap allocation, got {}", + acc.size() + ); + } } diff --git a/native/spark-expr/src/agg_funcs/hll_union_agg.rs b/native/spark-expr/src/agg_funcs/hll_union_agg.rs index 9d3658cef78..aee7c044a10 100644 --- a/native/spark-expr/src/agg_funcs/hll_union_agg.rs +++ b/native/spark-expr/src/agg_funcs/hll_union_agg.rs @@ -65,6 +65,10 @@ impl AggregateUDFImpl for HllUnionAgg { } } +/// Default `lgMaxK` used by Spark's `new Union()` when constructing the empty +/// union returned for a group that never absorbed any sketch. +const DEFAULT_LG_K: u8 = 12; + #[derive(Debug)] pub struct HllUnionAccumulator { // Spark's HllUnionAgg defers creating the Union until the first sketch is seen, @@ -119,18 +123,31 @@ impl Accumulator for HllUnionAccumulator { Ok(()) } fn evaluate(&mut self) -> Result { + // Spark's HllUnionAgg is declared non-nullable: an empty/all-null group + // still returns the serialized bytes of an empty `new Union()` (default + // lgMaxK), which estimates to 0, never NULL. match &self.union { Some(u) => Ok(ScalarValue::Binary(Some(u.to_sketch_bytes()))), - None => Ok(ScalarValue::Binary(None)), + None => Ok(ScalarValue::Binary(Some( + SparkHllUnion::new(DEFAULT_LG_K).to_sketch_bytes(), + ))), } } fn size(&self) -> usize { + // An HLL_8 sketch at lgConfigK=k can heap-allocate up to 1 << k bytes; + // account for that so memory reservation reflects actual usage. std::mem::size_of_val(self) + + self + .seen_lg_config_k + .map(|k| 1usize << k as usize) + .unwrap_or(0) } fn state(&mut self) -> Result> { match &self.union { Some(u) => Ok(vec![ScalarValue::Binary(Some(u.to_sketch_bytes()))]), - None => Ok(vec![ScalarValue::Binary(None)]), + None => Ok(vec![ScalarValue::Binary(Some( + SparkHllUnion::new(DEFAULT_LG_K).to_sketch_bytes(), + ))]), } } fn merge_batch(&mut self, states: &[ArrayRef]) -> Result<()> { @@ -174,4 +191,34 @@ mod tests { let est = crate::agg_funcs::estimate_from_bytes(&bytes).unwrap(); assert!((est - 1500).abs() <= 45, "union estimate {est}"); } + + /// Spark's `HllUnionAgg` is non-nullable: an empty/all-null group still + /// produces the serialized bytes of an empty union (estimate 0), not NULL. + #[test] + fn empty_group_evaluates_to_empty_sketch_not_null() { + let mut acc = HllUnionAccumulator::new(false); + let ScalarValue::Binary(Some(bytes)) = acc.evaluate().unwrap() else { + panic!("expected Binary(Some(_)) for an empty group, got NULL") + }; + let est = crate::agg_funcs::estimate_from_bytes(&bytes).unwrap(); + assert_eq!(est, 0, "empty union should estimate to 0, got {est}"); + } + + #[test] + fn size_accounts_for_sketch_heap() { + let mut a = SparkHllSketch::new(12); + for i in 0..10000i64 { + a.update_i64(i); + } + let arr = Arc::new(BinaryArray::from(vec![Some( + a.to_sketch_bytes().as_slice(), + )])); + let mut acc = HllUnionAccumulator::new(false); + acc.update_batch(&[arr]).unwrap(); + assert!( + acc.size() > 1000, + "size() should account for the sketch heap allocation, got {}", + acc.size() + ); + } } diff --git a/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala b/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala index 2b58a9e7913..9605f5dd160 100644 --- a/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala @@ -3452,4 +3452,24 @@ class CometExpressionSuite extends CometTestBase with AdaptiveSparkPlanHelper { } } + test("hll_sketch_agg over all-null input estimates to 0, not NULL") { + assume(isSpark40Plus) + // Spark's HllSketchAgg/HllSketchEstimate are declared non-nullable: an empty or + // all-null group still produces a serialized empty sketch, and hll_sketch_estimate + // reads that as 0, never NULL. Guard against regressing to Binary(None) here. + withSQLConf( + "spark.comet.expression.HllSketchAgg.allowIncompatible" -> "true", + "spark.comet.expression.HllSketchEstimate.allowIncompatible" -> "true") { + withParquetTable((0 until 100).map(_ => Tuple1(null.asInstanceOf[Integer])), "tbl") { + val df = sql("SELECT hll_sketch_estimate(hll_sketch_agg(_1)) FROM tbl") + checkCometOperators(stripAQEPlan(df.queryExecution.executedPlan)) + val row = df.collect().head + assert(!row.isNullAt(0), "expected a non-null estimate for an all-null group") + assert( + row.getLong(0) == 0, + s"expected estimate 0 for an all-null group, got ${row.getLong(0)}") + } + } + } + } From df37620d910bd817cb94c3f47da1a9a8b3579d26 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Thu, 2 Jul 2026 19:47:07 -0600 Subject: [PATCH 12/17] ci: run CI for HLL sketch functions From 32793513d86f8bfab548bf36d34deb0ffa9215ab Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Wed, 9 Sep 2026 08:24:47 -0600 Subject: [PATCH 13/17] fix: decode compact HLL sketches and keep empty union partials NULL --- native/spark-expr/src/agg_funcs/hll_sketch.rs | 55 ++++++++++- .../spark-expr/src/agg_funcs/hll_union_agg.rs | 98 ++++++++++++++++++- 2 files changed, 149 insertions(+), 4 deletions(-) diff --git a/native/spark-expr/src/agg_funcs/hll_sketch.rs b/native/spark-expr/src/agg_funcs/hll_sketch.rs index 68c8d039432..0a6497ce8ce 100644 --- a/native/spark-expr/src/agg_funcs/hll_sketch.rs +++ b/native/spark-expr/src/agg_funcs/hll_sketch.rs @@ -43,6 +43,58 @@ pub struct SparkHllSketch { inner: HllSketch, } +/// Byte offsets into the DataSketches HLL preamble, and the bits we need from it. +mod preamble { + /// Serialization flags. Bit 3 is COMPACT. + pub const FLAGS: usize = 5; + /// Mode byte. Low two bits are the current mode (0 LIST, 1 SET, 2 HLL). + pub const MODE: usize = 7; + pub const COMPACT_FLAG: u8 = 8; + pub const CUR_MODE_MASK: u8 = 0x3; + pub const CUR_MODE_HLL: u8 = 2; +} + +/// Work around a decoding bug in `datasketches` 0.3.0 for compact sketches in an HLL array mode. +/// +/// `Array4::deserialize` (and the `Array6` / `Array8` equivalents) skip the register block +/// entirely when the COMPACT flag is set, leaving every register zero: +/// +/// ```text +/// let mut data = vec![0u8; num_bytes]; +/// if !compact { +/// cursor.read_exact(&mut data)?; +/// } else { +/// cursor.advance(num_bytes as u64); +/// } +/// ``` +/// +/// The damage is quiet, which is what makes it worth guarding: the decoded sketch's own +/// `estimate()` still looks correct because it comes back from the HIP accumulator in the +/// preamble, but every union built from it is wrong. Two disjoint 1,000-value sketches union to +/// ~989 rather than ~1991. +/// +/// Clearing the flag is a correct parse rather than a guess. The register block is present in +/// both the compact and updatable forms, and the crate reads the HLL_4 auxiliary map as +/// `aux_count` coupons regardless of the flag - which is the compact layout. LIST and SET mode +/// compaction *is* a genuinely different layout, and the crate handles those correctly, so this +/// only touches HLL array mode. +/// +/// Returns `None` when the input needs no rewriting, so the common path does not copy. +/// +/// `compact_input_survives_a_union` pins the behaviour: if a future `datasketches` release fixes +/// the register read, that test is what tells us this can be deleted. +fn normalize_compact_hll_array(bytes: &[u8]) -> Option> { + if bytes.len() <= preamble::MODE + || bytes[preamble::MODE] & preamble::CUR_MODE_MASK != preamble::CUR_MODE_HLL + || bytes[preamble::FLAGS] & preamble::COMPACT_FLAG == 0 + { + return None; + } + let mut owned = bytes.to_vec(); + owned[preamble::FLAGS] &= !preamble::COMPACT_FLAG; + Some(owned) +} + impl SparkHllSketch { /// Create an empty HLL_8 sketch with the given `lgConfigK`. pub fn new(lg_config_k: u8) -> Self { @@ -89,7 +141,8 @@ impl SparkHllSketch { /// Deserialize a DataSketches sketch (either compact or updatable form). pub fn from_bytes(bytes: &[u8]) -> Result { - HllSketch::deserialize(bytes) + let normalized = normalize_compact_hll_array(bytes); + HllSketch::deserialize(normalized.as_deref().unwrap_or(bytes)) .map(|inner| Self { inner }) .map_err(|e| DataFusionError::Internal(format!("invalid HLL sketch bytes: {e}"))) } diff --git a/native/spark-expr/src/agg_funcs/hll_union_agg.rs b/native/spark-expr/src/agg_funcs/hll_union_agg.rs index aee7c044a10..ccd3a43077e 100644 --- a/native/spark-expr/src/agg_funcs/hll_union_agg.rs +++ b/native/spark-expr/src/agg_funcs/hll_union_agg.rs @@ -143,11 +143,16 @@ impl Accumulator for HllUnionAccumulator { .unwrap_or(0) } fn state(&mut self) -> Result> { + // Unlike `evaluate`, an empty partial emits NULL rather than an empty lgConfigK=12 + // sketch. `merge_batch` skips nulls, so the Final phase sees nothing at all from a + // partition that absorbed nothing - which is the point: emitting a concrete + // lgConfigK=12 sketch here would make it the first `lgConfigK` the Final accumulator + // sees, and every real sketch at a different k would then fail the "Sketches have + // different lgConfigK values" check. A partition with no input must not get a vote on + // the union's k. match &self.union { Some(u) => Ok(vec![ScalarValue::Binary(Some(u.to_sketch_bytes()))]), - None => Ok(vec![ScalarValue::Binary(Some( - SparkHllUnion::new(DEFAULT_LG_K).to_sketch_bytes(), - ))]), + None => Ok(vec![ScalarValue::Binary(None)]), } } fn merge_batch(&mut self, states: &[ArrayRef]) -> Result<()> { @@ -221,4 +226,91 @@ mod tests { acc.size() ); } + + /// An empty partial must emit NULL, not an empty lgConfigK=12 sketch. Emitting a concrete + /// sketch made the empty partition the first `lgConfigK` the Final accumulator saw, so a + /// partition holding only lgConfigK=10 sketches then failed the mismatch check. + #[test] + fn empty_partial_state_is_null() { + let mut acc = HllUnionAccumulator::new(false); + assert_eq!( + acc.state().unwrap(), + vec![ScalarValue::Binary(None)], + "an empty partial must not contribute a sketch to the final merge" + ); + } + + #[test] + fn empty_partial_does_not_fix_the_final_lg_config_k() { + // Partition A saw nothing; partition B saw only lgConfigK=10 sketches. Merging A's + // state before B's used to abort with "Sketches have different lgConfigK values: + // 12 and 10". + let mut empty_partial = HllUnionAccumulator::new(false); + let empty_state = empty_partial.state().unwrap(); + + let mut k10 = SparkHllSketch::new(10); + for i in 0..1000i64 { + k10.update_i64(i); + } + let mut k10_partial = HllUnionAccumulator::new(false); + k10_partial + .update_batch(&[Arc::new(BinaryArray::from(vec![Some( + k10.to_sketch_bytes().as_slice(), + )]))]) + .unwrap(); + let k10_state = k10_partial.state().unwrap(); + + let mut final_acc = HllUnionAccumulator::new(false); + for state in [empty_state, k10_state] { + let arrays: Vec = state + .into_iter() + .map(|sv| sv.to_array_of_size(1).unwrap()) + .collect(); + final_acc.merge_batch(&arrays).unwrap(); + } + + let ScalarValue::Binary(Some(bytes)) = final_acc.evaluate().unwrap() else { + panic!("expected Binary(Some(_))") + }; + let est = crate::agg_funcs::estimate_from_bytes(&bytes).unwrap(); + assert!((est - 1000).abs() <= 40, "union estimate {est}"); + } + + /// A compact-form sketch has to survive a union. `datasketches` 0.3.0 drops the register + /// block for compact HLL array modes, which leaves the decoded sketch's own estimate intact + /// (it is restored from the HIP accumulator) but makes every union built from it wrong. + #[test] + fn compact_input_survives_a_union() { + let mut a = SparkHllSketch::new(12); + for i in 0..1000i64 { + a.update_i64(i); + } + let mut b = SparkHllSketch::new(12); + for i in 1000..2000i64 { + b.update_i64(i); + } + // Set the COMPACT flag, which is what a DataSketches `toCompactByteArray()` sketch + // carries. For an HLL array mode the register block is present either way, so this is + // the same bytes with a different flag. + let compact = |s: &SparkHllSketch| { + let mut v = s.to_sketch_bytes(); + v[5] |= 8; + v + }; + let arr = Arc::new(BinaryArray::from(vec![ + Some(compact(&a).as_slice()), + Some(compact(&b).as_slice()), + ])); + let mut acc = HllUnionAccumulator::new(false); + acc.update_batch(&[arr]).unwrap(); + + let ScalarValue::Binary(Some(bytes)) = acc.evaluate().unwrap() else { + panic!("expected Binary(Some(_))") + }; + let est = crate::agg_funcs::estimate_from_bytes(&bytes).unwrap(); + assert!( + (est - 2000).abs() <= 80, + "union of two disjoint compact sketches estimated {est}, expected ~2000" + ); + } } From 1fa3f260d7a4d1289f30e0895006e2e6b3fb6f4e Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Wed, 9 Sep 2026 10:16:06 -0600 Subject: [PATCH 14/17] fix: honour a NULL allowDifferentLgConfigK and reject undecodable HLL_4 hll_union is a TernaryExpression, so a NULL third argument makes the whole call NULL. Separately, an updatable HLL_4 sketch carrying auxiliary-map exceptions is now rejected: the bundled decoder reads the compact aux layout in both forms, so those sketches would decode to a wrong estimate. --- native/spark-expr/src/agg_funcs/hll_sketch.rs | 114 ++++++++++++++++++ native/spark-expr/src/hll_scalar.rs | 48 +++++++- .../comet/serde/CometHllSketchEstimate.scala | 7 +- .../apache/comet/serde/CometHllUnion.scala | 7 +- .../apache/comet/serde/CometHllUnionAgg.scala | 7 +- .../apache/comet/CometExpressionSuite.scala | 7 ++ 6 files changed, 185 insertions(+), 5 deletions(-) diff --git a/native/spark-expr/src/agg_funcs/hll_sketch.rs b/native/spark-expr/src/agg_funcs/hll_sketch.rs index 0a6497ce8ce..58484ce2e42 100644 --- a/native/spark-expr/src/agg_funcs/hll_sketch.rs +++ b/native/spark-expr/src/agg_funcs/hll_sketch.rs @@ -49,9 +49,60 @@ mod preamble { pub const FLAGS: usize = 5; /// Mode byte. Low two bits are the current mode (0 LIST, 1 SET, 2 HLL). pub const MODE: usize = 7; + /// Number of auxiliary-map exceptions, for a sketch in an HLL array mode. + pub const AUX_COUNT: usize = 36; + /// Total preamble length for an HLL array mode, i.e. where the register block starts. + pub const HLL_SIZE: usize = 40; pub const COMPACT_FLAG: u8 = 8; pub const CUR_MODE_MASK: u8 = 0x3; pub const CUR_MODE_HLL: u8 = 2; + /// Target type lives in bits 2-3 of the mode byte: 0 HLL_4, 1 HLL_6, 2 HLL_8. + pub const TGT_HLL4: u8 = 0; + + pub fn tgt_type(mode_byte: u8) -> u8 { + (mode_byte >> 2) & 0x3 + } +} + +/// An error for the one input shape `datasketches` 0.3.0 decodes to silently wrong values: +/// an updatable HLL_4 sketch carrying auxiliary-map exceptions. +/// +/// `Array4::deserialize` reads exactly `aux_count` coupons from the aux region regardless of the +/// COMPACT flag. That is the *compact* aux layout. DataSketches-Java's *updatable* HLL_4 form +/// writes `1 << lgAuxArrInts` ints including the empty slots, so reading the first `aux_count` of +/// those pulls empty slots in as zero coupons and drops the real exceptions. Nothing errors, and +/// the estimate is quietly wrong. +/// +/// The exposure is real rather than theoretical: Comet always *writes* HLL_8, but +/// `hll_sketch_estimate` / `hll_union` / `hll_union_agg` accept any binary column, and +/// DataSketches-Java's default target type is HLL_4. So a sketch column produced elsewhere can +/// land here. An exception map only appears once some register exceeds `curMin + 15`, so ordinary +/// low-cardinality HLL_4 input still reads correctly and is deliberately still accepted - the +/// check is narrowed to the shape that is actually mis-decoded rather than rejecting HLL_4 +/// outright, which would refuse most third-party sketches for no reason. +fn reject_undecodable_hll4(bytes: &[u8]) -> Result<(), DataFusionError> { + if bytes.len() < preamble::HLL_SIZE + || bytes[preamble::MODE] & preamble::CUR_MODE_MASK != preamble::CUR_MODE_HLL + || preamble::tgt_type(bytes[preamble::MODE]) != preamble::TGT_HLL4 + || bytes[preamble::FLAGS] & preamble::COMPACT_FLAG != 0 + { + return Ok(()); + } + let aux_count = u32::from_le_bytes([ + bytes[preamble::AUX_COUNT], + bytes[preamble::AUX_COUNT + 1], + bytes[preamble::AUX_COUNT + 2], + bytes[preamble::AUX_COUNT + 3], + ]); + if aux_count == 0 { + return Ok(()); + } + Err(DataFusionError::Execution(format!( + "Cannot read an updatable HLL_4 sketch with {aux_count} auxiliary-map entries: the \ + bundled datasketches decoder reads the compact auxiliary layout in both forms, so this \ + sketch would decode to a silently wrong estimate. Convert it to HLL_8, or to the compact \ + HLL_4 form, before reading it with Comet." + ))) } /// Work around a decoding bug in `datasketches` 0.3.0 for compact sketches in an HLL array mode. @@ -141,6 +192,7 @@ impl SparkHllSketch { /// Deserialize a DataSketches sketch (either compact or updatable form). pub fn from_bytes(bytes: &[u8]) -> Result { + reject_undecodable_hll4(bytes)?; let normalized = normalize_compact_hll_array(bytes); HllSketch::deserialize(normalized.as_deref().unwrap_or(bytes)) .map(|inner| Self { inner }) @@ -268,4 +320,66 @@ mod tests { "estimate {est} of Spark-produced sketch not within 3% of 1000" ); } + + /// An HLL_4 sketch with auxiliary-map exceptions must be refused rather than silently + /// mis-decoded. At lgK=12 the exceptions appear somewhere above ~100k distinct values. + #[test] + fn updatable_hll4_with_aux_entries_is_rejected() { + use datasketches::hll::{HllSketch, HllType}; + let mut sketch = HllSketch::new(12, HllType::Hll4); + for i in 0..100_000i64 { + sketch.update(i); + } + let bytes = sketch.serialize(); + // Confirm the premise rather than assuming it: HLL array mode, HLL_4, not compact, and + // carrying at least one exception. + assert_eq!( + bytes[preamble::MODE] & preamble::CUR_MODE_MASK, + preamble::CUR_MODE_HLL + ); + assert_eq!( + preamble::tgt_type(bytes[preamble::MODE]), + preamble::TGT_HLL4 + ); + assert_eq!(bytes[preamble::FLAGS] & preamble::COMPACT_FLAG, 0); + let aux_count = u32::from_le_bytes([ + bytes[preamble::AUX_COUNT], + bytes[preamble::AUX_COUNT + 1], + bytes[preamble::AUX_COUNT + 2], + bytes[preamble::AUX_COUNT + 3], + ]); + assert!(aux_count > 0, "expected the sketch to carry aux entries"); + + let err = SparkHllSketch::from_bytes(&bytes).unwrap_err().to_string(); + assert!( + err.contains("auxiliary-map"), + "expected a clear rejection, got {err}" + ); + } + + /// The guard is narrowed to the shape that is actually mis-decoded, so ordinary HLL_4 input + /// with no exceptions still reads. Rejecting HLL_4 outright would refuse most third-party + /// sketches for no reason. + #[test] + fn hll4_without_aux_entries_still_reads() { + use datasketches::hll::{HllSketch, HllType}; + let mut sketch = HllSketch::new(12, HllType::Hll4); + for i in 0..1_000i64 { + sketch.update(i); + } + let bytes = sketch.serialize(); + let aux_count = u32::from_le_bytes([ + bytes[preamble::AUX_COUNT], + bytes[preamble::AUX_COUNT + 1], + bytes[preamble::AUX_COUNT + 2], + bytes[preamble::AUX_COUNT + 3], + ]); + assert_eq!(aux_count, 0, "this cardinality should need no exceptions"); + let read = SparkHllSketch::from_bytes(&bytes).unwrap(); + assert!( + (read.estimate() - 1000.0).abs() < 40.0, + "estimate {}", + read.estimate() + ); + } } diff --git a/native/spark-expr/src/hll_scalar.rs b/native/spark-expr/src/hll_scalar.rs index f9a9648b20c..b1f3e77167f 100644 --- a/native/spark-expr/src/hll_scalar.rs +++ b/native/spark-expr/src/hll_scalar.rs @@ -49,13 +49,18 @@ pub fn spark_hll_union(args: &[ColumnarValue]) -> Result { let allow = arrays[2].as_any().downcast_ref::().unwrap(); let mut out = arrow::array::BinaryBuilder::new(); for i in 0..a.len() { - if a.is_null(i) || b.is_null(i) { + // Spark's `HllUnion` is a `TernaryExpression` evaluated through `nullSafeEval`, and + // `TernaryExpression.eval` returns NULL when *any* of the three inputs is NULL - the + // `allowDifferentLgConfigK` flag included. `Literal(null, BooleanType)` is foldable, so + // the serde's `third.foldable` check lets it through and this loop is the only thing + // standing between a NULL flag and a non-null sketch. + if a.is_null(i) || b.is_null(i) || allow.is_null(i) { out.append_null(); continue; } let sa = SparkHllSketch::from_bytes(a.value(i))?; let sb = SparkHllSketch::from_bytes(b.value(i))?; - let allow_i = !allow.is_null(i) && allow.value(i); + let allow_i = allow.value(i); if !allow_i && sa.lg_config_k() != sb.lg_config_k() { return Err(DataFusionError::Execution(format!( "Sketches have different lgConfigK values: {} and {}. \ @@ -132,4 +137,43 @@ mod union_tests { .unwrap(); assert!((est - 1500).abs() <= 45, "estimate {est}"); } + + /// Spark's `HllUnion` is a `TernaryExpression`, so a NULL `allowDifferentLgConfigK` makes the + /// whole call NULL. `Literal(null, BooleanType)` is foldable and therefore passes the serde's + /// `third.foldable` gate, so this loop is what has to honour it. + #[test] + fn a_null_allow_flag_yields_null() { + let mut a = SparkHllSketch::new(12); + for i in 0..1000i64 { + a.update_i64(i); + } + let bytes = a.to_sketch_bytes(); + let column = || { + Arc::new(BinaryArray::from(vec![ + Some(bytes.as_slice()), + Some(bytes.as_slice()), + ])) + }; + // Row 0 has a NULL flag, row 1 a real one, so the non-null row proves the NULL is not + // just nulling the whole column. + let allow = Arc::new(BooleanArray::from(vec![None, Some(true)])); + let out = spark_hll_union(&[ + ColumnarValue::Array(column()), + ColumnarValue::Array(column()), + ColumnarValue::Array(allow), + ]) + .unwrap(); + let ColumnarValue::Array(arr) = out else { + panic!() + }; + let arr = arr.as_any().downcast_ref::().unwrap(); + assert!( + arr.is_null(0), + "a NULL allowDifferentLgConfigK must produce NULL" + ); + assert!( + !arr.is_null(1), + "a non-NULL flag must still produce a sketch" + ); + } } diff --git a/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchEstimate.scala b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchEstimate.scala index 3abf67fa857..a36121c2b24 100644 --- a/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchEstimate.scala +++ b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchEstimate.scala @@ -28,7 +28,12 @@ object CometHllSketchEstimate extends CometExpressionSerde[HllSketchEstimate] { private val incompatReason = "Comet uses a Rust DataSketches port; HLL estimates may differ slightly from Spark." - override def getIncompatibleReasons(): Seq[String] = Seq(incompatReason) + private val errorReason = + "Errors surface as plain Comet execution errors rather than Spark's SparkRuntimeException " + + "with condition HLL_INVALID_INPUT_SKETCH_BUFFER (sqlState 22000), so the message and " + + "error class differ even though both engines fail." + + override def getIncompatibleReasons(): Seq[String] = Seq(incompatReason, errorReason) override def getSupportLevel(expr: HllSketchEstimate): SupportLevel = Incompatible(Some(incompatReason)) diff --git a/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllUnion.scala b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllUnion.scala index d63472e605d..1c164177bc0 100644 --- a/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllUnion.scala +++ b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllUnion.scala @@ -31,7 +31,12 @@ object CometHllUnion extends CometExpressionSerde[HllUnion] { private val incompatReason = "Comet uses a Rust DataSketches port; HLL sketch bytes and estimates may differ slightly from Spark." - override def getIncompatibleReasons(): Seq[String] = Seq(incompatReason) + private val errorReason = + "Errors surface as plain Comet execution errors rather than Spark's SparkRuntimeException " + + "with condition HLL_UNION_DIFFERENT_LG_K / HLL_INVALID_INPUT_SKETCH_BUFFER (sqlState " + + "22000), so the message and error class differ even though both engines fail." + + override def getIncompatibleReasons(): Seq[String] = Seq(incompatReason, errorReason) override def getSupportLevel(expr: HllUnion): SupportLevel = { if (!expr.third.foldable) { diff --git a/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllUnionAgg.scala b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllUnionAgg.scala index 773342d2d0f..8ea080ee6a0 100644 --- a/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllUnionAgg.scala +++ b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllUnionAgg.scala @@ -35,7 +35,12 @@ object CometHllUnionAgg extends CometAggregateExpressionSerde[HllUnionAgg] { "Comet uses a Rust DataSketches port; HLL sketch bytes and estimates may differ slightly " + "from Spark." - override def getIncompatibleReasons(): Seq[String] = Seq(incompatReason) + private val errorReason = + "Errors surface as plain Comet execution errors rather than Spark's SparkRuntimeException " + + "with condition HLL_UNION_DIFFERENT_LG_K / HLL_INVALID_INPUT_SKETCH_BUFFER (sqlState " + + "22000), so the message and error class differ even though both engines fail." + + override def getIncompatibleReasons(): Seq[String] = Seq(incompatReason, errorReason) override def getSupportLevel(expr: HllUnionAgg): SupportLevel = { if (!expr.right.foldable) { diff --git a/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala b/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala index 4df5de9ba60..55e474735a1 100644 --- a/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala @@ -3617,6 +3617,13 @@ class CometExpressionSuite extends CometTestBase with AdaptiveSparkPlanHelper { val (sparkErr, cometErr) = checkSparkAnswerMaybeThrows(df) assert(sparkErr.isDefined, "expected Spark to throw on different lgConfigK") assert(cometErr.isDefined, "expected Comet to throw on different lgConfigK") + // Both arms raising is not enough: if Comet had fallen back for the whole plan, the + // second run would raise Spark's exception too and this test would pass without the + // native check ever running. The two messages differ - Spark raises + // HLL_UNION_DIFFERENT_LG_K - so asserting on Comet's own wording pins the native path. + assert( + cometErr.get.getMessage.contains("to enable unions of different lgConfigK"), + s"expected Comet's native lgConfigK error, got: ${cometErr.get.getMessage}") } } } From 406cc9d905968435ec751a1bcaa698e1bce51238 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Wed, 9 Sep 2026 12:52:53 -0600 Subject: [PATCH 15/17] fix: document HLL support and address review notes Remaining review feedback on the HLL sketch functions. `docs/source/user-guide/latest/expressions.md` still listed the `hll_*` family under "Not currently planned", so the page contradicted the feature. Narrow that bullet to the families that stay unplanned, add the four rows to the `agg_funcs` and `misc_funcs` tables, and explain why this is the one sketch family Comet accelerates. The Implementation cells were checked against a `generate-docs` run on spark-4.0 and spark-4.1 rather than filled in by hand. Minor notes from the same review: - Drop the dead `update_i32`/`update_i16`/`update_i8` helpers. The accumulator widens with `as i64`, which sign-extends identically. - Downcast once per batch in `HllSketchAccumulator::update_batch` instead of building a `ScalarValue` per row, which copied every string and binary value onto the heap only to hash it and drop it. `every_input_type_hashes_the_same_as_a_direct_update` pins the new dispatch to the old bytes exactly, per input type. - Invalid sketch bytes are user data, not a Comet invariant, so `from_bytes` reports `Execution` rather than `Internal`. This is also what the "surfaces as a plain Comet execution error" incompatibility note already promised. - Complete `getUnsupportedReasons()`: the lgConfigK range case was missing from `CometHllSketchAgg`, and `CometHllUnion` / `CometHllUnionAgg` documented no unsupported cases at all despite both returning `Unsupported` for a non-foldable flag. - Pin `datasketches` to `=0.3.0`. The compact-decode workaround and the HLL_4 aux guard are both written against that release's array code, so the version should not move without re-checking those two tests. Also record the HLL_4-with-aux-entries rejection as an incompatibility on the three serdes that read sketch bytes: Comet errors there where Spark returns an estimate, which belongs in front of anyone opting in. Add the end-to-end test comphead asked for on the NULL `allowDifferentLgConfigK` fix. Spark marks `HllUnion` null-intolerant, so `NullPropagation` folds a foldable NULL flag away before Comet sees it; the test excludes that rule so the native kernel is what answers, and asserts `hll_union` survives into the plan so it cannot pass on a folded or fallen-back plan. --- docs/source/user-guide/latest/expressions.md | 8 +- native/Cargo.toml | 7 +- native/spark-expr/src/agg_funcs/hll_sketch.rs | 22 +-- .../src/agg_funcs/hll_sketch_agg.rs | 161 +++++++++++++----- native/spark-expr/src/hll_scalar.rs | 13 +- .../comet/serde/CometHllSketchAgg.scala | 6 +- .../comet/serde/CometHllSketchEstimate.scala | 9 +- .../apache/comet/serde/CometHllUnion.scala | 16 +- .../apache/comet/serde/CometHllUnionAgg.scala | 16 +- .../apache/comet/CometExpressionSuite.scala | 43 +++++ 10 files changed, 231 insertions(+), 70 deletions(-) diff --git a/docs/source/user-guide/latest/expressions.md b/docs/source/user-guide/latest/expressions.md index 2cf0894ceb8..0da9ac5d6de 100644 --- a/docs/source/user-guide/latest/expressions.md +++ b/docs/source/user-guide/latest/expressions.md @@ -64,7 +64,7 @@ The Implementation column is auto-generated from the serde definitions in `Query Comet focuses acceleration on mainstream relational, string, datetime, math, and collection expressions. The following function families are **not currently planned** for native acceleration (they are not on the 1.0 roadmap): specialized functionality with narrow real-world analytics use and high implementation cost. They fall back to Spark and may be reconsidered based on demand: -- **Probabilistic sketches and approximate top-k** (`kll_sketch_*`, `hll_*`, `theta_*`, `count_min_sketch`, `bitmap_*`, `approx_top_k*`): specialized data structures with exact-correctness traps. +- **Probabilistic sketches and approximate top-k** (`kll_sketch_*`, `theta_*`, `count_min_sketch`, `bitmap_*`, `approx_top_k*`): specialized data structures with exact-correctness traps. - **Geospatial** (`st_*`): brand-new Spark 4.1 functionality, specialized. - **Avro / Protobuf codecs** (`from_avro`, `to_avro`, `from_protobuf`, `to_protobuf`, `schema_of_avro`): format conversion belongs at the IO layer, not expression evaluation. - **JVM reflection** (`java_method`, `reflect`): niche, and they invoke arbitrary JVM methods (a security concern). @@ -75,6 +75,8 @@ The file-metadata functions `input_file_name`, `input_file_block_start`, and `in Note that `median` and `mode` are planned: they are mainstream exact aggregates. `approx_count_distinct` is supported because Comet ports Spark's `HyperLogLogPlusPlus` exactly, so its result is bit-identical to Spark. +The Apache DataSketches HLL functions (`hll_sketch_agg`, `hll_union_agg`, `hll_sketch_estimate`, `hll_union`) are the one sketch family Comet does accelerate, on Spark 4.0+. The sketches are mutually readable with Spark's, but the point estimate can differ slightly after a merge, so the native path is off by default and opt-in per expression via `allowIncompatible` — see the rows below. + The tables below list every Spark built-in expression with its current status. ## agg_funcs @@ -104,6 +106,8 @@ The tables below list every Spark built-in expression with its current status. | `first_value` | ✅ | Native | | | `grouping` | ✅ | — | Grouping indicator for ROLLUP/CUBE/GROUPING SETS | | `grouping_id` | ✅ | — | Grouping indicator for ROLLUP/CUBE/GROUPING SETS | +| `hll_sketch_agg` | ✅ | Native | Spark 4.0+ only; falls back by default, the native path is opt-in via allowIncompatible ([details](compatibility/expressions/aggregate.md)) | +| `hll_union_agg` | ✅ | Native | Spark 4.0+ only; falls back by default, the native path is opt-in via allowIncompatible ([details](compatibility/expressions/aggregate.md)) | | `kurtosis` | 🔜 | — | Not yet implemented natively | | `last` | ✅ | Native | | | `last_value` | ✅ | Native | | @@ -498,6 +502,8 @@ The type-name conversion functions (`bigint`, `binary`, `boolean`, `date`, `deci | `current_schema` | ✅ | — | Alias of `current_database`; resolved to a literal by the analyzer | | `current_user` | ✅ | — | Resolved to a literal by the analyzer; same as `user` | | `equal_null` | ✅ | — | Lowers to `<=>` (`EqualNullSafe`) | +| `hll_sketch_estimate` | ✅ | Native | Spark 4.0+ only; falls back by default, the native path is opt-in via allowIncompatible ([details](compatibility/expressions/misc.md)) | +| `hll_union` | ✅ | Native | Spark 4.0+ only; falls back by default, the native path is opt-in via allowIncompatible ([details](compatibility/expressions/misc.md)) | | `is_variant_null` | 🔜 | — | Requires `VariantType` support | | `monotonically_increasing_id` | ✅ | Native | | | `parse_json` | 🔜 | — | Requires `VariantType` support | diff --git a/native/Cargo.toml b/native/Cargo.toml index 62a94237dd6..0b13b6c3978 100644 --- a/native/Cargo.toml +++ b/native/Cargo.toml @@ -53,7 +53,12 @@ datafusion-comet-jni-bridge = { path = "jni-bridge" } datafusion-comet-proto = { path = "proto" } datafusion-comet-shuffle = { path = "shuffle" } chrono = { version = "0.4", default-features = false, features = ["clock"] } -datasketches = { version = "0.3.0", features = ["hll"] } +# Pinned exactly rather than left to the default caret range. `hll_sketch.rs` carries a +# workaround for a compact-decoding bug in 0.3.0 and a guard against a mis-read aux layout, +# both written against that release's `Array4`/`Array6`/`Array8` code. A silent bump to a +# future 0.3.x could fix or move either one, so the version moves only with a deliberate +# change that re-checks those two tests. +datasketches = { version = "=0.3.0", features = ["hll"] } futures = "0.3.32" num = "0.4" rand = "0.10" diff --git a/native/spark-expr/src/agg_funcs/hll_sketch.rs b/native/spark-expr/src/agg_funcs/hll_sketch.rs index 58484ce2e42..433e8739344 100644 --- a/native/spark-expr/src/agg_funcs/hll_sketch.rs +++ b/native/spark-expr/src/agg_funcs/hll_sketch.rs @@ -34,7 +34,7 @@ //! ever read back by Comet. use datafusion::error::DataFusionError; -use datasketches::hash_value::{raw_bytes, sign_extend}; +use datasketches::hash_value::raw_bytes; use datasketches::hll::{HllSketch, HllType, HllUnion}; /// A DataSketches HLL_8 sketch configured to match Spark's `HllSketchAgg`. @@ -155,25 +155,13 @@ impl SparkHllSketch { } /// Update with a 64-bit integer. Spark widens narrower integrals to `long` - /// before hashing; callers should pass the already-widened value here. - /// Rust's `Hash` for `i64` writes 8 little-endian bytes with no prefix, - /// matching DataSketches-Java `update(long)`. + /// before hashing, so callers pass the already-widened value; Rust's `as i64` + /// sign-extends, matching Spark's `toLong`. Rust's `Hash` for `i64` writes 8 + /// little-endian bytes with no prefix, matching DataSketches-Java `update(long)`. pub fn update_i64(&mut self, v: i64) { self.inner.update(v); } - /// Update with a narrow signed integer, sign-extending to 64 bits exactly as - /// Spark's `toLong` does before hashing. - pub fn update_i32(&mut self, v: i32) { - self.inner.update(sign_extend::from_i32(v)); - } - pub fn update_i16(&mut self, v: i16) { - self.inner.update(sign_extend::from_i16(v)); - } - pub fn update_i8(&mut self, v: i8) { - self.inner.update(sign_extend::from_i8(v)); - } - /// Update with raw bytes (used for both StringType UTF-8 bytes and /// BinaryType), hashing without Rust's slice length prefix. Empty inputs are /// skipped, matching DataSketches (and Spark), which ignore empty values. @@ -196,7 +184,7 @@ impl SparkHllSketch { let normalized = normalize_compact_hll_array(bytes); HllSketch::deserialize(normalized.as_deref().unwrap_or(bytes)) .map(|inner| Self { inner }) - .map_err(|e| DataFusionError::Internal(format!("invalid HLL sketch bytes: {e}"))) + .map_err(|e| DataFusionError::Execution(format!("invalid HLL sketch bytes: {e}"))) } /// The configured `lgConfigK`. diff --git a/native/spark-expr/src/agg_funcs/hll_sketch_agg.rs b/native/spark-expr/src/agg_funcs/hll_sketch_agg.rs index 3453b646573..c0e716042b3 100644 --- a/native/spark-expr/src/agg_funcs/hll_sketch_agg.rs +++ b/native/spark-expr/src/agg_funcs/hll_sketch_agg.rs @@ -19,7 +19,11 @@ use crate::agg_funcs::hll_sketch::SparkHllSketch; use arrow::array::Array; use arrow::array::ArrayRef; use arrow::array::BinaryArray; -use arrow::datatypes::{DataType, Field, FieldRef}; +use arrow::array::{as_primitive_array, GenericByteArray, PrimitiveArray, StringArray}; +use arrow::datatypes::{ + ArrowPrimitiveType, ByteArrayType, DataType, Field, FieldRef, Int16Type, Int32Type, Int64Type, + Int8Type, +}; use datafusion::common::{downcast_value, ScalarValue}; use datafusion::error::{DataFusionError, Result}; use datafusion::logical_expr::function::{AccumulatorArgs, StateFieldsArgs}; @@ -85,6 +89,34 @@ impl HllSketchAccumulator { sketch: SparkHllSketch::new(lg_config_k), } } + + /// Spark widens every accepted integral to `long` before hashing, so all four widths + /// funnel through the same `i64` update. Nulls are ignored, matching `HllSketchAgg`. + fn update_ints(&mut self, arr: &PrimitiveArray) + where + T: ArrowPrimitiveType, + T::Native: Into, + { + for i in 0..arr.len() { + if !arr.is_null(i) { + self.sketch.update_i64(arr.value(i).into()); + } + } + } + + /// StringType hashes its UTF-8 bytes and BinaryType its bytes directly, so both share + /// this loop. + fn update_byte_slices(&mut self, arr: &GenericByteArray) + where + T: ByteArrayType, + for<'a> &'a T::Native: AsRef<[u8]>, + { + for i in 0..arr.len() { + if !arr.is_null(i) { + self.sketch.update_bytes(arr.value(i).as_ref()); + } + } + } } impl Accumulator for HllSketchAccumulator { @@ -93,41 +125,23 @@ impl Accumulator for HllSketchAccumulator { return Ok(()); } let arr = &values[0]; - (0..arr.len()).try_for_each(|i| { - match ScalarValue::try_from_array(arr, i)? { - ScalarValue::Int8(Some(v)) => { - self.sketch.update_i64(v as i64); - } - ScalarValue::Int16(Some(v)) => { - self.sketch.update_i64(v as i64); - } - ScalarValue::Int32(Some(v)) => { - self.sketch.update_i64(v as i64); - } - ScalarValue::Int64(Some(v)) => { - self.sketch.update_i64(v); - } - ScalarValue::Utf8(Some(v)) => { - self.sketch.update_bytes(v.as_bytes()); - } - ScalarValue::Binary(Some(v)) => { - self.sketch.update_bytes(&v); - } - // Spark's HllSketchAgg ignores null inputs. - ScalarValue::Int8(None) - | ScalarValue::Int16(None) - | ScalarValue::Int32(None) - | ScalarValue::Int64(None) - | ScalarValue::Utf8(None) - | ScalarValue::Binary(None) => {} - other => { - return Err(DataFusionError::Internal(format!( - "hll_sketch_agg received an unsupported input type: {other:?}" - ))) - } + // Downcast once per batch rather than going through `ScalarValue::try_from_array` per + // row: for the string and binary cases that copies every value onto the heap only to + // hash it and drop it again. + match arr.data_type() { + DataType::Int8 => self.update_ints(as_primitive_array::(arr)), + DataType::Int16 => self.update_ints(as_primitive_array::(arr)), + DataType::Int32 => self.update_ints(as_primitive_array::(arr)), + DataType::Int64 => self.update_ints(as_primitive_array::(arr)), + DataType::Utf8 => self.update_byte_slices(downcast_value!(arr, StringArray)), + DataType::Binary => self.update_byte_slices(downcast_value!(arr, BinaryArray)), + other => { + return Err(DataFusionError::Internal(format!( + "hll_sketch_agg received an unsupported input type: {other:?}" + ))) } - Ok(()) - }) + } + Ok(()) } fn evaluate(&mut self) -> Result { @@ -165,22 +179,91 @@ impl Accumulator for HllSketchAccumulator { #[cfg(test)] mod tests { use super::*; - use arrow::array::Int64Array; + use arrow::array::{Int32Array, Int64Array, Int8Array}; use datafusion::physical_plan::Accumulator; use std::sync::Arc; + fn sketch_bytes(acc: &mut HllSketchAccumulator) -> Vec { + let ScalarValue::Binary(Some(bytes)) = acc.evaluate().unwrap() else { + panic!("expected binary") + }; + bytes + } + #[test] fn accumulates_and_estimates() { let mut acc = HllSketchAccumulator::new(12); let arr = Arc::new(Int64Array::from((0..1000i64).collect::>())); acc.update_batch(&[arr]).unwrap(); - let ScalarValue::Binary(Some(bytes)) = acc.evaluate().unwrap() else { - panic!("expected binary") - }; + let bytes = sketch_bytes(&mut acc); let est = crate::agg_funcs::estimate_from_bytes(&bytes).unwrap(); assert!((est - 1000).abs() <= 30, "estimate {est}"); } + /// `update_batch` downcasts the whole array once instead of building a `ScalarValue` per + /// row. Every accepted input type has to keep hashing exactly as before, so compare the + /// accumulator's bytes against a sketch fed the same values directly - byte equality, not + /// an error bound, since any change in the hashed bytes would move registers. + #[test] + fn every_input_type_hashes_the_same_as_a_direct_update() { + // Narrow integrals are widened to i64 (sign-extending), so negatives matter here. + let mut acc = HllSketchAccumulator::new(12); + acc.update_batch(&[Arc::new(Int8Array::from(vec![ + Some(-128), + None, + Some(0), + Some(127), + ]))]) + .unwrap(); + let mut direct = SparkHllSketch::new(12); + for v in [-128i64, 0, 127] { + direct.update_i64(v); + } + assert_eq!(sketch_bytes(&mut acc), direct.to_sketch_bytes()); + + let mut acc = HllSketchAccumulator::new(12); + acc.update_batch(&[Arc::new(Int32Array::from(vec![ + Some(i32::MIN), + None, + Some(i32::MAX), + ]))]) + .unwrap(); + let mut direct = SparkHllSketch::new(12); + for v in [i32::MIN as i64, i32::MAX as i64] { + direct.update_i64(v); + } + assert_eq!(sketch_bytes(&mut acc), direct.to_sketch_bytes()); + + // Strings hash their UTF-8 bytes; the empty string is skipped by both paths. + let mut acc = HllSketchAccumulator::new(12); + acc.update_batch(&[Arc::new(StringArray::from(vec![ + Some("a"), + None, + Some(""), + Some("héllo"), + ]))]) + .unwrap(); + let mut direct = SparkHllSketch::new(12); + for v in ["a", "", "héllo"] { + direct.update_bytes(v.as_bytes()); + } + assert_eq!(sketch_bytes(&mut acc), direct.to_sketch_bytes()); + + let mut acc = HllSketchAccumulator::new(12); + acc.update_batch(&[Arc::new(BinaryArray::from(vec![ + Some(&b"\x00\xff"[..]), + None, + Some(&b""[..]), + Some(&b"xyz"[..]), + ]))]) + .unwrap(); + let mut direct = SparkHllSketch::new(12); + for v in [&b"\x00\xff"[..], &b""[..], &b"xyz"[..]] { + direct.update_bytes(v); + } + assert_eq!(sketch_bytes(&mut acc), direct.to_sketch_bytes()); + } + /// Spark's `HllSketchAgg` is non-nullable: an empty/all-null group still /// produces a serialized empty sketch (estimate 0), not NULL. #[test] diff --git a/native/spark-expr/src/hll_scalar.rs b/native/spark-expr/src/hll_scalar.rs index b1f3e77167f..51892896e65 100644 --- a/native/spark-expr/src/hll_scalar.rs +++ b/native/spark-expr/src/hll_scalar.rs @@ -51,9 +51,12 @@ pub fn spark_hll_union(args: &[ColumnarValue]) -> Result { for i in 0..a.len() { // Spark's `HllUnion` is a `TernaryExpression` evaluated through `nullSafeEval`, and // `TernaryExpression.eval` returns NULL when *any* of the three inputs is NULL - the - // `allowDifferentLgConfigK` flag included. `Literal(null, BooleanType)` is foldable, so - // the serde's `third.foldable` check lets it through and this loop is the only thing - // standing between a NULL flag and a non-null sketch. + // `allowDifferentLgConfigK` flag included. Spark also marks `HllUnion` null-intolerant, + // so with the default optimizer `NullPropagation` folds a foldable NULL flag away before + // Comet sees the expression; a user who excludes that rule reaches this loop instead, and + // it has to give the same answer. Note the flag is *not* coerced to false here: that + // coercion is what makes `HllUnionAgg` correct (Spark's `null.asInstanceOf[Boolean]`), + // and it does not apply to a `TernaryExpression`, where the NULL has to propagate. if a.is_null(i) || b.is_null(i) || allow.is_null(i) { out.append_null(); continue; @@ -139,8 +142,8 @@ mod union_tests { } /// Spark's `HllUnion` is a `TernaryExpression`, so a NULL `allowDifferentLgConfigK` makes the - /// whole call NULL. `Literal(null, BooleanType)` is foldable and therefore passes the serde's - /// `third.foldable` gate, so this loop is what has to honour it. + /// whole call NULL. `NullPropagation` normally removes that shape before Comet sees it, so + /// this pins the kernel for the case where the rule is excluded. #[test] fn a_null_allow_flag_yields_null() { let mut a = SparkHllSketch::new(12); diff --git a/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchAgg.scala b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchAgg.scala index d56a963c42c..076d4ec1e8d 100644 --- a/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchAgg.scala +++ b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchAgg.scala @@ -39,6 +39,8 @@ object CometHllSketchAgg extends CometAggregateExpressionSerde[HllSketchAgg] { private val nonLiteralLgConfigKReason = "The lgConfigK argument must be a foldable literal." + private val lgConfigKRangeReason = + s"The lgConfigK argument must be in the range [$MinLgConfigK, $MaxLgConfigK]." private val inputTypeReason = "Only int, long, string, and binary input types are supported." private val incompatReason = @@ -46,7 +48,7 @@ object CometHllSketchAgg extends CometAggregateExpressionSerde[HllSketchAgg] { "slightly from Spark." override def getUnsupportedReasons(): Seq[String] = - Seq(nonLiteralLgConfigKReason, inputTypeReason) + Seq(nonLiteralLgConfigKReason, lgConfigKRangeReason, inputTypeReason) override def getIncompatibleReasons(): Seq[String] = Seq(incompatReason) @@ -60,7 +62,7 @@ object CometHllSketchAgg extends CometAggregateExpressionSerde[HllSketchAgg] { case _ => return Unsupported(Some(nonLiteralLgConfigKReason)) } if (lgConfigK < MinLgConfigK || lgConfigK > MaxLgConfigK) { - return Unsupported(Some(s"lgConfigK must be in [$MinLgConfigK, $MaxLgConfigK]")) + return Unsupported(Some(lgConfigKRangeReason)) } expr.left.dataType match { case IntegerType | LongType | StringType | BinaryType => diff --git a/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchEstimate.scala b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchEstimate.scala index a36121c2b24..f5743bd1194 100644 --- a/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchEstimate.scala +++ b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllSketchEstimate.scala @@ -33,7 +33,14 @@ object CometHllSketchEstimate extends CometExpressionSerde[HllSketchEstimate] { "with condition HLL_INVALID_INPUT_SKETCH_BUFFER (sqlState 22000), so the message and " + "error class differ even though both engines fail." - override def getIncompatibleReasons(): Seq[String] = Seq(incompatReason, errorReason) + private val hll4AuxReason = + "An input sketch in the updatable HLL_4 form carrying auxiliary-map entries is rejected " + + "with an error, where Spark reads it: the bundled Rust decoder reads the compact " + + "auxiliary layout in both forms and would otherwise return a silently wrong estimate. " + + "Comet only ever writes HLL_8, so this affects sketch columns produced elsewhere." + + override def getIncompatibleReasons(): Seq[String] = + Seq(incompatReason, errorReason, hll4AuxReason) override def getSupportLevel(expr: HllSketchEstimate): SupportLevel = Incompatible(Some(incompatReason)) diff --git a/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllUnion.scala b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllUnion.scala index 1c164177bc0..ffe41ea2911 100644 --- a/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllUnion.scala +++ b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllUnion.scala @@ -31,16 +31,28 @@ object CometHllUnion extends CometExpressionSerde[HllUnion] { private val incompatReason = "Comet uses a Rust DataSketches port; HLL sketch bytes and estimates may differ slightly from Spark." + private val nonLiteralAllowReason = + "The allowDifferentLgConfigK argument must be a foldable literal." + private val errorReason = "Errors surface as plain Comet execution errors rather than Spark's SparkRuntimeException " + "with condition HLL_UNION_DIFFERENT_LG_K / HLL_INVALID_INPUT_SKETCH_BUFFER (sqlState " + "22000), so the message and error class differ even though both engines fail." - override def getIncompatibleReasons(): Seq[String] = Seq(incompatReason, errorReason) + private val hll4AuxReason = + "An input sketch in the updatable HLL_4 form carrying auxiliary-map entries is rejected " + + "with an error, where Spark reads it: the bundled Rust decoder reads the compact " + + "auxiliary layout in both forms and would otherwise return a silently wrong estimate. " + + "Comet only ever writes HLL_8, so this affects sketch columns produced elsewhere." + + override def getUnsupportedReasons(): Seq[String] = Seq(nonLiteralAllowReason) + + override def getIncompatibleReasons(): Seq[String] = + Seq(incompatReason, errorReason, hll4AuxReason) override def getSupportLevel(expr: HllUnion): SupportLevel = { if (!expr.third.foldable) { - Unsupported(Some("The allowDifferentLgConfigK argument must be a foldable literal.")) + Unsupported(Some(nonLiteralAllowReason)) } else { Incompatible(Some(incompatReason)) } diff --git a/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllUnionAgg.scala b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllUnionAgg.scala index 8ea080ee6a0..6029b67c16c 100644 --- a/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllUnionAgg.scala +++ b/spark/src/main/spark-4.x/org/apache/comet/serde/CometHllUnionAgg.scala @@ -35,16 +35,28 @@ object CometHllUnionAgg extends CometAggregateExpressionSerde[HllUnionAgg] { "Comet uses a Rust DataSketches port; HLL sketch bytes and estimates may differ slightly " + "from Spark." + private val nonLiteralAllowReason = + "The allowDifferentLgConfigK argument must be a foldable literal." + private val errorReason = "Errors surface as plain Comet execution errors rather than Spark's SparkRuntimeException " + "with condition HLL_UNION_DIFFERENT_LG_K / HLL_INVALID_INPUT_SKETCH_BUFFER (sqlState " + "22000), so the message and error class differ even though both engines fail." - override def getIncompatibleReasons(): Seq[String] = Seq(incompatReason, errorReason) + private val hll4AuxReason = + "An input sketch in the updatable HLL_4 form carrying auxiliary-map entries is rejected " + + "with an error, where Spark reads it: the bundled Rust decoder reads the compact " + + "auxiliary layout in both forms and would otherwise return a silently wrong estimate. " + + "Comet only ever writes HLL_8, so this affects sketch columns produced elsewhere." + + override def getUnsupportedReasons(): Seq[String] = Seq(nonLiteralAllowReason) + + override def getIncompatibleReasons(): Seq[String] = + Seq(incompatReason, errorReason, hll4AuxReason) override def getSupportLevel(expr: HllUnionAgg): SupportLevel = { if (!expr.right.foldable) { - Unsupported(Some("The allowDifferentLgConfigK argument must be a foldable literal.")) + Unsupported(Some(nonLiteralAllowReason)) } else { Incompatible(Some(incompatReason)) } diff --git a/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala b/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala index 55e474735a1..00830afe49a 100644 --- a/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala @@ -3628,6 +3628,49 @@ class CometExpressionSuite extends CometTestBase with AdaptiveSparkPlanHelper { } } + test("hll_union with a NULL allowDifferentLgConfigK returns NULL") { + assume(isSpark40Plus) + // HllUnion is a TernaryExpression evaluated through nullSafeEval, so a NULL in *any* of the + // three arguments - the allowDifferentLgConfigK flag included - makes the whole call NULL. + // Returning a sketch would be a categorically wrong answer rather than an approximation, + // which is not something the Incompatible opt-in covers. + // + // The serde only accepts a foldable third argument, and Spark marks HllUnion + // `nullIntolerant`, so with the default optimizer NullPropagation rewrites a foldable NULL + // flag to a NULL literal before Comet ever sees the expression. Excluding that rule is what + // makes the native kernel responsible for the NULL, which is the behaviour under test; a + // user who excludes NullPropagation must still get Spark's answer. + // + // (HllUnionAgg needs no equivalent guard: its `convert` falls back when `right.eval()` is + // not a Boolean, and Spark's `null.asInstanceOf[Boolean]` coerces to false, so the two agree + // whichever way the flag arrives.) + withSQLConf( + "spark.sql.optimizer.excludedRules" -> + "org.apache.spark.sql.catalyst.optimizer.NullPropagation", + "spark.comet.expression.HllSketchAgg.allowIncompatible" -> "true", + "spark.comet.expression.HllSketchEstimate.allowIncompatible" -> "true", + "spark.comet.expression.HllUnion.allowIncompatible" -> "true") { + withParquetTable((0 until 300).map(i => (i % 2, i)), "tbl") { + val query = + "SELECT hll_sketch_estimate(hll_union(a.s, b.s, cast(null as boolean))) FROM " + + "(SELECT hll_sketch_agg(_2) AS s FROM tbl WHERE _1 = 0) a, " + + "(SELECT hll_sketch_agg(_2) AS s FROM tbl WHERE _1 = 1) b" + val df = sql(query) + val plan = stripAQEPlan(df.queryExecution.executedPlan) + // Without these the test would pass on a plan where hll_union was folded away or fell + // back to Spark, neither of which exercises the native null check. + checkCometOperators(plan) + assert( + plan.toString().contains("hll_union"), + s"expected hll_union to survive into the native plan, got:\n$plan") + assert( + df.collect().head.isNullAt(0), + "a NULL allowDifferentLgConfigK must make hll_union return NULL") + checkSparkAnswer(query) + } + } + } + test("hll_sketch_agg over all-null input estimates to 0, not NULL") { assume(isSpark40Plus) // Spark's HllSketchAgg/HllSketchEstimate are declared non-nullable: an empty or From e70c3d4d845d289c472900d21be028be05de0e19 Mon Sep 17 00:00:00 2001 From: test Date: Wed, 9 Sep 2026 21:16:06 -0600 Subject: [PATCH 16/17] test: keep the HLL lgConfigK rejection test native on Spark 4.2 Spark 4.2 replaces MergeScalarSubqueries with MergeSubplans, which also merges non-grouping Aggregate nodes. That folded the two hll_sketch_agg branches of the test's UNION ALL into a single CTE projecting a struct with two fields both named 's'. Comet does not support CreateNamedStruct with duplicate field names, so the whole plan - hll_union_agg included - fell back to Spark and the test saw Spark's HLL_UNION_DIFFERENT_LG_K instead of Comet's native error. Materialize the lgConfigK=10 and lgConfigK=12 sketches into Parquet first and run the union over a plain scan, a shape that stays fully native on 3.4 through 4.2, and assert that with checkCometOperators. --- .../apache/comet/CometExpressionSuite.scala | 27 ++++++++++++------- 1 file changed, 18 insertions(+), 9 deletions(-) diff --git a/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala b/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala index 00830afe49a..a07a71553d8 100644 --- a/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala @@ -3604,16 +3604,25 @@ class CometExpressionSuite extends CometTestBase with AdaptiveSparkPlanHelper { test("hll_union_agg rejects different lgConfigK when not allowed") { assume(isSpark40Plus) - withSQLConf( - "spark.comet.expression.HllSketchAgg.allowIncompatible" -> "true", - "spark.comet.expression.HllUnionAgg.allowIncompatible" -> "true") { + withTempPath { dir => + val sketchPath = dir.getCanonicalPath + // Materialize one lgConfigK=10 and one lgConfigK=12 sketch into Parquet, written by Spark + // (the HllSketchAgg opt-in is deliberately absent here), so the query under test is a plain + // scan plus aggregate. Building the two sketches inline with UNION ALL instead is not + // version-stable: on Spark 4.2 MergeSubplans folds the two non-grouping aggregates into one + // CTE projecting a struct with duplicate field names, which Comet does not accelerate, so + // hll_union_agg falls back and the native check under test never runs. withParquetTable((0 until 100).map(i => Tuple1(i)), "tbl") { - // A lgConfigK=10 sketch unioned with a lgConfigK=12 sketch (allowDifferentLgConfigK - // defaults false) must throw in BOTH Spark and Comet. - val df = sql( - "SELECT hll_union_agg(s) FROM (" + - " SELECT hll_sketch_agg(_1, 10) AS s FROM tbl UNION ALL" + - " SELECT hll_sketch_agg(_1, 12) AS s FROM tbl)") + sql("SELECT hll_sketch_agg(_1, 10) AS s FROM tbl") + .union(sql("SELECT hll_sketch_agg(_1, 12) AS s FROM tbl")) + .write + .parquet(sketchPath) + } + withSQLConf("spark.comet.expression.HllUnionAgg.allowIncompatible" -> "true") { + // Unioning the two sketches (allowDifferentLgConfigK defaults false) must throw in BOTH + // Spark and Comet. + val df = spark.read.parquet(sketchPath).selectExpr("hll_union_agg(s)") + checkCometOperators(stripAQEPlan(df.queryExecution.executedPlan)) val (sparkErr, cometErr) = checkSparkAnswerMaybeThrows(df) assert(sparkErr.isDefined, "expected Spark to throw on different lgConfigK") assert(cometErr.isDefined, "expected Comet to throw on different lgConfigK") From 32b602559a3956737cd6da02e76e33643eea1fb2 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Fri, 25 Sep 2026 07:53:41 -0600 Subject: [PATCH 17/17] fix: charge HLL accumulators for the sketch they hold, not the dense maximum `size()` on both HLL accumulators charged the full 2^lgConfigK register array as soon as a group existed. A sketch stays in an 8-slot coupon list or a small coupon hash set until it has seen enough distinct values, so at lgConfigK 21 each singleton group was charged 2 MiB for 32 bytes. 64 such groups exhausted a 16 MiB pool in the Final aggregation even with spilling, where Spark runs the query fine. The wrapper now reports the representation the sketch actually holds. The datasketches crate keeps the mode private, but its estimate is never below the coupon count in LIST and SET mode and seeds the HIP accumulator on promotion, and a union with a dense input is tracked from that input's preamble. A test walks every SET resize and the promotion and compares against the layout the crate serializes: never below it, and exact at all but a handful of points. --- native/spark-expr/src/agg_funcs/hll_sketch.rs | 237 +++++++++++++++++- .../src/agg_funcs/hll_sketch_agg.rs | 40 ++- .../spark-expr/src/agg_funcs/hll_union_agg.rs | 52 +++- 3 files changed, 309 insertions(+), 20 deletions(-) diff --git a/native/spark-expr/src/agg_funcs/hll_sketch.rs b/native/spark-expr/src/agg_funcs/hll_sketch.rs index 433e8739344..ebb52156cd3 100644 --- a/native/spark-expr/src/agg_funcs/hll_sketch.rs +++ b/native/spark-expr/src/agg_funcs/hll_sketch.rs @@ -29,9 +29,9 @@ //! *compact* form, whereas Spark emits the *updatable* form. The bytes are //! therefore not byte-identical to Spark's output for small inputs, but //! DataSketches `deserialize` reads both forms, so estimates round-trip in both -//! directions. Comet must own both Partial and Final aggregation -//! (`supportsMixedPartialFinal = false`) so this compact intermediate is only -//! ever read back by Comet. +//! directions. Comet must own both Partial and Final aggregation (the HLL serdes leave +//! `supportsSparkPartialToNativeFinal` and `supportsNativePartialToSparkFinal` false) so this +//! compact intermediate is only ever read back by Comet. use datafusion::error::DataFusionError; use datasketches::hash_value::raw_bytes; @@ -41,6 +41,75 @@ use datasketches::hll::{HllSketch, HllType, HllUnion}; #[derive(Debug)] pub struct SparkHllSketch { inner: HllSketch, + /// Whether `inner` is known to hold the dense register array; see `layout`. Once set it + /// stays set, because a union's result estimates from its registers, and near the promotion + /// point that estimate can fall back below `layout::max_coupons`. + dense: bool, +} + +/// The in-memory layout `datasketches` 0.3.0 gives an HLL_8 sketch, which the crate keeps +/// private (`HllSketch::mode` is `pub(super)`). Accumulator `size()` reports it: most sketches +/// stay small, and charging every group the dense register array up front makes grouped +/// high-`lgConfigK` queries exhaust the memory pool. +/// +/// A sketch starts in LIST mode, an 8-slot coupon array. When that fills it moves to SET mode, a +/// hash table that starts at 32 slots and doubles once more than 3/4 full. At `2^(lgConfigK - 3)` +/// slots it is promoted to the register array instead, one byte per register. Below lgConfigK 8 +/// a full LIST goes straight to the register array. Modes only move forward, and a union's +/// gadget follows the same path. +/// +/// The crate exposes only the estimate, which is enough. In LIST and SET mode `estimate()` is +/// `max(couponCount, interpolation)`, so it is never below the coupon count. Promotion seeds the +/// HIP accumulator with that estimate, so from the moment the register array exists the estimate +/// exceeds `max_coupons`. The one other way into the register array is a union with a sketch +/// already in it, which the caller tracks as `dense`. +/// +/// `heap_size_follows_the_crates_layout` checks this against the preamble the crate serializes, +/// so a `datasketches` bump that changes the layout fails that test. +mod layout { + const COUPON_BYTES: usize = 4; + const LIST_BYTES: usize = COUPON_BYTES << 3; + const MIN_SET_SLOTS: usize = 1 << 5; + /// Below this lgConfigK there is no SET mode. + const MIN_LG_K_WITH_SET: u8 = 8; + + /// The most coupons a sketch at `lg_config_k` holds before the register array replaces them. + pub fn max_coupons(lg_config_k: u8) -> usize { + if lg_config_k < MIN_LG_K_WITH_SET { + 7 + } else { + // 3/4 of the largest SET, 2^(lgConfigK - 3) slots. + 3 << (lg_config_k - 5) + } + } + + /// Whether a sketch at `lg_config_k` holds the register array, given its `estimate()` and + /// whether it is already known to. + pub fn is_dense(lg_config_k: u8, known_dense: bool, estimate: f64) -> bool { + known_dense || estimate > max_coupons(lg_config_k) as f64 + } + + /// Heap bytes held by an HLL_8 sketch or union gadget, never less than the crate allocated. + /// It is exact except just below a SET resize or the promotion, where the estimate, which + /// corrects for coupon collisions, runs slightly ahead of the coupon count and this charges + /// the next size up. + pub fn heap_bytes(lg_config_k: u8, known_dense: bool, estimate: f64) -> usize { + let registers = 1usize << lg_config_k; + if lg_config_k < MIN_LG_K_WITH_SET { + // The larger of a LIST and the register array, 128 bytes at most. + return registers.max(LIST_BYTES); + } + if is_dense(lg_config_k, known_dense, estimate) { + return registers; + } + let coupons = estimate as usize; + if coupons < 8 { + LIST_BYTES + } else { + let slots = (4 * coupons).div_ceil(3).next_power_of_two(); + COUPON_BYTES * slots.max(MIN_SET_SLOTS) + } + } } /// Byte offsets into the DataSketches HLL preamble, and the bits we need from it. @@ -151,6 +220,7 @@ impl SparkHllSketch { pub fn new(lg_config_k: u8) -> Self { Self { inner: HllSketch::new(lg_config_k, HllType::Hll8), + dense: false, } } @@ -182,9 +252,11 @@ impl SparkHllSketch { pub fn from_bytes(bytes: &[u8]) -> Result { reject_undecodable_hll4(bytes)?; let normalized = normalize_compact_hll_array(bytes); - HllSketch::deserialize(normalized.as_deref().unwrap_or(bytes)) - .map(|inner| Self { inner }) - .map_err(|e| DataFusionError::Execution(format!("invalid HLL sketch bytes: {e}"))) + let inner = HllSketch::deserialize(normalized.as_deref().unwrap_or(bytes)) + .map_err(|e| DataFusionError::Execution(format!("invalid HLL sketch bytes: {e}")))?; + // `deserialize` has validated the preamble, so the mode byte is there. + let dense = bytes[preamble::MODE] & preamble::CUR_MODE_MASK == preamble::CUR_MODE_HLL; + Ok(Self { inner, dense }) } /// The configured `lgConfigK`. @@ -199,10 +271,19 @@ impl SparkHllSketch { /// Merge another sketch into this one via a union, keeping HLL_8 output. pub fn merge_sketch(&mut self, other: &SparkHllSketch) { - let mut u = HllUnion::new(self.lg_config_k()); - u.update(&self.inner); - u.update(&other.inner); - self.inner = u.to_sketch(HllType::Hll8); + let mut u = SparkHllUnion::new(self.lg_config_k()); + u.merge(self); + u.merge(other); + *self = u.to_sketch(); + } + + fn is_dense(&self) -> bool { + layout::is_dense(self.lg_config_k(), self.dense, self.inner.estimate()) + } + + /// Heap bytes the sketch currently holds, for accumulator `size()`. + pub fn heap_size(&self) -> usize { + layout::heap_bytes(self.lg_config_k(), self.dense, self.inner.estimate()) } } @@ -210,6 +291,9 @@ impl SparkHllSketch { #[derive(Debug)] pub struct SparkHllUnion { inner: HllUnion, + /// Whether the union's gadget is known to hold the dense register array. Merging a sketch + /// that holds one always promotes the gadget. See `SparkHllSketch::dense`. + dense: bool, } impl SparkHllUnion { @@ -217,18 +301,37 @@ impl SparkHllUnion { pub fn new(lg_max_k: u8) -> Self { Self { inner: HllUnion::new(lg_max_k), + dense: false, } } /// Merge a sketch into the union. pub fn merge(&mut self, sketch: &SparkHllSketch) { + self.dense = self.is_dense() || sketch.is_dense(); self.inner.update(&sketch.inner); } + /// The union result as an HLL_8 sketch. + fn to_sketch(&self) -> SparkHllSketch { + SparkHllSketch { + inner: self.inner.to_sketch(HllType::Hll8), + dense: self.is_dense(), + } + } + /// The union result as an HLL_8 sketch's serialized bytes. pub fn to_sketch_bytes(&self) -> Vec { self.inner.to_sketch(HllType::Hll8).serialize() } + + fn is_dense(&self) -> bool { + layout::is_dense(self.inner.lg_config_k(), self.dense, self.inner.estimate()) + } + + /// Heap bytes the union's gadget currently holds, for accumulator `size()`. + pub fn heap_size(&self) -> usize { + layout::heap_bytes(self.inner.lg_config_k(), self.dense, self.inner.estimate()) + } } /// Estimate the distinct count from serialized sketch bytes, rounded to the @@ -370,4 +473,118 @@ mod tests { read.estimate() ); } + + /// Heap bytes the crate allocated for an HLL_8 sketch, read back from the preamble it + /// serializes: `2^lgArr` four-byte slots in LIST and SET mode, one byte per register in HLL. + fn allocated(bytes: &[u8]) -> usize { + const LG_K: usize = 3; + const LG_ARR: usize = 4; + if bytes[preamble::MODE] & preamble::CUR_MODE_MASK == preamble::CUR_MODE_HLL { + 1 << bytes[LG_K] + } else { + 4 << bytes[LG_ARR] + } + } + + /// `heap_size` must never report less than the crate holds, and it should track the real + /// layout rather than the dense maximum. Walk each lgConfigK through every SET resize and the + /// promotion, with duplicates mixed in, and compare against what the crate serializes. + #[test] + fn heap_size_follows_the_crates_layout() { + for lg_k in [4u8, 7, 8, 11, 14] { + let mut sketch = SparkHllSketch::new(lg_k); + let steps = 2 * layout::max_coupons(lg_k) + 64; + let mut exact = 0; + for i in 0..steps as i64 { + sketch.update_i64(i); + sketch.update_i64(i / 2); + let actual = allocated(&sketch.to_sketch_bytes()); + let charged = sketch.heap_size(); + assert!( + charged >= actual, + "lgConfigK {lg_k}, value {i}: charged {charged} bytes, crate holds {actual}" + ); + // Off by at most one size step. Below lgConfigK 8 the charge is a flat 128 bytes + // at most. + assert!( + charged <= (2 * actual).max(128), + "lgConfigK {lg_k}, value {i}: charged {charged} bytes, crate holds {actual}" + ); + exact += usize::from(charged == actual); + } + if lg_k >= 8 { + assert!( + exact * 100 >= steps * 99, + "lgConfigK {lg_k}: exact for only {exact} of {steps} cardinalities" + ); + } + } + } + + /// At lgConfigK 21 the register array is 2 MiB, but a small group holds a few dozen bytes and + /// a large one only becomes dense after ~200,000 distinct values. + #[test] + fn high_lg_config_k_is_charged_for_what_it_holds() { + let mut sketch = SparkHllSketch::new(21); + let mut next = 0i64; + for checkpoint in [1i64, 100, 10_000, 150_000, 250_000] { + while next < checkpoint { + sketch.update_i64(next); + next += 1; + } + let actual = allocated(&sketch.to_sketch_bytes()); + assert_eq!(sketch.heap_size(), actual, "after {checkpoint} values"); + if checkpoint == 1 { + assert_eq!(actual, 32, "one value sits in an 8-slot LIST"); + } + } + assert_eq!(sketch.heap_size(), 1 << 21); + } + + /// A union's gadget stays small while it takes small sketches. Merging a sketch that holds the + /// register array promotes it however little that sketch has seen, so it is the preamble's + /// mode, not the estimate, that has to decide. + #[test] + fn union_heap_size_follows_the_crates_layout() { + let sketch = |values: std::ops::Range| { + let mut s = SparkHllSketch::new(16); + for v in values { + s.update_i64(v); + } + s + }; + let mut union = SparkHllUnion::new(16); + union.merge(&sketch(0..5)); + union.merge(&sketch(5..40)); + assert_eq!(union.heap_size(), allocated(&union.to_sketch_bytes())); + assert_eq!(union.heap_size(), 256, "40 coupons fit a 64-slot SET"); + + // An HLL-mode sketch with one register set: every estimate it or a union with it gives is + // tiny. Preamble offsets: HIP 8, KxQ0 16, KxQ1 24, zero-register count 32. + let mut bytes = sketch(0..20_000).to_sketch_bytes(); + assert_eq!(allocated(&bytes), 1 << 16); + let registers = &mut bytes[preamble::HLL_SIZE..]; + registers.fill(0); + registers[0] = 1; + let zeros = (1u32 << 16) - 1; + bytes[8..16].copy_from_slice(&1.0f64.to_le_bytes()); + bytes[16..24].copy_from_slice(&(f64::from(zeros) + 0.5).to_le_bytes()); + bytes[24..32].copy_from_slice(&0.0f64.to_le_bytes()); + bytes[32..36].copy_from_slice(&zeros.to_le_bytes()); + let barely_used = SparkHllSketch::from_bytes(&bytes).unwrap(); + assert!(barely_used.estimate() < 2.0); + assert_eq!(barely_used.heap_size(), 1 << 16); + + union.merge(&barely_used); + assert!(union.inner.estimate() < 100.0); + assert_eq!(allocated(&union.to_sketch_bytes()), 1 << 16); + assert_eq!(union.heap_size(), 1 << 16); + + // `merge_sketch` goes through a union too, and keeps the flag. + let mut small = sketch(0..5); + small.merge_sketch(&barely_used); + assert!(small.estimate() < 100.0); + assert_eq!(allocated(&small.to_sketch_bytes()), 1 << 16); + assert_eq!(small.heap_size(), 1 << 16); + } } diff --git a/native/spark-expr/src/agg_funcs/hll_sketch_agg.rs b/native/spark-expr/src/agg_funcs/hll_sketch_agg.rs index c0e716042b3..3bcb5d44e1c 100644 --- a/native/spark-expr/src/agg_funcs/hll_sketch_agg.rs +++ b/native/spark-expr/src/agg_funcs/hll_sketch_agg.rs @@ -151,9 +151,9 @@ impl Accumulator for HllSketchAccumulator { } fn size(&self) -> usize { - // An HLL_8 sketch at lgConfigK=k can heap-allocate up to 1 << k bytes; - // account for that so memory reservation reflects actual usage. - std::mem::size_of_val(self) + (1usize << self.sketch.lg_config_k() as usize) + // What the sketch holds now, not the 1 << lgConfigK dense maximum: a group stays in a + // few-dozen-byte coupon list until it has seen enough distinct values. + std::mem::size_of_val(self) + self.sketch.heap_size() } fn state(&mut self) -> Result> { @@ -287,4 +287,38 @@ mod tests { acc.size() ); } + + /// 64 singleton groups at lgConfigK=21, grouped through a real Partial/Final plan in a 16 MiB + /// pool. Each group's sketch holds one coupon in an 8-slot LIST, 32 bytes, and Spark runs the + /// same query without trouble. Charging every group the 2 MiB dense register array instead + /// asks for 128 MiB, and the Final aggregation fails even with spilling. + #[tokio::test] + async fn singleton_groups_at_high_lg_config_k_fit_a_small_pool() { + use arrow::array::RecordBatch; + use datafusion::execution::runtime_env::RuntimeEnvBuilder; + use datafusion::logical_expr::AggregateUDF; + use datafusion::prelude::{SessionConfig, SessionContext}; + + let runtime = RuntimeEnvBuilder::new() + .with_memory_limit(16 * 1024 * 1024, 1.0) + .build_arc() + .unwrap(); + let ctx = SessionContext::new_with_config_rt( + SessionConfig::new().with_target_partitions(2), + runtime, + ); + ctx.register_udaf(AggregateUDF::new_from_impl(HllSketchAgg::new(21))); + let ids: ArrayRef = Arc::new(Int64Array::from((0..64i64).collect::>())); + ctx.register_batch("t", RecordBatch::try_from_iter([("id", ids)]).unwrap()) + .unwrap(); + + let batches = ctx + .sql("SELECT id, hll_sketch_agg(id) FROM t GROUP BY id") + .await + .unwrap() + .collect() + .await + .unwrap(); + assert_eq!(batches.iter().map(|b| b.num_rows()).sum::(), 64); + } } diff --git a/native/spark-expr/src/agg_funcs/hll_union_agg.rs b/native/spark-expr/src/agg_funcs/hll_union_agg.rs index ccd3a43077e..73a2b49ce02 100644 --- a/native/spark-expr/src/agg_funcs/hll_union_agg.rs +++ b/native/spark-expr/src/agg_funcs/hll_union_agg.rs @@ -134,13 +134,8 @@ impl Accumulator for HllUnionAccumulator { } } fn size(&self) -> usize { - // An HLL_8 sketch at lgConfigK=k can heap-allocate up to 1 << k bytes; - // account for that so memory reservation reflects actual usage. - std::mem::size_of_val(self) - + self - .seen_lg_config_k - .map(|k| 1usize << k as usize) - .unwrap_or(0) + // What the union's gadget holds now, not the dense maximum; see `HllSketchAccumulator`. + std::mem::size_of_val(self) + self.union.as_ref().map_or(0, SparkHllUnion::heap_size) } fn state(&mut self) -> Result> { // Unlike `evaluate`, an empty partial emits NULL rather than an empty lgConfigK=12 @@ -313,4 +308,47 @@ mod tests { "union of two disjoint compact sketches estimated {est}, expected ~2000" ); } + + /// The union counterpart of the `hll_sketch_agg` test: 64 groups, each unioning a single + /// one-value lgConfigK=21 sketch, grouped through Partial/Final in a 16 MiB pool. + #[tokio::test] + async fn singleton_groups_at_high_lg_config_k_fit_a_small_pool() { + use arrow::array::{Int64Array, RecordBatch}; + use datafusion::execution::runtime_env::RuntimeEnvBuilder; + use datafusion::logical_expr::AggregateUDF; + use datafusion::prelude::{SessionConfig, SessionContext}; + + let runtime = RuntimeEnvBuilder::new() + .with_memory_limit(16 * 1024 * 1024, 1.0) + .build_arc() + .unwrap(); + let ctx = SessionContext::new_with_config_rt( + SessionConfig::new().with_target_partitions(2), + runtime, + ); + ctx.register_udaf(AggregateUDF::new_from_impl(HllUnionAgg::new(false))); + let sketches: Vec> = (0..64i64) + .map(|i| { + let mut s = SparkHllSketch::new(21); + s.update_i64(i); + s.to_sketch_bytes() + }) + .collect(); + let ids: ArrayRef = Arc::new(Int64Array::from((0..64i64).collect::>())); + let sketches: ArrayRef = Arc::new(BinaryArray::from_iter_values(sketches)); + ctx.register_batch( + "t", + RecordBatch::try_from_iter([("id", ids), ("s", sketches)]).unwrap(), + ) + .unwrap(); + + let batches = ctx + .sql("SELECT id, hll_union_agg(s) FROM t GROUP BY id") + .await + .unwrap() + .collect() + .await + .unwrap(); + assert_eq!(batches.iter().map(|b| b.num_rows()).sum::(), 64); + } }