use std::ops::Range;
use super::reindex::{locate_key_col, FoldCols, ReindexPacker};
use crate::repr::{Batch, MemBatch};
use crate::schema::key::{pk_width_dispatch, NarrowPkOpk, PkSortKey};
use crate::schema::{
ColumnLocator, DerivedSchema, ReduceOutKey, SchemaColumn, SchemaDescriptor, SchemaFacts, TypeCode,
};
use gnitz_wire::{ReduceOutSlot, NARROW_PK_MAX_BYTES};
use rustc_hash::FxHashMap;
use std::collections::hash_map::Entry;
pub(crate) enum GroupKey {
PkRange { at: usize, n: usize },
Image(ColumnLocator),
Packed(ReindexPacker),
Fold(FoldCols),
}
impl GroupKey {
pub(crate) fn new(schema: &SchemaDescriptor, group_cols: &[u32]) -> Result<Self, String> {
let locs: Vec<ColumnLocator> = group_cols
.iter()
.map(|&c| locate_key_col(schema, c, "group key"))
.collect::<Result<_, _>>()?;
Ok(match (schema.reduce_out_key(group_cols), &locs[..]) {
(ReduceOutKey::Natural, &[loc]) => Self::column(loc),
(ReduceOutKey::Natural, _) => GroupKey::PkRange { at: 0, n: schema.pk_stride() },
(ReduceOutKey::SyntheticFold, _) => match ReindexPacker::new_group_key(schema, group_cols, &[])?.0 {
p if p.packs_whole() && (1..=GROUP_PK_BYTES).contains(&p.out_stride) => Self::packed(p),
_ => GroupKey::Fold(FoldCols::new(locs)),
},
})
}
fn column(loc: ColumnLocator) -> Self {
match loc {
ColumnLocator::Pk { byte_off, size, .. } => GroupKey::PkRange { at: byte_off as usize, n: size as usize },
ColumnLocator::Payload { .. } => GroupKey::Image(loc),
}
}
pub(crate) fn packed(packer: ReindexPacker) -> Self {
match (packer.pk_range(), packer.identity_columns().as_deref()) {
(Some((at, n)), _) => GroupKey::PkRange { at, n },
(None, Some(&[loc])) => Self::column(loc),
_ => GroupKey::Packed(packer),
}
}
pub(crate) fn cells<'a>(&self, mb: &MemBatch<'a>) -> Option<KeyCells<'a>> {
let pk = |off: usize, width: usize| KeyCells {
region: mb.pk(),
stride: mb.pk_stride(),
off,
width,
opk: true,
signed: false,
};
match *self {
GroupKey::PkRange { at, n } if (1..=NARROW_PK_MAX_BYTES).contains(&n) => Some(pk(at, n)),
GroupKey::Image(ColumnLocator::Payload { slot, size, type_code }) => Some(KeyCells {
region: mb.col_data(slot as usize, size as usize),
stride: size as usize,
off: 0,
width: size as usize,
opk: false,
signed: type_code.is_signed_int(),
}),
GroupKey::Image(ColumnLocator::Pk { .. }) => unreachable!("a PK column keys as a PK range"),
GroupKey::PkRange { .. } | GroupKey::Packed(_) | GroupKey::Fold(_) => None,
}
}
}
pub(crate) struct KeyCells<'a> {
pub(crate) region: &'a [u8],
pub(crate) stride: usize,
pub(crate) off: usize,
pub(crate) width: usize,
pub(crate) opk: bool,
pub(crate) signed: bool,
}
impl KeyCells<'_> {
#[inline(always)]
fn cell<const W: usize>(&self, r: usize) -> &[u8; W] {
let at = r * self.stride + self.off;
self.region[at..at + W].try_into().unwrap()
}
}
#[inline(always)]
pub(crate) fn cell_image<const W: usize, const OPK: bool>(cell: &[u8; W], signed: bool) -> u128 {
if OPK {
return gnitz_wire::widen_pk_be(cell);
}
let mut le = [0u8; 16];
le[..W].copy_from_slice(cell);
u128::from_le_bytes(le) ^ ((signed as u128) << (8 * W - 1))
}
macro_rules! for_cell_width {
($width:expr, |$w:ident| $body:expr) => {
$crate::algebra::group_key::for_cell_width!(@arms $width, $w, $body, 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16)
};
(@arms $width:expr, $w:ident, $body:expr, $($n:literal)*) => {
match $width {
$($n => {
const $w: usize = $n;
$body
})*
_ => unreachable!("a key cell is 1..=16 bytes"),
}
};
}
pub(crate) use for_cell_width;
const GROUP_PK_COL: SchemaColumn = SchemaColumn::new(TypeCode::U128, false);
const GROUP_PK_BYTES: usize = GROUP_PK_COL.size() as usize;
pub(crate) enum OutPk<'a> {
Borrowed(&'a [u8]),
Narrow(NarrowPkOpk),
}
impl OutPk<'_> {
#[inline]
pub(crate) fn bytes(&self) -> &[u8] {
match self {
OutPk::Borrowed(b) => b,
OutPk::Narrow(k) => k.bytes(),
}
}
}
pub(crate) trait IdentityLoop {
type Out;
fn run(self, identity: impl Fn(usize) -> u128) -> Self::Out;
}
#[inline]
pub(crate) fn ground_pk() -> NarrowPkOpk {
NarrowPkOpk::new(gnitz_wire::global_group_key(), GROUP_PK_BYTES)
}
pub(crate) struct GroupOutKey {
key: GroupKey,
synthetic: bool,
carried: Vec<ColumnLocator>,
}
impl GroupOutKey {
pub(crate) fn new(
input: &SchemaDescriptor,
group_cols: &[u32],
row: impl IntoIterator<Item = u32>,
) -> Result<(Self, DerivedSchema), String> {
let key = GroupKey::new(input, group_cols)?;
let kind = input.reduce_out_key(group_cols);
let mut b = DerivedSchema::new();
let mut carried = Vec::new();
for slot in kind.output_layout(group_cols, row) {
match slot {
ReduceOutSlot::SyntheticKey => b.push_pk(GROUP_PK_COL),
ReduceOutSlot::Key(c) => b.push_pk(input.columns[c as usize]),
ReduceOutSlot::Carried(c) => {
carried.push(input.locate(c as usize));
b.push(input.columns[c as usize])
}
}
}
let synthetic = kind == ReduceOutKey::SyntheticFold;
Ok((GroupOutKey { key, synthetic, carried }, b))
}
#[inline]
pub(crate) fn is_global(&self) -> bool {
matches!(&self.key, GroupKey::Fold(f) if f.is_empty())
}
#[inline]
pub(crate) fn out_pk<'a>(&self, mb: &'a MemBatch, row: usize) -> OutPk<'a> {
let group_pk = |image| OutPk::Narrow(NarrowPkOpk::new(image, GROUP_PK_BYTES));
match &self.key {
&GroupKey::PkRange { at, n } if self.synthetic => {
group_pk(gnitz_wire::widen_pk_be(mb.get_pk_range(row, at, n)))
}
&GroupKey::PkRange { at, n } => OutPk::Borrowed(mb.get_pk_range(row, at, n)),
&GroupKey::Image(loc) => OutPk::Narrow(NarrowPkOpk::new(loc.opk_image(mb, row), loc.size())),
GroupKey::Packed(p) => group_pk(p.narrow_image(mb, row)),
GroupKey::Fold(f) => group_pk(f.key_row(mb, row, mb.get_null_word(row))),
}
}
#[inline]
pub(crate) fn with_identity<L: IdentityLoop>(&self, mb: &MemBatch, body: L) -> L::Out {
if let Some(cells) = self.key.cells(mb) {
return for_cell_width!(cells.width, |W| match cells.opk {
true => body.run(|r| cell_image::<W, true>(cells.cell::<W>(r), cells.signed)),
false => body.run(|r| cell_image::<W, false>(cells.cell::<W>(r), cells.signed)),
});
}
match &self.key {
GroupKey::Fold(f) => body.run(|r| f.key_row(mb, r, mb.get_null_word(r))),
GroupKey::Packed(p) => packed_identity(p, mb, body),
&GroupKey::PkRange { at, n } => body.run(|r| gnitz_wire::checksum_128(mb.get_pk_range(r, at, n))),
GroupKey::Image(_) => unreachable!("an image key is one cell"),
}
}
#[cfg(test)]
pub(crate) fn identity(&self, mb: &MemBatch, row: usize) -> u128 {
struct At(usize);
impl IdentityLoop for At {
type Out = u128;
fn run(self, identity: impl Fn(usize) -> u128) -> u128 {
identity(self.0)
}
}
self.with_identity(mb, At(row))
}
#[inline(always)]
pub(crate) fn carried(&self) -> &[ColumnLocator] {
&self.carried
}
pub(crate) fn runs(&self, batch: &Batch) -> Option<GroupRuns> {
let mb = &batch.as_mem_batch();
let n = mb.count;
match &self.key {
_ if n <= 1 || self.is_global() => Some(GroupRuns::of(n, |_| ())),
&GroupKey::PkRange { at: 0, n: w } if batch.is_consolidated() => Some(pk_width_dispatch!(w, |K| {
GroupRuns::of(n, |i| K::from_opk(mb.get_pk_range(i, 0, w)))
})),
_ => None,
}
}
pub(crate) fn ordinals(&self, batch: &Batch) -> GroupOrdinals {
match self.runs(batch) {
Some(runs) => GroupOrdinals::of_runs(&runs),
None => self.numbered(batch),
}
}
pub(crate) fn numbered(&self, batch: &Batch) -> GroupOrdinals {
let mb = &batch.as_mem_batch();
let n = mb.count;
let hashes = !matches!(self.key, GroupKey::PkRange { n: w, .. } if w > NARROW_PK_MAX_BYTES);
if let Some(groups) = hashes.then(|| self.with_identity(mb, Hashed { n })).flatten() {
return groups;
}
match &self.key {
&GroupKey::PkRange { at, n: w } => {
pk_width_dispatch!(w, |K| GroupOrdinals::sorted(n, |i| K::from_opk(
mb.get_pk_range(i, at, w)
)))
}
&GroupKey::Image(loc) if loc.size() <= 8 => GroupOrdinals::sorted(n, |i| loc.opk_image(mb, i) as u64),
&GroupKey::Image(loc) => GroupOrdinals::sorted(n, |i| loc.opk_image(mb, i)),
GroupKey::Packed(p) => {
let (keys, w) = (p.keys(mb), p.out_stride);
GroupOrdinals::sorted(n, |i| gnitz_wire::widen_pk_be(&keys[i * w..(i + 1) * w]))
}
GroupKey::Fold(f) => GroupOrdinals::sorted(n, |i| f.key_row(mb, i, mb.get_null_word(i))),
}
}
}
#[inline(never)]
fn packed_identity<L: IdentityLoop>(p: &ReindexPacker, mb: &MemBatch, body: L) -> L::Out {
let (keys, w) = (p.keys(mb), p.out_stride);
body.run(|r| gnitz_wire::widen_pk_be(&keys[r * w..(r + 1) * w]))
}
pub(crate) struct GroupRuns {
ends: Vec<u32>,
}
impl GroupRuns {
fn of<K: PartialEq>(n: usize, key: impl Fn(usize) -> K) -> Self {
let mut ends = Vec::new();
if n == 0 {
return GroupRuns { ends };
}
let mut prev = key(0);
for i in 1..n {
let k = key(i);
if k != prev {
ends.push(i as u32);
}
prev = k;
}
ends.push(n as u32);
GroupRuns { ends }
}
#[inline]
pub(crate) fn len(&self) -> usize {
self.ends.len()
}
#[inline]
pub(crate) fn last_start(&self) -> usize {
self.ends.iter().rev().nth(1).map_or(0, |&e| e as usize)
}
pub(crate) fn iter(&self) -> impl Iterator<Item = Range<usize>> + '_ {
let starts = std::iter::once(0).chain(self.ends.iter().map(|&e| e as usize));
starts.zip(self.ends.iter().map(|&e| e as usize)).map(|(s, e)| s..e)
}
}
const HASH_MIN_ROWS_PER_GROUP: usize = 4;
pub(crate) struct GroupOrdinals {
pub(crate) ord: Vec<u32>,
pub(crate) first: Vec<u32>,
pub(crate) by_pk: Vec<u32>,
}
impl GroupOrdinals {
#[inline]
pub(crate) fn len(&self) -> usize {
self.first.len()
}
pub(crate) fn of_runs(runs: &GroupRuns) -> Self {
let mut ord = Vec::with_capacity(runs.ends.last().map_or(0, |&e| e as usize));
let mut first = Vec::with_capacity(runs.len());
for (g, run) in runs.iter().enumerate() {
first.push(run.start as u32);
ord.extend(std::iter::repeat_n(g as u32, run.len()));
}
GroupOrdinals {
ord,
by_pk: (0..first.len() as u32).collect(),
first,
}
}
fn sorted<K: Ord>(n: usize, key: impl Fn(usize) -> K) -> Self {
let mut pairs: Vec<(K, u32)> = (0..n).map(|i| (key(i), i as u32)).collect();
pairs.sort_unstable();
let mut ord = vec![0u32; n];
let mut first: Vec<u32> = Vec::new();
for (p, (k, row)) in pairs.iter().enumerate() {
if p == 0 || *k != pairs[p - 1].0 {
first.push(*row);
}
ord[*row as usize] = (first.len() - 1) as u32;
}
GroupOrdinals {
ord,
by_pk: (0..first.len() as u32).collect(),
first,
}
}
}
#[derive(Default)]
pub(crate) struct GroupNumbers {
by_identity: FxHashMap<u128, u32>,
last: Option<(u128, u32)>,
}
impl GroupNumbers {
#[inline(always)]
pub(crate) fn ordinal<E>(&mut self, identity: u128, new: impl FnOnce() -> Result<u32, E>) -> Result<u32, E> {
if let Some((_, g)) = self.last.filter(|&(k, _)| k == identity) {
return Ok(g);
}
let g = match self.by_identity.entry(identity) {
Entry::Occupied(e) => *e.get(),
Entry::Vacant(e) => *e.insert(new()?),
};
self.last = Some((identity, g));
Ok(g)
}
}
struct Hashed {
n: usize,
}
impl IdentityLoop for Hashed {
type Out = Option<GroupOrdinals>;
fn run(self, identity: impl Fn(usize) -> u128) -> Self::Out {
let limit = self.n / HASH_MIN_ROWS_PER_GROUP;
let mut numbers = GroupNumbers::default();
let mut ord: Vec<u32> = Vec::with_capacity(self.n);
let mut first: Vec<u32> = Vec::new();
for row in 0..self.n {
let new = || {
let g = first.len();
first.push(row as u32);
(g < limit).then_some(g as u32).ok_or(())
};
ord.push(numbers.ordinal(identity(row), new).ok()?);
}
let mut by_pk: Vec<(u128, u32)> = first
.iter()
.zip(0..)
.map(|(&row, g)| (identity(row as usize), g))
.collect();
by_pk.sort_unstable();
Some(GroupOrdinals {
ord,
first,
by_pk: by_pk.into_iter().map(|(_, g)| g).collect(),
})
}
}
#[cfg(test)]
#[path = "tests/group_key.rs"]
mod tests;
#[cfg(test)]
#[path = "benches/group_key.rs"]
mod bench;