#!/usr/bin/env python3
"""Specialized solver for signed/deleted relativized pigeonhole CNFs.

The solver detects the RPHP skeleton structurally:

  P_i,j primary placement variables
  R_j   selected domain variables
  Q_j,c secondary placement variables

Some shuffled SAT instances are exact RPHP formulas with one secondary
collision clause deleted.  For those, the missing collision gives a direct
model.  Exact formulas are UNSAT; the solver emits a VeriPB counting proof
using extension atoms X_j,c = R_j & Q_j,c and cutting planes cardinality
derivations for the P and X cliques.
"""

from collections import defaultdict, deque
import os
import sys


class Unsupported(Exception):
    pass


def parse_dimacs(path):
    clauses = []
    n = m = None
    with open(path, "r", encoding="ascii", errors="strict") 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":
                    raise Unsupported("not DIMACS CNF")
                n, m = int(parts[2]), int(parts[3])
                continue
            lits = [int(x) for x in line.split() if x != "0"]
            if lits:
                clauses.append(lits)
    if n is None or m is None or m != len(clauses):
        raise Unsupported("bad DIMACS header")
    return n, clauses


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


def pb_terms(lits):
    return " ".join("+1 " + lit_name(l) for l in lits)


class Detector:
    def __init__(self, n, clauses):
        self.n = n
        self.clauses = clauses
        self.m = len(clauses)

    def detect(self):
        n = self.n
        clauses = self.clauses

        in_size4 = [False] * (n + 1)
        for c in clauses:
            if len(c) == 4:
                for lit in c:
                    in_size4[abs(lit)] = True

        pvars = {v for v in range(1, n + 1) if not in_size4[v]}
        p_clause_ids = []
        p_clause_lits = []
        for cid, c in enumerate(clauses, 1):
            if len(c) > 2 and all(abs(lit) in pvars for lit in c):
                p_clause_ids.append(cid)
                p_clause_lits.append(c)
        if not p_clause_lits:
            raise Unsupported("no primary clauses")

        p = len(p_clause_lits)
        r = len(p_clause_lits[0])
        if p < 2 or r < p or any(len(c) != r for c in p_clause_lits):
            raise Unsupported("bad primary dimensions")
        if len(pvars) != p * r or n != 2 * p * r:
            raise Unsupported("bad variable counts")

        canon_p = {}
        p_group = {}
        group_pvars = []
        for group, clause in enumerate(p_clause_lits):
            row = []
            for lit in clause:
                var = abs(lit)
                if var in canon_p:
                    raise Unsupported("primary variable reused")
                canon_p[var] = lit
                p_group[var] = group
                row.append(var)
            group_pvars.append(row)

        p_adj = defaultdict(list)
        p_edge_id = {}
        p_imp_id = {}
        rlit_for_p = {}
        for cid, c in enumerate(clauses, 1):
            if len(c) != 2:
                continue
            a, b = c
            ap = abs(a) in pvars
            bp = abs(b) in pvars
            if ap and bp:
                if a != -canon_p[abs(a)] or b != -canon_p[abs(b)]:
                    raise Unsupported("bad primary binary polarity")
                va, vb = abs(a), abs(b)
                p_adj[va].append(vb)
                p_adj[vb].append(va)
                p_edge_id[(min(va, vb), max(va, vb))] = cid
            elif ap ^ bp:
                plit, rlit = (a, b) if ap else (b, a)
                pvar = abs(plit)
                if plit != -canon_p[pvar] or pvar in p_imp_id:
                    raise Unsupported("bad primary implication")
                p_imp_id[pvar] = cid
                rlit_for_p[pvar] = rlit
            else:
                raise Unsupported("unexpected non-primary binary")

        seen = set()
        domains = []
        rlit_to_domain = {}
        for start in sorted(pvars):
            if start in seen:
                continue
            q = deque([start])
            seen.add(start)
            comp = []
            while q:
                var = q.popleft()
                comp.append(var)
                for nxt in p_adj[var]:
                    if nxt not in seen:
                        seen.add(nxt)
                        q.append(nxt)
            if len(comp) != p or len({p_group[v] for v in comp}) != p:
                raise Unsupported("bad primary component")
            rlits = {rlit_for_p.get(v) for v in comp}
            if None in rlits or len(rlits) != 1:
                raise Unsupported("bad R selector component")
            rlit = next(iter(rlits))
            dom = len(domains)
            if abs(rlit) in rlit_to_domain:
                raise Unsupported("R variable reused")
            rlit_to_domain[abs(rlit)] = dom
            domains.append({"pvars": comp, "rlit": rlit, "q_lits": [], "r_clause": None})
        if len(domains) != r:
            raise Unsupported("bad domain count")

        rvars = {abs(d["rlit"]) for d in domains}
        qvars = set(range(1, n + 1)) - pvars - rvars
        if len(qvars) != r * (p - 1):
            raise Unsupported("bad Q count")

        q_domain = {}
        q_canon = {}
        for cid, c in enumerate(clauses, 1):
            if len(c) != p or any(abs(lit) in pvars for lit in c):
                continue
            rs = [lit for lit in c if abs(lit) in rvars]
            if len(rs) != 1:
                continue
            dom = rlit_to_domain[abs(rs[0])]
            if rs[0] != -domains[dom]["rlit"]:
                continue
            qs = [lit for lit in c if abs(lit) in qvars]
            if len(qs) != p - 1 or domains[dom]["q_lits"]:
                continue
            domains[dom]["q_lits"] = qs
            domains[dom]["r_clause"] = cid
            for qlit in qs:
                qvar = abs(qlit)
                if qvar in q_domain:
                    raise Unsupported("Q variable reused")
                q_domain[qvar] = dom
                q_canon[qvar] = qlit
        if len(q_domain) != len(qvars) or any(not d["q_lits"] for d in domains):
            raise Unsupported("missing R-to-Q clauses")

        q_adj = defaultdict(list)
        q_edge_id = {}
        for cid, c in enumerate(clauses, 1):
            if len(c) != 4 or any(abs(lit) in pvars for lit in c):
                continue
            rs = [lit for lit in c if abs(lit) in rvars]
            qs = [lit for lit in c if abs(lit) in qvars]
            if len(rs) != 2 or len(qs) != 2:
                continue
            d1 = rlit_to_domain[abs(rs[0])]
            d2 = rlit_to_domain[abs(rs[1])]
            if d1 == d2:
                raise Unsupported("bad injection R pair")
            if rs[0] != -domains[d1]["rlit"] or rs[1] != -domains[d2]["rlit"]:
                raise Unsupported("bad injection R polarity")
            q1, q2 = abs(qs[0]), abs(qs[1])
            if qs[0] != -q_canon[q1] or qs[1] != -q_canon[q2]:
                raise Unsupported("bad injection Q polarity")
            if {q_domain[q1], q_domain[q2]} != {d1, d2}:
                raise Unsupported("injection Q domain mismatch")
            q_adj[q1].append(q2)
            q_adj[q2].append(q1)
            q_edge_id[(min(q1, q2), max(q1, q2))] = cid

        color_components = []
        seen.clear()
        for start in sorted(qvars):
            if start in seen:
                continue
            q = deque([start])
            seen.add(start)
            comp = []
            while q:
                var = q.popleft()
                comp.append(var)
                for nxt in q_adj[var]:
                    if nxt not in seen:
                        seen.add(nxt)
                        q.append(nxt)
            color_components.append(comp)
        if len(color_components) != p - 1:
            raise Unsupported("bad secondary color count")

        q_by_domain_color = [[None] * (p - 1) for _ in range(r)]
        for color, comp in enumerate(color_components):
            if len(comp) != r:
                raise Unsupported("bad secondary color component")
            used_domains = set()
            for qvar in comp:
                dom = q_domain[qvar]
                if dom in used_domains:
                    raise Unsupported("secondary color repeats domain")
                used_domains.add(dom)
                q_by_domain_color[dom][color] = qvar

        missing = []
        for j in range(r):
            for k in range(j + 1, r):
                for color in range(p - 1):
                    a = q_by_domain_color[j][color]
                    b = q_by_domain_color[k][color]
                    if (min(a, b), max(a, b)) not in q_edge_id:
                        missing.append((j, k, color))

        self.p = p
        self.r = r
        self.pvars = pvars
        self.canon_p = canon_p
        self.p_group = p_group
        self.group_pvars = group_pvars
        self.domains = domains
        self.rvars = rvars
        self.qvars = qvars
        self.q_canon = q_canon
        self.q_by_domain_color = q_by_domain_color
        self.p_clause_ids = p_clause_ids
        self.p_edge_id = p_edge_id
        self.p_imp_id = p_imp_id
        self.q_edge_id = q_edge_id
        self.missing = missing
        return self


class ProofWriter:
    def __init__(self, det, path):
        self.d = det
        self.f = open(path, "w", encoding="ascii")
        self.maxid = det.m
        self.write("pseudo-Boolean proof version 3.0")
        self.write("")
        self.write(f"f {det.m} ;")

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

    def write(self, line):
        self.f.write(line + "\n")

    def add_line(self, line):
        self.write(line)
        self.maxid += 1
        return self.maxid

    def red(self, lits, witness_var, witness_val):
        return self.add_line(f"red {pb_terms(lits)} >= 1 : x{witness_var} -> {witness_val} ;")

    def rup_clause(self, lits):
        return self.add_line(f"rup {pb_terms(lits)} >= 1 ;")

    def pol_sum(self, ids):
        if not ids:
            raise Unsupported("empty pol")
        expr = str(ids[0])
        for cid in ids[1:]:
            expr += f" {cid} +"
        return self.add_line(f"pol {expr} ;")

    def pbc_clique(self, atoms, edge_id):
        if len(atoms) == 1:
            return None
        if len(atoms) == 2:
            key = (min(abs(atoms[0]), abs(atoms[1])), max(abs(atoms[0]), abs(atoms[1])))
            return edge_id[key]

        groups = []
        i = 0
        while i + 1 < len(atoms):
            a, b = atoms[i], atoms[i + 1]
            groups.append({
                "atoms": [a, b],
                "amo": edge_id[(min(abs(a), abs(b)), max(abs(a), abs(b)))]
            })
            i += 2
        if i < len(atoms):
            groups.append({"atoms": [atoms[i]], "amo": None})

        while len(groups) > 1:
            groups.sort(key=lambda g: len(g["atoms"]))
            a = groups.pop(0)
            b = groups.pop(0)
            amo = self.merge_cliques(a, b, edge_id)
            groups.append({"atoms": a["atoms"] + b["atoms"], "amo": amo})
        return groups[0]["amo"]

    def merge_cliques(self, group_a, group_b, edge_id):
        if len(group_b["atoms"]) == 1 and len(group_a["atoms"]) > 1:
            group_a, group_b = group_b, group_a

        target_lits = [-x for x in (group_a["atoms"] + group_b["atoms"])]
        self.write(f"pbc {pb_terms(target_lits)} >= {len(target_lits) - 1} : subproof")
        old_max = self.maxid
        neg_id = old_max + 1
        local_max = old_max + 1

        def sub_pol(tokens):
            nonlocal local_max
            self.write("pol " + " ".join(str(t) for t in tokens) + " ;")
            local_max += 1
            return local_max

        if len(group_a["atoms"]) == 1:
            atom = group_a["atoms"][0]
            other = group_b
            if other["amo"] is None:
                raise Unsupported("bad singleton clique merge")
            atom_pos = sub_pol([neg_id, other["amo"], "+"])
            units = []
            for b in other["atoms"]:
                eid = edge_id[(min(abs(atom), abs(b)), max(abs(atom), abs(b)))]
                units.append(sub_pol([atom_pos, eid, "+"]))
            toks = [neg_id]
            for uid in units:
                toks += [uid, "+"]
            contradiction = sub_pol(toks)
        else:
            if len(group_b["atoms"]) < len(group_a["atoms"]):
                group_a, group_b = group_b, group_a
            b_pos = sub_pol([neg_id, group_a["amo"], "+"])
            a_pos = sub_pol([neg_id, group_b["amo"], "+"])
            units = []
            for a in group_a["atoms"]:
                toks = [b_pos]
                for b in group_b["atoms"]:
                    eid = edge_id[(min(abs(a), abs(b)), max(abs(a), abs(b)))]
                    toks += [eid, "+"]
                toks += [len(group_b["atoms"]), "d"]
                units.append(sub_pol(toks))
            toks = [a_pos]
            for uid in units:
                toks += [uid, "+"]
            contradiction = sub_pol(toks)

        self.write(f"qed pbc : {contradiction} ;")
        self.maxid = local_max + 1
        return self.maxid

    def prove_domain_link(self, dom_idx, p_amo_id):
        dom = self.d.domains[dom_idx]
        pvars = dom["pvars"]
        target = [-self.d.canon_p[v] for v in pvars] + [dom["rlit"]]
        self.write(f"pbc {pb_terms(target)} >= {self.d.p} : subproof")
        old_max = self.maxid
        neg_id = old_max + 1
        local_max = old_max + 1

        def sub_pol(tokens):
            nonlocal local_max
            self.write("pol " + " ".join(str(t) for t in tokens) + " ;")
            local_max += 1
            return local_max

        not_r = sub_pol([neg_id, p_amo_id, "+"])
        units = []
        for pvar in pvars:
            units.append(sub_pol([not_r, self.d.p_imp_id[pvar], "+"]))
        toks = [neg_id]
        for uid in units:
            toks += [uid, "+"]
        contradiction = sub_pol(toks)
        self.write(f"qed pbc : {contradiction} ;")
        self.maxid = local_max + 1
        return self.maxid

    def finish_unsat(self):
        d = self.d

        domain_constraints = []
        for dom_idx, dom in enumerate(d.domains):
            atoms = [d.canon_p[v] for v in dom["pvars"]]
            edge_ids = {}
            for i, a in enumerate(atoms):
                for b in atoms[i + 1:]:
                    edge_ids[(min(abs(a), abs(b)), max(abs(a), abs(b)))] = d.p_edge_id[
                        (min(abs(a), abs(b)), max(abs(a), abs(b)))
                    ]
            amo = self.pbc_clique(atoms, edge_ids)
            domain_constraints.append(self.prove_domain_link(dom_idx, amo))
        lower_r = self.pol_sum(d.p_clause_ids + domain_constraints)

        xvar = {}
        next_var = d.n + 1
        for j in range(d.r):
            for color in range(d.p - 1):
                xv = next_var
                next_var += 1
                xvar[(j, color)] = xv
                rlit = d.domains[j]["rlit"]
                qlit = d.q_canon[d.q_by_domain_color[j][color]]
                self.red([-xv, rlit], xv, 0)
                self.red([-xv, qlit], xv, 0)
                self.red([-rlit, -qlit, xv], xv, 1)

        x_edge_id = {}
        for color in range(d.p - 1):
            for j in range(d.r):
                for k in range(j + 1, d.r):
                    a = xvar[(j, color)]
                    b = xvar[(k, color)]
                    eid = self.rup_clause([-a, -b])
                    x_edge_id[(min(a, b), max(a, b))] = eid

        selected_to_x = []
        for j in range(d.r):
            lits = [-d.domains[j]["rlit"]] + [xvar[(j, c)] for c in range(d.p - 1)]
            selected_to_x.append(self.rup_clause(lits))

        color_amos = []
        for color in range(d.p - 1):
            atoms = [xvar[(j, color)] for j in range(d.r)]
            edge_ids = {}
            for i, a in enumerate(atoms):
                for b in atoms[i + 1:]:
                    edge_ids[(min(a, b), max(a, b))] = x_edge_id[(min(a, b), max(a, b))]
            color_amos.append(self.pbc_clique(atoms, edge_ids))

        upper_r = self.pol_sum(selected_to_x + color_amos)
        contradiction = self.pol_sum([lower_r, upper_r])
        self.write("")
        self.write("output NONE ;")
        self.write(f"conclusion UNSAT : {contradiction} ;")
        self.write("end pseudo-Boolean proof ;")
        return contradiction


def set_atom(values, atom_lit, truth):
    var = abs(atom_lit)
    values[var] = truth if atom_lit > 0 else not truth


def build_sat_model(det):
    if not det.missing:
        raise Unsupported("no missing collision")
    j, k, color = det.missing[0]
    if det.r < det.p:
        raise Unsupported("not enough domains for model")

    color_for_domain = {}
    color_for_domain[j] = color
    color_for_domain[k] = color
    chosen = {j, k}
    for c in range(det.p - 1):
        if c == color:
            continue
        for dom in range(det.r):
            if dom not in chosen:
                chosen.add(dom)
                color_for_domain[dom] = c
                break
    if len(chosen) != det.p:
        raise Unsupported("could not choose selected domains")

    selected_domains = list(color_for_domain.keys())
    values = [False] * (det.n + 1)

    for var, lit in det.canon_p.items():
        set_atom(values, lit, False)
    for dom in det.domains:
        set_atom(values, dom["rlit"], False)
    for var, lit in det.q_canon.items():
        set_atom(values, lit, False)

    for group, dom_idx in enumerate(selected_domains):
        for pvar in det.domains[dom_idx]["pvars"]:
            if det.p_group[pvar] == group:
                set_atom(values, det.canon_p[pvar], True)
                break
        else:
            raise Unsupported("missing primary variable for group")

    for dom_idx, c in color_for_domain.items():
        set_atom(values, det.domains[dom_idx]["rlit"], True)
        qvar = det.q_by_domain_color[dom_idx][c]
        set_atom(values, det.q_canon[qvar], True)

    return values


def verify_model(clauses, values):
    for clause in clauses:
        if not any(values[abs(lit)] == (lit > 0) for lit in clause):
            return False
    return True


def print_model(values):
    print("s SATISFIABLE")
    line = []
    for var in range(1, len(values)):
        line.append(str(var if values[var] else -var))
        if len(line) >= 20:
            print("v " + " ".join(line))
            line = []
    if line:
        print("v " + " ".join(line) + " 0")
    else:
        print("v 0")


def unknown(msg=None):
    if msg:
        print(f"c UNKNOWN: {msg}", file=sys.stderr)
    print("s UNKNOWN")
    return 0


def main():
    if len(sys.argv) != 3:
        print("usage: solver <formula.cnf> <out.pbp>", file=sys.stderr)
        return 1

    formula, proof_path = sys.argv[1], sys.argv[2]
    try:
        n, clauses = parse_dimacs(formula)
        det = Detector(n, clauses).detect()
        if det.missing:
            values = build_sat_model(det)
            if not verify_model(clauses, values):
                return unknown("constructed model failed verification")
            try:
                open(proof_path, "w", encoding="ascii").close()
            except OSError:
                pass
            print_model(values)
            return 10

        writer = ProofWriter(det, proof_path)
        try:
            writer.finish_unsat()
        finally:
            writer.close()
        print("s UNSATISFIABLE")
        return 20
    except Unsupported as exc:
        return unknown(str(exc))
    except Exception as exc:
        return unknown(f"internal error: {exc}")


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