use crate::prelude::*;
use crate::ast::{Net, Tree};
use core::ops::RangeFrom;
use ordered_float::OrderedFloat;
impl Net {
pub fn eta_reduce(&mut self) {
let mut phase1 = Phase1::default();
for tree in self.trees() {
phase1.walk_tree(tree);
}
let mut phase2 = Phase2 { nodes: phase1.nodes, index: 0 .. };
for tree in self.trees_mut() {
phase2.reduce_tree(tree);
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum NodeType {
Ctr(u16),
Var(isize),
Int(i64),
F32(OrderedFloat<f32>),
Era,
Other,
Hole,
}
#[derive(Default)]
struct Phase1<'a> {
vars: Map<&'a str, usize>,
nodes: Vec<NodeType>,
}
impl<'a> Phase1<'a> {
fn walk_tree(&mut self, tree: &'a Tree) {
match tree {
Tree::Ctr { lab, ports } => {
let last_port = ports.len() - 1;
for (idx, i) in ports.iter().enumerate() {
if idx != last_port {
self.nodes.push(NodeType::Ctr(*lab));
}
self.walk_tree(i);
}
}
Tree::Var { nam } => {
if let Some(i) = self.vars.get(&**nam) {
let j = self.nodes.len() as isize;
self.nodes.push(NodeType::Var(*i as isize - j));
self.nodes[*i] = NodeType::Var(j - *i as isize);
} else {
self.vars.insert(nam, self.nodes.len());
self.nodes.push(NodeType::Hole);
}
}
Tree::Era => self.nodes.push(NodeType::Era),
Tree::Int { val } => self.nodes.push(NodeType::Int(*val)),
Tree::F32 { val } => self.nodes.push(NodeType::F32(*val)),
_ => {
self.nodes.push(NodeType::Other);
for i in tree.children() {
self.walk_tree(i);
}
}
}
}
}
struct Phase2 {
nodes: Vec<NodeType>,
index: RangeFrom<usize>,
}
impl Phase2 {
fn reduce_ctr(&mut self, lab: u16, ports: &mut Vec<Tree>, skip: usize) -> NodeType {
if skip == ports.len() {
return NodeType::Other;
}
if skip == ports.len() - 1 {
return self.reduce_tree(&mut ports[skip]);
}
let head_index = self.index.next().unwrap();
let a = self.reduce_tree(&mut ports[skip]);
let b = self.reduce_ctr(lab, ports, skip + 1);
if a == b {
let reducible = match a {
NodeType::Var(delta) => self.nodes[head_index.wrapping_add_signed(delta)] == NodeType::Ctr(lab),
NodeType::Era | NodeType::Int(_) | NodeType::F32(_) => true,
_ => false,
};
if reducible {
ports.pop();
return a;
}
}
NodeType::Ctr(lab)
}
fn reduce_tree(&mut self, tree: &mut Tree) -> NodeType {
if let Tree::Ctr { lab, ports } = tree {
let ty = self.reduce_ctr(*lab, ports, 0);
if ports.len() == 1 {
*tree = ports.pop().unwrap();
}
ty
} else {
let index = self.index.next().unwrap();
for i in tree.children_mut() {
self.reduce_tree(i);
}
self.nodes[index]
}
}
}