From 9eec63fab890b5e5c1a442a284ae1805649ef600 Mon Sep 17 00:00:00 2001 From: Eli Rosenthal Date: Wed, 5 Aug 2026 00:50:34 -0700 Subject: [PATCH] Move Set containers to sequence storage --- src/sort/set.rs | 530 +++++++++++++++++- tests/container_rebuild.rs | 26 + tests/mixed-container-dirty-propagation.egg | 18 + ...hot_mixed_container_dirty_propagation.snap | 7 + 4 files changed, 568 insertions(+), 13 deletions(-) create mode 100644 tests/mixed-container-dirty-propagation.egg create mode 100644 tests/snapshots/files__shared_snapshot_mixed_container_dirty_propagation.snap diff --git a/src/sort/set.rs b/src/sort/set.rs index 654162d05..567e669a3 100644 --- a/src/sort/set.rs +++ b/src/sort/set.rs @@ -1,4 +1,5 @@ use super::*; +use crate::numeric_id::NumericId; use std::collections::BTreeSet; #[derive(Clone, Debug, PartialEq, Eq, Hash)] @@ -23,6 +24,77 @@ impl ContainerValue for SetContainer { } } +impl SequenceContainerValue for SetContainer { + fn encode_sequence(&self, out: &mut Vec) { + out.push(Value::from_usize(self.do_rebuild as usize)); + out.extend(self.data.iter().copied()); + } + + fn decode_sequence(sequence: &[Value]) -> Self { + let (&header, data) = sequence + .split_first() + .expect("serialized SetContainer must include its rebuild flag"); + assert!( + header == Value::from_usize(0) || header == Value::from_usize(1), + "serialized SetContainer has an invalid rebuild flag" + ); + debug_assert!(data.is_sorted()); + debug_assert!(data.windows(2).all(|pair| pair[0] != pair[1])); + Self { + do_rebuild: header == Value::from_usize(1), + data: data.iter().copied().collect(), + } + } + + fn sequence_values(sequence: &[Value]) -> &[Value] { + sequence + .get(1..) + .expect("serialized SetContainer must include its rebuild flag") + } + + fn visit_sequence_values(sequence: &[Value], visitor: &mut dyn FnMut(Value)) { + let (&header, data) = sequence + .split_first() + .expect("serialized SetContainer must include its rebuild flag"); + match header.index() { + 0 => {} + 1 => data.iter().copied().for_each(visitor), + _ => panic!("serialized SetContainer has an invalid rebuild flag"), + } + } + + fn rebuild_sequence( + sequence: &[Value], + rebuilder: &dyn ValueRebuilder, + out: &mut Vec, + ) -> bool { + let (&header, data) = sequence + .split_first() + .expect("serialized SetContainer must include its rebuild flag"); + if header == Value::from_usize(0) { + return false; + } + assert_eq!( + header, + Value::from_usize(1), + "serialized SetContainer has an invalid rebuild flag" + ); + + let mut rebuilt = data.to_vec(); + if !rebuilder.rebuild_slice(&mut rebuilt) { + return false; + } + rebuilt.sort_unstable(); + rebuilt.dedup(); + if rebuilt == data { + return false; + } + out.push(header); + out.extend(rebuilt); + true + } +} + /// The elements of a `(set-of e0 ...)` term as a Rust `BTreeSet` in AST /// order, matching `SetContainer`'s semantics; `None` for any other term. fn set_term_to_btreeset<'a>(termdag: &'a TermDag, term: TermId) -> Option>> { @@ -113,6 +185,10 @@ impl ContainerSort for SetSort { &self.name } + fn register_type(&self, backend: &mut egglog_bridge::EGraph) { + backend.register_sequence_container_ty::(); + } + fn inner_sorts(&self) -> Vec { vec![self.element.clone()] } @@ -229,19 +305,88 @@ impl ContainerSort for SetSort { data: xs.collect() } }, set_of_validator); - // No validator: `set-get` indexes the runtime `BTreeSet` order, - // which terms cannot reproduce, so it is unsupported in proof mode. - add_primitive!(eg, "set-get" = |xs: @SetContainer (arc), i: i64| -?> # (self.element()) { xs.data.iter().nth(i as usize).copied() }); - add_primitive_with_validator!(eg, "set-insert" = |mut xs: @SetContainer (arc), x: # (self.element())| -> @SetContainer (arc) {{ xs.data.insert( x); xs }}, set_insert_validator); - add_primitive_with_validator!(eg, "set-remove" = |mut xs: @SetContainer (arc), x: # (self.element())| -> @SetContainer (arc) {{ xs.data.remove(&x); xs }}, set_remove_validator); - - add_primitive_with_validator!(eg, "set-length" = |xs: @SetContainer (arc)| -> i64 { xs.data.len() as i64 }, set_length_validator); - add_primitive_with_validator!(eg, "set-contains" = |xs: @SetContainer (arc), x: # (self.element())| -?> () { ( xs.data.contains(&x)).then_some(()) }, set_contains_validator); - add_primitive_with_validator!(eg, "set-not-contains" = |xs: @SetContainer (arc), x: # (self.element())| -?> () { (!xs.data.contains(&x)).then_some(()) }, set_not_contains_validator); - - add_primitive_with_validator!(eg, "set-union" = |mut xs: @SetContainer (arc), ys: @SetContainer (arc)| -> @SetContainer (arc) {{ xs.data.extend(ys.data); xs }}, set_union_validator); - add_primitive_with_validator!(eg, "set-diff" = |mut xs: @SetContainer (arc), ys: @SetContainer (arc)| -> @SetContainer (arc) {{ xs.data.retain(|k| !ys.data.contains(k)); xs }}, set_diff_validator); - add_primitive_with_validator!(eg, "set-intersect" = |mut xs: @SetContainer (arc), ys: @SetContainer (arc)| -> @SetContainer (arc) {{ xs.data.retain(|k| ys.data.contains(k)); xs }}, set_intersect_validator); + // No validator: `set-get` indexes the runtime `Value` order, which + // terms cannot reproduce, so it is unsupported in proof mode. + eg.add_pure_primitive( + SetRead { + name: "set-get".into(), + set: arc.clone(), + element: self.element(), + op: SetReadOp::Get, + }, + None, + ); + eg.add_pure_primitive( + SetEdit { + name: "set-insert".into(), + set: arc.clone(), + element: self.element(), + op: SetEditOp::Insert, + }, + Some(Arc::new(set_insert_validator)), + ); + eg.add_pure_primitive( + SetEdit { + name: "set-remove".into(), + set: arc.clone(), + element: self.element(), + op: SetEditOp::Remove, + }, + Some(Arc::new(set_remove_validator)), + ); + for (name, op, validator) in [ + ( + "set-length", + SetReadOp::Length, + Arc::new(set_length_validator) as PrimitiveValidator, + ), + ( + "set-contains", + SetReadOp::Contains, + Arc::new(set_contains_validator) as PrimitiveValidator, + ), + ( + "set-not-contains", + SetReadOp::NotContains, + Arc::new(set_not_contains_validator) as PrimitiveValidator, + ), + ] { + eg.add_pure_primitive( + SetRead { + name: name.into(), + set: arc.clone(), + element: self.element(), + op, + }, + Some(validator), + ); + } + for (name, op, validator) in [ + ( + "set-union", + SetBinaryOp::Union, + Arc::new(set_union_validator) as PrimitiveValidator, + ), + ( + "set-diff", + SetBinaryOp::Diff, + Arc::new(set_diff_validator) as PrimitiveValidator, + ), + ( + "set-intersect", + SetBinaryOp::Intersect, + Arc::new(set_intersect_validator) as PrimitiveValidator, + ), + ] { + eg.add_pure_primitive( + SetBinary { + name: name.into(), + set: arc.clone(), + op, + }, + Some(validator), + ); + } } fn reconstruct_termdag( @@ -269,3 +414,362 @@ impl ContainerSort for SetSort { "set-of".to_owned() } } + +#[derive(Clone, Copy)] +enum SetReadOp { + Get, + Length, + Contains, + NotContains, +} + +#[derive(Clone)] +struct SetRead { + name: String, + set: ArcSort, + element: ArcSort, + op: SetReadOp, +} + +impl Primitive for SetRead { + fn name(&self) -> &str { + &self.name + } + + fn get_type_constraints(&self, span: &Span) -> Box { + let types = match self.op { + SetReadOp::Get => vec![self.set.clone(), I64Sort.to_arcsort(), self.element.clone()], + SetReadOp::Length => vec![self.set.clone(), I64Sort.to_arcsort()], + SetReadOp::Contains | SetReadOp::NotContains => vec![ + self.set.clone(), + self.element.clone(), + UnitSort.to_arcsort(), + ], + }; + SimpleTypeConstraint::new(self.name(), types, span.clone()).into_box() + } +} + +impl PurePrim for SetRead { + fn apply<'a, 'db>(&self, state: crate::PureState<'a, 'db>, args: &[Value]) -> Option { + let [set_id, rest @ ..] = args else { + return None; + }; + match self.op { + SetReadOp::Get => { + let [index] = rest else { return None }; + let index = usize::try_from(state.base_values().unwrap::(*index)).ok()?; + state + .with_container_sequence::(*set_id, |values| { + values.get(index).copied() + }) + .or_else(|| { + state + .value_to_owned_container::(*set_id) + .map(|set| set.data.iter().nth(index).copied()) + })? + } + SetReadOp::Length => { + if !rest.is_empty() { + return None; + } + let len = state + .with_container_sequence::(*set_id, <[Value]>::len) + .or_else(|| { + state + .value_to_owned_container::(*set_id) + .map(|set| set.data.len()) + })?; + Some(state.base_values().get::(len as i64)) + } + SetReadOp::Contains | SetReadOp::NotContains => { + let [needle] = rest else { return None }; + let contains = state + .with_container_sequence::(*set_id, |values| { + values.binary_search(needle).is_ok() + }) + .or_else(|| { + state + .value_to_owned_container::(*set_id) + .map(|set| set.data.contains(needle)) + })?; + let succeeds = match self.op { + SetReadOp::Contains => contains, + SetReadOp::NotContains => !contains, + _ => unreachable!(), + }; + succeeds.then(|| state.base_values().get::<()>(())) + } + } + } +} + +#[derive(Clone, Copy)] +enum SetEditOp { + Insert, + Remove, +} + +#[derive(Clone)] +struct SetEdit { + name: String, + set: ArcSort, + element: ArcSort, + op: SetEditOp, +} + +impl Primitive for SetEdit { + fn name(&self) -> &str { + &self.name + } + + fn get_type_constraints(&self, span: &Span) -> Box { + SimpleTypeConstraint::new( + self.name(), + vec![self.set.clone(), self.element.clone(), self.set.clone()], + span.clone(), + ) + .into_box() + } +} + +impl PurePrim for SetEdit { + fn apply<'a, 'db>( + &self, + mut state: crate::PureState<'a, 'db>, + args: &[Value], + ) -> Option { + let [set_id, needle] = args else { return None }; + let build_key = |values: &[Value]| { + let mut key = Vec::with_capacity(values.len() + 2); + key.push(Value::from_usize(self.set.is_eq_container_sort() as usize)); + match (self.op, values.binary_search(needle)) { + (SetEditOp::Insert, Ok(_)) | (SetEditOp::Remove, Err(_)) => { + key.extend_from_slice(values); + } + (SetEditOp::Insert, Err(index)) => { + key.extend_from_slice(&values[..index]); + key.push(*needle); + key.extend_from_slice(&values[index..]); + } + (SetEditOp::Remove, Ok(index)) => { + key.extend_from_slice(&values[..index]); + key.extend_from_slice(&values[index + 1..]); + } + } + key + }; + let key = state + .with_container_sequence::(*set_id, build_key) + .or_else(|| { + state + .value_to_owned_container::(*set_id) + .map(|set| build_key(&set.data.into_iter().collect::>())) + })?; + Some(state.register_container_sequence::(&key)) + } +} + +#[derive(Clone, Copy)] +enum SetBinaryOp { + Union, + Diff, + Intersect, +} + +#[derive(Clone)] +struct SetBinary { + name: String, + set: ArcSort, + op: SetBinaryOp, +} + +impl Primitive for SetBinary { + fn name(&self) -> &str { + &self.name + } + + fn get_type_constraints(&self, span: &Span) -> Box { + SimpleTypeConstraint::new( + self.name(), + vec![self.set.clone(), self.set.clone(), self.set.clone()], + span.clone(), + ) + .into_box() + } +} + +impl PurePrim for SetBinary { + fn apply<'a, 'db>( + &self, + mut state: crate::PureState<'a, 'db>, + args: &[Value], + ) -> Option { + let [left, right] = args else { return None }; + let read = |id| { + state + .with_container_sequence::(id, <[Value]>::to_vec) + .or_else(|| { + state + .value_to_owned_container::(id) + .map(|set| set.data.into_iter().collect()) + }) + }; + let left = read(*left)?; + let right = read(*right)?; + let mut key = Vec::with_capacity(left.len() + right.len() + 1); + key.push(Value::from_usize(self.set.is_eq_container_sort() as usize)); + merge_sets(&left, &right, self.op, &mut key); + Some(state.register_container_sequence::(&key)) + } +} + +fn merge_sets(left: &[Value], right: &[Value], op: SetBinaryOp, out: &mut Vec) { + let (mut l, mut r) = (0, 0); + while l < left.len() && r < right.len() { + match left[l].cmp(&right[r]) { + std::cmp::Ordering::Less => { + if matches!(op, SetBinaryOp::Union | SetBinaryOp::Diff) { + out.push(left[l]); + } + l += 1; + } + std::cmp::Ordering::Greater => { + if matches!(op, SetBinaryOp::Union) { + out.push(right[r]); + } + r += 1; + } + std::cmp::Ordering::Equal => { + if matches!(op, SetBinaryOp::Union | SetBinaryOp::Intersect) { + out.push(left[l]); + } + l += 1; + r += 1; + } + } + } + if matches!(op, SetBinaryOp::Union | SetBinaryOp::Diff) { + out.extend_from_slice(&left[l..]); + } + if matches!(op, SetBinaryOp::Union) { + out.extend_from_slice(&right[r..]); + } +} + +#[cfg(test)] +mod tests { + use super::*; + + struct Collapse { + from: Value, + to: Value, + } + + impl ValueRebuilder for Collapse { + fn rebuild_val(&self, value: Value) -> Value { + if value == self.from { self.to } else { value } + } + } + + fn value(index: usize) -> Value { + Value::from_usize(index) + } + + #[test] + fn sequence_codec_rebuilds_sorts_and_collapses() { + let set = SetContainer { + do_rebuild: true, + data: [value(2), value(4), value(6)].into_iter().collect(), + }; + let mut encoded = Vec::new(); + set.encode_sequence(&mut encoded); + assert_eq!(SetContainer::decode_sequence(&encoded), set); + + let mut rebuilt = Vec::new(); + assert!(SetContainer::rebuild_sequence( + &encoded, + &Collapse { + from: value(6), + to: value(2), + }, + &mut rebuilt, + )); + assert_eq!(rebuilt, vec![value(1), value(2), value(4)]); + assert_eq!( + SetContainer::sequence_values(&rebuilt), + &[value(2), value(4)] + ); + } + + #[test] + fn sequence_codec_skips_non_rebuildable_sets() { + let encoded = vec![value(0), value(2), value(4)]; + let mut rebuilt = Vec::new(); + assert!(!SetContainer::rebuild_sequence( + &encoded, + &Collapse { + from: value(2), + to: value(4), + }, + &mut rebuilt, + )); + assert!(rebuilt.is_empty()); + } + + #[test] + fn sequence_primitives_consume_local_predictions() { + let mut egraph = EGraph::default(); + egraph + .parse_and_run_program( + None, + r#" + (sort IntSet (Set i64)) + (check (= 4 + (set-length + (set-intersect + (set-insert (set-union (set-of 3 1) (set-of 2 4)) 4) + (set-of 4 3 2 1 1))))) + (check (= (set-of 1 3) + (set-diff (set-remove (set-of 1 2 3) 9) (set-of 2)))) + (check (set-contains (set-insert (set-empty) 7) 7)) + (check (set-not-contains (set-remove (set-of 7) 7) 7)) + "#, + ) + .unwrap(); + } + + #[test] + fn sorted_slice_algebra_matches_btree_sets_exhaustively() { + for left_mask in 0u8..16 { + for right_mask in 0u8..16 { + let left = (0..4) + .filter(|bit| left_mask & (1 << bit) != 0) + .map(|bit| value(bit + 10)) + .collect::>(); + let right = (0..4) + .filter(|bit| right_mask & (1 << bit) != 0) + .map(|bit| value(bit + 10)) + .collect::>(); + let left_set = left.iter().copied().collect::>(); + let right_set = right.iter().copied().collect::>(); + for op in [ + SetBinaryOp::Union, + SetBinaryOp::Diff, + SetBinaryOp::Intersect, + ] { + let mut actual = Vec::new(); + merge_sets(&left, &right, op, &mut actual); + let expected: Vec<_> = match op { + SetBinaryOp::Union => left_set.union(&right_set).copied().collect(), + SetBinaryOp::Diff => left_set.difference(&right_set).copied().collect(), + SetBinaryOp::Intersect => { + left_set.intersection(&right_set).copied().collect() + } + }; + assert_eq!(actual, expected); + } + } + } + } +} diff --git a/tests/container_rebuild.rs b/tests/container_rebuild.rs index 6e120f443..acde39ce0 100644 --- a/tests/container_rebuild.rs +++ b/tests/container_rebuild.rs @@ -401,3 +401,29 @@ fn map_rebuild_collapse_proof_mode() { ) .unwrap(); } + +/// Sequence-backed Set dependencies propagate through another Set. The outer +/// container must be revisited when rebuilding changes an inner identity. +#[test] +fn nested_set_rebuild_term_only() { + let mut egraph = EGraph::new_with_term_encoding(); + egraph + .parse_and_run_program( + None, + r#" + (sort Math) + (constructor A () Math) + (constructor B () Math) + (sort MathSet (Set Math)) + (sort NestedSet (Set MathSet)) + (constructor Holds (NestedSet) Math) + (Holds (set-of (set-of (A)))) + (Holds (set-of (set-of (B)))) + (union (A) (B)) + (run 1) + (check (= (Holds (set-of (set-of (A)))) + (Holds (set-of (set-of (B)))))) + "#, + ) + .unwrap(); +} diff --git a/tests/mixed-container-dirty-propagation.egg b/tests/mixed-container-dirty-propagation.egg new file mode 100644 index 000000000..c9424dc0e --- /dev/null +++ b/tests/mixed-container-dirty-propagation.egg @@ -0,0 +1,18 @@ +;; Rebuilding `w(b)` to `b` changes the Set in place. The enclosing Vec still +;; contains the same Set id, so its own row is unchanged too. Stable-id dirty +;; propagation must cross both sequence-container tables and revisit the +;; ordinary `p` row, allowing the opaque nested matcher to fire. +(sort E) +(sort SE (Set E)) +(sort VSE (Vec SE)) + +(constructor b () E) +(constructor w (E) E) +(constructor p (VSE) E) + +(rewrite (w x) x) +(rewrite (p (vec-of (set-of (b)))) (b)) + +(let $nested (p (vec-of (set-of (w (b)))))) +(run-schedule (saturate (run))) +(check (= $nested (b))) diff --git a/tests/snapshots/files__shared_snapshot_mixed_container_dirty_propagation.snap b/tests/snapshots/files__shared_snapshot_mixed_container_dirty_propagation.snap new file mode 100644 index 000000000..b9793ed2d --- /dev/null +++ b/tests/snapshots/files__shared_snapshot_mixed_container_dirty_propagation.snap @@ -0,0 +1,7 @@ +--- +source: tests/files.rs +expression: snapshot_content_across_treatments +--- +((b 1) + (p 1) + (w 1))