use std::cmp::Ordering;
use std::marker::PhantomData;
use std::ops::{ControlFlow, Range};
use super::batch::{Batch, FIXED_REGION_BYTES, REG_NULL_BMP, REG_PAYLOAD_START, REG_WEIGHT};
use super::loser_tree::{HeapNode, LoserTree};
use super::scatter::DecodedColumns;
use super::string_heap::{rebase_string_cells, row_long_bytes, BlobCache};
use super::writer::DirectWriter;
use crate::schema::key::{pack_pk_be, pk_width_dispatch, PkSortKey};
use crate::schema::payload_order::{compare_full_rows, with_payload_cmp, PayloadOrder};
use crate::schema::SchemaDescriptor;
use gnitz_expr::BatchView;
use gnitz_wire::read_u64_le;
use gnitz_wire::RowSource;
use gnitz_wire::NARROW_PK_MAX_BYTES;
#[derive(Clone, Copy)]
pub(crate) struct ColPtr {
pub base: *const u8,
pub stride: usize,
}
impl ColPtr {
#[inline(always)]
pub(crate) fn row_ptr(self, i: usize) -> *const u8 {
self.base.wrapping_add(i * self.stride)
}
#[inline(always)]
pub(crate) unsafe fn row<'a>(self, i: usize, len: usize) -> &'a [u8] {
std::slice::from_raw_parts(self.row_ptr(i), len)
}
pub(crate) unsafe fn copy_rows(self, start: usize, width: usize, dst: &mut [u8]) {
if self.stride != 0 {
return dst.copy_from_slice(std::slice::from_raw_parts(self.row_ptr(start), dst.len()));
}
if dst.is_empty() {
return;
}
dst[..width].copy_from_slice(self.row(0, width));
let mut filled = width;
while filled < dst.len() {
let n = filled.min(dst.len() - filled);
dst.copy_within(..n, filled);
filled += n;
}
}
}
#[derive(Clone, Copy)]
pub(crate) struct UnifiedSource<'a> {
pub pk: ColPtr,
pub null_bmp: ColPtr,
pub null_pad_mask: u64,
pub cols_off: usize,
pub blob: &'a [u8],
pub heap_at: Option<usize>,
}
pub(crate) fn mem_batch_to_unified<'a>(
mb: &MemBatch<'a>,
schema: &SchemaDescriptor,
cols: &mut Vec<ColPtr>,
) -> UnifiedSource<'a> {
let data_ptr = mb.data.as_ptr();
let cols_off = cols.len();
for (pi, col) in schema.payload_columns() {
cols.push(ColPtr {
base: unsafe { data_ptr.add(mb.region_start(REG_PAYLOAD_START + pi)) },
stride: col.size() as usize,
});
}
UnifiedSource {
pk: ColPtr { base: data_ptr, stride: mb.pk_stride() },
null_bmp: ColPtr {
base: unsafe { data_ptr.add(mb.region_start(REG_NULL_BMP)) },
stride: FIXED_REGION_BYTES,
},
null_pad_mask: 0,
cols_off,
blob: mb.blob,
heap_at: None,
}
}
#[derive(Clone)]
pub struct MemBatch<'a> {
pub(crate) data: &'a [u8],
pub(crate) schema: &'a SchemaDescriptor,
pub(crate) cap: usize,
pub(crate) blob: &'a [u8],
pub(crate) count: usize,
pub(in crate::repr) dead_heap: usize,
}
impl<'a> MemBatch<'a> {
#[inline]
pub fn len(&self) -> usize {
self.count
}
#[inline]
pub fn is_empty(&self) -> bool {
self.count == 0
}
#[inline(always)]
pub fn pk_stride(&self) -> usize {
self.schema.pk_stride()
}
#[inline(always)]
pub(crate) fn region_start(&self, r: usize) -> usize {
self.schema.region_start(r, self.cap)
}
#[inline(always)]
pub(crate) fn region(&self, r: usize) -> &'a [u8] {
let off = self.region_start(r);
&self.data[off..off + self.count * self.schema.region_stride(r)]
}
#[inline]
pub fn pk(&self) -> &'a [u8] {
&self.data[..self.count * self.pk_stride()]
}
#[inline]
pub(crate) fn weight(&self) -> &'a [u8] {
let off = self.region_start(REG_WEIGHT);
&self.data[off..off + self.count * 8]
}
#[inline]
pub fn sum_weights(&self, start: usize, end: usize) -> i64 {
self.weight()[start * 8..end * 8]
.as_chunks::<8>()
.0
.iter()
.fold(0i64, |a, w| a.wrapping_add(i64::from_le_bytes(*w)))
}
#[inline(always)]
pub(crate) fn null_bmp(&self) -> &'a [u8] {
let off = self.region_start(REG_NULL_BMP);
&self.data[off..off + self.count * 8]
}
#[inline(always)]
pub fn col_data(&self, pi: usize, stride: usize) -> &'a [u8] {
let off = self.region_start(REG_PAYLOAD_START + pi);
&self.data[off..off + self.count * stride]
}
#[inline(always)]
pub fn get_pk_bytes(&self, row: usize) -> &'a [u8] {
let stride = self.pk_stride();
let off = row * stride;
&self.data[off..off + stride]
}
#[inline(always)]
pub(crate) fn get_pk_range(&self, row: usize, at: usize, n: usize) -> &'a [u8] {
debug_assert!(at + n <= self.pk_stride());
let off = row * self.pk_stride() + at;
&self.data[off..off + n]
}
#[inline(always)]
pub fn get_weight(&self, row: usize) -> i64 {
gnitz_wire::read_i64_le(self.data, self.region_start(REG_WEIGHT) + row * 8)
}
#[inline(always)]
pub fn get_null_word(&self, row: usize) -> u64 {
read_u64_le(self.data, self.region_start(REG_NULL_BMP) + row * 8)
}
#[inline(always)]
pub(crate) fn get_col_ptr(&self, row: usize, payload_col: usize, col_size: usize) -> &'a [u8] {
let off = self.region_start(REG_PAYLOAD_START + payload_col) + row * col_size;
&self.data[off..off + col_size]
}
}
impl<'a> RowSource for MemBatch<'a> {
#[inline(always)]
fn get_pk_bytes(&self, row: usize) -> &[u8] {
MemBatch::get_pk_bytes(self, row)
}
#[inline(always)]
fn get_null_word(&self, row: usize) -> u64 {
MemBatch::get_null_word(self, row)
}
#[inline(always)]
fn get_col_ptr(&self, row: usize, payload_col: usize, col_size: usize) -> &[u8] {
MemBatch::get_col_ptr(self, row, payload_col, col_size)
}
#[inline(always)]
fn blob(&self) -> &[u8] {
self.blob
}
#[inline(always)]
fn row_count(&self) -> usize {
MemBatch::len(self)
}
}
impl<'a> BatchView for MemBatch<'a> {
#[inline(always)]
fn col_data(&self, payload_col: usize, col_size: usize) -> &[u8] {
MemBatch::col_data(self, payload_col, col_size)
}
#[inline(always)]
fn null_bmp(&self) -> &[u8] {
MemBatch::null_bmp(self)
}
#[inline(always)]
fn pk_region(&self) -> (&[u8], usize) {
(MemBatch::pk(self), self.pk_stride())
}
}
impl<'a> ColumnarSource for MemBatch<'a> {
#[inline(always)]
fn get_weight(&self, row: usize) -> i64 {
MemBatch::get_weight(self, row)
}
fn to_unified(
&self,
schema: &SchemaDescriptor,
cols: &mut Vec<ColPtr>,
_window: Range<usize>,
_decoded: &mut DecodedColumns,
) -> UnifiedSource<'_> {
mem_batch_to_unified(self, schema, cols)
}
#[inline(always)]
fn is_skeleton(&self) -> bool {
false
}
}
pub(crate) struct PosCursor {
pub(crate) position: usize,
pub(crate) count: usize,
}
impl PosCursor {
#[inline]
pub(crate) fn new(count: usize) -> Self {
Self::over(0..count)
}
#[inline]
pub(crate) fn over(rows: Range<usize>) -> Self {
debug_assert!(
rows.end < u32::MAX as usize,
"merge source exceeds the heap node's u32 row"
);
PosCursor { position: rows.start, count: rows.end }
}
#[inline]
pub(crate) fn is_valid(&self) -> bool {
self.position < self.count
}
#[inline]
pub(crate) fn advance(&mut self) {
self.position += 1;
}
}
pub(crate) fn run_merge_in<S: ColumnarSource>(
sources: &[S],
schema: &SchemaDescriptor,
windows: impl IntoIterator<Item = Range<usize>>,
emit: impl FnMut(usize, usize, i64),
) {
if sources.is_empty() {
return;
}
let mut cursors: Vec<PosCursor> = windows.into_iter().map(PosCursor::over).collect();
debug_assert_eq!(cursors.len(), sources.len(), "run_merge_in: one window per source");
with_payload_cmp!(schema, run_merge_body, sources, &mut cursors, schema, emit)
}
pub(crate) fn run_merge<S: ColumnarSource>(
sources: &[S],
schema: &SchemaDescriptor,
emit: impl FnMut(usize, usize, i64),
) {
run_merge_in(sources, schema, sources.iter().map(|s| 0..s.row_count()), emit)
}
pub fn merge_consolidated(sources: &[MemBatch<'_>], schema: &SchemaDescriptor) -> Batch {
let held = sources.iter().filter(|s| s.count > 0);
let ascending = held
.clone()
.zip(held.skip(1))
.all(|(a, b)| a.get_pk_bytes(a.count - 1) < b.get_pk_bytes(0));
let mut out = match ascending {
true => Batch::concat(schema, sources.iter().cloned()),
false => merge_rows(sources, schema),
};
out.certify_consolidated();
out
}
#[inline(never)]
fn merge_rows(sources: &[MemBatch<'_>], schema: &SchemaDescriptor) -> Batch {
let mut rows: Vec<(u32, u32, i64)> = Vec::with_capacity(sources.iter().map(|s| s.count).sum());
run_merge(sources, schema, |src, row, w| rows.push((src as u32, row as u32, w)));
super::scatter::materialize_carrying(sources, schema, &rows)
}
pub(crate) trait ColumnarSource: RowSource {
fn get_weight(&self, row: usize) -> i64;
fn to_unified(
&self,
schema: &SchemaDescriptor,
cols: &mut Vec<ColPtr>,
window: Range<usize>,
decoded: &mut DecodedColumns,
) -> UnifiedSource<'_>;
fn is_skeleton(&self) -> bool;
}
impl<T: ColumnarSource + ?Sized> ColumnarSource for &T {
#[inline(always)]
fn get_weight(&self, row: usize) -> i64 {
(**self).get_weight(row)
}
#[inline(always)]
fn to_unified(
&self,
schema: &SchemaDescriptor,
cols: &mut Vec<ColPtr>,
window: Range<usize>,
decoded: &mut DecodedColumns,
) -> UnifiedSource<'_> {
(**self).to_unified(schema, cols, window, decoded)
}
#[inline(always)]
fn is_skeleton(&self) -> bool {
(**self).is_skeleton()
}
}
#[inline]
pub(crate) fn merge_less<'a, S, P>(
schema: &'a SchemaDescriptor,
sources: &'a [S],
order: MergeOrder,
payload: P,
) -> impl Fn(&HeapNode, &HeapNode) -> bool + Copy + 'a
where
S: ColumnarSource,
P: PayloadOrder + 'a,
{
let wide = order.has_tail(schema.pk_stride());
move |a, b| {
let (a_src, a_row) = (a.source_idx as usize, a.row as usize);
let (b_src, b_row) = (b.source_idx as usize, b.row as usize);
let pk = match wide {
true => order
.tail(&sources[a_src], a_row)
.cmp(order.tail(&sources[b_src], b_row)),
false => Ordering::Equal,
};
match pk {
Ordering::Less => true,
Ordering::Greater => false,
Ordering::Equal => {
if order.coarsen {
let (sa, sb) = (sources[a_src].is_skeleton(), sources[b_src].is_skeleton());
if sa || sb {
return sa && !sb;
}
}
payload.compare(schema, &sources[a_src], a_row, &sources[b_src], b_row) == Ordering::Less
}
}
}
}
#[derive(Clone, Copy)]
pub(crate) struct MergeOrder {
wide_at: Option<usize>,
coarsen: bool,
}
impl MergeOrder {
pub(crate) fn leading(stride: usize, coarsen: bool) -> Self {
MergeOrder {
wide_at: (stride > NARROW_PK_MAX_BYTES).then_some(0),
coarsen,
}
}
fn past_shared_prefix<S: ColumnarSource>(sources: &[S], stride: usize) -> Self {
if stride <= NARROW_PK_MAX_BYTES {
return Self::leading(stride, false);
}
let mut bounds = sources
.iter()
.filter(|s| s.row_count() > 0)
.flat_map(|s| [s.get_pk_bytes(0), s.get_pk_bytes(s.row_count() - 1)]);
let first = bounds.next().unwrap_or(&[]);
let shared = |pk: &[u8]| first.iter().zip(pk).take_while(|(a, b)| a == b).count();
let at = bounds.map(shared).min().unwrap_or(0);
MergeOrder {
wide_at: Some(at.min(stride - NARROW_PK_MAX_BYTES)),
coarsen: false,
}
}
#[inline(always)]
fn has_tail(self, stride: usize) -> bool {
self.wide_at.is_some_and(|at| at + NARROW_PK_MAX_BYTES < stride)
}
#[inline(always)]
pub(crate) fn key<S: ColumnarSource>(self, sources: &[S], src: usize, row: usize) -> u128 {
let pk = sources[src].get_pk_bytes(row);
match self.wide_at {
None => pack_pk_be(pk),
Some(at) => u128::from_be_bytes(*pk[at..].first_chunk().unwrap()),
}
}
#[inline(always)]
fn tail<S: ColumnarSource>(self, source: &S, row: usize) -> &[u8] {
&source.get_pk_bytes(row)[self.wide_at.unwrap_or(0) + NARROW_PK_MAX_BYTES..]
}
}
#[inline(always)]
pub(crate) fn drive<S, P>(
tree: &mut LoserTree,
schema: &SchemaDescriptor,
sources: &[S],
order: MergeOrder,
cursors: &mut [PosCursor],
payload: P,
mut emit: impl FnMut(usize, usize, i64) -> ControlFlow<()>,
) where
S: ColumnarSource,
P: PayloadOrder,
{
let less = merge_less(schema, sources, order, payload);
let wide = order.has_tail(schema.pk_stride());
let same_group = |a_src: usize, a_row: usize, b_src: usize, b_row: usize| {
if wide && order.tail(&sources[a_src], a_row) != order.tail(&sources[b_src], b_row) {
return false;
}
if order.coarsen && (sources[a_src].is_skeleton() || sources[b_src].is_skeleton()) {
return true;
}
payload.compare(schema, &sources[a_src], a_row, &sources[b_src], b_row) == Ordering::Equal
};
macro_rules! step {
($src:expr) => {{
let c = &mut cursors[$src];
c.advance();
let next = c
.is_valid()
.then(|| (c.position as u32, order.key(sources, $src, c.position)));
tree.step_top(next, &less);
}};
}
while let Some(top) = tree.peek() {
let (group_src, group_row, group_key) = (top.source_idx as usize, top.row as usize, top.key());
let mut net_weight: i64 = sources[group_src].get_weight(group_row);
step!(group_src);
while let Some(top) = tree.peek() {
let (cur_src, cur_row) = (top.source_idx as usize, top.row as usize);
if top.key() != group_key || !same_group(group_src, group_row, cur_src, cur_row) {
break;
}
net_weight += sources[cur_src].get_weight(cur_row);
step!(cur_src);
}
if net_weight != 0 && emit(group_src, group_row, net_weight).is_break() {
return;
}
}
}
#[inline]
fn run_merge_body<S, P>(
sources: &[S],
cursors: &mut [PosCursor],
schema: &SchemaDescriptor,
mut emit: impl FnMut(usize, usize, i64),
payload: P,
) where
S: ColumnarSource,
P: PayloadOrder,
{
let order = MergeOrder::past_shared_prefix(sources, schema.pk_stride());
let mut tree = LoserTree::build(
cursors.len(),
|i| {
cursors[i]
.is_valid()
.then(|| (cursors[i].position as u32, order.key(sources, i, cursors[i].position)))
},
merge_less(schema, sources, order, payload),
);
drive(&mut tree, schema, sources, order, cursors, payload, |src, row, w| {
emit(src, row, w);
ControlFlow::Continue(())
});
}
impl Batch {
pub fn merged_consolidated(&self, other: &Batch, schema: &SchemaDescriptor) -> Batch {
debug_assert!(
self.stands_consolidated() && other.stands_consolidated(),
"merged_consolidated: both inputs must be consolidated",
);
let mut out = pk_width_dispatch!(schema.pk_stride(), |K| {
with_payload_cmp!(schema, merged_consolidated_body::<K, _>, self, other, schema)
});
out.certify_consolidated();
out
}
}
struct Keys<'a, K> {
pk: &'a [u8],
stride: usize,
count: usize,
key: PhantomData<K>,
}
impl<'a, K: PkSortKey<'a>> Keys<'a, K> {
fn of(batch: &'a Batch) -> Self {
Keys {
pk: batch.pk_data(),
stride: batch.schema().pk_stride(),
count: batch.count,
key: PhantomData,
}
}
#[inline(always)]
fn at(&self, row: usize) -> K {
assert!(row < self.count);
unsafe { self.at_unchecked(row) }
}
#[inline(always)]
unsafe fn at_unchecked(&self, row: usize) -> K {
debug_assert!(row < self.count);
K::from_opk(unsafe { self.pk.get_unchecked(row * self.stride..(row + 1) * self.stride) })
}
#[inline(always)]
fn group_end(&self, row: usize, key: K) -> usize {
let mut end = row + 1;
while end < self.count && unsafe { self.at_unchecked(end) } == key {
end += 1;
}
end
}
#[inline(always)]
fn skip_below(&self, row: usize, key: K) -> usize {
let (mut lo, mut step) = (row, 1);
while lo + step < self.count && unsafe { self.at_unchecked(lo + step) } < key {
lo += step;
step *= 2;
}
let (mut lo, mut hi) = (lo + 1, (lo + step).min(self.count));
while lo < hi {
let mid = lo + (hi - lo) / 2;
match unsafe { self.at_unchecked(mid) } < key {
true => lo = mid + 1,
false => hi = mid,
}
}
lo
}
}
struct MergeRegion<'w> {
src: [&'w [u8]; 2],
dst: &'w mut [u8],
stride: usize,
}
struct MergeSink<'w> {
srcs: &'w [MemBatch<'w>; 2],
heap_at: [Option<usize>; 2],
regions: Vec<MergeRegion<'w>>,
blob: &'w mut Vec<u8>,
cache: Option<&'w mut BlobCache>,
string_slots: u64,
rows: usize,
room: usize,
dead: usize,
}
impl<'w> MergeSink<'w> {
fn new(srcs: &'w [MemBatch<'w>; 2], heap_at: [Option<usize>; 2], writer: &'w mut DirectWriter<'_>) -> Self {
let (schema, room) = (writer.schema, writer.rows());
let (regions, blob, cache) = writer.split_mut();
let regions = regions.enumerate().map(|(r, dst)| MergeRegion {
src: [srcs[0].region(r), srcs[1].region(r)],
dst,
stride: schema.region_stride(r),
});
MergeSink {
srcs,
heap_at,
regions: regions.collect(),
blob,
cache,
string_slots: schema.string_payload_slots(),
rows: 0,
room,
dead: 0,
}
}
#[inline(always)]
fn push(&mut self, side: usize, start: usize, end: usize) {
if start == end {
return;
}
let at = self.rows;
assert!(start < end && end <= self.srcs[side].count && at + (end - start) <= self.room);
let n = end - start;
if n == 1 {
for region in &mut self.regions {
let w = region.stride;
unsafe {
copy_cell(
region.src[side].as_ptr().add(start * w),
region.dst.as_mut_ptr().add(at * w),
w,
)
};
}
} else {
for region in &mut self.regions {
let w = region.stride;
let src = unsafe { region.src[side].as_ptr().add(start * w) };
unsafe { std::ptr::copy_nonoverlapping(src, region.dst.as_mut_ptr().add(at * w), n * w) };
}
}
for pi in gnitz_wire::BitIter(self.string_slots) {
let cells = &mut self.regions[REG_PAYLOAD_START + pi].dst[at * 16..(at + n) * 16];
let (blob, heap_at) = (self.srcs[side].blob, self.heap_at[side]);
rebase_string_cells(cells, blob, self.blob, heap_at, self.cache.as_deref_mut());
}
self.rows += n;
}
#[inline(always)]
fn push_folded(&mut self, ia: usize, jb: usize, weight: i64) {
self.leave_out(1, jb);
if weight == 0 {
return self.leave_out(0, ia);
}
self.push(0, ia, ia + 1);
let weights = self.regions[REG_WEIGHT].dst.as_chunks_mut::<8>().0;
weights[self.rows - 1] = weight.to_le_bytes();
}
#[inline(always)]
fn leave_out(&mut self, side: usize, row: usize) {
if self.heap_at[side].is_some() {
self.dead += row_long_bytes(&self.srcs[side], self.string_slots, row);
}
}
}
#[inline(always)]
unsafe fn copy_cell(src: *const u8, dst: *mut u8, width: usize) {
use std::ptr::copy_nonoverlapping as copy;
unsafe {
match width {
8 => copy(src, dst, 8),
16 => copy(src, dst, 16),
4 => copy(src, dst, 4),
_ => copy(src, dst, width),
}
}
}
#[inline(never)]
fn merged_consolidated_body<'a, K, P>(a: &'a Batch, b: &'a Batch, schema: &SchemaDescriptor, payload: P) -> Batch
where
K: PkSortKey<'a>,
P: PayloadOrder,
{
let (n_a, n_b) = (a.count, b.count);
let srcs = [a.as_mem_batch(), b.as_mem_batch()];
let (keys_a, keys_b) = (Keys::<K>::of(a), Keys::<K>::of(b));
let mut out = Batch::with_capacity(schema, n_a + n_b);
let mut session = out.append_session(n_a + n_b);
let heap_at = [
session.carry(&srcs[0], &[(0, n_a)]),
session.carry(&srcs[1], &[(0, n_b)]),
];
let dead = session.write_at_most(n_a + n_b, |writer| {
let mut sink = MergeSink::new(&srcs, heap_at, writer);
let (mut ia, mut jb) = (0usize, 0usize);
while ia < n_a && jb < n_b {
let (ka, kb) = (keys_a.at(ia), keys_b.at(jb));
match ka.cmp(&kb) {
Ordering::Less => {
let start = ia;
ia = keys_a.skip_below(ia, kb);
sink.push(0, start, ia);
}
Ordering::Greater => {
let start = jb;
jb = keys_b.skip_below(jb, ka);
sink.push(1, start, jb);
}
Ordering::Equal => {
let (ga, gb) = (keys_a.group_end(ia, ka), keys_b.group_end(jb, ka));
while ia < ga && jb < gb {
match payload.compare(schema, &srcs[0], ia, &srcs[1], jb) {
Ordering::Less => {
sink.push(0, ia, ia + 1);
ia += 1;
}
Ordering::Greater => {
sink.push(1, jb, jb + 1);
jb += 1;
}
Ordering::Equal => {
sink.push_folded(ia, jb, a.get_weight(ia) + b.get_weight(jb));
ia += 1;
jb += 1;
}
}
}
sink.push(0, ia, ga);
sink.push(1, jb, gb);
(ia, jb) = (ga, gb);
}
}
}
sink.push(0, ia, n_a);
sink.push(1, jb, n_b);
(sink.rows, sink.dead)
});
out.charge_dead(dead);
out
}
pub(crate) fn consolidate_groups(batch: &MemBatch, schema: &SchemaDescriptor, out: &mut Vec<(u32, u32, i64)>) {
let n = batch.count;
if n == 0 {
return;
}
with_payload_cmp!(schema, consolidate_groups_inner, n, batch, schema, out)
}
pub(crate) fn in_consolidated_order(batch: &Batch) -> bool {
if batch.count == 0 {
return true;
}
let schema = batch.schema();
let stride = schema.pk_stride();
let ascending = pk_width_dispatch!(stride, |K| {
let mut keys = batch.pk_data().chunks_exact(stride).map(K::from_opk);
let mut prev = keys.next().expect("a batch of at least one row");
keys.zip(1..).all(|(key, i)| {
let below = Ord::cmp(&prev, &key).then_with(|| compare_full_rows(schema, batch, i - 1, batch, i));
prev = key;
below.is_lt()
})
});
ascending && !batch.has_ghost()
}
trait ArgEntry: Copy {
fn idx(self) -> usize;
fn same_pk(self, other: Self) -> bool;
}
#[derive(Copy, Clone)]
struct SortEntry<K> {
key: K,
idx: u32,
}
impl<K: Copy + Eq> ArgEntry for SortEntry<K> {
#[inline(always)]
fn idx(self) -> usize {
self.idx as usize
}
#[inline(always)]
fn same_pk(self, other: Self) -> bool {
self.key == other.key
}
}
#[derive(Copy, Clone, PartialEq, Eq, PartialOrd, Ord)]
struct PackedEntry([u128; 2]);
const PACKED_MAX_STRIDE: usize = size_of::<PackedEntry>() - size_of::<u32>();
impl PackedEntry {
#[inline(always)]
fn new(opk: &[u8], idx: u32) -> Self {
debug_assert!((17..=PACKED_MAX_STRIDE).contains(&opk.len()));
let [hi, lo] = <[u128; 2]>::from_opk(opk);
PackedEntry([hi, lo | idx as u128])
}
}
impl ArgEntry for PackedEntry {
#[inline(always)]
fn idx(self) -> usize {
self.0[1] as u32 as usize
}
#[inline(always)]
fn same_pk(self, other: Self) -> bool {
self.0[0] == other.0[0] && (self.0[1] ^ other.0[1]) >> 32 == 0
}
}
#[inline]
fn consolidate_groups_inner<P: PayloadOrder>(
n: usize,
batch: &MemBatch,
schema: &SchemaDescriptor,
out: &mut Vec<(u32, u32, i64)>,
payload: P,
) {
let stride = batch.pk_stride();
if (17..=PACKED_MAX_STRIDE).contains(&stride) {
let mut entries: Vec<PackedEntry> = (0..n as u32)
.map(|i| PackedEntry::new(batch.get_pk_bytes(i as usize), i))
.collect();
entries.sort_unstable();
if schema.num_payload_cols() > 0 {
for run in entries.chunk_by_mut(|a, b| a.same_pk(*b)).filter(|r| r.len() > 1) {
run.sort_unstable_by(|a, b| payload.compare(schema, batch, a.idx(), batch, b.idx()));
}
}
return drain_groups(&entries, batch, schema, payload, out);
}
pk_width_dispatch!(stride, |K| {
let mut entries: Vec<SortEntry<K>> = (0..n as u32)
.map(|i| SortEntry {
key: K::from_opk(batch.get_pk_bytes(i as usize)),
idx: i,
})
.collect();
entries.sort_unstable_by(|a, b| {
let (x, y) = (a.idx as usize, b.idx as usize);
Ord::cmp(&a.key, &b.key).then_with(|| payload.compare(schema, batch, x, batch, y))
});
drain_groups(&entries, batch, schema, payload, out);
})
}
#[inline]
fn drain_groups<E: ArgEntry, P: PayloadOrder>(
entries: &[E],
batch: &MemBatch,
schema: &SchemaDescriptor,
payload: P,
out: &mut Vec<(u32, u32, i64)>,
) {
let mut pending = entries[0];
let mut pending_weight = batch.get_weight(pending.idx());
for &cur in &entries[1..] {
let (pi, ci) = (pending.idx(), cur.idx());
if pending.same_pk(cur) && payload.compare(schema, batch, pi, batch, ci).is_eq() {
pending_weight += batch.get_weight(ci);
} else {
if pending_weight != 0 {
out.push((0, pi as u32, pending_weight));
}
pending = cur;
pending_weight = batch.get_weight(ci);
}
}
if pending_weight != 0 {
out.push((0, pending.idx() as u32, pending_weight));
}
}
#[cfg(test)]
#[path = "tests/merge.rs"]
mod tests;
#[cfg(test)]
#[path = "benches/merge.rs"]
mod bench;