diff --git a/native/spark-expr/src/agg_funcs/approx_percentile.rs b/native/spark-expr/src/agg_funcs/approx_percentile.rs index cb67a1b6d0..366f18922c 100644 --- a/native/spark-expr/src/agg_funcs/approx_percentile.rs +++ b/native/spark-expr/src/agg_funcs/approx_percentile.rs @@ -15,15 +15,21 @@ // specific language governing permissions and limitations // under the License. -use super::quantile_summaries::QuantileSummaries; -use arrow::array::{Array, ArrayRef, BinaryArray, Float64Array, ListArray}; +use super::quantile_summaries::{QuantileSummaries, QuantileSummariesScratch}; +use arrow::array::{ + new_empty_array, Array, ArrayRef, BinaryArray, BinaryBuilder, BooleanArray, Float64Array, + ListArray, +}; use arrow::datatypes::{DataType, Field, FieldRef}; use datafusion::common::utils::SingleRowListArrayBuilder; use datafusion::common::{downcast_value, Result, ScalarValue}; use datafusion::logical_expr::function::{AccumulatorArgs, StateFieldsArgs}; use datafusion::logical_expr::Volatility::Immutable; -use datafusion::logical_expr::{Accumulator, AggregateUDFImpl, Signature}; +use datafusion::logical_expr::{ + Accumulator, AggregateUDFImpl, EmitTo, GroupsAccumulator, Signature, +}; use datafusion::physical_expr::expressions::format_state_name; +use std::mem::{size_of, size_of_val}; use std::sync::Arc; /// Native implementation of Spark's `approx_percentile` / `percentile_approx`, @@ -119,14 +125,27 @@ impl AggregateUDFImpl for ApproxPercentile { ))]) } - fn groups_accumulator_supported(&self, _args: AccumulatorArgs) -> bool { - false + fn groups_accumulator_supported(&self, args: AccumulatorArgs) -> bool { + !args.is_distinct + } + + fn create_groups_accumulator( + &self, + _args: AccumulatorArgs, + ) -> Result> { + Ok(Box::new(ApproxPercentileGroupsAccumulator::new( + self.percentiles.clone(), + self.accuracy, + self.input_type.clone(), + self.return_array, + ))) } } #[derive(Debug)] struct ApproxPercentileAccumulator { summary: QuantileSummaries, + scratch: QuantileSummariesScratch, percentiles: Vec, input_type: DataType, return_array: bool, @@ -140,40 +159,62 @@ impl ApproxPercentileAccumulator { QuantileSummaries::DEFAULT_COMPRESS_THRESHOLD, relative_error, ), + scratch: QuantileSummariesScratch::default(), percentiles, input_type, return_array, } } +} - /// Cast a double quantile back to Spark's output type. GK always returns an - /// actual inserted value (never an interpolation), so for the supported - /// numeric types this round-trips exactly and is always in range. - fn cast_back(&self, d: f64) -> ScalarValue { - match &self.input_type { - DataType::Int8 => ScalarValue::Int8(Some(d as i8)), - DataType::Int16 => ScalarValue::Int16(Some(d as i16)), - DataType::Int32 => ScalarValue::Int32(Some(d as i32)), - DataType::Int64 => ScalarValue::Int64(Some(d as i64)), - DataType::Float32 => ScalarValue::Float32(Some(d as f32)), - DataType::Float64 => ScalarValue::Float64(Some(d)), - // The serde only marks byte/short/int/long/float/double as - // supported, so no other type reaches the accumulator. - other => unreachable!("unsupported approx_percentile input type: {other}"), - } +/// Cast a double quantile back to Spark's output type. GK always returns an +/// inserted value, so every supported numeric type round-trips exactly. +fn cast_back(input_type: &DataType, value: f64) -> ScalarValue { + match input_type { + DataType::Int8 => ScalarValue::Int8(Some(value as i8)), + DataType::Int16 => ScalarValue::Int16(Some(value as i16)), + DataType::Int32 => ScalarValue::Int32(Some(value as i32)), + DataType::Int64 => ScalarValue::Int64(Some(value as i64)), + DataType::Float32 => ScalarValue::Float32(Some(value as f32)), + DataType::Float64 => ScalarValue::Float64(Some(value)), + // The serde only marks byte/short/int/long/float/double as supported. + other => unreachable!("unsupported approx_percentile input type: {other}"), } +} - /// The null Spark produces for an empty result: a typed null scalar, or a - /// null list when the call returns an array of percentiles. - fn null_result(&self) -> Result { - if self.return_array { - Ok(ScalarValue::List(Arc::new(ListArray::new_null( - Arc::new(Field::new("item", self.input_type.clone(), false)), - 1, - )))) - } else { - Ok(ScalarValue::try_from(&self.input_type)?) - } +fn null_result(input_type: &DataType, return_array: bool) -> Result { + if return_array { + Ok(ScalarValue::List(Arc::new(ListArray::new_null( + Arc::new(Field::new("item", input_type.clone(), false)), + 1, + )))) + } else { + Ok(ScalarValue::try_from(input_type)?) + } +} + +fn evaluate_summary( + summary: &mut QuantileSummaries, + scratch: &mut QuantileSummariesScratch, + percentiles: &[f64], + input_type: &DataType, + return_array: bool, +) -> Result { + summary.compress(scratch); + let results = match summary.query(percentiles) { + Some(results) if !results.is_empty() => results, + _ => return null_result(input_type, return_array), + }; + let scalars = results + .into_iter() + .map(|value| cast_back(input_type, value)); + if return_array { + let values = ScalarValue::iter_to_array(scalars)?; + Ok(SingleRowListArrayBuilder::new(values) + .with_nullable(false) + .build_list_scalar()) + } else { + Ok(scalars.into_iter().next().unwrap()) } } @@ -184,11 +225,11 @@ impl Accumulator for ApproxPercentileAccumulator { if arr.null_count() == 0 { // Fast path: no validity checks needed, iterate the raw values. for &v in arr.values() { - self.summary.insert(v); + self.summary.insert(v, &mut self.scratch); } } else { for v in arr.iter().flatten() { - self.summary.insert(v); + self.summary.insert(v, &mut self.scratch); } } Ok(()) @@ -196,7 +237,7 @@ impl Accumulator for ApproxPercentileAccumulator { fn merge_batch(&mut self, states: &[ArrayRef]) -> Result<()> { let digests = downcast_value!(&states[0], BinaryArray); - self.summary.compress(); + self.summary.compress(&mut self.scratch); for i in 0..digests.len() { if digests.is_null(i) { continue; @@ -206,45 +247,204 @@ impl Accumulator for ApproxPercentileAccumulator { digests.value(i), ); if self.summary.count() == 0 { - // Empty self: `merge` would return a clone of the (potentially - // large) peer, so move the owned peer in and skip the clone. + // Move the already-owned first digest into the accumulator. self.summary = peer; } else { - self.summary = self.summary.merge(&peer); + self.summary.merge(&peer, &mut self.scratch); } } Ok(()) } fn state(&mut self) -> Result> { - self.summary.compress(); + self.summary.compress(&mut self.scratch); Ok(vec![ScalarValue::Binary(Some(self.summary.to_bytes()))]) } fn evaluate(&mut self) -> Result { - self.summary.compress(); - // Spark returns null whenever the result would be empty, i.e. no rows - // were aggregated (`query` returns `None`) or the percentage argument - // was an empty array (`query` returns `Some([])`). - let results = match self.summary.query(&self.percentiles) { - Some(r) if !r.is_empty() => r, - _ => return self.null_result(), - }; - let scalars: Vec = results.into_iter().map(|d| self.cast_back(d)).collect(); + evaluate_summary( + &mut self.summary, + &mut self.scratch, + &self.percentiles, + &self.input_type, + self.return_array, + ) + } + + fn size(&self) -> usize { + size_of_val(self) + + self.summary.heap_size() + + self.scratch.heap_size() + + self.percentiles.capacity() * size_of::() + } +} + +#[derive(Debug)] +struct ApproxPercentileGroupsAccumulator { + summaries: Vec, + scratch: QuantileSummariesScratch, + summaries_heap_size: usize, + relative_error: f64, + percentiles: Vec, + input_type: DataType, + return_array: bool, +} + +impl ApproxPercentileGroupsAccumulator { + fn new(percentiles: Vec, accuracy: i64, input_type: DataType, return_array: bool) -> Self { + Self { + summaries: Vec::new(), + scratch: QuantileSummariesScratch::default(), + summaries_heap_size: 0, + relative_error: 1.0 / accuracy as f64, + percentiles, + input_type, + return_array, + } + } + + fn resize(&mut self, total_num_groups: usize) { + let relative_error = self.relative_error; + self.summaries.resize_with(total_num_groups, || { + QuantileSummaries::new( + QuantileSummaries::DEFAULT_COMPRESS_THRESHOLD, + relative_error, + ) + }); + } + + fn take_needed(&mut self, emit_to: EmitTo) -> Vec { + let summaries = emit_to.take_needed(&mut self.summaries); + let emitted_size = summaries + .iter() + .map(QuantileSummaries::heap_size) + .sum::(); + self.summaries_heap_size = self.summaries_heap_size.saturating_sub(emitted_size); + summaries + } + + fn output_type(&self) -> DataType { if self.return_array { - let values = ScalarValue::iter_to_array(scalars)?; - Ok(SingleRowListArrayBuilder::new(values) - .with_nullable(false) - .build_list_scalar()) + DataType::List(Arc::new(Field::new("item", self.input_type.clone(), false))) } else { - Ok(scalars.into_iter().next().unwrap()) + self.input_type.clone() } } +} + +fn selected(filter: Option<&BooleanArray>, row: usize) -> bool { + match filter { + Some(filter) => filter.is_valid(row) && filter.value(row), + None => true, + } +} + +fn adjust_size(total: &mut usize, before: usize, after: usize) { + if after >= before { + *total += after - before; + } else { + *total = total.saturating_sub(before - after); + } +} + +impl GroupsAccumulator for ApproxPercentileGroupsAccumulator { + fn update_batch( + &mut self, + values: &[ArrayRef], + group_indices: &[usize], + opt_filter: Option<&BooleanArray>, + total_num_groups: usize, + ) -> Result<()> { + let values = downcast_value!(&values[0], Float64Array); + self.resize(total_num_groups); + let (summaries, scratch, summaries_heap_size) = ( + &mut self.summaries, + &mut self.scratch, + &mut self.summaries_heap_size, + ); + for (row, &group_index) in group_indices.iter().enumerate() { + if !selected(opt_filter, row) || values.is_null(row) { + continue; + } + let summary = &mut summaries[group_index]; + let before = summary.heap_size(); + summary.insert(values.value(row), scratch); + adjust_size(summaries_heap_size, before, summary.heap_size()); + } + Ok(()) + } + + fn merge_batch( + &mut self, + states: &[ArrayRef], + group_indices: &[usize], + opt_filter: Option<&BooleanArray>, + total_num_groups: usize, + ) -> Result<()> { + let digests = downcast_value!(&states[0], BinaryArray); + self.resize(total_num_groups); + let (summaries, scratch, summaries_heap_size) = ( + &mut self.summaries, + &mut self.scratch, + &mut self.summaries_heap_size, + ); + for (row, &group_index) in group_indices.iter().enumerate() { + if !selected(opt_filter, row) || digests.is_null(row) { + continue; + } + let peer = QuantileSummaries::from_bytes( + QuantileSummaries::DEFAULT_COMPRESS_THRESHOLD, + digests.value(row), + ); + let summary = &mut summaries[group_index]; + let before = summary.heap_size(); + summary.compress(scratch); + if summary.count() == 0 { + *summary = peer; + } else { + summary.merge(&peer, scratch); + } + adjust_size(summaries_heap_size, before, summary.heap_size()); + } + Ok(()) + } + + fn state(&mut self, emit_to: EmitTo) -> Result> { + let summaries = self.take_needed(emit_to); + let mut builder = BinaryBuilder::new(); + for mut summary in summaries { + summary.compress(&mut self.scratch); + builder.append_value(summary.to_bytes()); + } + Ok(vec![Arc::new(builder.finish())]) + } + + fn evaluate(&mut self, emit_to: EmitTo) -> Result { + let summaries = self.take_needed(emit_to); + if summaries.is_empty() { + return Ok(new_empty_array(&self.output_type())); + } + let results = summaries + .into_iter() + .map(|mut summary| { + evaluate_summary( + &mut summary, + &mut self.scratch, + &self.percentiles, + &self.input_type, + self.return_array, + ) + }) + .collect::>>()?; + ScalarValue::iter_to_array(results) + } fn size(&self) -> usize { - std::mem::size_of_val(self) - + self.summary.heap_size() - + self.percentiles.capacity() * std::mem::size_of::() + size_of_val(self) + + self.summaries.capacity() * size_of::() + + self.summaries_heap_size + + self.scratch.heap_size() + + self.percentiles.capacity() * size_of::() } } @@ -256,6 +456,17 @@ mod tests { Arc::new(Float64Array::from(v)) } + fn grouped_values(start: i32) -> (ArrayRef, Vec) { + let values = (start..start + 100) + .map(|value| value as f64) + .chain((start + 1000..start + 1100).map(|value| value as f64)) + .collect(); + ( + f64_array(values), + vec![0; 100].into_iter().chain(vec![1; 100]).collect(), + ) + } + #[test] fn scalar_median_of_int_column() { let mut acc = ApproxPercentileAccumulator::new(vec![0.5], 10000, DataType::Int32, false); @@ -340,4 +551,83 @@ mod tests { } } } + + #[test] + fn grouped_update_filter_and_partial_emit() { + let mut acc = + ApproxPercentileGroupsAccumulator::new(vec![0.5], 10000, DataType::Float64, false); + let values: ArrayRef = Arc::new(Float64Array::from(vec![ + Some(1.0), + Some(3.0), + None, + Some(10.0), + Some(20.0), + ])); + let filter = BooleanArray::from(vec![true, true, true, true, false]); + acc.update_batch(&[values], &[0, 0, 1, 2, 2], Some(&filter), 4) + .unwrap(); + + let first = acc.evaluate(EmitTo::First(2)).unwrap(); + let first = first.as_any().downcast_ref::().unwrap(); + assert!((1.0..=3.0).contains(&first.value(0))); + assert!(first.is_null(1)); + + let rest = acc.evaluate(EmitTo::All).unwrap(); + let rest = rest.as_any().downcast_ref::().unwrap(); + assert_eq!(rest.value(0), 10.0); + assert!(rest.is_null(1)); + } + + #[test] + fn grouped_merge_reuses_one_scratch_buffer() { + let mut left = + ApproxPercentileGroupsAccumulator::new(vec![0.5], 10000, DataType::Float64, false); + let (values, groups) = grouped_values(1); + left.update_batch(&[values], &groups, None, 2).unwrap(); + let left_state = left.state(EmitTo::All).unwrap(); + + let mut right = + ApproxPercentileGroupsAccumulator::new(vec![0.5], 10000, DataType::Float64, false); + let (values, groups) = grouped_values(101); + right.update_batch(&[values], &groups, None, 2).unwrap(); + let right_state = right.state(EmitTo::All).unwrap(); + + let mut merged = + ApproxPercentileGroupsAccumulator::new(vec![0.5], 10000, DataType::Float64, false); + merged.merge_batch(&left_state, &[0, 1], None, 2).unwrap(); + let filter = BooleanArray::from(vec![Some(true), None]); + merged + .merge_batch(&right_state, &[0, 1], Some(&filter), 2) + .unwrap(); + + assert!(merged.scratch.heap_size() > 0); + assert_eq!( + merged.summaries_heap_size, + merged + .summaries + .iter() + .map(QuantileSummaries::heap_size) + .sum::() + ); + + let result = merged.evaluate(EmitTo::All).unwrap(); + let result = result.as_any().downcast_ref::().unwrap(); + assert!((90.0..=110.0).contains(&result.value(0))); + assert!((1040.0..=1060.0).contains(&result.value(1))); + } + + #[test] + fn grouped_array_result() { + let mut acc = ApproxPercentileGroupsAccumulator::new( + vec![0.25, 0.75], + 10000, + DataType::Float64, + true, + ); + acc.update_batch(&[f64_array(vec![1.0, 2.0, 3.0])], &[0, 0, 0], None, 1) + .unwrap(); + let result = acc.evaluate(EmitTo::All).unwrap(); + let result = result.as_any().downcast_ref::().unwrap(); + assert_eq!(result.value_length(0), 2); + } } diff --git a/native/spark-expr/src/agg_funcs/quantile_summaries.rs b/native/spark-expr/src/agg_funcs/quantile_summaries.rs index 2ed6b32004..ceb6fe9bcf 100644 --- a/native/spark-expr/src/agg_funcs/quantile_summaries.rs +++ b/native/spark-expr/src/agg_funcs/quantile_summaries.rs @@ -44,6 +44,18 @@ pub struct Stats { pub delta: i64, } +/// Reusable workspace owned by an accumulator rather than by each summary. +#[derive(Debug, Default)] +pub(crate) struct QuantileSummariesScratch { + sampled: Vec, +} + +impl QuantileSummariesScratch { + pub(crate) fn heap_size(&self) -> usize { + self.sampled.capacity() * std::mem::size_of::() + } +} + #[derive(Debug, Clone)] pub struct QuantileSummaries { compress_threshold: usize, @@ -87,18 +99,18 @@ impl QuantileSummaries { self.head_sampled.reserve(additional); } - pub fn insert(&mut self, x: f64) { + pub(crate) fn insert(&mut self, x: f64, scratch: &mut QuantileSummariesScratch) { self.head_sampled.push(x); self.compressed = false; if self.head_sampled.len() >= Self::DEFAULT_HEAD_SIZE { - self.with_head_buffer_inserted(); + self.with_head_buffer_inserted(scratch); if self.sampled.len() >= self.compress_threshold { - self.compress(); + self.compress(scratch); } } } - fn with_head_buffer_inserted(&mut self) { + fn with_head_buffer_inserted(&mut self, scratch: &mut QuantileSummariesScratch) { if self.head_sampled.is_empty() { return; } @@ -109,7 +121,8 @@ impl QuantileSummaries { // irrelevant; `total_cmp` gives a deterministic total order. sorted.sort_unstable_by(|a, b| a.total_cmp(b)); - let mut new_samples: Vec = Vec::with_capacity(self.sampled.len() + sorted.len()); + scratch.sampled.clear(); + scratch.sampled.reserve(self.sampled.len() + sorted.len()); let mut sample_idx = 0usize; let mut ops_idx = 0usize; while ops_idx < sorted.len() { @@ -117,11 +130,11 @@ impl QuantileSummaries { while sample_idx < self.sampled.len() && self.sampled[sample_idx].value <= current_sample { - new_samples.push(self.sampled[sample_idx]); + scratch.sampled.push(self.sampled[sample_idx]); sample_idx += 1; } current_count += 1; - let delta = if new_samples.is_empty() + let delta = if scratch.sampled.is_empty() || (sample_idx == self.sampled.len() && ops_idx == sorted.len() - 1) { 0 @@ -130,7 +143,7 @@ impl QuantileSummaries { // (verified `.toLong` in 3.4/3.5/4.0/4.1), matching our i64. (2.0 * self.relative_error * current_count as f64).floor() as i64 }; - new_samples.push(Stats { + scratch.sampled.push(Stats { value: current_sample, g: 1, delta, @@ -138,14 +151,15 @@ impl QuantileSummaries { ops_idx += 1; } while sample_idx < self.sampled.len() { - new_samples.push(self.sampled[sample_idx]); + scratch.sampled.push(self.sampled[sample_idx]); sample_idx += 1; } - self.sampled = new_samples; + std::mem::swap(&mut self.sampled, &mut scratch.sampled); + scratch.sampled.clear(); self.count = current_count; } - pub fn compress(&mut self) { + pub(crate) fn compress(&mut self, scratch: &mut QuantileSummariesScratch) { // Already compressed and the head buffer is empty (insert clears the // flag whenever it stages a value), so there is nothing to do. This // mirrors Spark's `PercentileDigest.isCompressed` guard, which also @@ -153,19 +167,21 @@ impl QuantileSummaries { if self.compressed { return; } - self.with_head_buffer_inserted(); + self.with_head_buffer_inserted(scratch); let merge_threshold = 2.0 * self.relative_error * self.count as f64; - self.sampled = Self::compress_immut(&self.sampled, merge_threshold); + Self::compress_immut(&self.sampled, merge_threshold, &mut scratch.sampled); + std::mem::swap(&mut self.sampled, &mut scratch.sampled); + scratch.sampled.clear(); self.compressed = true; } - fn compress_immut(current_samples: &[Stats], merge_threshold: f64) -> Vec { + fn compress_immut(current_samples: &[Stats], merge_threshold: f64, res: &mut Vec) { + res.clear(); if current_samples.is_empty() { - return Vec::new(); + return; } // Spark prepends into a `ListBuffer`; we push in the same order and // reverse once, which yields an identical sequence. - let mut res: Vec = Vec::with_capacity(current_samples.len()); let mut head = current_samples[current_samples.len() - 1]; // Traverse backward from size-2 down to index 1 (index 0 is preserved // separately so the minimum is always kept). @@ -186,17 +202,27 @@ impl QuantileSummaries { res.push(curr_head); } res.reverse(); - res } - pub fn merge(&self, other: &QuantileSummaries) -> QuantileSummaries { + pub(crate) fn merge( + &mut self, + other: &QuantileSummaries, + scratch: &mut QuantileSummariesScratch, + ) { debug_assert!(self.head_sampled.is_empty(), "compress before merge"); debug_assert!(other.head_sampled.is_empty(), "compress before merge"); if other.count == 0 { - return self.clone(); + return; } if self.count == 0 { - return other.clone(); + self.compress_threshold = other.compress_threshold; + self.relative_error = other.relative_error; + self.sampled.clear(); + self.sampled.extend_from_slice(&other.sampled); + self.count = other.count; + self.compressed = other.compressed; + self.head_sampled.clear(); + return; } let merged_relative_error = self.relative_error.max(other.relative_error); let merged_count = self.count + other.count; @@ -204,8 +230,10 @@ impl QuantileSummaries { (2.0 * other.relative_error * other.count as f64).floor() as i64; let additional_other_delta = (2.0 * self.relative_error * self.count as f64).floor() as i64; - let mut merged_sampled: Vec = - Vec::with_capacity(self.sampled.len() + other.sampled.len()); + scratch.sampled.clear(); + scratch + .sampled + .reserve(self.sampled.len() + other.sampled.len()); let mut self_idx = 0usize; let mut other_idx = 0usize; while self_idx < self.sampled.len() && other_idx < other.sampled.len() { @@ -233,28 +261,26 @@ impl QuantileSummaries { ) }; next_sample.delta += additional_delta; - merged_sampled.push(next_sample); + scratch.sampled.push(next_sample); } while self_idx < self.sampled.len() { - merged_sampled.push(self.sampled[self_idx]); + scratch.sampled.push(self.sampled[self_idx]); self_idx += 1; } while other_idx < other.sampled.len() { - merged_sampled.push(other.sampled[other_idx]); + scratch.sampled.push(other.sampled[other_idx]); other_idx += 1; } - let comp = Self::compress_immut( - &merged_sampled, + Self::compress_immut( + &scratch.sampled, 2.0 * merged_relative_error * merged_count as f64, + &mut self.sampled, ); - QuantileSummaries { - compress_threshold: other.compress_threshold, - relative_error: merged_relative_error, - sampled: comp, - count: merged_count, - compressed: true, - head_sampled: Vec::new(), - } + scratch.sampled.clear(); + self.compress_threshold = other.compress_threshold; + self.relative_error = merged_relative_error; + self.count = merged_count; + self.compressed = true; } pub fn query(&self, percentiles: &[f64]) -> Option> { @@ -373,11 +399,19 @@ mod tests { const EPS: f64 = 1.0 / 10000.0; fn summary_of(values: &[f64]) -> QuantileSummaries { - let mut qs = QuantileSummaries::new(QuantileSummaries::DEFAULT_COMPRESS_THRESHOLD, EPS); + summary_of_with_error(values, EPS) + } + + fn summary_of_with_error(values: &[f64], relative_error: f64) -> QuantileSummaries { + let mut qs = QuantileSummaries::new( + QuantileSummaries::DEFAULT_COMPRESS_THRESHOLD, + relative_error, + ); + let mut scratch = QuantileSummariesScratch::default(); for &v in values { - qs.insert(v); + qs.insert(v, &mut scratch); } - qs.compress(); + qs.compress(&mut scratch); qs } @@ -435,19 +469,65 @@ mod tests { } #[test] - fn merge_is_within_bound() { + fn repeated_merges_are_within_bound() { let left: Vec = (1..=5000).map(|i| i as f64).collect(); - let right: Vec = (5001..=10000).map(|i| i as f64).collect(); - let a = summary_of(&left); - let b = summary_of(&right); - let merged = a.merge(&b); - let mut all: Vec = left.iter().chain(right.iter()).cloned().collect(); + let middle: Vec = (5001..=10000).map(|i| i as f64).collect(); + let right: Vec = (10001..=15000).map(|i| i as f64).collect(); + let mut merged = summary_of(&left); + let mut scratch = QuantileSummariesScratch::default(); + merged.merge(&summary_of(&middle), &mut scratch); + merged.merge(&summary_of(&right), &mut scratch); + let mut all: Vec = left + .iter() + .chain(middle.iter()) + .chain(right.iter()) + .cloned() + .collect(); all.sort_by(|x, y| x.total_cmp(y)); let got = merged.query(&[0.5]).unwrap()[0]; let exact = exact_percentile(&all, 0.5); assert!((got - exact).abs() <= EPS * all.len() as f64 + 1.0); } + #[test] + fn flush_reuses_shared_sample_buffer_and_drops_head() { + let mut scratch = QuantileSummariesScratch::default(); + let mut summary = summary_of(&(1..=100).map(|i| i as f64).collect::>()); + summary.insert(50.5, &mut scratch); + scratch + .sampled + .reserve(summary.sampled.len() + summary.head_sampled.len()); + let sampled_ptr = summary.sampled.as_ptr(); + let buffer_ptr = scratch.sampled.as_ptr(); + + summary.with_head_buffer_inserted(&mut scratch); + + assert_eq!(summary.sampled.as_ptr(), buffer_ptr); + assert_eq!(scratch.sampled.as_ptr(), sampled_ptr); + assert_eq!(summary.head_sampled.capacity(), 0); + } + + #[test] + fn merge_keeps_uncompressed_capacity_in_shared_scratch() { + let error = 0.01; + let mut summary = + summary_of_with_error(&(1..=10_000).map(|i| i as f64).collect::>(), error); + let other = summary_of_with_error( + &(10_001..=20_000).map(|i| i as f64).collect::>(), + error, + ); + let merged_len = summary.sampled.len() + other.sampled.len(); + let mut scratch = QuantileSummariesScratch::default(); + scratch.sampled.reserve(merged_len); + let buffer_ptr = scratch.sampled.as_ptr(); + + summary.merge(&other, &mut scratch); + + assert_eq!(scratch.sampled.as_ptr(), buffer_ptr); + assert!(scratch.heap_size() >= merged_len * std::mem::size_of::()); + assert!(summary.sampled.capacity() < merged_len); + } + #[test] fn extreme_percentiles_hit_short_circuits() { let values: Vec = (1..=1000).map(|i| i as f64).collect();