use std::collections::HashMap;
use std::collections::hash_map::Entry;
use crate::allocator::{Allocator, Atom, NodePtr, SExp};
use crate::error::Result;
use super::bytes32::Bytes32;
use super::object_cache::{ObjectCache, treehash};
#[derive(Debug)]
pub struct InternedTree {
pub allocator: Allocator,
pub root: NodePtr,
pub atoms: Vec<NodePtr>,
pub pairs: Vec<NodePtr>,
}
impl InternedTree {
pub fn tree_hash(&self) -> [u8; 32] {
let mut cache: ObjectCache<Bytes32> = ObjectCache::new(treehash);
*cache
.get_or_calculate(&self.allocator, &self.root, None)
.expect("treehash should not fail on valid tree")
}
}
pub fn intern_tree_limited(
source: &Allocator,
node: NodePtr,
heap_limit: usize,
) -> Result<InternedTree> {
let mut new_allocator = Allocator::new_limited(heap_limit);
let mut atoms: Vec<NodePtr> = Vec::new();
let mut pairs: Vec<NodePtr> = Vec::new();
let mut node_to_interned: HashMap<NodePtr, NodePtr> = HashMap::new();
let mut atom_to_interned: HashMap<Atom, NodePtr> = HashMap::new();
let mut pair_to_interned: HashMap<(NodePtr, NodePtr), NodePtr> = HashMap::new();
let mut stack = vec![node];
while let Some(current) = stack.pop() {
if node_to_interned.contains_key(¤t) {
continue;
}
match source.sexp(current) {
SExp::Atom => {
let atom = source.atom(current);
let interned = match atom_to_interned.entry(atom) {
Entry::Occupied(o) => *o.get(),
Entry::Vacant(v) => {
let new_node = new_allocator.new_atom(atom.as_ref())?;
v.insert(new_node);
atoms.push(new_node);
new_node
}
};
node_to_interned.insert(current, interned);
}
SExp::Pair(left, right) => {
let left_interned = node_to_interned.get(&left);
let right_interned = node_to_interned.get(&right);
if let (Some(l), Some(r)) = (left_interned, right_interned) {
let interned = match pair_to_interned.entry((*l, *r)) {
Entry::Occupied(o) => *o.get(),
Entry::Vacant(v) => {
let new_node = new_allocator.new_pair(*l, *r)?;
v.insert(new_node);
pairs.push(new_node);
new_node
}
};
node_to_interned.insert(current, interned);
} else {
stack.push(current);
if right_interned.is_none() {
stack.push(right);
}
if left_interned.is_none() {
stack.push(left);
}
}
}
}
}
let root = node_to_interned[&node];
Ok(InternedTree {
allocator: new_allocator,
root,
atoms,
pairs,
})
}
pub fn intern_tree(source: &Allocator, node: NodePtr) -> Result<InternedTree> {
intern_tree_limited(source, node, u32::MAX as usize)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_intern_single_atom() {
let mut allocator = Allocator::new();
let node = allocator.new_atom(&[1, 2, 3]).unwrap();
let tree = intern_tree(&allocator, node).unwrap();
assert_eq!(tree.atoms.len(), 1);
assert_eq!(tree.pairs.len(), 0);
assert_eq!(tree.allocator.atom(tree.root).as_ref(), &[1, 2, 3]);
}
#[test]
fn test_intern_simple_pair() {
let mut allocator = Allocator::new();
let left = allocator.new_atom(&[1]).unwrap();
let right = allocator.new_atom(&[2]).unwrap();
let node = allocator.new_pair(left, right).unwrap();
let tree = intern_tree(&allocator, node).unwrap();
assert_eq!(tree.atoms.len(), 2);
assert_eq!(tree.pairs.len(), 1);
}
#[test]
fn test_intern_deduplicates_atoms() {
let mut allocator = Allocator::new();
let a1 = allocator.new_atom(&[42]).unwrap();
let a2 = allocator.new_atom(&[42]).unwrap(); let node = allocator.new_pair(a1, a2).unwrap();
let tree = intern_tree(&allocator, node).unwrap();
assert_eq!(tree.atoms.len(), 1);
assert_eq!(tree.pairs.len(), 1);
}
#[test]
fn test_intern_deduplicates_pairs() {
let mut allocator = Allocator::new();
let a = allocator.new_atom(&[1]).unwrap();
let b = allocator.new_atom(&[2]).unwrap();
let p1 = allocator.new_pair(a, b).unwrap();
let p2 = allocator.new_pair(a, b).unwrap(); let node = allocator.new_pair(p1, p2).unwrap();
let tree = intern_tree(&allocator, node).unwrap();
assert_eq!(tree.atoms.len(), 2);
assert_eq!(tree.pairs.len(), 2); }
#[test]
fn test_tree_hash_deterministic() {
let mut alloc1 = Allocator::new();
let a1 = alloc1.new_atom(&[1, 2, 3]).unwrap();
let b1 = alloc1.new_atom(&[4, 5, 6]).unwrap();
let node1 = alloc1.new_pair(a1, b1).unwrap();
let mut alloc2 = Allocator::new();
let a2 = alloc2.new_atom(&[1, 2, 3]).unwrap();
let b2 = alloc2.new_atom(&[4, 5, 6]).unwrap();
let node2 = alloc2.new_pair(a2, b2).unwrap();
let tree1 = intern_tree(&alloc1, node1).unwrap();
let tree2 = intern_tree(&alloc2, node2).unwrap();
assert_eq!(tree1.tree_hash(), tree2.tree_hash());
}
#[test]
fn test_pairs_in_post_order() {
let mut allocator = Allocator::new();
let a = allocator.new_atom(&[1]).unwrap();
let b = allocator.new_atom(&[2]).unwrap();
let c = allocator.new_atom(&[3]).unwrap();
let inner = allocator.new_pair(b, c).unwrap();
let outer = allocator.new_pair(a, inner).unwrap();
let tree = intern_tree(&allocator, outer).unwrap();
assert_eq!(tree.pairs.len(), 2);
let inner_pair = tree.pairs[0];
let outer_pair = tree.pairs[1];
let SExp::Pair(left, right) = tree.allocator.sexp(inner_pair) else {
panic!("Expected inner_pair to be a pair");
};
assert_eq!(tree.allocator.atom(left).as_ref(), &[2]);
assert_eq!(tree.allocator.atom(right).as_ref(), &[3]);
let SExp::Pair(left, right) = tree.allocator.sexp(outer_pair) else {
panic!("Expected outer_pair to be a pair");
};
assert_eq!(tree.allocator.atom(left).as_ref(), &[1]);
assert_eq!(
right, inner_pair,
"Outer pair's right child should be the inner pair"
);
}
}