#!/usr/bin/env python3
"""Clean-room verifier for two frozen Hales--Jewett lower-bound certificates.

For a word in [t]^n, record the counts of its t letters. A symmetric coloring
depends only on this count vector. A combinatorial line with k wildcard
coordinates and fixed-letter count vector v has the t count vectors
v + k*e_j. Exhaustively rejecting monochromatic tuples of this form therefore
rejects every monochromatic combinatorial line in the lifted grid coloring.

The certificate is treated only as a proposed coloring. This verifier derives
the simplex, every corner tuple, the lower-bound implication, and two negative
controls independently from the definitions.
"""

from __future__ import annotations

import argparse
import hashlib
import json
from pathlib import Path


CASES = {
    "33": {
        "alphabet_size": 3,
        "dimension": 21,
        "color_count": 3,
        "bound": 22,
        "certificate_sha256": "876eabf0c351c36f96160a8ef3baf4b131a1e8a7664be21ebacce39ffdbbdbb5",
        "expected_cells": 253,
        "expected_corners": 1771,
    },
    "42": {
        "alphabet_size": 4,
        "dimension": 13,
        "color_count": 2,
        "bound": 14,
        "certificate_sha256": "ed602c972fc2e3556eacd3b2176aff81420e04736ace49174dc651f94aaf7de4",
        "expected_cells": 560,
        "expected_corners": 1820,
    },
}


def compositions(total: int, parts: int):
    if parts == 1:
        yield (total,)
        return
    for first in range(total + 1):
        for rest in compositions(total - first, parts - 1):
            yield (first,) + rest


def parse_certificate(payload: bytes, *, parts: int, colors: int):
    coloring = {}
    for line_number, raw_line in enumerate(payload.decode("utf-8").splitlines(), 1):
        line = raw_line.strip()
        if not line:
            continue
        tokens = line.split()
        if len(tokens) != parts + 1:
            raise ValueError(f"line {line_number} does not contain {parts + 1} integers")
        try:
            values = tuple(int(token) for token in tokens)
        except ValueError as error:
            raise ValueError(f"line {line_number} contains a non-integer token") from error
        cell, color = values[:-1], values[-1]
        if any(coordinate < 0 for coordinate in cell):
            raise ValueError(f"line {line_number} contains a negative coordinate")
        if not 0 <= color < colors:
            raise ValueError(f"line {line_number} contains an out-of-range color")
        if cell in coloring:
            raise ValueError(f"line {line_number} duplicates cell {cell}")
        coloring[cell] = color
    return coloring


def audit_coloring(coloring, *, parts: int, dimension: int):
    expected = set(compositions(dimension, parts))
    actual = set(coloring)
    missing = expected - actual
    extra = actual - expected
    if missing or extra:
        return {
            "complete": False,
            "missing_cells": len(missing),
            "extra_cells": len(extra),
            "corners_checked": 0,
            "monochromatic_corners": 0,
        }

    checked = 0
    monochromatic = 0
    for wildcard_count in range(1, dimension + 1):
        for fixed_counts in compositions(dimension - wildcard_count, parts):
            corner = [
                tuple(
                    fixed_counts[index] + (wildcard_count if index == letter else 0)
                    for index in range(parts)
                )
                for letter in range(parts)
            ]
            checked += 1
            if len({coloring[cell] for cell in corner}) == 1:
                monochromatic += 1
    return {
        "complete": True,
        "missing_cells": 0,
        "extra_cells": 0,
        "corners_checked": checked,
        "monochromatic_corners": monochromatic,
    }


def verify(case_key: str, certificate: Path):
    case = CASES[case_key]
    payload = certificate.read_bytes()
    digest = hashlib.sha256(payload).hexdigest()
    if digest != case["certificate_sha256"]:
        raise ValueError("certificate SHA-256 does not match the frozen source artifact")

    coloring = parse_certificate(
        payload,
        parts=case["alphabet_size"],
        colors=case["color_count"],
    )
    audit = audit_coloring(
        coloring,
        parts=case["alphabet_size"],
        dimension=case["dimension"],
    )
    if not audit["complete"]:
        raise ValueError("certificate does not cover the discrete simplex exactly")
    if len(coloring) != case["expected_cells"]:
        raise ValueError("simplex cell count differs from the frozen audit target")
    if audit["corners_checked"] != case["expected_corners"]:
        raise ValueError("corner count differs from the frozen audit target")
    if audit["monochromatic_corners"]:
        raise ValueError("certificate contains a monochromatic combinatorial line")

    incomplete = dict(coloring)
    incomplete.pop(next(iter(incomplete)))
    incomplete_audit = audit_coloring(
        incomplete,
        parts=case["alphabet_size"],
        dimension=case["dimension"],
    )
    if incomplete_audit["complete"]:
        raise AssertionError("negative control failed to reject an incomplete coloring")

    monochromatic = dict(coloring)
    pure_cells = [
        tuple(case["dimension"] if index == letter else 0 for index in range(case["alphabet_size"]))
        for letter in range(case["alphabet_size"])
    ]
    for cell in pure_cells:
        monochromatic[cell] = 0
    monochromatic_audit = audit_coloring(
        monochromatic,
        parts=case["alphabet_size"],
        dimension=case["dimension"],
    )
    if monochromatic_audit["monochromatic_corners"] < 1:
        raise AssertionError("negative control failed to detect a monochromatic line")

    return {
        "alphabet_size": case["alphabet_size"],
        "bound": case["bound"],
        "case": case_key,
        "certificate_sha256": digest,
        "colors": case["color_count"],
        "dimension": case["dimension"],
        "negative_controls": {
            "incomplete_coloring_rejected": True,
            "monochromatic_diagonal_rejected": True,
        },
        "result": "VERIFIED",
        "simplex_cells": len(coloring),
        "corner_tuples_checked": audit["corners_checked"],
        "monochromatic_corner_tuples": audit["monochromatic_corners"],
        "theorem": f"HJ({case['alphabet_size']},{case['color_count']}) >= {case['bound']}",
    }


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("case", choices=sorted(CASES))
    parser.add_argument("certificate", type=Path)
    arguments = parser.parse_args()
    print(json.dumps(verify(arguments.case, arguments.certificate), sort_keys=True, separators=(",", ":")))


if __name__ == "__main__":
    main()
