use crate::schema::ColumnTable;
use std::cmp::Ordering;
use std::ops::Range;
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, SchemaFacts, MAX_PK_BYTES};
use crate::algebra::MapPlan;
use gnitz_expr::{ColCopy, ColumnLocator, NullPerm};
use gnitz_wire::RowSource;
use gnitz_wire::{null_word_at, JoinKind, PkBuf, RangeRel, TypeCode};
pub struct JoinPlan {
pub probe: JoinProbe,
pub out_schema: SchemaDescriptor,
}
pub struct JoinProbe {
walk: Walk,
d_first: u16,
t_cols: Box<[ColCopy]>,
t_nulls: NullPerm,
trace_leads: bool,
trace_ordered: bool,
}
impl JoinProbe {
pub fn probes_delta_keys(&self) -> bool {
matches!(self.walk, Walk::Equi)
}
}
#[derive(Clone, Copy)]
struct Slots {
start: u16,
end: u16,
}
impl Slots {
fn new(start: usize, end: usize) -> Slots {
debug_assert!(start <= end && end <= u16::MAX as usize);
Slots { start: start as u16, end: end as u16 }
}
#[inline]
fn range(self) -> Range<usize> {
self.start as usize..self.end as usize
}
}
#[derive(Clone, Copy)]
enum Walk {
Equi,
Range(RangeProbe),
Cross { pk_len: u16, d_key: Slots, t_key: Slots },
}
impl JoinPlan {
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 halves = |l: usize, total: usize| match delta_is_right {
true => (Slots::new(l, total), Slots::new(0, l)),
false => (Slots::new(0, l), Slots::new(l, total)),
};
let (d_slots, t_slots) = halves(left.num_payload_cols(), out_schema.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 => {
let (d_key, t_key) = halves(left.pk_stride(), out_schema.pk_stride());
Walk::Cross {
pk_len: out_schema.pk_stride() as u16,
d_key,
t_key,
}
}
};
let t_cols: Box<[ColCopy]> = (t_slots.start as usize..)
.zip(t_cols)
.map(|(slot, src)| ColCopy { src, slot, width: src.size() })
.collect();
let t_nulls = NullPerm::new(&t_cols, walked.nullable_payload_slots());
let probe = JoinProbe {
walk,
d_first: d_slots.start,
t_cols,
t_nulls,
trace_leads: delta_is_right,
trace_ordered,
};
Ok(JoinPlan { probe, out_schema })
}
}
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_columns()
.take(trace.pk_cols().len() - 1)
.map(|(_, c)| c.size() as usize)
.sum();
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,
cursor: &mut ReadCursor,
out_schema: &SchemaDescriptor,
probe: &JoinProbe,
) -> Batch {
debug_assert!(delta.is_consolidated());
let n = delta.count;
if n == 0 {
return Batch::empty_with_schema(out_schema);
}
let mut pairs: Vec<Pairing> = Vec::new();
let mut rows = 0usize;
let mut ordered = matches!(probe.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 = probe.trace_ordered && (probe.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 probe.walk {
Walk::Equi => equi_merge_walk(delta, cursor, emit),
Walk::Range(range) => range_merge_walk(delta, cursor, range, emit),
Walk::Cross { .. } => {
cursor.rewind();
cursor.for_each_row_while(|_| true, |c| emit(0, n, c));
}
}
let mut out = write_pairings(delta, cursor, out_schema, probe, &pairs, rows);
if ordered && !out.is_empty() {
out.certify_consolidated();
}
out
}
fn write_pairings(
delta: &Batch,
cursor: &ReadCursor,
out_schema: &SchemaDescriptor,
probe: &JoinProbe,
pairs: &[Pairing],
rows: usize,
) -> Batch {
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 = probe.d_first as usize;
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 probe.walk {
Walk::Cross { pk_len, d_key, t_key } => {
let (d_key, t_key) = (d_key.range(), t_key.range());
let mut keys = pk.chunks_exact_mut(pk_len as usize);
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 = probe.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() as usize,
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 &probe.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;