///|
/// Analysis payload supporting numeric and boolean constants plus free-variable sets.
pub(all) enum Value {
  Num(Float)
  Bool(Bool)
} derive(Eq, Show)

///|
/// Per-eclass analysis data.
pub(all) struct Data {
  free : Map[Id, Bool]
  constant : Value?
} derive(Show)

///|
pub fn data_num(d : Data?) -> Float? {
  match d {
    Some(payload) =>
      match payload.constant {
        Some(Value::Num(n)) => Some(n)
        _ => None
      }
    None => None
  }
}

///|
pub fn data_bool(d : Data?) -> Bool? {
  match d {
    Some(payload) =>
      match payload.constant {
        Some(Value::Bool(b)) => Some(b)
        _ => None
      }
    None => None
  }
}

///|
fn empty_data() -> Data {
  let free : Map[Id, Bool] = Map::new()
  Data::{ free, constant: None }
}

///|
pub fn intersect_free(a : Map[Id, Bool], b : Map[Id, Bool]) -> Map[Id, Bool] {
  let result = Map::new()
  for pair in a.iter() {
    let (k, v) = pair
    if v && b.contains(k) {
      result.set(k, true)
    }
  }
  result
}

///|
fn union_free(children : Array[Data?]) -> Map[Id, Bool] {
  let free : Map[Id, Bool] = Map::new()
  for child_data in children {
    match child_data {
      Some(payload) =>
        for pair in payload.free.iter() {
          let (k, v) = pair
          if v {
            free.set(k, true)
          }
        }
      None => ignore(())
    }
  }
  free
}

///|
pub fn merge_data(left : Data, right : Data) -> Data {
  let constant = match (left.constant, right.constant) {
    (Some(a), Some(b)) => if a == b { Some(a) } else { None }
    (Some(a), None) => Some(a)
    (None, Some(b)) => Some(b)
    (None, None) => None
  }
  Data::{ free: intersect_free(left.free, right.free), constant }
}

///|
/// Keep only leaf nodes inside a constant class to curb blowup.
fn prune_to_leaves(egraph : EGraph, id : Id) -> Unit {
  let root = egraph.find(id)
  match egraph.class_for(root) {
    Some(class) => {
      let leaf_nodes = class.nodes.filter(idx => egraph.nodes[idx].children.is_empty())
      egraph.classes.set(root, EClass::{
        id: root,
        nodes: leaf_nodes,
        data: class.data,
      })
    }
    None => ()
  }
}

///|
/// Per-eclass analysis callbacks.
pub(all) struct Analysis {
  make : (ENode, (Id) -> Data?) -> Data
  merge : (Data, Data) -> Data
  modify : (EGraph, Id) -> Unit
}

///|
pub fn default_analysis() -> Analysis {
  Analysis::{
    make: (_, _) => empty_data(),
    merge: (a, b) => merge_data(a, b),
    modify: (_, _) => (),
  }
}

///|
pub fn constant_analysis() -> Analysis {
  Analysis::{
    make: (node, child_data) => {
      let free = union_free(node.children.map(child => child_data(child)))
      let constant = match node.op {
        NodeOp::Number(n) => Some(Value::Num(n))
        NodeOp::Symbol(name) if name == "true" => Some(Value::Bool(true))
        NodeOp::Symbol(name) if name == "false" => Some(Value::Bool(false))
        NodeOp::Name(op) =>
          if op == "add" || op == "+" {
            match (child_data(node.children[0]), child_data(node.children[1])) {
              (Some(payload_a), Some(payload_b)) =>
                match (payload_a.constant, payload_b.constant) {
                  (Some(Value::Num(a)), Some(Value::Num(b))) =>
                    Some(Value::Num(a + b))
                  _ => None
                }
              _ => None
            }
          } else if op == "sub" || op == "-" {
            match (child_data(node.children[0]), child_data(node.children[1])) {
              (Some(payload_a), Some(payload_b)) =>
                match (payload_a.constant, payload_b.constant) {
                  (Some(Value::Num(a)), Some(Value::Num(b))) =>
                    Some(Value::Num(a - b))
                  _ => None
                }
              _ => None
            }
          } else if op == "mul" || op == "*" {
            match (child_data(node.children[0]), child_data(node.children[1])) {
              (Some(payload_a), Some(payload_b)) =>
                match (payload_a.constant, payload_b.constant) {
                  (Some(Value::Num(a)), Some(Value::Num(b))) =>
                    Some(Value::Num(a * b))
                  _ => None
                }
              _ => None
            }
          } else if op == "div" || op == "/" {
            match (child_data(node.children[0]), child_data(node.children[1])) {
              (Some(payload_a), Some(payload_b)) =>
                match (payload_a.constant, payload_b.constant) {
                  (Some(Value::Num(a)), Some(Value::Num(b))) =>
                    if b != 0.0 {
                      Some(Value::Num(a / b))
                    } else {
                      None
                    }
                  _ => None
                }
              _ => None
            }
          } else if op == "=" {
            match (child_data(node.children[0]), child_data(node.children[1])) {
              (Some(payload_a), Some(payload_b)) =>
                match (payload_a.constant, payload_b.constant) {
                  (Some(Value::Num(a)), Some(Value::Num(b))) =>
                    Some(Value::Bool(a == b))
                  (Some(Value::Bool(a)), Some(Value::Bool(b))) =>
                    Some(Value::Bool(a == b))
                  _ => None
                }
              _ => None
            }
          } else {
            None
          }
        _ => None
      }
      Data::{ free, constant }
    },
    merge: (a, b) => merge_data(a, b),
    modify: (egraph, id) => match egraph.data(id) {
      Some(payload) =>
        match payload.constant {
          Some(Value::Num(n)) => {
            let num = egraph.add(make_number(n))
            ignore(egraph.union(id, num))
            prune_to_leaves(egraph, id)
          }
          Some(Value::Bool(true)) => {
            let t = egraph.add(make_symbol("true"))
            ignore(egraph.union(id, t))
            prune_to_leaves(egraph, id)
          }
          Some(Value::Bool(false)) => {
            let f = egraph.add(make_symbol("false"))
            ignore(egraph.union(id, f))
            prune_to_leaves(egraph, id)
          }
          None => ()
        }
      None => ()
    },
  }
}

///|
/// Boolean constant folding for propositional logic tests.
pub fn bool_analysis() -> Analysis {
  Analysis::{
    make: (node, child_data) => {
      let free = union_free(node.children.map(child => child_data(child)))
      let constant = match node.op {
        NodeOp::Symbol(sym) if sym == "true" => Some(Value::Bool(true))
        NodeOp::Symbol(sym) if sym == "false" => Some(Value::Bool(false))
        NodeOp::Name(op) =>
          if op == "&" {
            match (child_data(node.children[0]), child_data(node.children[1])) {
              (Some(payload_a), Some(payload_b)) =>
                match (payload_a.constant, payload_b.constant) {
                  (Some(Value::Bool(a)), Some(Value::Bool(b))) =>
                    Some(Value::Bool(a && b))
                  _ => None
                }
              _ => None
            }
          } else if op == "|" {
            match (child_data(node.children[0]), child_data(node.children[1])) {
              (Some(payload_a), Some(payload_b)) =>
                match (payload_a.constant, payload_b.constant) {
                  (Some(Value::Bool(a)), Some(Value::Bool(b))) =>
                    Some(Value::Bool(a || b))
                  _ => None
                }
              _ => None
            }
          } else if op == "~" {
            match child_data(node.children[0]) {
              Some(payload) =>
                match payload.constant {
                  Some(Value::Bool(v)) => Some(Value::Bool(!v))
                  _ => None
                }
              None => None
            }
          } else if op == "->" {
            match (child_data(node.children[0]), child_data(node.children[1])) {
              (Some(payload_a), Some(payload_b)) =>
                match (payload_a.constant, payload_b.constant) {
                  (Some(Value::Bool(a)), Some(Value::Bool(b))) =>
                    Some(Value::Bool(!a || b))
                  _ => None
                }
              _ => None
            }
          } else {
            None
          }
        _ => None
      }
      Data::{ free, constant }
    },
    merge: (a, b) => merge_data(a, b),
    modify: (egraph, id) => match egraph.data(id) {
      Some(payload) =>
        match payload.constant {
          Some(Value::Bool(true)) => {
            let t = egraph.add(make_symbol("true"))
            ignore(egraph.union(id, t))
          }
          Some(Value::Bool(false)) => {
            let f = egraph.add(make_symbol("false"))
            ignore(egraph.union(id, f))
          }
          _ => ()
        }
      None => ()
    },
  }
}

///|
/// Lambda analysis tracking free variables and simple constant folding.
pub fn lambda_analysis() -> Analysis {
  Analysis::{
    make: (node, child_data) => {
      let free = union_free(node.children.map(child => child_data(child)))
      let constant = match node.op {
        NodeOp::Number(n) => Some(Value::Num(n))
        NodeOp::Symbol(sym) if sym == "true" => Some(Value::Bool(true))
        NodeOp::Symbol(sym) if sym == "false" => Some(Value::Bool(false))
        NodeOp::Name(op) =>
          if op == "+" {
            match (child_data(node.children[0]), child_data(node.children[1])) {
              (Some(a), Some(b)) =>
                match (a.constant, b.constant) {
                  (Some(Value::Num(x)), Some(Value::Num(y))) =>
                    Some(Value::Num(x + y))
                  _ => None
                }
              _ => None
            }
          } else if op == "=" {
            match (child_data(node.children[0]), child_data(node.children[1])) {
              (Some(a), Some(b)) =>
                match (a.constant, b.constant) {
                  (Some(Value::Num(x)), Some(Value::Num(y))) =>
                    Some(Value::Bool(x == y))
                  (Some(Value::Bool(x)), Some(Value::Bool(y))) =>
                    Some(Value::Bool(x == y))
                  _ => None
                }
              _ => None
            }
          } else {
            None
          }
        _ => None
      }
      // adjust free variables for binders
      match node.op {
        NodeOp::Name(op) if op == "var" =>
          if node.children.length() > 0 {
            free.set(node.children[0], true)
          }
        NodeOp::Name(op) if op == "lam" =>
          if node.children.length() > 0 {
            free.remove(node.children[0])
          }
        NodeOp::Name(op) if op == "fix" =>
          if node.children.length() > 0 {
            free.remove(node.children[0])
          }
        NodeOp::Name(op) if op == "let" =>
          if node.children.length() > 0 {
            free.remove(node.children[0])
          }
        _ => ignore(())
      }
      Data::{ free, constant }
    },
    merge: (a, b) => merge_data(a, b),
    modify: (egraph, id) => {
      match egraph.data(id) {
        Some(payload) =>
          match payload.constant {
            Some(Value::Num(n)) => {
              let num = egraph.add(make_number(n))
              ignore(egraph.union(id, num))
              prune_to_leaves(egraph, id)
            }
            Some(Value::Bool(true)) => {
              let t = egraph.add(make_symbol("true"))
              ignore(egraph.union(id, t))
              prune_to_leaves(egraph, id)
            }
            Some(Value::Bool(false)) => {
              let f = egraph.add(make_symbol("false"))
              ignore(egraph.union(id, f))
              prune_to_leaves(egraph, id)
            }
            None => ()
          }
        None => ()
      }
      match egraph.class_for(id) {
        Some(class) =>
          for node_idx in class.nodes {
            let node = egraph.nodes[node_idx]
            match node.op {
              NodeOp::Name(op) if op == "if" && node.children.length() == 3 =>
                match egraph.data(node.children[0]) {
                  Some(payload) =>
                    match payload.constant {
                      Some(Value::Bool(true)) =>
                        ignore(egraph.union(id, node.children[1]))
                      Some(Value::Bool(false)) =>
                        ignore(egraph.union(id, node.children[2]))
                      _ => ()
                    }
                  None => ()
                }
              _ => ()
            }
          }
        None => ()
      }
    },
  }
}