From c3ea022de5826350edd7da11937a173f29a94d20 Mon Sep 17 00:00:00 2001 From: "Lei, HUANG" <6406592+v0y4g3r@users.noreply.github.com> Date: Mon, 7 Sep 2026 06:34:43 +0000 Subject: [PATCH] perf(mito2): blazing-fast tournament tree merger (#8989) * perf(mito2): optimize flat merge heap and primary-key interleave Replace the per-row BinaryHeap pop/push cycle in FlatMerge with an in-place root mutation plus a single sift-down repair on a custom RootHeap, keeping the cold heap and direct-batch fast path unchanged. Fallible or awaiting batch transitions move the hot node out of the heap first, preserving error and cancellation semantics. Exploit the globally sorted merge output to build the internal Dictionary primary-key column with a one-pass ordered gather: append a Binary value only when the PK changes and reuse the current key for adjacent equal PKs, bypassing Arrow dictionary masks, hash interning and key remapping. Non-PK columns still use Arrow interleave. Also cache the current primary-key byte range in RowCursor to avoid repeated dictionary range decoding during comparisons, and add a setup-free Criterion benchmark with exact output-row assertions. 32-way/1 row-per-series/40-tag improves 955.79ms -> 562.36ms (-41.2%); 0-tag -39.5%, 64 rows/series -79.8%, 8-way -56.1%, single-iterator control +0.3%. Signed-off-by: Lei, HUANG * test(mito2): add rows-per-series sweep to flat merge bench Add 32-way/40-tag shapes for 1, 10, 100, 1000 and 10000 rows per series, and allow FLAT_MERGE_BENCH_SHAPE to match shape name prefixes so the whole sweep can run in one invocation. Signed-off-by: Lei, HUANG * test(mito2): add oracle-based correctness tests for RootHeap Drive RootHeap and a std BinaryHeap oracle with the same seeded op sequence (push / pop / mutate-root + repair) and assert peek, len, best_child and the full drain order after every operation. A second run with a tiny value range makes duplicates dominate, covering the equal-key branches of sift_up/sift_down and best_child. Signed-off-by: Lei, HUANG * perf(mito2): replace hot heap with a tournament tree in flat merge Replace the hot RootHeap with a fixed-capacity tournament (winner) tree over per-node slots: every internal node caches the champion of its subtree, so advancing the winner only replays the ~log2(k) nodes on its leaf-to-root path with one compare per level, instead of the heap's two-compares-per-level sift that also re-compares the same node pairs on every row. Two fast paths keep dense shapes at O(1) per row: - champion retention: after mutating the winner in place, skip the replay entirely when it still beats the runner-up (its path caches are unchanged by construction); - a second-best slot cache, invalidated on any structural change, so the retention check costs a single compare without walking the tree. The cold heap, hot/cold overlap window, direct-batch fast path and the remove-before-fallible-fetch batch transition semantics are unchanged. Vs the RootHeap version: 1rps/32way/40tag -19.7%, 0tag -34.4%, 8way -15.9%, 64rps -30.5%, sweep 10/100/1000/10000rps -29~32%; vs the original BinaryHeap baseline the main shape is -52.8%. The single-iterator control is +8% (+50ns one-time construction allocation, no merge work). Signed-off-by: Lei, HUANG * fix(mito2): support generic schemas in flat merge Signed-off-by: Lei, HUANG * fix(mito2): satisfy clippy in flat merge benchmark Signed-off-by: Lei, HUANG * perf(mito2): cache flat merge primary key index Compute the internal primary-key column index once when constructing BatchBuilder and reuse it for every output batch. Preserve the column-name gate for generic schemas. Signed-off-by: Lei, HUANG * test(mito2): benchmark high-fan-in flat merges Add sparse 64, 128, 256, and 512-way merge shapes while keeping the total input fixed at 3.2 million rows. Compared with the merge-base heap implementation, median time improves by 56.0%, 60.8%, 55.6%, and 56.1%, respectively. Signed-off-by: Lei, HUANG --------- Signed-off-by: Lei, HUANG --- src/mito2/Cargo.toml | 5 + src/mito2/benches/bench_flat_merge.rs | 281 ++++++ src/mito2/src/read/flat_merge.rs | 1218 +++++++++++++++++++++++-- 3 files changed, 1430 insertions(+), 74 deletions(-) create mode 100644 src/mito2/benches/bench_flat_merge.rs diff --git a/src/mito2/Cargo.toml b/src/mito2/Cargo.toml index bf8e1cdaec..116c4066ae 100644 --- a/src/mito2/Cargo.toml +++ b/src/mito2/Cargo.toml @@ -130,6 +130,11 @@ name = "bench_compaction_picker" harness = false required-features = ["testing"] +[[bench]] +name = "bench_flat_merge" +harness = false +required-features = ["testing"] + [[bench]] name = "simple_bulk_memtable" harness = false diff --git a/src/mito2/benches/bench_flat_merge.rs b/src/mito2/benches/bench_flat_merge.rs new file mode 100644 index 0000000000..621e1550f4 --- /dev/null +++ b/src/mito2/benches/bench_flat_merge.rs @@ -0,0 +1,281 @@ +// Copyright 2023 Greptime Team +// +// Licensed 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 std::hint::black_box; +use std::sync::Arc; + +use criterion::{BenchmarkId, Criterion, criterion_group, criterion_main}; +use datatypes::arrow::array::builder::BinaryDictionaryBuilder; +use datatypes::arrow::array::{ + ArrayRef, Float64Array, TimestampMillisecondArray, UInt8Array, UInt64Array, +}; +use datatypes::arrow::datatypes::{DataType, Field, Schema, SchemaRef, TimeUnit, UInt32Type}; +use datatypes::arrow::record_batch::RecordBatch; +use mito_codec::row_converter::SparsePrimaryKeyCodec; +use mito2::memtable::BoxedRecordBatchIterator; +use mito2::read::flat_merge::FlatMergeIterator; + +struct Shape { + name: &'static str, + num_iters: usize, + rows_per_iter: usize, + rows_per_series: usize, + num_pk_tags: u32, +} + +const TABLE_ID: u32 = 1024; + +fn make_common_tag_suffix(codec: &SparsePrimaryKeyCodec, num_pk_tags: u32) -> Vec { + let mut suffix = Vec::new(); + codec + .encode_raw_tag_value( + (1..num_pk_tags).map(|column_id| (column_id + 1, b"tagvalue".as_slice())), + &mut suffix, + ) + .unwrap(); + suffix +} + +fn make_key( + codec: &SparsePrimaryKeyCodec, + series: u64, + num_pk_tags: u32, + common_tag_suffix: &[u8], +) -> Vec { + let mut key = Vec::new(); + codec.encode_internal(TABLE_ID, series, &mut key).unwrap(); + if num_pk_tags > 0 { + let series_tag = series.to_be_bytes(); + codec + .encode_raw_tag_value(std::iter::once((1, series_tag.as_slice())), &mut key) + .unwrap(); + key.extend_from_slice(common_tag_suffix); + } + key +} + +fn build_input(shape: &Shape) -> (SchemaRef, Vec) { + let fields = vec![ + Field::new("value", DataType::Float64, true), + Field::new( + "ts", + DataType::Timestamp(TimeUnit::Millisecond, None), + false, + ), + Field::new( + "__primary_key", + DataType::Dictionary(Box::new(DataType::UInt32), Box::new(DataType::Binary)), + false, + ), + Field::new("__sequence", DataType::UInt64, false), + Field::new("__op_type", DataType::UInt8, false), + ]; + let schema = Arc::new(Schema::new(fields)); + assert_eq!( + schema.fields().len(), + 5, + "sparse SST input must not contain raw tag columns" + ); + + let num_series = shape.rows_per_iter / shape.rows_per_series; + let num_rows = num_series * shape.rows_per_series; + let codec = SparsePrimaryKeyCodec::schemaless(); + let common_tag_suffix = make_common_tag_suffix(&codec, shape.num_pk_tags); + let keys: Vec<_> = (0..num_series as u64) + .map(|series| make_key(&codec, series, shape.num_pk_tags, &common_tag_suffix)) + .collect(); + for series in [0, keys.len() - 1] { + let key = &keys[series]; + let (table_id, tsid) = codec + .decode_ids(key) + .expect("benchmark primary keys must use sparse encoding"); + assert_eq!(TABLE_ID, table_id); + assert_eq!(series as u64, tsid); + } + let mut batches = Vec::with_capacity(shape.num_iters); + for iter_idx in 0..shape.num_iters { + let mut primary_key = BinaryDictionaryBuilder::::new(); + let mut timestamps = Vec::with_capacity(num_rows); + for key in &keys { + for row in 0..shape.rows_per_series { + primary_key.append(key).unwrap(); + timestamps.push(((iter_idx * shape.rows_per_series + row) as i64) * 1000); + } + } + + let mut columns = Vec::with_capacity(schema.fields().len()); + columns.push(Arc::new(Float64Array::from(vec![1.0; num_rows])) as ArrayRef); + columns.push(Arc::new(TimestampMillisecondArray::from(timestamps)) as ArrayRef); + columns.push(Arc::new(primary_key.finish()) as ArrayRef); + columns.push(Arc::new(UInt64Array::from(vec![1; num_rows])) as ArrayRef); + columns.push(Arc::new(UInt8Array::from(vec![1; num_rows])) as ArrayRef); + + let batch = RecordBatch::try_new(Arc::clone(&schema), columns).unwrap(); + batches.push(batch); + } + + (schema, batches) +} + +fn run_merge( + schema: SchemaRef, + iters: Vec, + expected_rows: usize, +) -> usize { + let iter = FlatMergeIterator::new(schema, iters, 8192).unwrap(); + let output_rows = iter.map(|batch| batch.unwrap().num_rows()).sum(); + assert_eq!(expected_rows, output_rows); + black_box(output_rows) +} + +fn bench_merge(c: &mut Criterion) { + let mut group = c.benchmark_group("flat_merge"); + group.sample_size(10); + + let shapes = [ + Shape { + name: "sparse_1rps_32way_40tags", + num_iters: 32, + rows_per_iter: 100_000, + rows_per_series: 1, + num_pk_tags: 40, + }, + Shape { + name: "sparse_1rps_64way_40tags", + num_iters: 64, + rows_per_iter: 50_000, + rows_per_series: 1, + num_pk_tags: 40, + }, + Shape { + name: "sparse_1rps_128way_40tags", + num_iters: 128, + rows_per_iter: 25_000, + rows_per_series: 1, + num_pk_tags: 40, + }, + Shape { + name: "sparse_1rps_256way_40tags", + num_iters: 256, + rows_per_iter: 12_500, + rows_per_series: 1, + num_pk_tags: 40, + }, + Shape { + name: "sparse_1rps_512way_40tags", + num_iters: 512, + rows_per_iter: 6_250, + rows_per_series: 1, + num_pk_tags: 40, + }, + Shape { + name: "sparse_64rps_32way_40tags", + num_iters: 32, + rows_per_iter: 100_000, + rows_per_series: 64, + num_pk_tags: 40, + }, + Shape { + name: "sparse_1rps_32way_0tags", + num_iters: 32, + rows_per_iter: 100_000, + rows_per_series: 1, + num_pk_tags: 0, + }, + Shape { + name: "sparse_1rps_8way_40tags", + num_iters: 8, + rows_per_iter: 400_000, + rows_per_series: 1, + num_pk_tags: 40, + }, + Shape { + name: "single_iter_1rps_40tags", + num_iters: 1, + rows_per_iter: 3_200_000, + rows_per_series: 1, + num_pk_tags: 40, + }, + // Rows-per-series sweep, all 32-way with 40 encoded tags. + Shape { + name: "sweep_1rps_32way_40tags", + num_iters: 32, + rows_per_iter: 100_000, + rows_per_series: 1, + num_pk_tags: 40, + }, + Shape { + name: "sweep_10rps_32way_40tags", + num_iters: 32, + rows_per_iter: 100_000, + rows_per_series: 10, + num_pk_tags: 40, + }, + Shape { + name: "sweep_100rps_32way_40tags", + num_iters: 32, + rows_per_iter: 100_000, + rows_per_series: 100, + num_pk_tags: 40, + }, + Shape { + name: "sweep_1000rps_32way_40tags", + num_iters: 32, + rows_per_iter: 100_000, + rows_per_series: 1000, + num_pk_tags: 40, + }, + Shape { + name: "sweep_10000rps_32way_40tags", + num_iters: 32, + rows_per_iter: 100_000, + rows_per_series: 10000, + num_pk_tags: 40, + }, + ]; + + let selected_shape = std::env::var("FLAT_MERGE_BENCH_SHAPE").ok(); + for shape in &shapes { + if selected_shape + .as_deref() + .is_some_and(|selected| !shape.name.starts_with(selected)) + { + continue; + } + let (schema, batches) = build_input(shape); + let expected_rows = + shape.num_iters * (shape.rows_per_iter / shape.rows_per_series) * shape.rows_per_series; + group.bench_function(BenchmarkId::from_parameter(shape.name), |b| { + b.iter_batched( + || { + let iters = batches + .iter() + .cloned() + .map(|batch| { + Box::new(std::iter::once(Ok(batch))) as BoxedRecordBatchIterator + }) + .collect(); + (Arc::clone(&schema), iters) + }, + |(schema, iters)| run_merge(schema, iters, expected_rows), + criterion::BatchSize::LargeInput, + ); + }); + } + + group.finish(); +} + +criterion_group!(benches, bench_merge); +criterion_main!(benches); diff --git a/src/mito2/src/read/flat_merge.rs b/src/mito2/src/read/flat_merge.rs index 2562f36430..f1c30dfbfd 100644 --- a/src/mito2/src/read/flat_merge.rs +++ b/src/mito2/src/read/flat_merge.rs @@ -12,15 +12,20 @@ // See the License for the specific language governing permissions and // limitations under the License. +#[cfg(test)] +use std::cell::Cell; use std::cmp::Ordering; use std::collections::BinaryHeap; use std::fmt; +use std::ops::Range; use std::sync::Arc; use std::time::{Duration, Instant}; use async_stream::try_stream; use common_telemetry::debug; -use datatypes::arrow::array::{Array, AsArray, Int64Array, UInt64Array}; +use datatypes::arrow::array::{ + Array, ArrayRef, AsArray, BinaryBuilder, Int64Array, UInt32Array, UInt64Array, +}; use datatypes::arrow::compute::interleave; use datatypes::arrow::datatypes::{ArrowNativeType, BinaryType, DataType, SchemaRef, Utf8Type}; use datatypes::arrow::error::ArrowError; @@ -30,6 +35,7 @@ use datatypes::timestamp::timestamp_array_to_primitive; use futures::{Stream, TryStreamExt}; use snafu::ResultExt; use store_api::storage::SequenceNumber; +use store_api::storage::consts::PRIMARY_KEY_COLUMN_NAME; use crate::error::{ComputeArrowSnafu, Result}; use crate::memtable::BoxedRecordBatchIterator; @@ -96,6 +102,97 @@ fn check_interleave_overflow( Ok(()) } +/// Interleaves the non-null internal primary-key column from globally sorted rows. +fn interleave_primary_key( + arrays: &[&dyn Array], + indices: &[(usize, usize)], +) -> std::result::Result { + if arrays.is_empty() { + return Err(ArrowError::InvalidArgumentError( + "interleave requires input of at least one array".to_string(), + )); + } + + let dictionaries = arrays + .iter() + .map(|array| { + let dictionary = array + .as_any() + .downcast_ref::() + .ok_or_else(|| { + ArrowError::CastError(format!( + "expected Dictionary(UInt32, Binary) primary key, got {}", + array.data_type() + )) + })?; + let values = dictionary + .values() + .as_any() + .downcast_ref::() + .ok_or_else(|| { + ArrowError::CastError(format!( + "expected Binary primary-key dictionary values, got {}", + dictionary.values().data_type() + )) + })?; + Ok((dictionary, values)) + }) + .collect::, ArrowError>>()?; + + let mut keys = Vec::with_capacity(indices.len()); + let mut values = BinaryBuilder::with_capacity(indices.len(), 0); + let mut previous_primary_key = None; + let mut current_key = 0; + let mut num_dictionary_values = 0_usize; + let mut value_bytes = 0_usize; + + for &(array_idx, row_idx) in indices { + let (dictionary, dictionary_values) = dictionaries.get(array_idx).ok_or_else(|| { + ArrowError::InvalidArgumentError(format!( + "primary-key source index {array_idx} is out of bounds for {} arrays", + dictionaries.len() + )) + })?; + if row_idx >= dictionary.len() { + return Err(ArrowError::InvalidArgumentError(format!( + "primary-key row index {row_idx} is out of bounds for array of length {}", + dictionary.len() + ))); + } + let source_key = dictionary.key(row_idx).ok_or_else(|| { + ArrowError::InvalidArgumentError( + "internal primary-key dictionary contains a null key".to_string(), + ) + })?; + if dictionary_values.is_null(source_key) { + return Err(ArrowError::InvalidArgumentError( + "internal primary-key dictionary contains a null dictionary value".to_string(), + )); + } + let primary_key = dictionary_values.value(source_key); + + if previous_primary_key != Some(primary_key) { + current_key = u32::try_from(num_dictionary_values) + .map_err(|_| ArrowError::DictionaryKeyOverflowError)?; + value_bytes = value_bytes.checked_add(primary_key.len()).ok_or_else(|| { + ArrowError::ArithmeticOverflow( + "primary-key dictionary value length overflow".to_string(), + ) + })?; + if value_bytes > i32::MAX as usize { + return Err(ArrowError::OffsetOverflowError(value_bytes)); + } + values.append_value(primary_key); + num_dictionary_values += 1; + previous_primary_key = Some(primary_key); + } + keys.push(current_key); + } + + let dictionary = PrimaryKeyArray::try_new(UInt32Array::from(keys), Arc::new(values.finish()))?; + Ok(Arc::new(dictionary)) +} + /// Keeps track of the current position in a batch #[derive(Debug, Copy, Clone, Default)] struct BatchCursor { @@ -191,6 +288,9 @@ pub struct BatchBuilder { /// The schema of the RecordBatches yielded by this stream schema: SchemaRef, + /// Index of the internal primary key column, if present. + primary_key_column_idx: Option, + /// Maintain a list of [`RecordBatch`] and their corresponding stream batches: Vec<(usize, RecordBatch)>, @@ -205,8 +305,12 @@ pub struct BatchBuilder { impl BatchBuilder { /// Create a new [`BatchBuilder`] with the provided `stream_count` and `batch_size` pub fn new(schema: SchemaRef, stream_count: usize, batch_size: usize) -> Self { + let primary_key_column_idx = (schema.fields.len() >= 3) + .then(|| primary_key_column_index(schema.fields.len())) + .filter(|&column_idx| schema.field(column_idx).name() == PRIMARY_KEY_COLUMN_NAME); Self { schema, + primary_key_column_idx, batches: Vec::with_capacity(stream_count * 2), cursors: vec![BatchCursor::default(); stream_count], indices: Vec::with_capacity(batch_size), @@ -265,7 +369,11 @@ impl BatchBuilder { .iter() .map(|(_, batch)| batch.column(column_idx).as_ref()) .collect(); - interleave(&arrays, &self.indices).context(ComputeArrowSnafu) + if Some(column_idx) == self.primary_key_column_idx { + interleave_primary_key(&arrays, &self.indices).context(ComputeArrowSnafu) + } else { + interleave(&arrays, &self.indices).context(ComputeArrowSnafu) + } }) .collect::>>()?; @@ -322,6 +430,199 @@ impl BatchBuilder { } } +/// Sentinel for an empty slot in a [TournamentTree]. +const EMPTY_SLOT: usize = usize::MAX; + +/// A tournament tree over a fixed set of slots. +/// +/// Each node occupies one slot; a slot is empty while its node is not in the +/// tree (i.e. it is in the cold heap or reached EOF). An empty slot always +/// loses a match. The number of leaves is padded to a power of two; leaves +/// beyond the capacity never participate, so the tree works for arbitrary +/// capacities. +/// +/// Invariant: every internal tree node caches the champion (hottest slot, per +/// `Ord`) of its subtree, so the root always holds the hottest occupied slot +/// and updating a leaf only requires recomputing the ~log2(capacity) internal +/// nodes on its path to the root. We cache champions ("winner tree") so that +/// each internal node is a pure function of its children — insertion, removal +/// and mutation share one replay path — and the runner-up is directly +/// available from the sibling champions on the winner's path, which the +/// hot/cold transition check needs after mutating the winner in place. +struct TournamentTree { + /// Slot storage, one slot per node. `None` means the slot is empty. + nodes: Vec>, + /// Tree of champions: `tree[i]` is the slot index of the champion of the + /// subtree rooted at `i` (or [EMPTY_SLOT]). Leaves for slot `s` live at + /// index `leaves + s`; the root is index 1. + tree: Vec, + /// Number of occupied slots. + len: usize, + /// Number of leaves, padded to a power of two. + leaves: usize, + /// Cached `(winner slot, second-best slot)`, with [EMPTY_SLOT] standing + /// for no second-best. Invalidated by [TournamentTree::replay], the single + /// choke point of every structural change (push, pop, winner replay). + /// + /// While the tree is structurally unchanged and the winner keeps its slot + /// the cache stays valid: only the winner's node may be mutated in place + /// and the second-best slot is never the winner's slot, so the cached + /// slot's node is untouched. + second_best_cache: Option<(usize, usize)>, +} + +impl TournamentTree { + fn with_capacity(capacity: usize) -> Self { + let leaves = capacity.next_power_of_two().max(1); + let mut nodes = Vec::new(); + nodes.resize_with(capacity, || None); + Self { + nodes, + tree: vec![EMPTY_SLOT; leaves * 2], + len: 0, + leaves, + second_best_cache: None, + } + } + + fn is_empty(&self) -> bool { + self.len == 0 + } + + fn len(&self) -> usize { + self.len + } + + /// Returns the winner (greatest element among occupied slots). + fn peek(&self) -> Option<&T> { + self.winner_slot() + .and_then(|slot| self.nodes[slot].as_ref()) + } + + /// Returns the winner mutably. Call [TournamentTree::replay_winner] after + /// mutating it, or remove it with `pop()` if it reached EOF or moved cold. + fn winner_mut(&mut self) -> Option<&mut T> { + let slot = self.winner_slot()?; + self.nodes[slot].as_mut() + } + + /// Returns the second greatest element among occupied slots. + /// + /// The runner-up is the champion of one of the sibling subtrees on the + /// winner's path to the root. + #[cfg(test)] + fn second_best(&mut self) -> Option<&T> { + let slot = self.second_best_slot()?; + Some(self.nodes[slot].as_ref().unwrap()) + } + + /// Returns the winner and the second greatest element among occupied slots. + fn winner_and_second_best(&mut self) -> (Option<&T>, Option<&T>) { + let second = self.second_best_slot(); + let winner = self.winner_slot(); + ( + winner.map(|slot| self.nodes[slot].as_ref().unwrap()), + second.map(|slot| self.nodes[slot].as_ref().unwrap()), + ) + } + + /// Inserts `value` into a free slot and replays its path to the root. + /// + /// Scans for a free slot instead of keeping a free list to keep the tree + /// cheap to construct. This is O(capacity), but pushes only happen on + /// batch transitions, never per row. + /// + /// # Panics + /// Panics if the tree is full. + fn push(&mut self, value: T) { + let slot = self + .nodes + .iter() + .position(Option::is_none) + .expect("tournament tree is full"); + self.nodes[slot] = Some(value); + self.tree[self.leaves + slot] = slot; + self.len += 1; + self.replay(slot); + } + + /// Removes and returns the winner, leaving its slot empty. + fn pop(&mut self) -> Option { + let slot = self.winner_slot()?; + let value = self.nodes[slot].take(); + debug_assert!(value.is_some()); + self.tree[self.leaves + slot] = EMPTY_SLOT; + self.len -= 1; + self.replay(slot); + value + } + + /// Replays the winner's path after its value was mutated in place. + fn replay_winner(&mut self) { + if let Some(slot) = self.winner_slot() { + self.replay(slot); + } + } + + /// Returns the slot of the champion at the root, if any slot is occupied. + fn winner_slot(&self) -> Option { + (self.tree[1] != EMPTY_SLOT).then_some(self.tree[1]) + } + + /// Returns the slot of the second greatest element among occupied slots. + fn second_best_slot(&mut self) -> Option { + let winner = self.winner_slot()?; + if let Some((cached_winner, cached)) = self.second_best_cache + && cached_winner == winner + { + return (cached != EMPTY_SLOT).then_some(cached); + } + let second = self.compute_second_best_slot(winner).unwrap_or(EMPTY_SLOT); + self.second_best_cache = Some((winner, second)); + (second != EMPTY_SLOT).then_some(second) + } + + /// Scans the champions of the sibling subtrees on `winner`'s path to the + /// root for the hottest one. + fn compute_second_best_slot(&self, winner: usize) -> Option { + let mut node = self.leaves + winner; + let mut best = EMPTY_SLOT; + while node > 1 { + let challenger = self.tree[node ^ 1]; + if self.wins(challenger, best) { + best = challenger; + } + node /= 2; + } + (best != EMPTY_SLOT).then_some(best) + } + + /// Returns true if slot `a` wins its match against slot `b`, i.e. `a` is + /// the greater element. An occupied slot always beats an empty slot + /// ([EMPTY_SLOT] or a slot whose node was removed). + fn wins(&self, a: usize, b: usize) -> bool { + match (self.nodes.get(a), self.nodes.get(b)) { + (Some(Some(x)), Some(Some(y))) => x >= y, + (Some(Some(_)), _) => true, + _ => false, + } + } + + /// Recomputes the internal nodes on the path from `slot`'s leaf to the + /// root (~log2(capacity) comparisons). + fn replay(&mut self, slot: usize) { + debug_assert!(slot < self.nodes.len()); + self.second_best_cache = None; + let mut node = (self.leaves + slot) / 2; + while node > 0 { + let left = self.tree[node * 2]; + let right = self.tree[node * 2 + 1]; + self.tree[node] = if self.wins(left, right) { left } else { right }; + node /= 2; + } + } +} + /// A comparable node of the heap. trait NodeCmp: Eq + Ord { /// Returns whether the node still has batch to read. @@ -336,13 +637,13 @@ trait NodeCmp: Eq + Ord { } /// Common algorithm of merging sorted batches from multiple nodes. -struct MergeAlgo { +struct MergeAlgo { /// Holds nodes whose key range of current batch **is** overlapped with the merge window. /// Each node yields batches from a `source`. /// - /// Node in this heap **MUST** not be empty. A `merge window` is the (primary key, timestamp) - /// range of the **root node** in the `hot` heap. - hot: BinaryHeap, + /// Node in this tree **MUST** not be empty. A `merge window` is the (primary key, timestamp) + /// range of the **winner node** in the `hot` tree. + hot: TournamentTree, /// Holds nodes whose key range of current batch **isn't** overlapped with the merge window. /// /// Nodes in this heap **MUST** not be empty. @@ -356,7 +657,7 @@ impl MergeAlgo { fn new(mut nodes: Vec) -> Self { // Skips EOF nodes. nodes.retain(|node| !node.is_eof()); - let hot = BinaryHeap::with_capacity(nodes.len()); + let hot = TournamentTree::with_capacity(nodes.len()); let cold = BinaryHeap::from(nodes); let mut algo = MergeAlgo { hot, cold }; @@ -367,7 +668,7 @@ impl MergeAlgo { } /// Moves nodes in `cold` heap, whose key range is overlapped with current merge - /// window to `hot` heap. + /// window to `hot` tree. fn refill_hot(&mut self) { while !self.cold.is_empty() { if let Some(merge_window) = self.hot.peek() { @@ -385,40 +686,61 @@ impl MergeAlgo { } } - /// Push the node popped from `hot` back to a proper heap. - fn reheap(&mut self, node: T) { - if node.is_eof() { - // If the node is EOF, don't put it into the heap again. - // The merge window would be updated, need to refill the hot heap. - self.refill_hot(); - } else { - // Find a proper heap for this node. - let node_is_cold = if let Some(hottest) = self.hot.peek() { - // If key range of this node is behind the hottest node's then we can - // push it to the cold heap. Otherwise we should push it to the hot heap. - node.is_behind(hottest) - } else { - // The hot heap is empty, but we don't known whether the current - // batch of this node is still the hottest. - true - }; - - if node_is_cold { - self.cold.push(node); - } else { - self.hot.push(node); - } - // Anyway, the merge window has been changed, we need to refill the hot heap. - self.refill_hot(); - } + /// Returns the hottest node mutably. + fn hottest_mut(&mut self) -> Option<&mut T> { + self.hot.winner_mut() } - /// Pops the hottest node. - fn pop_hot(&mut self) -> Option { + /// Removes the hottest node before a transition that can fetch a batch. + fn pop_hot_for_batch_transition(&mut self) -> Option { self.hot.pop() } - /// Returns true if there are rows in the hot heap. + /// Returns a node to the appropriate heap after a batch transition. + fn reheap_after_batch_transition(&mut self, node: T) { + if node.is_eof() { + self.refill_hot(); + return; + } + + let node_is_cold = self + .hot + .peek() + .is_none_or(|hottest| node.is_behind(hottest)); + if node_is_cold { + self.cold.push(node); + } else { + self.hot.push(node); + } + self.refill_hot(); + } + + /// Repairs the hot tree after mutating its winner and refills the merge window. + fn repair_hot_root(&mut self) { + if self.hot.peek().is_some_and(NodeCmp::is_eof) { + self.hot.pop(); + } else { + let (winner, second_best) = self.hot.winner_and_second_best(); + let Some(winner) = winner else { + self.refill_hot(); + return; + }; + let root_is_cold = second_best.is_some_and(|best| winner.is_behind(best)); + // If the winner still wins (tie included), every cached champion on its + // path is unchanged and the tree invariant already holds, so the replay + // can be skipped. + let winner_lost = second_best.is_some_and(|best| winner < best); + if root_is_cold { + self.cold.push(self.hot.pop().unwrap()); + } else if winner_lost { + self.hot.replay_winner(); + } + } + + self.refill_hot(); + } + + /// Returns true if there are rows in the hot tree. fn has_rows(&self) -> bool { !self.hot.is_empty() } @@ -433,8 +755,11 @@ impl MergeAlgo { /// Columns to compare for a [RecordBatch]. struct SortColumns { primary_key: PrimaryKeyArray, + primary_key_values: BinaryArray, timestamp: Int64Array, sequence: UInt64Array, + #[cfg(test)] + primary_key_lookups: Cell, } impl SortColumns { @@ -450,6 +775,12 @@ impl SortColumns { .downcast_ref::() .unwrap() .clone(); + let primary_key_values = primary_key + .values() + .as_any() + .downcast_ref::() + .unwrap() + .clone(); let timestamp = batch.column(time_index_column_index(num_columns)); let (timestamp, _unit) = timestamp_array_to_primitive(timestamp).unwrap(); let sequence = batch @@ -461,20 +792,31 @@ impl SortColumns { Self { primary_key, + primary_key_values, timestamp, sequence, + #[cfg(test)] + primary_key_lookups: Cell::new(0), } } fn primary_key_at(&self, index: usize) -> &[u8] { - let key = self.primary_key.keys().value(index); - let binary_values = self - .primary_key - .values() - .as_any() - .downcast_ref::() - .unwrap(); - binary_values.value(key as usize) + let range = self.primary_key_range_at(index); + &self.primary_key_values.value_data()[range] + } + + fn primary_key_range_at(&self, index: usize) -> Range { + #[cfg(test)] + self.primary_key_lookups + .set(self.primary_key_lookups.get() + 1); + let key = self.primary_key.keys().value(index) as usize; + let offsets = self.primary_key_values.value_offsets(); + offsets[key].as_usize()..offsets[key + 1].as_usize() + } + + #[cfg(test)] + fn primary_key_lookups(&self) -> usize { + self.primary_key_lookups.get() } fn timestamp_at(&self, index: usize) -> i64 { @@ -497,6 +839,8 @@ impl SortColumns { struct RowCursor { /// Current row offset. offset: usize, + /// Byte range of the current primary key in the dictionary values. + primary_key_range: Range, /// Keys of the batch. columns: SortColumns, } @@ -504,8 +848,13 @@ struct RowCursor { impl RowCursor { fn new(columns: SortColumns) -> Self { debug_assert!(columns.num_rows() > 0); + let primary_key_range = columns.primary_key_range_at(0); - Self { offset: 0, columns } + Self { + offset: 0, + primary_key_range, + columns, + } } fn is_finished(&self) -> bool { @@ -519,10 +868,13 @@ impl RowCursor { fn advance(&mut self) { self.offset += 1; + if !self.is_finished() { + self.primary_key_range = self.columns.primary_key_range_at(self.offset); + } } fn first_primary_key(&self) -> &[u8] { - self.columns.primary_key_at(self.offset) + &self.columns.primary_key_values.value_data()[self.primary_key_range.clone()] } fn first_timestamp(&self) -> i64 { @@ -640,25 +992,29 @@ impl FlatMergeIterator { debug_assert!(self.in_progress.is_empty()); // Safety: next_batch() ensures the heap is not empty. - let mut hottest = self.algo.pop_hot().unwrap(); + let mut hottest = self.algo.pop_hot_for_batch_transition().unwrap(); debug_assert!(!hottest.current_cursor().is_finished()); + let node_index = hottest.node_index; let next = hottest.advance_batch()?; // The node is the heap is not empty, so it must have existing rows in the builder. - let batch = self - .in_progress - .take_remaining_rows(hottest.node_index, next); + let batch = self.in_progress.take_remaining_rows(node_index, next); Self::maybe_output_batch(batch, &mut self.output_batch); - self.algo.reheap(hottest); + self.algo.reheap_after_batch_transition(hottest); Ok(()) } /// Fetches a row from the hottest node. fn fetch_row_from_hottest(&mut self) -> Result<()> { - // Safety: next_batch() ensures the heap has more than 1 element. - let mut hottest = self.algo.pop_hot().unwrap(); - debug_assert!(!hottest.current_cursor().is_finished()); - self.in_progress.push_row(hottest.node_index); + let (node_index, at_batch_boundary) = { + // Safety: next_batch() ensures the heap has more than 1 element. + let hottest = self.algo.hottest_mut().unwrap(); + debug_assert!(!hottest.current_cursor().is_finished()); + (hottest.node_index, hottest.current_cursor().is_last_row()) + }; + let mut boundary_node = + at_batch_boundary.then(|| self.algo.pop_hot_for_batch_transition().unwrap()); + self.in_progress.push_row(node_index); if self.in_progress.len() >= self.batch_size { // We buffered enough rows. if let Some(output) = self.in_progress.build_record_batch()? { @@ -666,11 +1022,20 @@ impl FlatMergeIterator { } } - if let Some(next) = hottest.advance_row()? { - self.in_progress.push_batch(hottest.node_index, next); + let next = if let Some(hottest) = &mut boundary_node { + hottest.advance_row()? + } else { + self.algo.hottest_mut().unwrap().advance_row()? + }; + if let Some(next) = next { + self.in_progress.push_batch(node_index, next); } - self.algo.reheap(hottest); + if let Some(hottest) = boundary_node { + self.algo.reheap_after_batch_transition(hottest); + } else { + self.algo.repair_hot_root(); + } Ok(()) } @@ -796,27 +1161,31 @@ impl FlatMergeReader { debug_assert!(self.in_progress.is_empty()); // Safety: next_batch() ensures the heap is not empty. - let mut hottest = self.algo.pop_hot().unwrap(); + let mut hottest = self.algo.pop_hot_for_batch_transition().unwrap(); debug_assert!(!hottest.current_cursor().is_finished()); + let node_index = hottest.node_index; let start = Instant::now(); let next = hottest.advance_batch().await?; self.metrics.fetch_cost += start.elapsed(); // The node is the heap is not empty, so it must have existing rows in the builder. - let batch = self - .in_progress - .take_remaining_rows(hottest.node_index, next); + let batch = self.in_progress.take_remaining_rows(node_index, next); Self::maybe_output_batch(batch, &mut self.output_batch); - self.algo.reheap(hottest); + self.algo.reheap_after_batch_transition(hottest); Ok(()) } /// Fetches a row from the hottest node. async fn fetch_row_from_hottest(&mut self) -> Result<()> { - // Safety: next_batch() ensures the heap has more than 1 element. - let mut hottest = self.algo.pop_hot().unwrap(); - debug_assert!(!hottest.current_cursor().is_finished()); - self.in_progress.push_row(hottest.node_index); + let (node_index, at_batch_boundary) = { + // Safety: next_batch() ensures the heap has more than 1 element. + let hottest = self.algo.hottest_mut().unwrap(); + debug_assert!(!hottest.current_cursor().is_finished()); + (hottest.node_index, hottest.current_cursor().is_last_row()) + }; + let mut boundary_node = + at_batch_boundary.then(|| self.algo.pop_hot_for_batch_transition().unwrap()); + self.in_progress.push_row(node_index); if self.in_progress.len() >= self.batch_size { // We buffered enough rows. if let Some(output) = self.in_progress.build_record_batch()? { @@ -824,18 +1193,24 @@ impl FlatMergeReader { } } - // Only read the clock when advancing will attempt to fetch the next batch. - // Check first because `None` means either no fetch or EOF after a fetch attempt. - let start = hottest.current_cursor().is_last_row().then(Instant::now); - let next = hottest.advance_row().await?; + let start = at_batch_boundary.then(Instant::now); + let next = if let Some(hottest) = &mut boundary_node { + hottest.advance_row().await? + } else { + self.algo.hottest_mut().unwrap().advance_row().await? + }; if let Some(start) = start { self.metrics.fetch_cost += start.elapsed(); } if let Some(next) = next { - self.in_progress.push_batch(hottest.node_index, next); + self.in_progress.push_batch(node_index, next); } - self.algo.reheap(hottest); + if let Some(hottest) = boundary_node { + self.algo.reheap_after_batch_transition(hottest); + } else { + self.algo.repair_hot_root(); + } Ok(()) } @@ -1012,15 +1387,462 @@ impl GenericNode { #[cfg(test)] mod tests { + use std::cmp::Reverse; + use std::rc::Rc; use std::sync::Arc; + use std::sync::atomic::{AtomicBool, Ordering as AtomicOrdering}; + use std::task::Poll; use api::v1::OpType; use datatypes::arrow::array::builder::BinaryDictionaryBuilder; use datatypes::arrow::array::{Int64Array, TimestampMillisecondArray, UInt8Array, UInt64Array}; use datatypes::arrow::datatypes::{DataType, Field, Schema, TimeUnit, UInt32Type}; use datatypes::arrow::record_batch::RecordBatch; + use futures::FutureExt; use super::*; + use crate::error::UnexpectedSnafu; + + #[derive(Debug, Eq, PartialEq)] + struct TestNode { + id: usize, + current_rank: Option, + end_rank: usize, + } + + impl TestNode { + fn new(id: usize, current_rank: usize, end_rank: usize) -> Self { + Self { + id, + current_rank: Some(current_rank), + end_rank, + } + } + } + + impl NodeCmp for TestNode { + fn is_eof(&self) -> bool { + self.current_rank.is_none() + } + + fn is_behind(&self, other: &Self) -> bool { + self.current_rank.unwrap() > other.end_rank + } + } + + impl PartialOrd for TestNode { + fn partial_cmp(&self, other: &Self) -> Option { + Some(self.cmp(other)) + } + } + + impl Ord for TestNode { + fn cmp(&self, other: &Self) -> Ordering { + Reverse((self.current_rank, self.id)).cmp(&Reverse((other.current_rank, other.id))) + } + } + + #[test] + fn test_merge_algo_repairs_overlapping_hot_root_in_place() { + let mut algo = MergeAlgo::new(vec![ + TestNode::new(0, 0, 10), + TestNode::new(1, 5, 15), + TestNode::new(2, 20, 25), + ]); + assert_eq!(0, algo.hot.peek().unwrap().id); + assert_eq!((2, 1), (algo.hot.len(), algo.cold.len())); + + algo.hottest_mut().unwrap().current_rank = Some(6); + algo.repair_hot_root(); + + assert_eq!(1, algo.hot.peek().unwrap().id); + assert_eq!((2, 1), (algo.hot.len(), algo.cold.len())); + } + + #[test] + fn test_merge_algo_moves_root_beyond_remaining_hot_range_to_cold() { + let mut algo = MergeAlgo::new(vec![ + TestNode::new(0, 0, 10), + TestNode::new(1, 5, 7), + TestNode::new(2, 20, 25), + ]); + + algo.hottest_mut().unwrap().current_rank = Some(8); + algo.repair_hot_root(); + + assert_eq!(1, algo.hot.peek().unwrap().id); + assert_eq!((1, 2), (algo.hot.len(), algo.cold.len())); + } + + #[test] + fn test_merge_algo_removes_eof_root_and_refills_hot() { + let mut algo = MergeAlgo::new(vec![TestNode::new(0, 0, 4), TestNode::new(1, 10, 14)]); + assert_eq!((1, 1), (algo.hot.len(), algo.cold.len())); + + algo.hottest_mut().unwrap().current_rank = None; + algo.repair_hot_root(); + + assert_eq!(1, algo.hot.peek().unwrap().id); + assert_eq!((1, 0), (algo.hot.len(), algo.cold.len())); + } + + #[test] + fn test_merge_algo_single_hot_node_can_fetch_batch() { + let algo = MergeAlgo::new(vec![TestNode::new(0, 0, 4), TestNode::new(1, 10, 14)]); + + assert_eq!(0, algo.hot.peek().unwrap().id); + assert_eq!((1, 1), (algo.hot.len(), algo.cold.len())); + assert!(algo.can_fetch_batch()); + } + + /// A merge node that counts its `Ord::cmp` invocations, to assert how many + /// comparisons a repair performs. Unlike [TestNode], nodes with equal + /// `current_rank` compare equal (no id tie-break), like rows with equal + /// (primary key, timestamp, sequence). + #[derive(Debug)] + struct CountedNode { + id: usize, + current_rank: Option, + end_rank: usize, + compares: Rc>, + } + + impl CountedNode { + fn new(id: usize, current_rank: usize, compares: &Rc>) -> Self { + Self { + id, + current_rank: Some(current_rank), + // Never behind, so nodes never move to the cold heap. + end_rank: usize::MAX, + compares: Rc::clone(compares), + } + } + } + + impl NodeCmp for CountedNode { + fn is_eof(&self) -> bool { + self.current_rank.is_none() + } + + fn is_behind(&self, other: &Self) -> bool { + self.current_rank.unwrap() > other.end_rank + } + } + + impl PartialEq for CountedNode { + fn eq(&self, other: &Self) -> bool { + self.current_rank == other.current_rank + } + } + + impl Eq for CountedNode {} + + impl PartialOrd for CountedNode { + fn partial_cmp(&self, other: &Self) -> Option { + Some(self.cmp(other)) + } + } + + impl Ord for CountedNode { + fn cmp(&self, other: &Self) -> Ordering { + self.compares.set(self.compares.get() + 1); + Reverse(self.current_rank).cmp(&Reverse(other.current_rank)) + } + } + + fn drain_merge_algo(algo: &mut MergeAlgo) -> Vec { + let mut nodes = Vec::with_capacity(algo.hot.len()); + while let Some(node) = algo.pop_hot_for_batch_transition() { + nodes.push(node); + } + nodes + } + + #[test] + fn test_merge_algo_skips_tree_replay_when_winner_stays_hottest() { + let compares = Rc::new(Cell::new(0)); + let mut algo = MergeAlgo::new(vec![ + CountedNode::new(0, 10, &compares), + CountedNode::new(1, 50, &compares), + CountedNode::new(2, 40, &compares), + CountedNode::new(3, 30, &compares), + ]); + assert_eq!(0, algo.hot.peek().unwrap().id); + assert_eq!((4, 0), (algo.hot.len(), algo.cold.len())); + + // Advance the winner within its batch: it stays hotter than every + // other node, so it remains the champion. + algo.hottest_mut().unwrap().current_rank = Some(20); + compares.set(0); + algo.repair_hot_root(); + + // The second-best scan (1 compare) plus the champion-retention check + // (1 compare) suffice; a full replay would cost 2 more compares. + assert_eq!(2, compares.get()); + assert_eq!(0, algo.hot.peek().unwrap().id); + assert_eq!((4, 0), (algo.hot.len(), algo.cold.len())); + + // The tree still drains in merge order afterwards. + let drained: Vec<_> = drain_merge_algo(&mut algo) + .into_iter() + .map(|node| node.id) + .collect(); + assert_eq!(vec![0, 3, 2, 1], drained); + } + + #[test] + fn test_merge_algo_retains_winner_tied_with_second_best() { + let compares = Rc::new(Cell::new(0)); + let mut algo = MergeAlgo::new(vec![ + CountedNode::new(0, 10, &compares), + CountedNode::new(1, 20, &compares), + CountedNode::new(2, 20, &compares), + ]); + let winner_id = algo.hot.peek().unwrap().id; + + // The winner drops to exactly tie the second hottest node. + algo.hottest_mut().unwrap().current_rank = Some(20); + compares.set(0); + algo.repair_hot_root(); + + // A tie retains the champion without a replay: the same node stays + // the winner and nothing moves to the cold heap. + assert_eq!(2, compares.get()); + assert_eq!(winner_id, algo.hot.peek().unwrap().id); + assert_eq!((3, 0), (algo.hot.len(), algo.cold.len())); + + let mut drained: Vec<_> = drain_merge_algo(&mut algo) + .into_iter() + .map(|node| node.id) + .collect(); + drained.sort_unstable(); + assert_eq!(vec![0, 1, 2], drained); + } + + #[test] + fn test_merge_algo_caches_second_best_across_retention_repairs() { + let compares = Rc::new(Cell::new(0)); + let mut algo = MergeAlgo::new(vec![ + CountedNode::new(0, 10, &compares), + CountedNode::new(1, 50, &compares), + CountedNode::new(2, 40, &compares), + CountedNode::new(3, 30, &compares), + ]); + + // First retention repair computes the second-best slot. + algo.hottest_mut().unwrap().current_rank = Some(20); + algo.repair_hot_root(); + assert_eq!(0, algo.hot.peek().unwrap().id); + + // While the winner keeps its slot and the tree is structurally + // unchanged, repairs reuse the cached second-best slot: only the + // retention check itself (1 compare) runs per repair. + compares.set(0); + for rank in 21..=23 { + algo.hottest_mut().unwrap().current_rank = Some(rank); + algo.repair_hot_root(); + } + assert_eq!(3, compares.get()); + assert_eq!(0, algo.hot.peek().unwrap().id); + assert_eq!((4, 0), (algo.hot.len(), algo.cold.len())); + + // The tree still drains in merge order afterwards. + let drained: Vec<_> = drain_merge_algo(&mut algo) + .into_iter() + .map(|node| node.id) + .collect(); + assert_eq!(vec![0, 3, 2, 1], drained); + } + + fn drain_tournament_tree(tree: &mut TournamentTree) -> Vec { + let mut values = Vec::with_capacity(tree.len()); + while let Some(value) = tree.pop() { + values.push(value); + } + values + } + + #[test] + fn test_tournament_tree_empty() { + let mut tree = TournamentTree::::with_capacity(0); + + assert!(tree.is_empty()); + assert_eq!(0, tree.len()); + assert_eq!(None, tree.peek()); + assert_eq!(None, tree.winner_mut()); + assert_eq!(None, tree.second_best()); + assert_eq!(None, tree.pop()); + tree.replay_winner(); + } + + #[test] + fn test_tournament_tree_single_element() { + let mut tree = TournamentTree::with_capacity(1); + tree.push(7); + + assert!(!tree.is_empty()); + assert_eq!(1, tree.len()); + assert_eq!(Some(&7), tree.peek()); + assert_eq!(None, tree.second_best()); + assert_eq!(Some(7), tree.pop()); + assert!(tree.is_empty()); + } + + #[test] + fn test_tournament_tree_drains_in_descending_order() { + let mut tree = TournamentTree::with_capacity(8); + for value in [3, 1, 4, 1, 5, 9, 2, 6] { + tree.push(value); + } + + assert_eq!(Some(&9), tree.peek()); + assert_eq!(Some(&6), tree.second_best()); + assert_eq!( + vec![9, 6, 5, 4, 3, 2, 1, 1], + drain_tournament_tree(&mut tree) + ); + } + + #[test] + fn test_tournament_tree_non_power_of_two_capacity() { + let mut tree = TournamentTree::with_capacity(5); + for value in [40, 10, 50, 20, 30] { + tree.push(value); + } + + assert_eq!(Some(&50), tree.peek()); + assert_eq!(Some(&40), tree.second_best()); + assert_eq!(vec![50, 40, 30, 20, 10], drain_tournament_tree(&mut tree)); + } + + #[test] + fn test_tournament_tree_replays_winner_after_mutation() { + let mut tree = TournamentTree::with_capacity(4); + for value in [7, 3, 9, 5] { + tree.push(value); + } + + *tree.winner_mut().unwrap() = 1; + tree.replay_winner(); + + assert_eq!(Some(&7), tree.peek()); + assert_eq!(Some(&5), tree.second_best()); + assert_eq!(vec![7, 5, 3, 1], drain_tournament_tree(&mut tree)); + } + + #[test] + fn test_tournament_tree_mutated_winner_can_stay_winner() { + let mut tree = TournamentTree::with_capacity(3); + for value in [1, 2, 9] { + tree.push(value); + } + + *tree.winner_mut().unwrap() = 8; + tree.replay_winner(); + + assert_eq!(Some(&8), tree.peek()); + assert_eq!(vec![8, 2, 1], drain_tournament_tree(&mut tree)); + } + + #[test] + fn test_tournament_tree_remove_and_reinsert() { + let mut tree = TournamentTree::with_capacity(3); + for value in [5, 9, 7] { + tree.push(value); + } + + // Remove the winner; its slot is freed for a later reinsert. + assert_eq!(Some(9), tree.pop()); + tree.push(8); + assert_eq!(Some(&8), tree.peek()); + assert_eq!(vec![8, 7, 5], drain_tournament_tree(&mut tree)); + + // Refill after the tree was drained to empty. + tree.push(4); + tree.push(6); + assert_eq!(Some(&6), tree.peek()); + assert_eq!(vec![6, 4], drain_tournament_tree(&mut tree)); + } + + /// Drives a TournamentTree and a std BinaryHeap oracle with the same seeded op + /// sequence (push / pop winner / mutate winner + replay) and compares + /// observable behavior after every op. The number of live elements never + /// exceeds `capacity`, mirroring how MergeAlgo uses the tree. + fn assert_tournament_tree_matches_oracle( + seed: u64, + value_range: u32, + capacity: usize, + num_ops: usize, + ) { + use rand::rngs::StdRng; + use rand::{Rng, SeedableRng}; + + let mut rng = StdRng::seed_from_u64(seed); + let mut tree = TournamentTree::::with_capacity(capacity); + let mut oracle = BinaryHeap::::new(); + let mut next_value = 0_u32; + + for _ in 0..num_ops { + match rng.random_range(0..3) { + 0 if tree.len() < capacity => { + let pushed_value = next_value % value_range; + next_value += 1; + tree.push(pushed_value); + oracle.push(pushed_value); + } + 1 => { + assert_eq!(oracle.pop(), tree.pop()); + } + _ => { + let new_value = rng.random_range(0..value_range); + if let Some(winner) = tree.winner_mut() { + *winner = new_value; + tree.replay_winner(); + + oracle.pop(); + oracle.push(new_value); + } + } + } + + assert_eq!(oracle.peek(), tree.peek()); + assert_eq!(oracle.len(), tree.len()); + let oracle_second_best = { + let mut rest = oracle.clone(); + rest.pop(); + rest.peek().copied() + }; + assert_eq!(oracle_second_best, tree.second_best().copied()); + } + + // Both structures must drain in the same non-increasing order. + let mut oracle_values = Vec::with_capacity(oracle.len()); + while let Some(value) = oracle.pop() { + oracle_values.push(value); + } + assert_eq!(oracle_values, drain_tournament_tree(&mut tree)); + } + + #[test] + fn test_tournament_tree_matches_binary_heap_oracle() { + for seed in [0x5eed, 0xdead_beef, 42] { + assert_tournament_tree_matches_oracle(seed, 1000, 13, 2000); + } + } + + #[test] + fn test_tournament_tree_matches_oracle_with_duplicate_heavy_values() { + // A tiny value range makes duplicates dominate, which exercises the + // tie-breaking branches of the tree matches. + assert_tournament_tree_matches_oracle(0xc0ffee, 3, 8, 2000); + } + + #[test] + fn test_tournament_tree_matches_oracle_with_tiny_capacities() { + for capacity in 1..=3 { + assert_tournament_tree_matches_oracle(0xbeef, 100, capacity, 500); + } + } /// Creates a test RecordBatch with the specified data. fn create_test_record_batch( @@ -1074,6 +1896,38 @@ mod tests { Box::new(batches.into_iter().map(Ok)) } + fn boundary_test_batches() -> (RecordBatch, RecordBatch, RecordBatch) { + let first = create_test_record_batch( + &[b"k1", b"k1"], + &[1000, 2000], + &[1, 2], + &[OpType::Put, OpType::Put], + &[10, 12], + ); + let second = create_test_record_batch( + &[b"k1", b"k1", b"k1"], + &[1500, 2000, 2500], + &[1, 1, 1], + &[OpType::Put, OpType::Put, OpType::Put], + &[11, 13, 14], + ); + let pending = create_test_record_batch( + &[b"k1", b"k1", b"k1"], + &[1000, 1500, 2000], + &[1, 1, 2], + &[OpType::Put, OpType::Put, OpType::Put], + &[10, 11, 12], + ); + (first, second, pending) + } + + fn test_source_error() -> crate::error::Error { + UnexpectedSnafu { + reason: "test source failed".to_string(), + } + .build() + } + #[test] fn test_row_cursor_last_row() { let batch = create_test_record_batch( @@ -1260,6 +2114,76 @@ mod tests { assert_record_batches_eq(&expected, &result); } + #[test] + fn test_merge_iterator_retry_after_row_boundary_error_removes_source() { + let (first, second, pending) = boundary_test_batches(); + let schema = first.schema(); + let first_source = Box::new(vec![Ok(first), Err(test_source_error())].into_iter()) + as BoxedRecordBatchIterator; + let second_source = new_test_iter(vec![second.clone()]); + let mut merge = + FlatMergeIterator::new(schema, vec![first_source, second_source], 1024).unwrap(); + + assert!(merge.next_batch().is_err()); + assert_eq!(pending, merge.next_batch().unwrap().unwrap()); + assert_eq!(second.slice(1, 2), merge.next_batch().unwrap().unwrap()); + assert!(merge.next_batch().unwrap().is_none()); + } + + #[tokio::test] + async fn test_merge_reader_retry_after_row_boundary_error_removes_source() { + let (first, second, pending) = boundary_test_batches(); + let schema = first.schema(); + let first_source = Box::pin(futures::stream::iter(vec![ + Ok(first), + Err(test_source_error()), + ])) as BoxedRecordBatchStream; + let second_source = + Box::pin(futures::stream::iter(vec![Ok(second.clone())])) as BoxedRecordBatchStream; + let mut merge = FlatMergeReader::new(schema, vec![first_source, second_source], 1024, None) + .await + .unwrap(); + + assert!(merge.next_batch().await.is_err()); + assert_eq!(pending, merge.next_batch().await.unwrap().unwrap()); + assert_eq!( + second.slice(1, 2), + merge.next_batch().await.unwrap().unwrap() + ); + assert!(merge.next_batch().await.unwrap().is_none()); + } + + #[tokio::test] + async fn test_merge_reader_cancelled_row_boundary_fetch_removes_source() { + let (first, second, pending) = boundary_test_batches(); + let schema = first.schema(); + let fetch_pending = Arc::new(AtomicBool::new(false)); + let fetch_pending_on_poll = Arc::clone(&fetch_pending); + let mut first_batch = Some(first); + let first_source = Box::pin(futures::stream::poll_fn(move |_cx| { + if let Some(batch) = first_batch.take() { + Poll::Ready(Some(Ok(batch))) + } else { + fetch_pending_on_poll.store(true, AtomicOrdering::Relaxed); + Poll::Pending + } + })) as BoxedRecordBatchStream; + let second_source = + Box::pin(futures::stream::iter(vec![Ok(second.clone())])) as BoxedRecordBatchStream; + let mut merge = FlatMergeReader::new(schema, vec![first_source, second_source], 1024, None) + .await + .unwrap(); + + assert!(Box::pin(merge.next_batch()).now_or_never().is_none()); + assert!(fetch_pending.load(AtomicOrdering::Relaxed)); + assert_eq!(pending, merge.next_batch().await.unwrap().unwrap()); + assert_eq!( + second.slice(1, 2), + merge.next_batch().await.unwrap().unwrap() + ); + assert!(merge.next_batch().await.unwrap().is_none()); + } + #[test] fn test_batch_builder_basic() { let schema = Arc::new(Schema::new(vec![ @@ -1294,6 +2218,137 @@ mod tests { assert_eq!(result_batch.num_rows(), 2); } + #[test] + fn test_batch_builder_generic_three_column_schema() { + let schema = Arc::new(Schema::new(vec![ + Field::new("field1", DataType::Int64, false), + Field::new("field2", DataType::Int64, false), + Field::new("field3", DataType::Int64, false), + ])); + let batch = RecordBatch::try_new( + Arc::clone(&schema), + vec![ + Arc::new(Int64Array::from(vec![1, 2])), + Arc::new(Int64Array::from(vec![3, 4])), + Arc::new(Int64Array::from(vec![5, 6])), + ], + ) + .unwrap(); + let mut builder = BatchBuilder::new(schema, 1, 2); + builder.push_batch(0, batch.clone()); + builder.push_row(0); + builder.push_row(0); + + let output_batch = builder.build_record_batch().unwrap().unwrap(); + + assert_eq!(batch, output_batch); + } + + fn assert_primary_key_dictionary( + array: &dyn Array, + expected_decoded: &[&[u8]], + expected_values: &[&[u8]], + ) { + let dictionary = array.as_any().downcast_ref::().unwrap(); + let values = dictionary + .values() + .as_any() + .downcast_ref::() + .unwrap(); + let decoded: Vec<_> = dictionary + .keys() + .iter() + .map(|key| values.value(key.unwrap() as usize)) + .collect(); + let dictionary_values: Vec<_> = values.iter().map(Option::unwrap).collect(); + + assert_eq!(expected_decoded, decoded); + assert_eq!(expected_values, dictionary_values); + } + + #[test] + fn test_interleave_primary_key_deduplicates_separate_dictionaries() { + let batch0 = create_test_record_batch( + &[b"k1", b"k2"], + &[1000, 2000], + &[1, 1], + &[OpType::Put, OpType::Put], + &[10, 20], + ); + let batch1 = create_test_record_batch( + &[b"k1", b"k2"], + &[1000, 2000], + &[1, 1], + &[OpType::Put, OpType::Put], + &[11, 21], + ); + let pk_idx = primary_key_column_index(batch0.num_columns()); + let arrays: Vec<_> = [&batch0, &batch1] + .into_iter() + .map(|batch| batch.column(pk_idx).as_ref()) + .collect(); + + let output = interleave_primary_key(&arrays, &[(0, 0), (1, 0), (0, 1), (1, 1)]).unwrap(); + + assert_primary_key_dictionary( + output.as_ref(), + &[b"k1", b"k1", b"k2", b"k2"], + &[b"k1", b"k2"], + ); + } + + #[test] + fn test_interleave_primary_key_rejects_null_dictionary_value() { + let primary_key = PrimaryKeyArray::try_new( + UInt32Array::from(vec![0]), + Arc::new(BinaryArray::from(vec![None::<&[u8]>])), + ) + .unwrap(); + + let error = interleave_primary_key(&[&primary_key], &[(0, 0)]).unwrap_err(); + + assert!(error.to_string().contains("null dictionary value")); + } + + #[test] + fn test_batch_builder_primary_key_has_no_state_between_builds() { + let long_k1 = vec![b'a'; 4096]; + let long_k2 = vec![b'b'; 8192]; + let batch0 = + create_test_record_batch(&[long_k1.as_slice()], &[1000], &[1], &[OpType::Put], &[10]); + let batch1 = + create_test_record_batch(&[long_k1.as_slice()], &[1000], &[1], &[OpType::Put], &[11]); + let mut builder = BatchBuilder::new(batch0.schema(), 2, 4); + builder.push_batch(0, batch0); + builder.push_batch(1, batch1); + builder.push_row(0); + builder.push_row(1); + + let first = builder.build_record_batch().unwrap().unwrap(); + let pk_idx = primary_key_column_index(first.num_columns()); + assert_primary_key_dictionary( + first.column(pk_idx).as_ref(), + &[long_k1.as_slice(), long_k1.as_slice()], + &[long_k1.as_slice()], + ); + + let batch0 = + create_test_record_batch(&[long_k2.as_slice()], &[2000], &[2], &[OpType::Put], &[20]); + let batch1 = + create_test_record_batch(&[long_k2.as_slice()], &[2000], &[2], &[OpType::Put], &[21]); + builder.push_batch(0, batch0); + builder.push_batch(1, batch1); + builder.push_row(0); + builder.push_row(1); + + let second = builder.build_record_batch().unwrap().unwrap(); + assert_primary_key_dictionary( + second.column(pk_idx).as_ref(), + &[long_k2.as_slice(), long_k2.as_slice()], + &[long_k2.as_slice()], + ); + } + #[test] fn test_row_cursor_comparison() { // Create test batches for cursor comparison @@ -1322,4 +2377,19 @@ mod tests { // cursor1 has sequence 22, cursor2 has sequence 23, so cursor2 < cursor1 (higher sequence comes first) assert!(cursor2 < cursor1); } + + #[test] + fn test_row_cursor_caches_current_primary_key() { + let batch1 = create_test_record_batch(&[b"k1"], &[1000], &[1], &[OpType::Put], &[11]); + let batch2 = create_test_record_batch(&[b"k2"], &[1000], &[1], &[OpType::Put], &[12]); + let cursor1 = RowCursor::new(SortColumns::new(&batch1)); + let cursor2 = RowCursor::new(SortColumns::new(&batch2)); + + for _ in 0..5 { + assert_eq!(Ordering::Less, cursor1.cmp(&cursor2)); + } + + assert_eq!(1, cursor1.columns.primary_key_lookups()); + assert_eq!(1, cursor2.columns.primary_key_lookups()); + } }