use std::cmp::Ordering;
use crate::repr::{
copy_runs, pk_group_end, pk_prefix_group_end, relocate_german_string_vec, runs_where, should_relocate_blob,
width_dispatch, Batch, ReadCursor,
};
use crate::schema::key::{compare_pk_ordering, key_range_between_cuts, KeyCut};
use crate::schema::{DerivedSchema, SchemaDescriptor, MAX_PK_BYTES};
use crate::algebra::MapPlan;
use crate::stream::OpenAt;
use gnitz_expr::{ColCopy, ColumnLocator, NullPerm};
use gnitz_wire::RowSource;
use gnitz_wire::{null_word_at, JoinKind, PkBuf, RangeRel, TypeCode};
pub struct JoinPlan {
out_schema: SchemaDescriptor,
walk: Walk,
t_cols: Box<[ColCopy]>,
t_nulls: NullPerm,
trace_leads: bool,
trace_ordered: bool,
}
#[derive(Clone, Copy)]
enum Walk {
Equi,
Range(RangeProbe),
Cross,
}
impl JoinPlan {
pub fn out_schema(&self) -> &SchemaDescriptor {
&self.out_schema
}
pub fn from_wire(
kind: JoinKind,
delta_is_right: bool,
delta: &SchemaDescriptor,
trace: &SchemaDescriptor,
) -> Result<JoinPlan, String> {
Self::new(kind, delta_is_right, delta, trace, Walked::Trace)
}
pub fn over_source(
delta_is_right: bool,
delta: &SchemaDescriptor,
source: &SchemaDescriptor,
rekey: &MapPlan,
) -> Result<JoinPlan, String> {
let cols = rekey
.rekeys_onto_pk_prefix()
.ok_or("join: the trace is not its source re-keyed onto leading primary-key columns")?;
Self::new(
JoinKind::Equi,
delta_is_right,
delta,
rekey.out_schema(),
Walked::Source(source, cols),
)
}
fn new(
kind: JoinKind,
delta_is_right: bool,
delta: &SchemaDescriptor,
trace: &SchemaDescriptor,
walked: Walked<'_>,
) -> Result<JoinPlan, String> {
let trace_ordered = matches!(walked, Walked::Trace);
let (walked, t_cols) = match walked {
Walked::Trace => (trace, trace.payload_locators()),
Walked::Source(source, cols) => (source, cols),
};
let (left, right) = match delta_is_right {
true => (trace, delta),
false => (delta, trace),
};
let mut b = DerivedSchema::new();
b.push_pk_of(left);
if kind == JoinKind::Cross {
b.push_pk_of(right);
}
b.push_payload_of(left);
b.push_payload_of(right);
let out_schema = b.finish().map_err(|e| format!("join: merged schema {e}"))?;
let t_first = match delta_is_right {
true => 0,
false => delta.num_payload_cols(),
};
let walk = match kind {
JoinKind::Equi => {
same_pk_types(delta, trace)?;
Walk::Equi
}
JoinKind::Range { rel } => {
same_pk_types(delta, trace)?;
Walk::Range(RangeProbe::new(trace, rel, delta_is_right))
}
JoinKind::Cross => Walk::Cross,
};
let t_cols: Box<[ColCopy]> = (t_first..)
.zip(t_cols)
.map(|(slot, src)| ColCopy { src, slot, width: src.size() })
.collect();
let t_nulls = NullPerm::new(&t_cols, walked.nullable_payload_slots());
Ok(JoinPlan {
out_schema,
walk,
t_cols,
t_nulls,
trace_leads: delta_is_right,
trace_ordered,
})
}
}
enum Walked<'a> {
Trace,
Source(&'a SchemaDescriptor, Vec<ColumnLocator>),
}
fn same_pk_types(delta: &SchemaDescriptor, trace: &SchemaDescriptor) -> Result<(), String> {
fn types(s: &SchemaDescriptor) -> impl Iterator<Item = TypeCode> + '_ {
s.pk_columns().map(|(_, c)| c.type_code)
}
if types(delta).eq(types(trace)) {
return Ok(());
}
Err("join: delta and trace PK column types differ (both sides must reindex at the pair's common type)".into())
}
#[derive(Clone, Copy)]
struct RangeProbe {
eq_size: usize,
above: bool,
cuts_below: bool,
}
impl RangeProbe {
fn new(trace: &SchemaDescriptor, rel: RangeRel, delta_is_right: bool) -> RangeProbe {
let eq_size = trace.pk_prefix_stride(trace.pk_cols().len() - 1);
let rel = match delta_is_right {
true => rel,
false => rel.converse(),
};
RangeProbe {
eq_size,
above: rel.bounds_below(),
cuts_below: rel.bounds_below() == rel.admits_equal(),
}
}
fn cut_points(&self, pk: &[u8]) -> Option<(PkBuf, Option<PkBuf>)> {
let group = &pk[..self.eq_size];
let slot = KeyCut::new(pk, !self.cuts_below);
let (start, end) = match self.above {
true => (slot, KeyCut::above(group)),
false => (KeyCut::min_of(group), slot),
};
key_range_between_cuts(start, end, pk.len())
}
#[inline]
fn last_before_split(&self) -> Ordering {
match self.cuts_below {
true => Ordering::Equal,
false => Ordering::Less,
}
}
}
#[derive(Clone, Copy)]
struct Pairing {
rs: u32,
re: u32,
src: u32,
row: u32,
w_trace: i64,
}
impl Pairing {
#[inline(always)]
fn delta_rows(&self) -> std::ops::Range<usize> {
self.rs as usize..self.re as usize
}
#[inline(always)]
fn len(&self) -> usize {
(self.re - self.rs) as usize
}
}
pub fn op_join_delta_trace(delta: &Batch, trace: OpenAt<'_>, plan: &JoinPlan) -> Batch {
let out_schema = &plan.out_schema;
debug_assert!(delta.is_consolidated());
let n = delta.count;
if n == 0 {
return Batch::empty_with_schema(out_schema);
}
let prefix = match plan.walk {
Walk::Equi => delta.schema().pk_stride(),
Walk::Range(range) => range.eq_size,
Walk::Cross => 0,
};
let cursor = &mut trace(&delta.get_pk_bytes(0)[..prefix], &delta.get_pk_bytes(n - 1)[..prefix]);
let mut pairs: Vec<Pairing> = Vec::new();
let mut rows = 0usize;
let mut ordered = matches!(plan.walk, Walk::Equi);
let mut last_run = usize::MAX;
let mut emit = |rs: usize, re: usize, c: &ReadCursor| {
if rs == re {
return;
}
if ordered && rs == last_run {
ordered = plan.trace_ordered && (plan.trace_leads || re - rs == 1);
}
last_run = rs;
let (src, row) = c.current_position();
pairs.push(Pairing {
rs: rs as u32,
re: re as u32,
src: src as u32,
row: row as u32,
w_trace: c.current_weight,
});
rows += re - rs;
};
match plan.walk {
Walk::Equi => equi_merge_walk(delta, cursor, emit),
Walk::Range(range) => range_merge_walk(delta, cursor, range, emit),
Walk::Cross => cursor.for_each_row_while(|_| true, |c| emit(0, n, c)),
}
let mut out = write_pairings(delta, cursor, plan, &pairs, rows);
if ordered && !out.is_empty() {
out.certify_consolidated();
}
out
}
fn write_pairings(delta: &Batch, cursor: &ReadCursor, plan: &JoinPlan, pairs: &[Pairing], rows: usize) -> Batch {
let out_schema = &plan.out_schema;
if rows == 0 {
return Batch::empty_with_schema(out_schema);
}
let d_schema = delta.schema();
let mut out = Batch::with_capacity(out_schema, rows);
let strings = d_schema.string_payload_slots();
let heap_at = match strings != 0 && !should_relocate_blob(delta.blob().len(), delta.count, rows) {
true => {
let mut emitted = vec![false; delta.count];
pairs.iter().for_each(|p| emitted[p.delta_rows()].fill(true));
let kept = runs_where(delta.count, |row| emitted[row]);
out.carry_heap(&delta.as_mem_batch(), strings, &kept)
}
false => None,
};
let d_first = if plan.trace_leads { plan.t_cols.len() } else { 0 };
let trace_row = |p: &Pairing| (cursor.source_at(p.src as usize), p.row as usize);
let any_ghost = out.append_session(rows).write(rows, |w| {
let mut any_ghost = false;
let (pk, weights, nulls) = w.fixed_mut();
match plan.walk {
Walk::Cross => {
let (pk_len, d_len) = (out_schema.pk_stride(), d_schema.pk_stride());
let (d_key, t_key) = match plan.trace_leads {
true => (pk_len - d_len..pk_len, 0..pk_len - d_len),
false => (0..d_len, d_len..pk_len),
};
let mut keys = pk.chunks_exact_mut(pk_len);
for p in pairs {
let (t_src, t_row) = trace_row(p);
let t_pk = t_src.get_pk_bytes(t_row);
for i in p.delta_rows() {
let key = keys.next().expect("one key per output row");
key[d_key.clone()].copy_from_slice(delta.get_pk_bytes(i));
key[t_key.clone()].copy_from_slice(t_pk);
}
}
}
_ => width_dispatch!(d_schema.pk_stride(), copy_runs, delta.pk_data(), pk, delta_runs(pairs)),
}
{
let src = delta.weight_data().as_chunks::<8>().0;
let mut dst = weights.as_chunks_mut::<8>().0.iter_mut();
for p in pairs {
for (w_delta, w_out) in src[p.delta_rows()].iter().zip(&mut dst) {
let w = i64::from_le_bytes(*w_delta).wrapping_mul(p.w_trace);
any_ghost |= w == 0;
*w_out = w.to_le_bytes();
}
}
}
{
let src = delta.null_bmp_data().as_chunks::<8>().0;
let mut dst = nulls.as_chunks_mut::<8>().0.iter_mut();
for p in pairs {
let (t_src, t_row) = trace_row(p);
let t_bits = plan.t_nulls.apply(t_src.get_null_word(t_row));
for (d_null, word) in src[p.delta_rows()].iter().zip(&mut dst) {
*word = (t_bits | null_word_at(u64::from_le_bytes(*d_null), d_first)).to_le_bytes();
}
}
}
for (pi, col) in d_schema.payload_columns() {
let slot = d_first + pi;
width_dispatch!(
col.size(),
copy_runs,
delta.col_data(pi),
w.col_mut(slot),
delta_runs(pairs)
);
if col.type_code.is_german_string() {
w.rebase_string_col(slot, delta.blob(), heap_at);
}
}
for &ColCopy { src, slot, .. } in &plan.t_cols {
match src {
ColumnLocator::Payload { slot: pi, type_code, .. } if type_code.is_german_string() => {
let (dst, dst_blob, mut cache) = w.string_col_mut(slot);
let mut cells = dst.as_chunks_mut::<16>().0.iter_mut();
for p in pairs {
let (t_src, t_row) = trace_row(p);
let cell = relocate_german_string_vec(
t_src.get_col_ptr(t_row, pi as usize, 16),
t_src.blob(),
dst_blob,
cache.as_deref_mut(),
);
cells.by_ref().take(p.len()).for_each(|c| *c = cell);
}
}
ColumnLocator::Payload { slot: pi, size, .. } => {
width_dispatch!(size as usize, repeat_cells, cursor, pi as usize, w.col_mut(slot), pairs)
}
ColumnLocator::Pk { byte_off, size, type_code } => width_dispatch!(
size as usize,
repeat_pk_cells,
cursor,
byte_off as usize,
type_code.is_signed_int(),
w.col_mut(slot),
pairs
),
}
}
any_ghost
});
if any_ghost {
let live = runs_where(rows, |row| out.get_weight(row) != 0);
return Batch::from_ranges(&out, &live, 0);
}
out
}
#[inline(always)]
fn delta_runs(pairs: &[Pairing]) -> impl Iterator<Item = (usize, usize)> + '_ {
pairs.iter().map(|p| (p.rs as usize, p.re as usize))
}
#[inline(always)]
fn repeat_cells<const N: usize>(cursor: &ReadCursor, pi: usize, dst: &mut [u8], pairs: &[Pairing], width: usize) {
assert!(N != 0, "a payload column is 1, 2, 4, 8 or 16 bytes, not {width}");
let mut cells = dst.as_chunks_mut::<N>().0.iter_mut();
for p in pairs {
let src = cursor.source_at(p.src as usize);
let cell: [u8; N] = src.get_col_ptr(p.row as usize, pi, N).try_into().unwrap();
cells.by_ref().take(p.len()).for_each(|c| *c = cell);
}
}
#[inline(always)]
fn repeat_pk_cells<const N: usize>(
cursor: &ReadCursor,
off: usize,
signed: bool,
dst: &mut [u8],
pairs: &[Pairing],
width: usize,
) {
assert!(N != 0, "a PK column is 1, 2, 4, 8 or 16 bytes, not {width}");
let mut cells = dst.as_chunks_mut::<N>().0.iter_mut();
for p in pairs {
let pk = cursor.source_at(p.src as usize).get_pk_bytes(p.row as usize);
let mut cell = [0u8; N];
gnitz_wire::decode_pk_cell(&pk[off..off + N], signed, &mut cell);
cells.by_ref().take(p.len()).for_each(|c| *c = cell);
}
}
fn equi_merge_walk(delta: &Batch, m: &mut ReadCursor, mut emit: impl FnMut(usize, usize, &ReadCursor)) {
let (n, k) = (delta.count, delta.schema().pk_stride());
let mut i = 0;
while i < n {
let dk = delta.get_pk_bytes(i);
if m.seek_pk_group_ascending(dk) {
let j = pk_group_end(delta, i); m.for_each_pk_group_row(dk, |c| emit(i, j, c));
i = j;
} else if m.valid {
i = delta.advance_to(&m.current_pk_bytes()[..k], i); } else {
break;
}
}
}
fn range_merge_walk(
delta: &Batch,
cursor: &mut ReadCursor,
probe: RangeProbe,
mut emit: impl FnMut(usize, usize, &ReadCursor),
) {
let (eq_size, above) = (probe.eq_size, probe.above);
let (stride, last) = (delta.schema().pk_stride(), probe.last_before_split());
let mut group_lb = [0u8; MAX_PK_BYTES];
let mut lo = 0;
while lo < delta.count {
let hi = pk_prefix_group_end(delta, lo, eq_size);
let cut_row = if above { lo } else { hi - 1 };
let Some((start, end)) = probe.cut_points(delta.get_pk_bytes(cut_row)) else {
lo = hi;
continue;
};
cursor.advance_to(start.pk_bytes());
let end = end.as_ref().map(PkBuf::pk_bytes);
let mut ptr = lo; cursor.for_each_row_while(
|pk| end.is_none_or(|e| compare_pk_ordering(pk, e).is_lt()),
|c| {
let s = &c.current_pk_bytes()[eq_size..];
while ptr < hi && compare_pk_ordering(&delta.get_pk_bytes(ptr)[eq_size..], s) <= last {
ptr += 1;
}
let (rs, re) = if above { (lo, ptr) } else { (ptr, hi) };
emit(rs, re, c);
},
);
if !cursor.valid {
return;
}
let trace_group = &cursor.current_pk_bytes()[..eq_size];
lo = hi;
if hi < delta.count && compare_pk_ordering(&delta.get_pk_bytes(hi)[..eq_size], trace_group).is_lt() {
group_lb[..eq_size].copy_from_slice(trace_group);
lo = delta.advance_to(&group_lb[..stride], hi).max(hi);
}
}
}
#[cfg(test)]
#[path = "tests/join.rs"]
mod tests;
#[cfg(test)]
#[path = "benches/join.rs"]
mod bench;