use crate::codec::{decode_all, Reader, Wire, Writer};
use crate::{TypeCode, MAX_COLUMNS};
wire_enum! {
enum Opcode: u8 {
ScanDelta = 1,
Filter = 2,
MapProj = 3,
MapExpr = 4,
MapHashRow = 5,
MapReindex = 6,
Negate = 7,
Union = 8,
Distinct = 9,
PositivePart = 10,
Reduce = 11,
JoinEqui = 12,
JoinRange = 13,
JoinCross = 14,
IntegrateSink = 15,
ExchangeShard = 16,
NullExtend = 17,
WorkerFilter = 18,
TopN = 19,
}
}
pub(crate) const CIRCUIT_VERSION: u8 = 9;
wire_enum! {
pub enum AggFunc: u8 {
Count = 1,
Sum = 2,
Min = 3,
Max = 4,
CountNonNull = 5,
}
}
impl AggFunc {
pub const fn is_linear(self) -> bool {
match self {
AggFunc::Count | AggFunc::CountNonNull | AggFunc::Sum => true,
AggFunc::Min | AggFunc::Max => false,
}
}
pub const fn merge_op(self) -> AggFunc {
if self.is_linear() {
AggFunc::Sum
} else {
self
}
}
pub const fn raw_output_nullable(self, src_nullable: bool, ungrouped: bool) -> bool {
!self.is_linear() && (src_nullable || ungrouped)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct AggDescriptor {
pub col_idx: u32,
pub agg_op: AggFunc,
}
impl AggDescriptor {
pub const COUNT_STAR: AggDescriptor = AggDescriptor { col_idx: 0, agg_op: AggFunc::Count };
}
impl Wire for AggDescriptor {
fn write(&self, w: &mut Writer) {
w.put(&self.agg_op).u32(self.col_idx);
}
fn read(r: &mut Reader) -> Result<Self, String> {
Ok(AggDescriptor { agg_op: r.get()?, col_idx: r.u32()? })
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ComputeMap {
pub program: Vec<u8>,
pub out_cols: Vec<(TypeCode, bool)>,
}
impl Wire for ComputeMap {
fn write(&self, w: &mut Writer) {
w.list(&self.out_cols).bytes32(&self.program);
}
fn read(r: &mut Reader) -> Result<Self, String> {
let out_cols = r.list("compute map", MAX_COLUMNS)?;
Ok(ComputeMap { program: r.bytes32()?.to_vec(), out_cols })
}
}
pub const fn agg_output_type(func: AggFunc, src_tc: TypeCode) -> Option<TypeCode> {
match func {
AggFunc::Count | AggFunc::CountNonNull => Some(TypeCode::I64),
AggFunc::Sum => {
if crate::ScalarKind::from_type_code(src_tc).is_some() && !src_tc.is_temporal() {
Some(src_tc.register_image())
} else {
None
}
}
AggFunc::Min | AggFunc::Max => Some(src_tc),
}
}
wire_enum! {
pub enum RangeRel: u8 {
Lt = 0,
Le = 1,
Gt = 2,
Ge = 3,
}
}
impl RangeRel {
pub fn converse(self) -> RangeRel {
match self {
RangeRel::Lt => RangeRel::Gt,
RangeRel::Le => RangeRel::Ge,
RangeRel::Gt => RangeRel::Lt,
RangeRel::Ge => RangeRel::Le,
}
}
pub fn bounds_below(self) -> bool {
matches!(self, RangeRel::Gt | RangeRel::Ge)
}
pub fn admits_equal(self) -> bool {
matches!(self, RangeRel::Ge | RangeRel::Le)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ReduceOutKey {
Natural,
SyntheticFold,
}
impl ReduceOutKey {
pub fn for_group_cols(pk_cols: &[u32], group_cols: &[u32], col: impl Fn(u32) -> (TypeCode, bool)) -> Self {
let natural = group_cols == pk_cols
|| matches!(*group_cols, [c] if {
let (type_code, nullable) = col(c);
!nullable && type_code.is_pk_eligible()
});
if natural {
ReduceOutKey::Natural
} else {
ReduceOutKey::SyntheticFold
}
}
pub fn output_layout(self, group_cols: &[u32], row: impl IntoIterator<Item = u32>) -> Vec<ReduceOutSlot> {
let (lead, spelled): (Vec<ReduceOutSlot>, &[u32]) = match self {
ReduceOutKey::SyntheticFold => (vec![ReduceOutSlot::SyntheticKey], &[]),
ReduceOutKey::Natural => (group_cols.iter().map(|&c| ReduceOutSlot::Key(c)).collect(), group_cols),
};
lead.into_iter()
.chain(
row.into_iter()
.filter(|c| !spelled.contains(c))
.map(ReduceOutSlot::Carried),
)
.collect()
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ReduceOutSlot {
SyntheticKey,
Key(u32),
Carried(u32),
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum JoinKind {
Equi,
Range {
rel: RangeRel,
},
Cross,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ClampKind {
Distinct,
PositivePart,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ReindexRole {
Auxiliary,
ScatterKey { source_key: Vec<ReindexSlot> },
}
wire_enum! {
pub enum NullKeys: u8 {
Keep = 0,
Drop = 1,
}
}
pub type ReindexSlot = (u32, TypeCode);
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum MapKind {
Projection(Vec<u32>),
Compute(ComputeMap),
Reindex {
keep: Vec<u32>,
key: Vec<ReindexSlot>,
role: ReindexRole,
nulls: NullKeys,
},
HashRow { cols: Vec<ReindexSlot> },
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum OpNode {
ScanDelta {
source: u64,
bound: crate::ReadBound,
},
Filter(Vec<u8>),
Map(MapKind),
Negate,
Union,
WeightClamp(ClampKind),
Reduce {
group_cols: Vec<u32>,
agg: Vec<AggDescriptor>,
global_ground: bool,
},
Join {
kind: JoinKind,
delta_is_right: bool,
},
IntegrateSink,
ExchangeShard {
shard_cols: Vec<u32>,
},
NullExtend {
type_codes: Vec<TypeCode>,
nulls_first: bool,
},
WorkerFilter,
TopN {
group_cols: Vec<u32>,
order: Vec<crate::OrderKey>,
limit: u64,
offset: u64,
},
}
impl OpNode {
pub const fn arity(&self) -> usize {
match self {
OpNode::ScanDelta { .. } => 0,
OpNode::Union | OpNode::Join { .. } => 2,
OpNode::Filter(_)
| OpNode::Map(_)
| OpNode::Negate
| OpNode::WeightClamp(_)
| OpNode::Reduce { .. }
| OpNode::IntegrateSink
| OpNode::ExchangeShard { .. }
| OpNode::NullExtend { .. }
| OpNode::WorkerFilter
| OpNode::TopN { .. } => 1,
}
}
fn check(&self) -> Result<(), String> {
match self {
OpNode::Map(MapKind::Reindex { key, role, .. }) => {
if key.is_empty() {
return Err("a reindex names no key columns".into());
}
if let ReindexRole::ScatterKey { source_key } = role {
if !source_key.iter().map(|s| s.1).eq(key.iter().map(|s| s.1)) {
return Err("a scatter key's slot types are not its reindex key's".into());
}
}
}
OpNode::Map(MapKind::HashRow { cols }) if cols.is_empty() => {
return Err("a hash-row map names no columns".into());
}
OpNode::Reduce { group_cols, global_ground: true, .. } if !group_cols.is_empty() => {
return Err("a global-ground reduce over a non-empty group set".into());
}
_ => {}
}
Ok(())
}
fn write(&self, w: &mut Writer) {
match self {
OpNode::ScanDelta { source, bound } => w.put(&Opcode::ScanDelta).u64(*source).put(bound),
OpNode::Filter(program) => w.put(&Opcode::Filter).bytes32(program),
OpNode::Map(MapKind::Projection(cols)) => w.put(&Opcode::MapProj).list(cols),
OpNode::Map(MapKind::Compute(map)) => w.put(&Opcode::MapExpr).put(map),
OpNode::Map(MapKind::Reindex { keep, key, role, nulls }) => {
w.put(&Opcode::MapReindex).put(nulls).list(key).list(keep);
match role {
ReindexRole::Auxiliary => w.bool(false),
ReindexRole::ScatterKey { source_key } => w.bool(true).list(source_key),
}
}
OpNode::Map(MapKind::HashRow { cols }) => w.put(&Opcode::MapHashRow).list(cols),
OpNode::Negate => w.put(&Opcode::Negate),
OpNode::Union => w.put(&Opcode::Union),
OpNode::WeightClamp(ClampKind::Distinct) => w.put(&Opcode::Distinct),
OpNode::WeightClamp(ClampKind::PositivePart) => w.put(&Opcode::PositivePart),
OpNode::Reduce { group_cols, agg, global_ground } => {
w.put(&Opcode::Reduce).bool(*global_ground).list(group_cols).list(agg)
}
OpNode::Join { kind, delta_is_right } => match kind {
JoinKind::Equi => w.put(&Opcode::JoinEqui),
JoinKind::Range { rel } => w.put(&Opcode::JoinRange).put(rel),
JoinKind::Cross => w.put(&Opcode::JoinCross),
}
.bool(*delta_is_right),
OpNode::IntegrateSink => w.put(&Opcode::IntegrateSink),
OpNode::ExchangeShard { shard_cols } => w.put(&Opcode::ExchangeShard).list(shard_cols),
OpNode::NullExtend { type_codes, nulls_first } => {
w.put(&Opcode::NullExtend).bool(*nulls_first).list(type_codes)
}
OpNode::WorkerFilter => w.put(&Opcode::WorkerFilter),
OpNode::TopN { group_cols, order, limit, offset } => w
.put(&Opcode::TopN)
.u64(*limit)
.u64(*offset)
.list(group_cols)
.list(order),
};
}
fn read(r: &mut Reader) -> Result<OpNode, String> {
fn cols(r: &mut Reader) -> Result<Vec<u32>, String> {
r.list("column list", MAX_COLUMNS)
}
fn slots(r: &mut Reader) -> Result<Vec<ReindexSlot>, String> {
r.list("slot list", MAX_COLUMNS)
}
Ok(match r.get::<Opcode>()? {
Opcode::ScanDelta => OpNode::ScanDelta { source: r.u64()?, bound: r.get()? },
Opcode::Filter => OpNode::Filter(r.bytes32()?.to_vec()),
Opcode::MapProj => OpNode::Map(MapKind::Projection(cols(r)?)),
Opcode::MapExpr => OpNode::Map(MapKind::Compute(r.get()?)),
Opcode::MapReindex => {
let nulls = r.get()?;
let key = slots(r)?;
let keep = cols(r)?;
let role = match r.bool()? {
false => ReindexRole::Auxiliary,
true => ReindexRole::ScatterKey { source_key: slots(r)? },
};
OpNode::Map(MapKind::Reindex { keep, key, role, nulls })
}
Opcode::MapHashRow => OpNode::Map(MapKind::HashRow { cols: slots(r)? }),
Opcode::Negate => OpNode::Negate,
Opcode::Union => OpNode::Union,
Opcode::Distinct => OpNode::WeightClamp(ClampKind::Distinct),
Opcode::PositivePart => OpNode::WeightClamp(ClampKind::PositivePart),
Opcode::Reduce => {
let global_ground = r.bool()?;
let group_cols = cols(r)?;
let agg = r.list("aggregate list", MAX_COLUMNS)?;
OpNode::Reduce { group_cols, agg, global_ground }
}
Opcode::JoinEqui => OpNode::Join {
kind: JoinKind::Equi,
delta_is_right: r.bool()?,
},
Opcode::JoinRange => OpNode::Join {
kind: JoinKind::Range { rel: r.get()? },
delta_is_right: r.bool()?,
},
Opcode::JoinCross => OpNode::Join {
kind: JoinKind::Cross,
delta_is_right: r.bool()?,
},
Opcode::IntegrateSink => OpNode::IntegrateSink,
Opcode::ExchangeShard => OpNode::ExchangeShard { shard_cols: cols(r)? },
Opcode::NullExtend => {
let nulls_first = r.bool()?;
let type_codes = r.list("type list", MAX_COLUMNS)?;
OpNode::NullExtend { type_codes, nulls_first }
}
Opcode::WorkerFilter => OpNode::WorkerFilter,
Opcode::TopN => {
let limit = r.u64()?;
let offset = r.u64()?;
let group_cols = cols(r)?;
let order = r.list("order keys", crate::MAX_ORDER_KEYS)?;
OpNode::TopN { group_cols, order, limit, offset }
}
})
}
}
pub type NodeId = usize;
pub const MAX_CIRCUIT_NODES: usize = 16_384;
const ARITY_MISMATCH: &str = "node's inputs do not match its operator's arity";
const EARLIER_NODE: &str = "a node's input is not an earlier node";
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Node {
pub op: OpNode,
inputs: [NodeId; 2],
}
impl Node {
pub fn inputs(&self) -> &[NodeId] {
&self.inputs[..self.op.arity()]
}
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct Circuit {
nodes: Vec<Node>,
}
impl Circuit {
pub fn nodes(&self) -> &[Node] {
&self.nodes
}
pub fn push(&mut self, op: OpNode, inputs: &[NodeId]) -> Result<NodeId, String> {
op.check()?;
if op.arity() != inputs.len() {
return Err(ARITY_MISMATCH.into());
}
if inputs.iter().any(|&p| p >= self.nodes.len()) {
return Err(EARLIER_NODE.into());
}
let mut slots = [0; 2];
slots[..inputs.len()].copy_from_slice(inputs);
self.nodes.push(Node { op, inputs: slots });
Ok(self.nodes.len() - 1)
}
pub fn encode(&self) -> Vec<u8> {
let mut w = Writer::new();
w.count(self.nodes.len());
for node in &self.nodes {
node.op.write(&mut w);
for &input in node.inputs() {
w.u16(input as u16);
}
}
w.into_vec()
}
pub fn decode(buf: &[u8]) -> Result<Circuit, String> {
decode_all(buf, "circuit", |r| {
let mut circuit = Circuit::default();
for _ in 0..r.count("nodes", MAX_CIRCUIT_NODES)? {
let op = OpNode::read(r)?;
let mut inputs = [0; 2];
let inputs = &mut inputs[..op.arity()];
for slot in inputs.iter_mut() {
*slot = r.u16()? as NodeId;
}
circuit.push(op, inputs)?;
}
Ok(circuit)
})
}
pub fn sources(&self) -> impl Iterator<Item = u64> + '_ {
self.nodes.iter().filter_map(|n| match n.op {
OpNode::ScanDelta { source, .. } => Some(source),
_ => None,
})
}
pub fn sources_mut(&mut self) -> impl Iterator<Item = &mut u64> {
self.nodes.iter_mut().filter_map(|n| match &mut n.op {
OpNode::ScanDelta { source, .. } => Some(source),
_ => None,
})
}
fn add(&mut self, op: OpNode, inputs: &[NodeId]) -> NodeId {
self.push(op, inputs)
.expect("a builder wires a well-formed operator on its arity to earlier nodes")
}
pub fn input_delta(&mut self, source: u64, bound: crate::ReadBound) -> NodeId {
self.add(OpNode::ScanDelta { source, bound }, &[])
}
pub fn filter(&mut self, input: NodeId, program: Vec<u8>) -> NodeId {
self.add(OpNode::Filter(program), &[input])
}
pub fn map_expr(&mut self, input: NodeId, map: ComputeMap) -> NodeId {
self.add(OpNode::Map(MapKind::Compute(map)), &[input])
}
pub fn map_reindex(
&mut self,
input: NodeId,
key: &[ReindexSlot],
keep: &[u32],
role: ReindexRole,
nulls: NullKeys,
) -> NodeId {
let op = OpNode::Map(MapKind::Reindex {
keep: keep.to_vec(),
key: key.to_vec(),
role,
nulls,
});
self.add(op, &[input])
}
pub fn map_hash_row(&mut self, input: NodeId, cols: &[ReindexSlot]) -> NodeId {
let map = self.add(OpNode::Map(MapKind::HashRow { cols: cols.to_vec() }), &[input]);
self.shard(map, &[0])
}
pub fn map(&mut self, input: NodeId, projection: &[u32]) -> NodeId {
self.add(OpNode::Map(MapKind::Projection(projection.to_vec())), &[input])
}
pub fn negate(&mut self, input: NodeId) -> NodeId {
self.add(OpNode::Negate, &[input])
}
pub fn union(&mut self, a: NodeId, b: NodeId) -> NodeId {
self.add(OpNode::Union, &[a, b])
}
pub fn distinct(&mut self, input: NodeId) -> NodeId {
self.add(OpNode::WeightClamp(ClampKind::Distinct), &[input])
}
pub fn difference(&mut self, minuend: NodeId, subtrahend: NodeId) -> NodeId {
let neg = self.negate(subtrahend);
self.union(neg, minuend)
}
pub fn positive_diff(&mut self, minuend: NodeId, subtrahend: NodeId) -> NodeId {
let diff = self.difference(minuend, subtrahend);
self.add(OpNode::WeightClamp(ClampKind::PositivePart), &[diff])
}
pub fn join(&mut self, delta: NodeId, integrand: NodeId, kind: JoinKind, delta_is_right: bool) -> NodeId {
self.add(OpNode::Join { kind, delta_is_right }, &[delta, integrand])
}
pub fn join_terms(&mut self, [da, db]: [NodeId; 2], [ia, ib]: [NodeId; 2], kind: JoinKind) -> NodeId {
let ab = self.join(da, ib, kind, false);
let ba = self.join(db, ia, kind, true);
self.union(ab, ba)
}
pub fn worker_filter(&mut self, input: NodeId) -> NodeId {
self.add(OpNode::WorkerFilter, &[input])
}
pub fn reduce_multi(&mut self, input: NodeId, group_cols: &[u32], agg_specs: &[AggDescriptor]) -> NodeId {
let sharded = self.shard(input, group_cols);
self.reduce(sharded, group_cols, agg_specs, group_cols.is_empty())
}
pub fn reduce_multi_local(&mut self, input: NodeId, group_cols: &[u32], agg_specs: &[AggDescriptor]) -> NodeId {
self.reduce(input, group_cols, agg_specs, false)
}
fn reduce(&mut self, input: NodeId, group_cols: &[u32], agg: &[AggDescriptor], global_ground: bool) -> NodeId {
let op = OpNode::Reduce {
group_cols: group_cols.to_vec(),
agg: agg.to_vec(),
global_ground,
};
self.add(op, &[input])
}
pub fn top_n(
&mut self,
input: NodeId,
group_cols: &[u32],
order: &[crate::OrderKey],
limit: u64,
offset: u64,
) -> NodeId {
let sharded = self.shard(input, group_cols);
let op = OpNode::TopN {
group_cols: group_cols.to_vec(),
order: order.to_vec(),
limit,
offset,
};
self.add(op, &[sharded])
}
pub fn shard(&mut self, input: NodeId, shard_cols: &[u32]) -> NodeId {
let op = OpNode::ExchangeShard { shard_cols: shard_cols.to_vec() };
self.add(op, &[input])
}
pub fn null_extend(&mut self, input: NodeId, type_codes: &[TypeCode], nulls_first: bool) -> NodeId {
let op = OpNode::NullExtend {
type_codes: type_codes.to_vec(),
nulls_first,
};
self.add(op, &[input])
}
pub fn sink(&mut self, input: NodeId) -> NodeId {
self.add(OpNode::IntegrateSink, &[input])
}
}
#[cfg(test)]
#[path = "tests/circuit.rs"]
mod tests;