"""Adversarial tests against the independent verifier, with hashes resealed."""
import copy
import csv
import json
import shutil
import sys
import tempfile
import unittest
from pathlib import Path

ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT/"src"))
from verify_independent import (CertificateError, sha256, verify_all,
                                verify_eigenline, verify_logs, verify_sum)


def write_json(path, value):
    path.write_text(json.dumps(value, indent=2, sort_keys=True)+"\n", encoding="utf-8")


def prepare(directory, eigen, logs, summation):
    directory.mkdir(parents=True, exist_ok=True)
    write_json(directory/"eigenline_certificate.json", eigen)
    write_json(directory/"log_tables.json", logs)
    shutil.copyfile(ROOT/"results/kurihara_terms.csv", directory/"kurihara_terms.csv")
    reseal(directory, summation)


def reseal(directory, summation):
    for field, filename in (("eigenline_sha256", "eigenline_certificate.json"),
                            ("log_tables_sha256", "log_tables.json"),
                            ("terms_csv_sha256", "kurihara_terms.csv")):
        summation[field] = sha256(directory/filename)
    write_json(directory/"kurihara_certificate.json", summation)


def rewrite_terms(directory, summation, mode):
    """Create internally consistent altered arithmetic, including all block sums."""
    source, target = directory/"kurihara_terms.csv", directory/"replacement.csv"
    for block in summation["blocks"]:
        block["raw_product_sum"] = 0
        block["unit_count"] = 0
    count = raw_total = 0
    changed = False
    with source.open(newline="", encoding="utf-8") as inp, target.open("w", newline="", encoding="utf-8") as out:
        reader, writer = csv.reader(inp), csv.writer(out, lineterminator="\n")
        writer.writerow(next(reader))
        for row in reader:
            a, symbol, l1, l2, term = map(int, row)
            if mode == "scale2":
                symbol = 2*symbol % 11
            elif not changed and l1*l2 != 0:
                symbol = (symbol+1) % 11
                changed = True
            raw = symbol*l1*l2
            writer.writerow([a, symbol, l1, l2, raw % 11])
            raw_total += raw
            count += 1
            block = summation["blocks"][(a-1)//10000]
            block["unit_count"] += 1
            block["raw_product_sum"] += raw
    target.replace(source)
    for block in summation["blocks"]:
        block["residue"] = block["raw_product_sum"] % 11
    summation["unit_count"] = count
    summation["raw_product_sum"] = raw_total
    summation["residue"] = raw_total % 11
    reseal(directory, summation)


class IndependentCertificateTests(unittest.TestCase):
    @classmethod
    def setUpClass(cls):
        cls.eigen = json.loads((ROOT/"results/eigenline_certificate.json").read_text())
        cls.logs = json.loads((ROOT/"results/log_tables.json").read_text())
        cls.summation = json.loads((ROOT/"results/kurihara_certificate.json").read_text())

    def test_complete_independent_replay(self):
        report = verify_all(ROOT/"results")
        self.assertEqual(report["verdict"], "PASS")
        self.assertEqual(report["summation"]["unit_count"], 392040)
        self.assertNotEqual(report["summation"]["residue"], 0)
        self.assertEqual(report["summation"]["residue"], self.summation["residue"])

    def test_resealed_mutations_and_explicit_unit_equivalence(self):
        cases = [
            ("vector_coordinate", "eigenvector:relation_residual"),
            ("relation_rank", "relation_rank"),
            ("relation_rref", "relation_rref"),
            ("hecke_operator", "hecke:7:whole_quotient"),
            ("log_entry", "logs:397:entry"),
            ("sum_residue", "sum:residue"),
            ("consistent_wrong_csv_symbol", "terms:symbol"),
            ("unit_scaled_vector_fixed_rule", "normalization:scalar"),
            ("unit_scaled_vector_changed_rule_default_protocol", "normalization:rule"),
        ]
        outcomes = []
        with tempfile.TemporaryDirectory(prefix="kurihara-independent-attacks-") as temporary:
            base = Path(temporary)
            for name, expected_step in cases:
                eigen, logs, summation = copy.deepcopy((self.eigen, self.logs, self.summation))
                data = eigen["linear_algebra"]
                if name == "vector_coordinate":
                    data["eigenvector"][0] = (data["eigenvector"][0]+1) % 11
                elif name == "relation_rank":
                    data["relation_rank"] -= 1
                elif name == "relation_rref":
                    data["relation_rref_nonzero"][0][0] = (data["relation_rref_nonzero"][0][0]+1) % 11
                elif name == "hecke_operator":
                    eigen["hecke_matrices"]["7"][0][5] = (eigen["hecke_matrices"]["7"][0][5]+1) % 11
                elif name == "log_entry":
                    logs["tables"]["397"]["values"][2] = (logs["tables"]["397"]["values"][2]+1) % 11
                elif name == "sum_residue":
                    summation["residue"] = (summation["residue"]+1) % 11
                elif name.startswith("unit_scaled"):
                    data["eigenvector"] = [2*x % 11 for x in data["eigenvector"]]
                    if "changed_rule" in name:
                        data["normalization_rule"] = "first nonzero coordinate in the fixed generator order equals 2"
                directory = base/name
                prepare(directory, eigen, logs, summation)
                if name == "consistent_wrong_csv_symbol":
                    rewrite_terms(directory, summation, "one_symbol")
                elif name.startswith("unit_scaled"):
                    rewrite_terms(directory, summation, "scale2")
                try:
                    verify_all(directory)
                except CertificateError as error:
                    self.assertEqual(error.step, expected_step, str(error))
                    self.assertNotIn("hash", error.step)
                    outcomes.append({"mutation": name, "verdict": "REJECTED",
                                     "failed_step": error.step, "detail": error.detail,
                                     "all_binding_hashes_resealed": True})
                else:
                    self.fail(f"Mutated certificate was accepted: {name}")

            # A declared unit scale is mathematically valid. It remains outside
            # the fixed producer protocol, and is accepted here only by explicit
            # opt-in of both the rule and the verifier's expected scalar.
            eigen, logs, summation = copy.deepcopy((self.eigen, self.logs, self.summation))
            data = eigen["linear_algebra"]
            data["eigenvector"] = [2*x % 11 for x in data["eigenvector"]]
            data["normalization_rule"] = "first nonzero coordinate in the fixed generator order equals 2"
            directory = base/"declared_scale2_mathematical_equivalence"
            prepare(directory, eigen, logs, summation)
            rewrite_terms(directory, summation, "scale2")
            state = verify_eigenline(eigen, expected_normalization_scalar=2)
            recomputed_logs = verify_logs(logs)
            recomputed_sum = verify_sum(directory, summation, state["vector"], recomputed_logs)
            self.assertEqual(recomputed_sum["residue"], 2*self.summation["residue"] % 11)
            equivalence = {"verdict": "PASS_WITH_EXPLICIT_TEST_ONLY_SCALAR_2",
                           "base_residue": self.summation["residue"],
                           "scaled_residue": recomputed_sum["residue"],
                           "all_scaled_csv_rows_independently_recomputed": True,
                           "default_fixed_normalization_protocol_accepts": False,
                           "real_producer_certificate_modified": False}
        report = {"schema": "kurihara-independent-negative-checks-v1", "verdict": "PASS",
                  "rejected_mutation_count": len(outcomes), "mutations": outcomes,
                  "declared_unit_equivalence": equivalence}
        write_json(ROOT/"results/independent_negative_checks.json", report)


if __name__ == "__main__":
    unittest.main()
