#!/usr/bin/env python3
import itertools
import math
import sys


def read_dimacs(path):
    nvars = None
    nclauses = None
    clauses = []
    with open(path, "rt", encoding="ascii") as f:
        for line in f:
            if not line or line[0] == "c":
                continue
            if line[0] == "p":
                parts = line.split()
                if len(parts) < 4 or parts[1] != "cnf":
                    return None
                nvars = int(parts[2])
                nclauses = int(parts[3])
                continue
            vals = [int(x) for x in line.split()]
            if not vals:
                continue
            if vals[-1] == 0:
                vals.pop()
            clauses.append(tuple(vals))
    if nvars is None:
        return None
    if nclauses is None:
        nclauses = len(clauses)
    return nvars, nclauses, clauses


def triangular_n(m):
    d = 1 + 8 * m
    r = math.isqrt(d)
    if r * r != d or (1 + r) % 2:
        return None
    n = (1 + r) // 2
    return n if n * (n - 1) // 2 == m else None


def edge_map(n):
    idx = {}
    k = 1
    for i in range(n):
        for j in range(i + 1, n):
            idx[i, j] = k
            k += 1
    return idx


def is_r44(n, nvars, clauses, idx):
    if n is None or n < 18 or nvars != n * (n - 1) // 2:
        return False
    if len(clauses) != 2 * math.comb(n, 4):
        return False
    pos = set()
    neg = set()
    for c in clauses:
        if len(c) != 6:
            return False
        if all(x > 0 for x in c):
            pos.add(tuple(sorted(c)))
        elif all(x < 0 for x in c):
            neg.add(tuple(sorted(-x for x in c)))
        else:
            return False
    if len(pos) != math.comb(n, 4) or len(neg) != math.comb(n, 4):
        return False
    for verts in itertools.combinations(range(n), 4):
        edges = []
        for a, b in itertools.combinations(verts, 2):
            edges.append(idx[a, b])
        key = tuple(sorted(edges))
        if key not in pos or key not in neg:
            return False
    return True


def lit_term(lit):
    if lit > 0:
        return f"1 x{lit}"
    return f"1 ~x{-lit}"


class Proof:
    def __init__(self, path, n_original):
        self.f = open(path, "w", encoding="ascii")
        self.next_id = n_original + 1
        self.f.write("pseudo-Boolean proof version 3.0\n")

    def close(self):
        self.f.close()

    def constraint(self, lits, rhs):
        seen = set()
        out = []
        for lit in lits:
            if lit not in seen:
                seen.add(lit)
                out.append(lit)
        if out:
            return " ".join(lit_term(x) for x in out) + f" >= {rhs}"
        return f">= {rhs}"

    def rup(self, lits, rhs=1):
        self.f.write("rup " + self.constraint(lits, rhs) + " ;\n")
        cid = self.next_id
        self.next_id += 1
        return cid

    def red(self, lits, rhs, witness):
        parts = []
        for var, lit in witness:
            if lit > 0:
                parts.append(f"x{var} -> x{lit}")
            elif lit < 0:
                parts.append(f"x{var} -> ~x{-lit}")
            else:
                parts.append(f"x{var} -> 0")
        self.f.write(
            "red "
            + self.constraint(lits, rhs)
            + " : "
            + " ".join(parts)
            + " ;\n"
        )
        cid = self.next_id
        self.next_id += 1
        return cid

    def pol_sum_div(self, ids, div=1):
        assert ids
        expr = [str(ids[0])]
        for cid in ids[1:]:
            expr.append(str(cid))
            expr.append("+")
        if div != 1:
            expr.append(str(div))
            expr.append("d")
        expr.append("s")
        self.f.write("pol " + " ".join(expr) + " ;\n")
        cid = self.next_id
        self.next_id += 1
        return cid

    def finish_unsat(self, cid):
        self.f.write("output NONE ;\n")
        self.f.write(f"conclusion UNSAT : {cid} ;\n")
        self.f.write("end pseudo-Boolean proof ;\n")
        self.close()


def perm_witness(n, idx, perm, complement=False, only_changed=True):
    out = []
    for a in range(n):
        for b in range(a + 1, n):
            var = idx[a, b]
            pa, pb = perm[a], perm[b]
            if pa > pb:
                pa, pb = pb, pa
            image = idx[pa, pb]
            lit = -image if complement else image
            if (not only_changed) or lit != var:
                out.append((var, lit))
    return out


def lift_at_least(proof, elems, q, base_ids):
    prev = {tuple(s): base_ids[tuple(s)] for s in itertools.combinations(elems, q)}
    for k in range(q + 1, len(elems) + 1):
        cur = {}
        for s in itertools.combinations(elems, k):
            ids = [prev[tuple(x for x in s if x != drop)] for drop in s]
            cur[tuple(s)] = proof.pol_sum_div(ids, k - 1)
        prev = cur
    return prev[tuple(elems)]


def prove_r33(proof, vertices, context, idx, red_true=True):
    def rlit(a, b):
        e = idx[min(a, b), max(a, b)]
        return e if red_true else -e

    def blit(a, b):
        return -rlit(a, b)

    pivot = vertices[0]
    neigh = list(vertices[1:])
    red_edges = [rlit(pivot, u) for u in neigh]
    blue_edges = [blit(pivot, u) for u in neigh]

    base_blue = {}
    for comb in itertools.combinations(range(len(neigh)), 3):
        key = tuple(blue_edges[i] for i in comb)
        base_blue[key] = proof.rup(list(context) + list(key))
    blue_full = lift_at_least(proof, blue_edges, 3, base_blue)

    base_red = {}
    for comb in itertools.combinations(range(len(neigh)), 3):
        key = tuple(red_edges[i] for i in comb)
        base_red[key] = proof.rup(list(context) + list(key))
    red_full = lift_at_least(proof, red_edges, 3, base_red)

    return proof.pol_sum_div([blue_full, red_full])


def prove_r34(proof, vertices, context, idx, red_true=True):
    def rlit(a, b):
        e = idx[min(a, b), max(a, b)]
        return e if red_true else -e

    def blit(a, b):
        return -rlit(a, b)

    lower_red = []
    upper_blue = []
    for v in vertices:
        neigh = [u for u in vertices if u != v]
        red_edges = [rlit(v, u) for u in neigh]
        blue_edges = [blit(v, u) for u in neigh]

        base = {}
        for comb in itertools.combinations(range(len(neigh)), 4):
            key = tuple(blue_edges[i] for i in comb)
            base[key] = proof.rup(list(context) + list(key))
        upper_blue.append(lift_at_least(proof, blue_edges, 4, base))

        base = {}
        for comb in itertools.combinations(range(len(neigh)), 6):
            sub = [neigh[i] for i in comb]
            key = tuple(red_edges[i] for i in comb)
            base[key] = prove_r33(proof, sub, list(context) + list(key), idx, red_true)
        lower_red.append(lift_at_least(proof, red_edges, 6, base))

    ge14 = proof.pol_sum_div(lower_red, 2)
    ge23_blue = proof.pol_sum_div(upper_blue, 2)
    return proof.pol_sum_div([ge14, ge23_blue])


def prove_r44(n, nclauses, out_path, idx):
    proof = Proof(out_path, nclauses)

    # Use the induced 18-vertex subformula on vertices 0..17.  The same proof
    # is valid in larger R(4,4,n) formulas because those clauses are present.
    for i in range(1, 17):
        perm = list(range(n))
        perm[1] = i + 1
        for j in range(2, i + 2):
            perm[j] = j - 1
        proof.red([idx[0, i], -idx[0, i + 1]], 1, perm_witness(n, idx, perm))

    perm = list(range(n))
    for i in range(1, 18):
        perm[i] = 18 - i
    proof.red(
        [idx[0, 9]],
        1,
        perm_witness(n, idx, perm, complement=True, only_changed=False),
    )

    context = [-idx[0, i] for i in range(1, 10)]
    prove_r34(proof, list(range(1, 10)), context, idx, True)
    empty = proof.rup([], 1)
    proof.finish_unsat(empty)


def main():
    if len(sys.argv) != 3:
        print("usage: solver <formula.cnf> <out.pbp>", file=sys.stderr)
        return 1
    parsed = read_dimacs(sys.argv[1])
    if parsed is None:
        print("s UNKNOWN")
        return 0
    nvars, nclauses, clauses = parsed
    n = triangular_n(nvars)
    if n is None:
        print("s UNKNOWN")
        return 0
    idx = edge_map(n)
    if is_r44(n, nvars, clauses, idx):
        prove_r44(n, nclauses, sys.argv[2], idx)
        print("s UNSATISFIABLE")
        return 20
    print("s UNKNOWN")
    return 0


if __name__ == "__main__":
    sys.exit(main())
