///|
pub enum NodeOp {
  Name(String)
  Symbol(String)
  Number(Float)
} derive(Eq, Hash, Show)

///|
pub fn NodeOp::label(self : NodeOp) -> String {
  match self {
    Name(s) => s
    Symbol(s) => s
    Number(n) => n.to_string()
  }
}

///|
struct NodeKey {
  op : NodeOp
  children : Array[Id]
} derive(Eq, Hash)

///|
pub struct ENode {
  op : NodeOp
  children : Array[Id]
} derive(Eq, Hash, Show)

///|
pub fn make_enode(op : String, children : Array[Id]) -> ENode {
  ENode::{ op: NodeOp::Name(op), children }
}

///|
pub fn make_symbol(sym : String) -> ENode {
  ENode::{ op: NodeOp::Symbol(sym), children: [] }
}

///|
pub fn make_number(n : Float) -> ENode {
  ENode::{ op: NodeOp::Number(n), children: [] }
}

///|
fn parse_number_literal(token : String) -> Float? {
  if token.length() == 0 {
    return None
  }
  let mut idx = 0
  let mut neg = false
  if token.get_char(0) == Some('-') {
    neg = true
    idx = 1
  }
  if idx >= token.length() {
    return None
  }
  let mut acc : Int = 0
  while idx < token.length() {
    match token.get_char(idx) {
      Some(c) if c >= '0' && c <= '9' => {
        let digit : Int = c.to_int() - '0'.to_int()
        acc = acc * 10 + digit
      }
      _ => return None
    }
    idx = idx + 1
  }
  let val = if neg { -acc } else { acc }
  Some(Float::from_int(val))
}

///|
pub struct EClass {
  id : Id
  nodes : Array[Int]
  data : Data
}

///|
pub struct EGraph {
  uf : UnionFind
  mut classes : Map[Id, EClass]
  mut memo : Map[NodeKey, Id]
  nodes : Array[ENode]
  node_classes : Array[Id]
  mut dirty : Bool
  analysis : Analysis
  mut allow_cycles : Bool
}

///|
pub type SimpleGraph = EGraph

///|
pub fn EGraph::new_with(analysis : Analysis) -> EGraph {
  EGraph::{
    uf: UnionFind::new(),
    classes: Map::new(),
    memo: Map::new(),
    nodes: Array::new(),
    node_classes: Array::new(),
    dirty: false,
    analysis,
    allow_cycles: true,
  }
}

///|
pub fn EGraph::new() -> SimpleGraph {
  EGraph::new_with(default_analysis())
}

///|
pub fn EGraph::set_allow_cycles(self : EGraph, allow : Bool) -> Unit {
  self.allow_cycles = allow
}

///|
pub fn EGraph::allow_cycles(self : EGraph) -> Bool {
  self.allow_cycles
}

///|
fn EGraph::canonicalize_children(
  self : EGraph,
  children : Array[Id],
) -> Array[Id] {
  children.map(child => self.uf.find(child))
}

///|
pub fn EGraph::data(self : EGraph, id : Id) -> Data? {
  match self.classes.get(self.uf.find_read(id)) {
    Some(class) => Some(class.data)
    None => None
  }
}

///|
pub fn EGraph::find(self : EGraph, id : Id) -> Id {
  self.uf.find(id)
}

///|
/// Lookup an existing enode without mutating the e-graph.
pub fn EGraph::lookup(self : EGraph, enode : ENode) -> Id? {
  let canon_children = self.canonicalize_children(enode.children)
  let key = NodeKey::{ op: enode.op, children: canon_children }
  self.memo.get(key)
}

///|
pub fn EGraph::find_read(self : EGraph, id : Id) -> Id {
  self.uf.find_read(id)
}

///|
pub fn EGraph::add(self : EGraph, enode : ENode) -> Id {
  let canon_children = self.canonicalize_children(enode.children)
  let key = NodeKey::{ op: enode.op, children: canon_children }
  match self.memo.get(key) {
    Some(existing) => self.find(existing)
    None => {
      let id = self.uf.make_set()
      let stored = ENode::{ op: key.op, children: key.children }
      let data = (self.analysis.make)(stored, child => self.data(child))
      self.nodes.push(stored)
      self.node_classes.push(id)
      self.classes.set(id, EClass::{
        id,
        nodes: [self.nodes.length() - 1],
        data,
      })
      self.memo.set(key, id)
      self.dirty = true
      id
    }
  }
}

///|
pub fn EGraph::add_expr(self : EGraph, expr : Expr) -> Id {
  match expr {
    Expr::Leaf(name) =>
      match parse_number_literal(name) {
        Some(n) => self.add(make_number(n))
        None => self.add(make_symbol(name))
      }
    Expr::Node(op, children) => {
      let child_ids = children.map(child => self.add_expr(child))
      self.add(make_enode(op, child_ids))
    }
  }
}

///|
pub fn EGraph::union(self : EGraph, a : Id, b : Id) -> Id {
  let root_a = self.find(a)
  let root_b = self.find(b)
  if root_a == root_b {
    return root_a
  }
  let merged_nodes = {
    let left_nodes = match self.classes.get(root_a) {
      Some(class) => class.nodes
      None => []
    }
    let right_nodes = match self.classes.get(root_b) {
      Some(class) => class.nodes
      None => []
    }
    let combined = left_nodes
    combined.append(right_nodes.op_as_view())
    combined
  }
  let merged_data = {
    let data_a = self.classes.get(root_a).unwrap().data
    let data_b = self.classes.get(root_b).unwrap().data
    (self.analysis.merge)(data_a, data_b)
  }
  let leader = self.uf.union(root_a, root_b)
  self.classes.remove(root_b)
  self.classes.set(leader, EClass::{
    id: leader,
    nodes: merged_nodes,
    data: merged_data,
  })
  self.dirty = true
  leader
}

///|
pub fn EGraph::rebuild(self : EGraph) -> Unit {
  if !self.dirty {
    return
  }
  loop () {
    _ => {
      let mut merged = false
      let seen : Map[NodeKey, Id] = Map::new()
      for idx in 0.. {
            let before = self.find(owner)
            let after = self.union(existing, before)
            merged = merged || after != before
          }
          None => seen.set(key, owner)
        }
      }
      self.memo = seen
      if merged {
        continue ()
      }
      break ()
    }
  }
  let next_classes : Map[Id, EClass] = Map::new()
  let next_memo : Map[NodeKey, Id] = Map::new()
  for idx in 0.. {
      let child_root = self.uf.find_read(child)
      match next_classes.get(child_root) {
        Some(class) => Some(class.data)
        None =>
          match self.classes.get(child_root) {
            Some(class) => Some(class.data)
            None => None
          }
      }
    })
    next_classes.update(root, existing => match existing {
      Some(class) => {
        let merged_nodes = class.nodes
        merged_nodes.push(idx)
        let merged_data = (self.analysis.merge)(class.data, node_data)
        Some(EClass::{ id: root, nodes: merged_nodes, data: merged_data })
      }
      None => Some(EClass::{ id: root, nodes: [idx], data: node_data })
    })
    next_memo.set(
      NodeKey::{ op: canon_node.op, children: canon_node.children },
      root,
    )
    self.node_classes[idx] = root
    self.nodes[idx] = canon_node
  }
  self.classes = next_classes
  self.memo = next_memo
  self.dirty = false
  for class_id in self.class_ids() {
    let root = self.find_read(class_id)
    if root != class_id {
      continue
    }
    (self.analysis.modify)(self, root)
  }
  if self.dirty {
    self.rebuild()
  }
}

///|
pub fn EGraph::are_equivalent(self : EGraph, a : Id, b : Id) -> Bool {
  self.find(a) == self.find(b)
}

///|
pub fn EGraph::class_ids(self : EGraph) -> Array[Id] {
  Array::from_iter(self.classes.keys())
}

///|
pub fn EGraph::class_for(self : EGraph, id : Id) -> EClass? {
  self.classes.get(self.find_read(id))
}