Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
40 changes: 27 additions & 13 deletions Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -11,18 +11,19 @@ repository = "https://github.com/egraphs-good/egg"
version = "0.11.0"

[dependencies]
env_logger = {version = "0.9.0", default-features = false}
hashbrown = "0.15.2"
indexmap = "2.7.0"
log = "0.4.17"
num-bigint = "0.4"
num-traits = "0.2"
quanta = "0.12"
rustc-hash = "2.0.0"
env_logger = {version = "0.9.0", default-features = false, optional = true}
hashbrown = {version = "0.15.2", default-features = false}
indexmap = {version = "2.7.0", default-features = false}
log = {version = "0.4.17", default-features = false}
num-bigint = {version = "0.4", default-features = false}
num-traits = {version = "0.2", default-features = false}
quanta = {version = "0.12", optional = true}
rustc-hash = {version = "2.0.0", default-features = false}
smallvec = {version = "1.8.0", features = ["union", "const_generics"]}
symbol_table = {version = "0.4.0", features = ["global"]}
symbolic_expressions = "5.0.3"
thiserror = "1.0.31"
spin = {version = "0.9", default-features = false, features = ["spin_mutex", "lazy"]}
symbol_table = {git = "https://github.com/mwillsey/symbol_table", rev = "a716c39a4aa4757c910c93fc10c5c7fa8f467133", default-features = false, features = ["global"]}
symbolic_expressions = {version = "5.0.3", optional = true}
thiserror = "2"

# for the lp feature
good_lp = { version = "1", optional = true }
Expand All @@ -38,9 +39,22 @@ serde_json = {version = "1.0.81", optional = true}
ordered-float = "3.0.0"

[features]
default = ["std"]
std = [
"dep:env_logger",
"dep:quanta",
"dep:symbolic_expressions",
"hashbrown/default",
"indexmap/std",
"num-bigint/std",
"num-traits/std",
"log/std",
"rustc-hash/std",
"symbol_table/std",
]
# forces the use of indexmaps over hashmaps
deterministic = []
lp = ["good_lp"]
lp = ["std", "good_lp"]
reports = ["serde-1", "serde_json"]
serde-1 = [
"serde",
Expand All @@ -49,7 +63,7 @@ serde-1 = [
"symbol_table/serde",
"vectorize",
]
wasm-bindgen = []
wasm-bindgen = ["std"]

# private features for testing
test-explanations = []
Expand Down
4 changes: 3 additions & 1 deletion Makefile
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,9 @@ test:
cargo test --release --features=lp
# don't run examples in proof-production mode
cargo test --release --features "test-explanations"

# verify no_std build
cargo test --no-default-features


.PHONY: nits
nits:
Expand Down
15 changes: 6 additions & 9 deletions src/dot.rs
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@ Use the [`Dot`] struct to visualize an [`EGraph`]

use std::ffi::OsStr;
use std::fmt::{self, Debug, Display, Formatter};
use std::io::{Error, ErrorKind, Result, Write};
use std::io::{Error, Result, Write};
use std::path::Path;

use crate::{Analysis, Language, egraph::EGraph};
Expand Down Expand Up @@ -140,14 +140,11 @@ where
write!(stdin, "{}", self)?;
match child.wait()?.code() {
Some(0) => Ok(()),
Some(e) => Err(Error::new(
ErrorKind::Other,
format!("dot program returned error code {}", e),
)),
None => Err(Error::new(
ErrorKind::Other,
"dot program was killed by a signal",
)),
Some(e) => Err(Error::other(format!(
"dot program returned error code {}",
e
))),
None => Err(Error::other("dot program was killed by a signal")),
}
}

Expand Down
5 changes: 3 additions & 2 deletions src/eclass.rs
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
use std::fmt::Debug;
use std::iter::ExactSizeIterator;
use crate::no_std_prelude::*;
use core::fmt::Debug;
use core::iter::ExactSizeIterator;

use crate::*;

Expand Down
30 changes: 16 additions & 14 deletions src/egraph.rs
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
use crate::no_std_prelude::*;
use crate::*;
use std::{
use core::{
borrow::BorrowMut,
fmt::{self, Debug, Display},
marker::PhantomData,
Expand Down Expand Up @@ -577,6 +578,7 @@ impl<L: Language, N: Analysis<L>> EGraph<L, N> {
}

/// Creates a [`Dot`] to visualize this egraph. See [`Dot`].
#[cfg(feature = "std")]
pub fn dot(&self) -> Dot<'_, L, N> {
Dot {
egraph: self,
Expand Down Expand Up @@ -792,7 +794,7 @@ where
}

/// Given an `Id` using the `egraph[id]` syntax, retrieve the e-class.
impl<L: Language, N: Analysis<L>> std::ops::Index<Id> for EGraph<L, N> {
impl<L: Language, N: Analysis<L>> core::ops::Index<Id> for EGraph<L, N> {
type Output = EClass<L, N::Data>;
fn index(&self, id: Id) -> &Self::Output {
let id = self.find(id);
Expand All @@ -804,7 +806,7 @@ impl<L: Language, N: Analysis<L>> std::ops::Index<Id> for EGraph<L, N> {

/// Given an `Id` using the `&mut egraph[id]` syntax, retrieve a mutable
/// reference to the e-class.
impl<L: Language, N: Analysis<L>> std::ops::IndexMut<Id> for EGraph<L, N> {
impl<L: Language, N: Analysis<L>> core::ops::IndexMut<Id> for EGraph<L, N> {
fn index_mut(&mut self, id: Id) -> &mut Self::Output {
let id = self.find_mut(id);
self.classes
Expand Down Expand Up @@ -1144,7 +1146,7 @@ impl<L: Language, N: Analysis<L>> EGraph<L, N> {
#[track_caller]
pub fn union(&mut self, id1: Id, id2: Id) -> bool {
if self.explain.is_some() {
let caller = std::panic::Location::caller();
let caller = core::panic::Location::caller();
self.union_trusted(id1, id2, caller.to_string())
} else {
self.perform_union(id1, id2, None)
Expand All @@ -1158,18 +1160,18 @@ impl<L: Language, N: Analysis<L>> EGraph<L, N> {
let mut id1 = self.find_mut(enode_id1);
let mut id2 = self.find_mut(enode_id2);
if id1 == id2 {
if let Some(Justification::Rule(_)) = rule {
if let Some(explain) = &mut self.explain {
explain.alternate_rewrite(enode_id1, enode_id2, rule.unwrap());
}
if let Some(Justification::Rule(_)) = rule
&& let Some(explain) = &mut self.explain
{
explain.alternate_rewrite(enode_id1, enode_id2, rule.unwrap());
}
return false;
}
// make sure class2 has fewer parents
let class1_parents = self.classes[&id1].parents.len();
let class2_parents = self.classes[&id2].parents.len();
if class1_parents < class2_parents {
std::mem::swap(&mut id1, &mut id2);
core::mem::swap(&mut id1, &mut id2);
}

if let Some(explain) = &mut self.explain {
Expand All @@ -1180,6 +1182,7 @@ impl<L: Language, N: Analysis<L>> EGraph<L, N> {
self.unionfind.union(id1, id2);

assert_ne!(id1, id2);
#[allow(deprecated)]
let class2 = self.classes.remove(&id2).unwrap();
let class1 = self.classes.get_mut(&id1).unwrap();
assert_eq!(id1, class1.id);
Expand Down Expand Up @@ -1232,10 +1235,10 @@ impl<L: Language + Display, N: Analysis<L>> EGraph<L, N> {
/// Useful for testing.
pub fn check_goals(&self, id: Id, goals: &[Pattern<L>]) {
let (cost, best) = Extractor::new(self, AstSize).find_best(id);
println!("End ({}): {}", cost, best.pretty(80));
log::info!("End ({}): {}", cost, best.pretty(80));

for (i, goal) in goals.iter().enumerate() {
println!("Trying to prove goal {}: {}", i, goal.pretty(40));
log::info!("Trying to prove goal {}: {}", i, goal.pretty(40));
let matches = goal.search_eclass(self, id);
if matches.is_none() {
let best = Extractor::new(self, AstSize).find_best(id).1;
Expand All @@ -1257,7 +1260,7 @@ impl<L: Language + Display, N: Analysis<L>> EGraph<L, N> {
impl<L: Language, N: Analysis<L>> EGraph<L, N> {
#[inline(never)]
fn rebuild_classes(&mut self) -> usize {
let mut classes_by_op = std::mem::take(&mut self.classes_by_op);
let mut classes_by_op = core::mem::take(&mut self.classes_by_op);
classes_by_op.values_mut().for_each(|ids| ids.clear());

let mut trimmed = 0;
Expand Down Expand Up @@ -1398,7 +1401,6 @@ impl<L: Language, N: Analysis<L>> EGraph<L, N> {
/// let y = egraph.add(S::leaf("y"));
/// let ax = egraph.add_expr(&"(+ a x)".parse().unwrap());
/// let ay = egraph.add_expr(&"(+ a y)".parse().unwrap());

/// // Union x and y
/// egraph.union(x, y);
/// // Classes: [x y] [ax] [ay] [a]
Expand Down Expand Up @@ -1467,7 +1469,7 @@ impl<'a, L: Language, N: Analysis<L>> Debug for EGraphDump<'a, L, N> {
}
}

#[cfg(test)]
#[cfg(all(test, feature = "std"))]
mod tests {

use super::*;
Expand Down
Loading
Loading