///|
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
}
}