use gnitz_expr::{ColCopy, ExprValidateErr, LogicalProgram, MapEval};
use super::reindex::{locate_key_col, FoldCols, ReindexPacker};
use crate::repr::{Batch, DirectWriter};
use crate::schema::{ColumnLocator, DerivedSchema, SchemaColumn, SchemaDescriptor, SchemaFacts, TypeCode};
use gnitz_wire::{zip_cells, FixedInt};
#[derive(Clone, Copy)]
struct RowWindow {
src: usize,
dst: usize,
n: usize,
}
const COMPACT_RUN_LEN: usize = 128;
const PACK_COMPACT_RUN_LEN: usize = 24;
enum PkSource {
Inherit,
Pack(ReindexPacker),
HashRow(FoldCols),
}
fn copy_column(
in_batch: &Batch,
output: &mut DirectWriter<'_>,
&ColCopy {
src: src_loc,
slot: dst_payload,
width: stride,
}: &ColCopy,
w: RowWindow,
) {
let RowWindow { src: src_start, dst: dst_base, n } = w;
match src_loc {
ColumnLocator::Pk { byte_off, size, type_code } => {
let pk_stride = in_batch.schema().pk_stride();
let pk = &in_batch.pk_data()[src_start * pk_stride..(src_start + n) * pk_stride];
let dst = &mut output.col_mut(dst_payload)[dst_base * stride..(dst_base + n) * stride];
let (off, src_stride) = (byte_off as usize, size as usize);
if src_stride == stride {
gnitz_wire::decode_pk_cells(pk, pk_stride, off, stride, type_code.is_signed_int(), dst);
} else {
widen_column(pk, pk_stride, off, type_code, true, stride, dst);
}
}
ColumnLocator::Payload { slot, size, type_code } => {
let in_pi = slot as usize;
let src_stride = size as usize; debug_assert!(
!type_code.is_german_string() || (src_stride, stride) == (16, 16),
"German-string column moved at a non-16-byte stride",
);
debug_assert!(
(src_start + n) * src_stride <= in_batch.col_data(in_pi).len(),
"copy_column: source column {in_pi} is shorter than rows [{src_start}, {}) at stride {src_stride}",
src_start + n
);
let src = &in_batch.col_data(in_pi)[src_start * src_stride..(src_start + n) * src_stride];
let dst = &mut output.col_mut(dst_payload)[dst_base * stride..(dst_base + n) * stride];
if src_stride == stride {
dst.copy_from_slice(src);
} else {
widen_column(src, src_stride, 0, type_code, false, stride, dst);
}
}
}
}
fn widen_column(src: &[u8], src_stride: usize, off: usize, tc: TypeCode, pk: bool, dw: usize, dst: &mut [u8]) {
let fi = FixedInt::from_type_code(tc).expect("a widened column is a fixed int");
match dw {
2 => widen_cells::<2>(src, src_stride, off, fi, pk, dst),
4 => widen_cells::<4>(src, src_stride, off, fi, pk, dst),
8 => widen_cells::<8>(src, src_stride, off, fi, pk, dst),
other => unreachable!("a widened slot is 2/4/8 bytes, got {other}"),
}
}
fn widen_cells<const DW: usize>(src: &[u8], src_stride: usize, off: usize, fi: FixedInt, pk: bool, dst: &mut [u8]) {
let out = dst.as_chunks_mut::<DW>().0.iter_mut();
gnitz_wire::for_each_fixed_int!(fi, |FI| {
const W: usize = FI.width();
let store = |v: i64, d: &mut [u8; DW]| *d = v.to_le_bytes()[..DW].try_into().unwrap();
match pk {
true => zip_cells::<W, _>(src, src_stride, off, out, |c, d| {
store(gnitz_wire::decode_opk_i64(c, FI), d)
}),
false => zip_cells::<W, _>(src, src_stride, off, out, |c, d| store(FI.decode_le_i64(c), d)),
}
});
}
pub struct MapPlan {
ev: MapEval,
pk_source: PkSource,
copied_string_slots: u64,
null_key_mask: u64,
out_schema: SchemaDescriptor,
}
fn reindex_hash_row(output: &mut DirectWriter<'_>, fold: &FoldCols) {
let (first, n) = (output.first_row(), output.rows());
const KEY_BYTES: usize = std::mem::size_of::<u128>();
assert_eq!(output.schema.pk_stride(), KEY_BYTES, "a hash-row PK is one U128 column");
const CHUNK: usize = 256;
let mut keys = [0u128; CHUNK];
let mut start = 0;
while start < n {
let end = (start + CHUNK).min(n);
{
let mb = output.written();
for (key, row) in keys.iter_mut().zip(first + start..first + end) {
*key = fold.key_row(&mb, row, mb.get_null_word(row));
}
}
let pk = &mut output.pk_mut()[start * KEY_BYTES..end * KEY_BYTES];
for (key, dst) in keys.iter().zip(pk.as_chunks_mut::<KEY_BYTES>().0) {
*dst = key.to_be_bytes();
}
start = end;
}
}
fn compute_map_output_schema(
in_schema: &SchemaDescriptor,
out_cols: &[(TypeCode, bool)],
) -> Result<SchemaDescriptor, String> {
let mut b = DerivedSchema::new();
b.push_pk_of(in_schema);
for &(tc, nullable) in out_cols {
b.push(SchemaColumn::new(tc, nullable));
}
b.finish().map_err(|e| format!("compute map: output {e}"))
}
fn hashrow_output_schema(
in_schema: &SchemaDescriptor,
cols: &[gnitz_wire::ReindexSlot],
) -> Result<SchemaDescriptor, String> {
let mut b = DerivedSchema::new();
b.push_pk(SchemaColumn::new(crate::schema::TypeCode::U128, false));
for &(c, t) in cols {
locate_key_col(in_schema, c, "hash-row map")?;
let src = in_schema.columns[c as usize];
b.push(SchemaColumn::new(t, src.nullable));
}
b.finish().map_err(|e| format!("hash-row map: output {e}"))
}
impl MapPlan {
pub fn from_wire(in_schema: &SchemaDescriptor, mk: &gnitz_wire::MapKind) -> Result<Self, String> {
let mut null_key_mask = 0;
let (out_schema, prog, pk_source) = match mk {
gnitz_wire::MapKind::Compute(map) => return Self::from_compute_map(in_schema, map),
gnitz_wire::MapKind::Reindex { keep, key, nulls, .. } => {
let packer = ReindexPacker::new(in_schema, key)?;
let out_schema = packer.output_schema(in_schema, keep)?;
if *nulls == gnitz_wire::NullKeys::Drop {
null_key_mask = key
.iter()
.filter_map(|&(c, _)| in_schema.payload_slot(c as usize))
.fold(0, |mask, slot| mask | 1u64 << slot)
& in_schema.nullable_payload_slots();
}
(out_schema, LogicalProgram::copy_cols(keep), PkSource::Pack(packer))
}
gnitz_wire::MapKind::HashRow { cols } => {
let out_schema = hashrow_output_schema(in_schema, cols)?;
let proj: Vec<u32> = cols.iter().map(|&(c, _)| c).collect();
let fold = FoldCols::new(out_schema.payload_locators());
(out_schema, LogicalProgram::copy_cols(&proj), PkSource::HashRow(fold))
}
gnitz_wire::MapKind::Projection(cols) => {
let out_schema =
crate::schema::project_schema(in_schema, cols).map_err(|e| format!("projection map: {e}"))?;
(out_schema, LogicalProgram::copy_cols(cols), PkSource::Inherit)
}
};
let plan = Self::from_map(prog, in_schema, &out_schema, pk_source)
.map_err(|e| format!("map: program/schema mismatch: {e}"))?;
Ok(MapPlan { null_key_mask, ..plan })
}
pub(crate) fn from_compute_map(in_schema: &SchemaDescriptor, map: &gnitz_wire::ComputeMap) -> Result<Self, String> {
let out_schema = compute_map_output_schema(in_schema, &map.out_cols)?;
let prog = LogicalProgram::from_blob(&map.program).map_err(|e| format!("map: invalid program: {e}"))?;
Self::from_map(prog, in_schema, &out_schema, PkSource::Inherit)
.map_err(|e| format!("map: program/schema mismatch: {e}"))
}
fn from_map(
logical: LogicalProgram,
in_schema: &SchemaDescriptor,
out_schema: &SchemaDescriptor,
pk_source: PkSource,
) -> Result<Self, ExprValidateErr> {
let ev = logical.resolve_map(in_schema, out_schema)?;
let copied_string_slots = ev.copies().iter().fold(0u64, |slots, c| match c.src {
ColumnLocator::Payload { slot, type_code, .. } if type_code.is_german_string() => slots | 1u64 << slot,
_ => slots,
});
Ok(MapPlan {
ev,
pk_source,
copied_string_slots,
null_key_mask: 0,
out_schema: *out_schema,
})
}
pub fn out_schema(&self) -> &SchemaDescriptor {
&self.out_schema
}
pub fn rekeys_onto_pk_prefix(&self) -> Option<Vec<ColumnLocator>> {
let PkSource::Pack(packer) = &self.pk_source else {
return None;
};
packer.pk_range().filter(|&(at, _)| at == 0)?;
debug_assert!(self.null_key_mask == 0 && !self.ev.emits_anything());
debug_assert!(self.ev.copies().iter().all(|c| c.width == c.src.size()));
Some(self.ev.copies().iter().map(|c| c.src).collect())
}
pub fn drops_null_keys(&self) -> bool {
self.null_key_mask != 0
}
pub fn is_identity(&self) -> bool {
matches!(self.pk_source, PkSource::Inherit) && self.ev.is_identity()
}
pub fn evaluate_map_batch(&mut self, in_batch: &Batch) -> Batch {
let whole = [(0, in_batch.count)];
let runs;
let ranges: &[(usize, usize)] = match self.null_key_mask {
0 => &whole,
mask => {
runs = in_batch.runs_without_nulls(mask);
&runs
}
};
let n = crate::repr::range_rows(ranges);
if n == 0 {
return Batch::empty_with_schema(&self.out_schema);
}
let mut output = Batch::with_capacity(&self.out_schema, n);
self.append_map_ranges(in_batch, &mut output, ranges);
output
}
pub(crate) fn append_map_ranges(&mut self, src: &Batch, out: &mut Batch, ranges: &[(usize, usize)]) {
let total = crate::repr::range_rows(ranges);
if total == 0 {
return;
}
let compact_below = match (self.ev.emits_anything(), &self.pk_source) {
(true, _) => COMPACT_RUN_LEN,
(false, PkSource::Pack(_)) => PACK_COMPACT_RUN_LEN,
(false, _) => 0,
};
let starves_kernel = ranges.len() > 1 && total < ranges.len() * compact_below;
let compacted = starves_kernel.then(|| Batch::from_ranges(src, ranges, 0));
let whole = [(0, total)];
let (src, ranges) = match &compacted {
Some(c) => (c, &whole[..]),
None => (src, ranges),
};
let heap_at = out.carry_heap(&src.as_mem_batch(), self.copied_string_slots, ranges);
if heap_at.is_none() {
out.reserve_blob(crate::repr::prorated_blob_cap(src.blob().len(), src.count, total));
}
out.append_session(total).write(total, |out| {
let mut dst = 0;
for &(start, end) in ranges {
let w = RowWindow { src: start, dst, n: end - start };
self.map_rows_into(src, out, w);
dst += w.n;
}
for c in self.ev.copies() {
if matches!(c.src, ColumnLocator::Payload { type_code, .. } if type_code.is_german_string()) {
out.rebase_string_col(c.slot, src.blob(), heap_at);
}
}
if let PkSource::HashRow(fold) = &self.pk_source {
reindex_hash_row(out, fold);
}
});
}
#[inline(always)]
fn map_rows_into(&mut self, in_batch: &Batch, output: &mut DirectWriter<'_>, w: RowWindow) {
let RowWindow { src: src_start, dst: dst_base, n } = w;
match &self.pk_source {
PkSource::Inherit => {
let pk_st = in_batch.schema().pk_stride();
debug_assert_eq!(
pk_st,
output.schema.pk_stride(),
"PkSource::Inherit: PK stride mismatch"
);
output.pk_mut()[dst_base * pk_st..(dst_base + n) * pk_st]
.copy_from_slice(&in_batch.pk_data()[src_start * pk_st..(src_start + n) * pk_st]);
}
PkSource::Pack(packer) => {
debug_assert_eq!(output.schema.pk_stride(), packer.out_stride);
let stride = packer.out_stride;
let pk = &mut output.pk_mut()[dst_base * stride..];
packer.pack_rows(pk, stride, &in_batch.as_mem_batch(), &[(src_start, src_start + n)]);
}
PkSource::HashRow(_) => {}
}
output.weight_mut()[dst_base * 8..(dst_base + n) * 8]
.copy_from_slice(&in_batch.weight_data()[src_start * 8..(src_start + n) * 8]);
for c in self.ev.copies() {
copy_column(in_batch, output, c, w);
}
self.ev
.write_computed(&in_batch.as_mem_batch(), src_start, n, output, dst_base);
}
}
#[cfg(test)]
#[path = "tests/map.rs"]
mod tests;
#[cfg(test)]
#[path = "benches/map.rs"]
mod bench;