"""Arithmetic and falsification witnesses; not a proof assistant."""
import sys
import unittest
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))
from exact_arithmetic import Curve, certificate, rank_mod


class ExactTests(unittest.TestCase):
    @classmethod
    def setUpClass(cls):
        cls.result = certificate()

    def test_counts_by_two_algorithms_and_generator_orbits(self):
        for ell, expected in ((397, 374), (991, 1045)):
            curve = Curve(ell)
            self.assertEqual(curve.count(), expected)
            self.assertEqual(len(curve.points()), expected)
            self.assertEqual(set(curve.orbit((12, 108))), set(curve.points()))

    def test_published_local_witness_points(self):
        a, b = self.result["local_results"]
        self.assertEqual((a["multiplied_P"], a["multiplied_Q"]), ((281, 236), (11, 334)))
        self.assertEqual((b["multiplied_P"], b["multiplied_Q"]), ((39, 97), (865, 243)))
        self.assertEqual((a["Q_norm_coordinate"], b["Q_norm_coordinate"]), (2, 4))

    def test_global_source_mod_11_is_detected(self):
        self.assertEqual(self.result["source_images_count"], 121)
        self.assertEqual(self.result["determinant_mod_11"], 2)

    def test_simultaneous_lattice(self):
        value = self.result["simultaneous_reduction_lattice"]
        self.assertEqual(value["congruence_residues"], [[0, 0], [0, 0]])
        self.assertEqual(value["index"], 390830)
        self.assertEqual(value["index"], value["target_order"])
        self.assertNotEqual((356-244) % 11, 0)

    def test_group_edge_cases_and_small_complete_associativity(self):
        curve = Curve(5, 1, 1)
        points = curve.points()
        for a in points:
            self.assertEqual(curve.add(a, None), a)
            self.assertEqual(curve.add(a, curve.neg(a)), None)
            for b in points:
                self.assertTrue(curve.on_curve(curve.add(a, b)))
                for c in points:
                    self.assertEqual(curve.add(curve.add(a, b), c), curve.add(a, curve.add(b, c)))

    def test_scalar_algorithm_against_repeated_addition(self):
        curve = Curve(397)
        point = (12, 108)
        for n, expected in enumerate(curve.orbit(point)):
            self.assertEqual(curve.mul(n, point), expected)
            self.assertEqual(curve.mul(-n, point), curve.neg(expected))

    def test_primitivity_inference_has_counterexample(self):
        self.assertNotEqual(11 % 121, 0)
        self.assertEqual(11 % 11, 0)

    def test_corrupted_row_loses_detection(self):
        self.assertEqual(rank_mod([[1, 2], [1, 2]]), 1)
        images = {((a+2*b) % 11, (a+2*b) % 11) for a in range(11) for b in range(11)}
        self.assertEqual(len(images), 11)

    def test_extra_selmer_directions_are_invisible(self):
        matrix = [[1, 2, 0, 0], [1, 4, 0, 0]]
        self.assertEqual(rank_mod(matrix), 2)
        kernel = [(a,b,c,d) for a in range(11) for b in range(11)
                  for c in range(11) for d in range(11)
                  if (a+2*b) % 11 == 0 and (a+4*b) % 11 == 0]
        self.assertEqual(len(kernel), 121)
        self.assertTrue(all(a == b == 0 for a,b,c,d in kernel))

    def test_square_zero_lifts_have_different_first_differentials(self):
        # Coefficients in the ordered basis (X,Y) of J/J^2.
        zero = [[(0,0),(0,0)],[(0,0),(0,0)]]
        lift = [[(1,0),(2,0)],[(0,1),(0,4)]]
        self.assertNotEqual(zero, lift)
        self.assertEqual((lift[0][0][0]*lift[1][1][1] -
                          lift[0][1][0]*lift[1][0][1]) % 11, 2)
        # Degree-two information belongs to Sym^2(J/J^2), not the zero product in R.

    def test_good_ordinary_and_singular_model_guard(self):
        self.assertEqual(self.result["good_at_11"], {"group_order":16, "a_11":-4})
        with self.assertRaises(ValueError):
            Curve(389)

    def test_minimal_model_and_residual_irreducibility_witness(self):
        d = self.result["minimal_invariants"]
        self.assertEqual((d["discriminant"], d["c4"]), (389,112))
        self.assertTrue(d["discriminant_is_prime"])
        self.assertEqual(self.result["minimal_model_point_counts"], {"2":5, "5":9, "11":16})
        witness = self.result["residual_irreducibility_witness"]
        self.assertEqual(witness["a_2"], -2)
        self.assertEqual(witness["roots_mod_11"], [])
        self.assertNotIn(7, witness["square_residues_mod_11"])


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