///|
pub(all) struct Var {
  id : Int
  name : String
} derive(Debug, Eq)

///|
pub(all) struct Lit {
  var_id : Int
  negated : Bool
} derive(Debug, Eq)

///|
pub(all) struct Clause {
  lits : Array[Lit]
} derive(Debug)

///|
pub(all) struct Cnf {
  var_count : Int
  clauses : Array[Clause]
} derive(Debug)

///|
pub(all) struct CnfBuilder {
  names : Array[String]
  clauses : Array[Clause]
} derive(Debug)

///|
pub(all) struct UnitSearch {
  found : Bool
  conflict : Bool
  lit : Lit
} derive(Debug, Eq)

///|
pub(all) struct SolveResult {
  sat : Bool
  decided : Bool
  assignment : Array[Int]
  trace : Array[String]
} derive(Debug)

///|
pub(all) struct CnfStats {
  vars : Int
  clauses : Int
  literals : Int
  max_clause_len : Int
} derive(Debug, Eq)

///|
pub fn Var::new(id : Int, name : String) -> Var {
  { id: normalize_positive(id, 1), name }
}

///|
pub fn Lit::pos(v : Var) -> Lit {
  { var_id: v.id, negated: false }
}

///|
pub fn Lit::neg(v : Var) -> Lit {
  { var_id: v.id, negated: true }
}

///|
pub fn Lit::not(self : Lit) -> Lit {
  { var_id: self.var_id, negated: !self.negated }
}

///|
pub fn Lit::to_dimacs(self : Lit) -> Int {
  if self.negated {
    -self.var_id
  } else {
    self.var_id
  }
}

///|
pub fn Clause::new(lits : Array[Lit]) -> Clause {
  { lits, }
}

///|
pub fn Clause::unit(lit : Lit) -> Clause {
  { lits: [lit] }
}

///|
pub fn Clause::len(self : Clause) -> Int {
  self.lits.length()
}

///|
pub fn Cnf::new(var_count : Int, clauses : Array[Clause]) -> Cnf {
  { var_count: normalize_positive(var_count, 0), clauses }
}

///|
pub fn Cnf::clause_count(self : Cnf) -> Int {
  self.clauses.length()
}

///|
pub fn CnfBuilder::new() -> CnfBuilder {
  { names: [], clauses: [] }
}

///|
pub fn CnfBuilder::new_var(self : CnfBuilder, name : String) -> Var {
  self.names.push(name)
  Var::new(self.names.length(), name)
}

///|
pub fn CnfBuilder::add_clause(self : CnfBuilder, lits : Array[Lit]) -> Unit {
  self.clauses.push(Clause::new(lits))
}

///|
pub fn CnfBuilder::require_literal(self : CnfBuilder, lit : Lit) -> Unit {
  self.clauses.push(Clause::unit(lit))
}

///|
pub fn CnfBuilder::at_least_one(self : CnfBuilder, vars : Array[Var]) -> Unit {
  let lits : Array[Lit] = []
  for v in vars {
    lits.push(Lit::pos(v))
  }
  self.add_clause(lits)
}

///|
pub fn CnfBuilder::at_most_one_pairwise(
  self : CnfBuilder,
  vars : Array[Var],
) -> Unit {
  for i = 0; i < vars.length(); i = i + 1 {
    for j = i + 1; j < vars.length(); j = j + 1 {
      self.add_clause([Lit::neg(vars[i]), Lit::neg(vars[j])])
    }
  }
}

///|
pub fn CnfBuilder::exactly_one(self : CnfBuilder, vars : Array[Var]) -> Unit {
  self.at_least_one(vars)
  self.at_most_one_pairwise(vars)
}

///|
pub fn CnfBuilder::implies(self : CnfBuilder, a : Var, b : Var) -> Unit {
  self.add_clause([Lit::neg(a), Lit::pos(b)])
}

///|
pub fn CnfBuilder::equivalent(self : CnfBuilder, a : Var, b : Var) -> Unit {
  self.implies(a, b)
  self.implies(b, a)
}

///|
pub fn CnfBuilder::and_gate(
  self : CnfBuilder,
  out : Var,
  inputs : Array[Var],
) -> Unit {
  let backward : Array[Lit] = [Lit::pos(out)]
  for v in inputs {
    self.add_clause([Lit::neg(out), Lit::pos(v)])
    backward.push(Lit::neg(v))
  }
  self.add_clause(backward)
}

///|
pub fn CnfBuilder::or_gate(
  self : CnfBuilder,
  out : Var,
  inputs : Array[Var],
) -> Unit {
  let backward : Array[Lit] = [Lit::neg(out)]
  for v in inputs {
    self.add_clause([Lit::neg(v), Lit::pos(out)])
    backward.push(Lit::pos(v))
  }
  self.add_clause(backward)
}

///|
pub fn CnfBuilder::at_most_k_sequential(
  self : CnfBuilder,
  vars : Array[Var],
  k : Int,
  prefix? : String = "s",
) -> Unit {
  if k < 0 {
    self.add_clause([])
    return
  }
  if k == 0 {
    for v in vars {
      self.require_literal(Lit::neg(v))
    }
    return
  }
  if vars.length() <= k {
    return
  }
  let counters : Array[Var] = []
  for i = 0; i < vars.length(); i = i + 1 {
    for j = 1; j <= k; j = j + 1 {
      counters.push(self.new_var("\{prefix}_\{i + 1}_\{j}"))
    }
  }
  for i = 0; i < vars.length(); i = i + 1 {
    self.add_clause([Lit::neg(vars[i]), Lit::pos(counter_at(counters, k, i, 1))])
    if i > 0 {
      for j = 1; j <= k; j = j + 1 {
        self.add_clause([
          Lit::neg(counter_at(counters, k, i - 1, j)),
          Lit::pos(counter_at(counters, k, i, j)),
        ])
      }
      for j = 2; j <= k; j = j + 1 {
        self.add_clause([
          Lit::neg(vars[i]),
          Lit::neg(counter_at(counters, k, i - 1, j - 1)),
          Lit::pos(counter_at(counters, k, i, j)),
        ])
      }
      self.add_clause([
        Lit::neg(vars[i]),
        Lit::neg(counter_at(counters, k, i - 1, k)),
      ])
    }
  }
}

///|
fn counter_at(counters : Array[Var], width : Int, row : Int, col : Int) -> Var {
  counters[row * width + col - 1]
}

///|
pub fn CnfBuilder::cnf(self : CnfBuilder) -> Cnf {
  Cnf::new(self.names.length(), self.clauses)
}

///|
pub fn Cnf::stats(self : Cnf) -> CnfStats {
  let mut literals = 0
  let mut max_len = 0
  for clause in self.clauses {
    literals = literals + clause.lits.length()
    if clause.lits.length() > max_len {
      max_len = clause.lits.length()
    }
  }
  {
    vars: self.var_count,
    clauses: self.clauses.length(),
    literals,
    max_clause_len: max_len,
  }
}

///|
pub fn CnfStats::to_json(self : CnfStats) -> String {
  "{\"vars\":\{self.vars},\"clauses\":\{self.clauses},\"literals\":\{self.literals},\"max_clause_len\":\{self.max_clause_len}}"
}

///|
pub fn Cnf::to_dimacs(self : Cnf) -> String {
  let buf = StringBuilder()
  buf.write_string("p cnf \{self.var_count} \{self.clauses.length()}\n")
  for clause in self.clauses {
    for lit in clause.lits {
      buf.write_string("\{lit.to_dimacs()} ")
    }
    buf.write_string("0\n")
  }
  buf.to_string()
}

///|
pub fn parse_dimacs(input : String) -> Cnf {
  let clauses : Array[Clause] = []
  let mut declared_vars = 0
  let mut current : Array[Lit] = []
  let tokens = split_ascii_tokens(input)
  let mut i = 0
  while i < tokens.length() {
    let token = tokens[i]
    if token == "c" {
      while i < tokens.length() && tokens[i] != "\n" {
        i = i + 1
      }
    } else if token == "p" && i + 3 < tokens.length() {
      declared_vars = parse_int_token(tokens[i + 2])
      i = i + 4
      continue
    } else {
      let value = parse_int_token(token)
      if value == 0 {
        if current.length() > 0 {
          clauses.push(Clause::new(current))
          current = []
        }
      } else if value < 0 {
        current.push({ var_id: -value, negated: true })
      } else {
        current.push({ var_id: value, negated: false })
      }
    }
    i = i + 1
  }
  if current.length() > 0 {
    clauses.push(Clause::new(current))
  }
  Cnf::new(declared_vars, clauses)
}

///|
fn split_ascii_tokens(input : String) -> Array[String] {
  let result : Array[String] = []
  let mut start = 0
  for i = 0; i < input.length(); i = i + 1 {
    match input.get_char(i) {
      Some(' ') | Some('\t') | Some('\r') | Some('\n') => {
        if start < i {
          result.push(input.unsafe_substring(start~, end=i))
        }
        match input.get_char(i) {
          Some('\n') => result.push("\n")
          _ => ()
        }
        start = i + 1
      }
      _ => ()
    }
  }
  if start < input.length() {
    result.push(input.unsafe_substring(start~, end=input.length()))
  }
  result
}

///|
fn parse_int_token(token : String) -> Int {
  if token.length() == 0 {
    return 0
  }
  let mut sign = 1
  let mut index = 0
  if token.get_char(0) == Some('-') {
    sign = -1
    index = 1
  }
  let mut value = 0
  while index < token.length() {
    let ch = token.get_char(index)
    match ch {
      Some(c) => if c >= '0' && c <= '9' { value = value * 10 + char_digit(c) }
      None => ()
    }
    index = index + 1
  }
  sign * value
}

///|
fn char_digit(ch : Char) -> Int {
  if ch == '0' {
    0
  } else if ch == '1' {
    1
  } else if ch == '2' {
    2
  } else if ch == '3' {
    3
  } else if ch == '4' {
    4
  } else if ch == '5' {
    5
  } else if ch == '6' {
    6
  } else if ch == '7' {
    7
  } else if ch == '8' {
    8
  } else if ch == '9' {
    9
  } else {
    0
  }
}

///|
fn lit_value(lit : Lit, assignment : Array[Int]) -> Int {
  let index = lit.var_id - 1
  if index < 0 || index >= assignment.length() {
    return -1
  }
  let value = assignment[index]
  if value < 0 {
    -1
  } else if lit.negated {
    if value == 0 {
      1
    } else {
      0
    }
  } else {
    value
  }
}

///|
fn clause_value(clause : Clause, assignment : Array[Int]) -> Int {
  let mut unknown = false
  for lit in clause.lits {
    let value = lit_value(lit, assignment)
    if value == 1 {
      return 1
    }
    if value < 0 {
      unknown = true
    }
  }
  if unknown {
    -1
  } else {
    0
  }
}

///|
pub fn Cnf::is_satisfied(self : Cnf, assignment : Array[Int]) -> Bool {
  for clause in self.clauses {
    if clause_value(clause, assignment) != 1 {
      return false
    }
  }
  true
}

///|
fn find_unit(cnf : Cnf, assignment : Array[Int]) -> UnitSearch {
  for clause in cnf.clauses {
    let mut unknown_count = 0
    let mut unit = { var_id: 1, negated: false }
    let mut satisfied = false
    for lit in clause.lits {
      let value = lit_value(lit, assignment)
      if value == 1 {
        satisfied = true
      } else if value < 0 {
        unknown_count = unknown_count + 1
        unit = lit
      }
    }
    if !satisfied && unknown_count == 0 {
      return { found: false, conflict: true, lit: unit }
    }
    if !satisfied && unknown_count == 1 {
      return { found: true, conflict: false, lit: unit }
    }
  }
  { found: false, conflict: false, lit: { var_id: 1, negated: false } }
}

///|
fn copy_assignment(assignment : Array[Int]) -> Array[Int] {
  let result : Array[Int] = []
  for item in assignment {
    result.push(item)
  }
  result
}

///|
fn first_unassigned(assignment : Array[Int]) -> Int {
  for i = 0; i < assignment.length(); i = i + 1 {
    if assignment[i] < 0 {
      return i
    }
  }
  -1
}

///|
fn assign_lit(assignment : Array[Int], lit : Lit) -> Bool {
  let index = lit.var_id - 1
  if index < 0 || index >= assignment.length() {
    return false
  }
  let value = if lit.negated { 0 } else { 1 }
  if assignment[index] >= 0 && assignment[index] != value {
    false
  } else {
    assignment[index] = value
    true
  }
}

///|
fn dpll(cnf : Cnf, assignment : Array[Int], trace : Array[String]) -> Bool {
  let mut keep_propagating = true
  while keep_propagating {
    let unit = find_unit(cnf, assignment)
    if unit.conflict {
      trace.push("conflict")
      return false
    }
    if unit.found {
      trace.push("unit \{unit.lit.to_dimacs()}")
      if !assign_lit(assignment, unit.lit) {
        trace.push("conflict assignment")
        return false
      }
    } else {
      keep_propagating = false
    }
  }
  if cnf.is_satisfied(assignment) {
    trace.push("satisfied")
    return true
  }
  let next = first_unassigned(assignment)
  if next < 0 {
    trace.push("conflict complete")
    return false
  }
  let positive = copy_assignment(assignment)
  positive[next] = 1
  trace.push("decide \{next + 1}")
  if dpll(cnf, positive, trace) {
    for i = 0; i < assignment.length(); i = i + 1 {
      assignment[i] = positive[i]
    }
    return true
  }
  let negative = copy_assignment(assignment)
  negative[next] = 0
  trace.push("decide -\{next + 1}")
  if dpll(cnf, negative, trace) {
    for i = 0; i < assignment.length(); i = i + 1 {
      assignment[i] = negative[i]
    }
    return true
  }
  false
}

///|
pub fn solve(cnf : Cnf) -> SolveResult {
  let assignment = Array::make(cnf.var_count, -1)
  let trace : Array[String] = []
  let sat = dpll(cnf, assignment, trace)
  { sat, decided: true, assignment, trace }
}

///|
pub fn SolveResult::value_of(self : SolveResult, v : Var) -> Int {
  let index = v.id - 1
  if index < 0 || index >= self.assignment.length() {
    -1
  } else {
    self.assignment[index]
  }
}

///|
pub fn SolveResult::to_json(self : SolveResult) -> String {
  let buf = StringBuilder()
  buf.write_string("{\"sat\":\{self.sat},\"assignment\":[")
  for i = 0; i < self.assignment.length(); i = i + 1 {
    if i > 0 {
      buf.write_char(',')
    }
    buf.write_string("\{self.assignment[i]}")
  }
  buf.write_string("],\"trace\":[")
  for i = 0; i < self.trace.length(); i = i + 1 {
    if i > 0 {
      buf.write_char(',')
    }
    buf.write_string("\"" + self.trace[i] + "\"")
  }
  buf.write_string("]}")
  buf.to_string()
}

///|
pub fn SolveResult::assignment_dimacs(self : SolveResult) -> String {
  let buf = StringBuilder()
  if self.sat {
    buf.write_string("s SATISFIABLE\nv ")
    for i = 0; i < self.assignment.length(); i = i + 1 {
      let lit = if self.assignment[i] == 1 { i + 1 } else { -(i + 1) }
      buf.write_string("\{lit} ")
    }
    buf.write_string("0\n")
  } else {
    buf.write_string("s UNSATISFIABLE\n")
  }
  buf.to_string()
}

///|
pub fn encode_n_queens(size : Int) -> Cnf {
  let n = normalize_positive(size, 1)
  let builder = CnfBuilder::new()
  let vars : Array[Var] = []
  for row = 0; row < n; row = row + 1 {
    for col = 0; col < n; col = col + 1 {
      vars.push(builder.new_var("q_\{row}_\{col}"))
    }
  }
  for row = 0; row < n; row = row + 1 {
    let row_vars : Array[Var] = []
    for col = 0; col < n; col = col + 1 {
      row_vars.push(queen_at(vars, n, row, col))
    }
    builder.exactly_one(row_vars)
  }
  for col = 0; col < n; col = col + 1 {
    let col_vars : Array[Var] = []
    for row = 0; row < n; row = row + 1 {
      col_vars.push(queen_at(vars, n, row, col))
    }
    builder.at_most_one_pairwise(col_vars)
  }
  for a = 0; a < vars.length(); a = a + 1 {
    let row_a = a / n
    let col_a = a % n
    for b = a + 1; b < vars.length(); b = b + 1 {
      let row_b = b / n
      let col_b = b % n
      if abs_int(row_a - row_b) == abs_int(col_a - col_b) {
        builder.add_clause([Lit::neg(vars[a]), Lit::neg(vars[b])])
      }
    }
  }
  builder.cnf()
}

///|
fn queen_at(vars : Array[Var], size : Int, row : Int, col : Int) -> Var {
  vars[row * size + col]
}

///|
fn abs_int(value : Int) -> Int {
  if value < 0 {
    -value
  } else {
    value
  }
}

///|
fn normalize_positive(value : Int, fallback : Int) -> Int {
  if value > 0 {
    value
  } else {
    fallback
  }
}