use crate::circuit::{AggDescriptor, ComputeMap};
use std::num::NonZeroU64;
use std::ops::Range;
use crate::codec::{decode_all, Reader, Wire, Writer};
use crate::range::KeyRange;
use crate::{MAX_COLUMNS, MAX_PK_BYTES};
pub const MAX_ORDER_KEYS: usize = 16;
const BOUND_NONE: u8 = 0;
const BOUND_RANGE: u8 = 1;
const BOUND_PK_SET: u8 = 2;
const SINK_ROWS: u8 = 0;
const SINK_FOLD: u8 = 1;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct OrderKey {
pub col: u16,
pub desc: bool,
pub nulls_first: bool,
}
const ORDER_DESC: u8 = 1 << 0;
const ORDER_NULLS_FIRST: u8 = 1 << 1;
impl Wire for OrderKey {
fn write(&self, w: &mut Writer) {
let mut flags = 0u8;
if self.desc {
flags |= ORDER_DESC;
}
if self.nulls_first {
flags |= ORDER_NULLS_FIRST;
}
w.u16(self.col).u8(flags);
}
fn read(r: &mut Reader) -> Result<Self, String> {
let col = r.u16()?;
let flags = r.flags(ORDER_DESC | ORDER_NULLS_FIRST)?;
Ok(OrderKey {
col,
desc: flags & ORDER_DESC != 0,
nulls_first: flags & ORDER_NULLS_FIRST != 0,
})
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct AggReadSpec {
pub group_cols: Vec<u32>,
pub aggs: Vec<AggDescriptor>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ReadSink {
pub map: Option<ComputeMap>,
pub kind: SinkKind,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RowsCut {
pub k: NonZeroU64,
pub order: Vec<OrderKey>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum SinkKind {
Rows { cut: Option<RowsCut> },
Fold(AggReadSpec),
}
impl ReadSink {
pub fn all_rows() -> Self {
ReadSink {
map: None,
kind: SinkKind::Rows { cut: None },
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PkKeys {
stride: u8,
bytes: Vec<u8>,
}
impl PkKeys {
pub fn from_keys<'a>(stride: usize, keys: impl IntoIterator<Item = &'a [u8]>) -> Self {
assert!((1..=MAX_PK_BYTES).contains(&stride), "PkKeys: stride {stride}");
let mut ks: Vec<&[u8]> = keys
.into_iter()
.inspect(|k| assert_eq!(k.len(), stride, "PkKeys: key width"))
.collect();
ks.sort_unstable();
ks.dedup();
PkKeys { stride: stride as u8, bytes: ks.concat() }
}
pub fn from_sorted(stride: usize, bytes: Vec<u8>) -> Self {
assert!((1..=MAX_PK_BYTES).contains(&stride), "PkKeys: stride {stride}");
assert_eq!(bytes.len() % stride, 0, "PkKeys: key width");
Self::checked(stride, bytes).expect("PkKeys")
}
pub fn checked(stride: usize, bytes: Vec<u8>) -> Result<Self, String> {
match strictly_ascending(&bytes, stride) {
true => Ok(PkKeys { stride: stride as u8, bytes }),
false => Err("keys are not strictly ascending".to_string()),
}
}
pub fn stride(&self) -> usize {
self.stride as usize
}
pub fn len(&self) -> usize {
self.bytes.len() / self.stride()
}
pub fn is_empty(&self) -> bool {
self.bytes.is_empty()
}
pub fn as_bytes(&self) -> &[u8] {
&self.bytes
}
pub fn into_bytes(self) -> Vec<u8> {
self.bytes
}
pub fn iter(&self) -> std::slice::ChunksExact<'_, u8> {
self.bytes.chunks_exact(self.stride())
}
pub fn contains(&self, key: &[u8]) -> bool {
debug_assert_eq!(key.len(), self.stride());
let (mut lo, mut hi) = (0, self.len());
while lo < hi {
let mid = lo + (hi - lo) / 2;
match self.bytes[mid * self.stride()..][..self.stride()].cmp(key) {
std::cmp::Ordering::Less => lo = mid + 1,
std::cmp::Ordering::Greater => hi = mid,
std::cmp::Ordering::Equal => return true,
}
}
false
}
pub fn bounds(&self) -> Option<(&[u8], &[u8])> {
let first = self.iter().next()?;
Some((first, &self.bytes[self.bytes.len() - self.stride()..]))
}
}
pub fn strictly_ascending(bytes: &[u8], stride: usize) -> bool {
bytes.chunks_exact(stride).is_sorted_by(|a, b| a < b)
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ReadBound {
None,
Range(KeyRange),
PkSet(PkKeys),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ReadSpec {
pub bound: ReadBound,
pub predicate: Vec<u8>,
pub sink: ReadSink,
}
impl ReadSpec {
pub fn all_rows(bound: ReadBound) -> Self {
ReadSpec {
bound,
predicate: Vec::new(),
sink: ReadSink::all_rows(),
}
}
pub fn is_whole(&self) -> bool {
matches!(self.bound, ReadBound::None)
&& self.predicate.is_empty()
&& matches!(
self.sink,
ReadSink {
map: None,
kind: SinkKind::Rows { cut: None }
}
)
}
pub fn encode(&self) -> Vec<u8> {
let keys = match &self.bound {
ReadBound::PkSet(k) => k.as_bytes().len(),
_ => 0,
};
let map = self.sink.map.as_ref().map_or(0, |m| m.program.len());
let mut w = Writer::with_capacity(64 + self.predicate.len() + keys + map);
w.put(&self.bound).bytes32(&self.predicate);
match &self.sink.map {
Some(m) => w.bool(true).put(m),
None => w.bool(false),
};
match &self.sink.kind {
SinkKind::Rows { cut: None } => w.u8(SINK_ROWS).u64(0),
SinkKind::Rows { cut: Some(cut) } => w.u8(SINK_ROWS).u64(cut.k.get()).list(&cut.order),
SinkKind::Fold(agg) => w.u8(SINK_FOLD).list(&agg.group_cols).list(&agg.aggs),
};
w.into_vec()
}
pub fn decode(buf: &[u8]) -> Result<ReadSpec, String> {
decode_all(buf, "read_spec", |r| {
let bound = r.get()?;
let predicate = r.bytes32()?.to_vec();
let map = if r.bool()? { Some(r.get()?) } else { None };
let kind = match r.u8()? {
SINK_ROWS => {
let cut = match NonZeroU64::new(r.u64()?) {
None => None,
Some(k) => Some(RowsCut {
k,
order: r.list("order keys", MAX_ORDER_KEYS)?,
}),
};
SinkKind::Rows { cut }
}
SINK_FOLD => {
let group_cols = r.list("column list", MAX_COLUMNS)?;
let aggs = r.list("aggregate list", MAX_COLUMNS)?;
SinkKind::Fold(AggReadSpec { group_cols, aggs })
}
other => return Err(format!("unknown sink tag {other}")),
};
Ok(ReadSpec {
bound,
predicate,
sink: ReadSink { map, kind },
})
})
}
}
fn write_pk_set(w: &mut Writer, stride: usize, keys: &[u8]) {
w.u8(stride as u8).u32((keys.len() / stride) as u32).raw(keys);
}
fn read_pk_set<'a>(r: &mut Reader<'a>) -> Result<(usize, &'a [u8]), String> {
let stride = r.u8()? as usize;
if !(1..=MAX_PK_BYTES).contains(&stride) {
return Err(format!("PkSet stride {stride} outside 1..={MAX_PK_BYTES}"));
}
let count = r.u32()? as usize;
let keys = r.take(count * stride)?;
Ok((stride, keys))
}
impl Wire for ReadBound {
fn write(&self, w: &mut Writer) {
match self {
ReadBound::None => {
w.u8(BOUND_NONE);
}
ReadBound::Range(range) => {
w.u8(BOUND_RANGE).put(range);
}
ReadBound::PkSet(keys) => {
w.u8(BOUND_PK_SET);
write_pk_set(w, keys.stride(), keys.as_bytes());
}
}
}
fn read(r: &mut Reader) -> Result<Self, String> {
Ok(match r.u8()? {
BOUND_NONE => ReadBound::None,
BOUND_RANGE => ReadBound::Range(r.get()?),
BOUND_PK_SET => {
let (stride, bytes) = read_pk_set(r)?;
ReadBound::PkSet(PkKeys::checked(stride, bytes.to_vec()).map_err(|e| format!("PkSet {e}"))?)
}
other => return Err(format!("unknown bound kind {other}")),
})
}
}
pub enum BoundPeek<'a> {
Range(KeyRange),
PkSet(PkSetPeek<'a>),
}
pub struct PkSetPeek<'a> {
blob: &'a [u8],
span: Range<usize>,
pub stride: usize,
pub keys: &'a [u8],
}
impl PkSetPeek<'_> {
pub fn with_keys(&self, keys: &[u8]) -> Vec<u8> {
debug_assert!(
strictly_ascending(keys, self.stride),
"a subsequence of the peeked keys"
);
let mut w = Writer::with_capacity(self.blob.len() - self.keys.len() + keys.len());
w.raw(&self.blob[..self.span.start]);
write_pk_set(&mut w, self.stride, keys);
w.raw(&self.blob[self.span.end..]);
w.into_vec()
}
}
pub fn peek_bound(blob: &[u8]) -> Option<BoundPeek<'_>> {
let mut r = Reader::new(blob);
match r.u8().ok()? {
BOUND_RANGE => r.get().ok().map(BoundPeek::Range),
BOUND_PK_SET => {
let start = r.pos();
let (stride, keys) = read_pk_set(&mut r).ok()?;
Some(BoundPeek::PkSet(PkSetPeek { blob, span: start..r.pos(), stride, keys }))
}
_ => None,
}
}
#[cfg(test)]
#[path = "tests/read_spec.rs"]
mod tests;