// Expansion of `inherit` statements.

///|
fn[T] keep_last_defined(get_name : (T) -> String, l : Array[T]) -> Array[T] {
  let seen : Map[String, Unit] = Map([])
  let rev = []
  for i = l.length() - 1; i >= 0; i = i - 1 {
    let k = get_name(l[i])
    if !seen.contains(k) {
      seen[k] = ()
      rev.push(l[i])
    }
  }
  rev.rev()
}

///|
fn get_field_name(f : Field) -> String {
  match f {
    Field(sf) => sf.name
    Inherit(_) => abort("assertion failed")
  }
}

///|
fn get_variant_name(v : Variant) -> String {
  match v {
    Variant(_, k, _, _) => k
    Inherit(_) => abort("assertion failed")
  }
}

///|
priv struct InheritExpander {
  tbl : PredefTable
  inherit_fields : Bool
  inherit_variants : Bool
}

///|
fn InheritExpander::subst(
  self : InheritExpander,
  deref : Bool,
  param : Array[(String, TypeExpr)],
  t : TypeExpr,
) -> TypeExpr raise AtdError {
  match t {
    Sum(loc, vl, a) => {
      let vl2 = []
      for v in vl {
        vl2.append(self.subst_variant(param, v))
      }
      let vl2 = if self.inherit_variants {
        keep_last_defined(get_variant_name, vl2)
      } else {
        vl2
      }
      Sum(loc, vl2, a)
    }
    Record(loc, fl, a) => {
      let fl2 = []
      for f in fl {
        fl2.append(self.subst_field(param, f))
      }
      let fl2 = if self.inherit_fields {
        keep_last_defined(get_field_name, fl2)
      } else {
        fl2
      }
      Record(loc, fl2, a)
    }
    Tuple(loc, tl, a) => {
      let cells = []
      for c in tl {
        cells.push({ ..c, expr: self.subst(false, param, c.expr), })
      }
      Tuple(loc, cells, a)
    }
    List(loc, t, a)
    | Name(loc, { name: { path: ["list"], }, args: [t], .. }, a) =>
      List(loc, self.subst(false, param, t), a)
    Option(loc, t, a)
    | Name(loc, { name: { path: ["option"], }, args: [t], .. }, a) =>
      Option(loc, self.subst(false, param, t), a)
    Nullable(loc, t, a)
    | Name(loc, { name: { path: ["nullable"], }, args: [t], .. }, a) =>
      Nullable(loc, self.subst(false, param, t), a)
    Shared(loc, t, a)
    | Name(loc, { name: { path: ["shared"], }, args: [t], .. }, a) =>
      Shared(loc, self.subst(false, param, t), a)
    Wrap(loc, t, a)
    | Name(loc, { name: { path: ["wrap"], }, args: [t], .. }, a) =>
      Wrap(loc, self.subst(false, param, t), a)
    Tvar(_, s) =>
      match param.iter().find_first(x => x.0 == s) {
        Some((_, v)) => v
        None => t
      }
    Name(loc, { loc: loc2, name: k, args, }, a) => {
      let expanded_args = []
      for x in args {
        expanded_args.push(self.subst(false, param, x))
      }
      if deref {
        let (vars, t) = match self.tbl.get(k) {
          Some((_, Some(x))) => (x.param, x.value)
          Some((_, None)) => error("Cannot inherit from type \{k}")
          None => error("Missing type definition for \{k}")
        }
        if vars.length() != expanded_args.length() {
          error("Invalid_argument(\"List.combine\")")
        }
        let param = vars.mapi((i, v) => (v, expanded_args[i]))
        self.subst(true, param, t)
      } else {
        Name(loc, { loc: loc2, name: k, args: expanded_args, }, a)
      }
    }
  }
}

///|
fn InheritExpander::subst_field(
  self : InheritExpander,
  param : Array[(String, TypeExpr)],
  f : Field,
) -> Array[Field] raise AtdError {
  match f {
    Field(sf) => [Field({ ..sf, expr: self.subst(false, param, sf.expr), })]
    Inherit(_, t) as x =>
      match self.subst(true, param, t) {
        Record(_, vl, _) => if self.inherit_fields { vl } else { [x] }
        _ => error("Not a record type")
      }
  }
}

///|
fn InheritExpander::subst_variant(
  self : InheritExpander,
  param : Array[(String, TypeExpr)],
  v : Variant,
) -> Array[Variant] raise AtdError {
  match v {
    Variant(loc, k, a, opt_t) as x =>
      match opt_t {
        None => [x]
        Some(t) => [Variant(loc, k, a, Some(self.subst(false, param, t)))]
      }
    Inherit(_, t) as x =>
      match self.subst(true, param, t) {
        Sum(_, vl, _) => if self.inherit_variants { vl } else { [x] }
        _ => error("Not a sum type")
      }
  }
}

///|
/// Expand the `inherit` statements of all the type definitions.
pub fn expand_inherit(
  defs : Array[TypeDef],
  inherit_fields? : Bool = true,
  inherit_variants? : Bool = true,
) -> Array[TypeDef] raise AtdError {
  let tbl = make_predef_table(defs)
  let e : InheritExpander = { tbl, inherit_fields, inherit_variants, }
  let res = []
  for x in defs {
    res.push({ ..x, value: e.subst(false, [], x.value), })
  }
  res
}