import sys
import unittest
from pathlib import Path
import numpy as np
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))
from manin import Manin, nullspace, a_ell, discrete_logs, build_eigenline


class FiniteCertificateTests(unittest.TestCase):
    @classmethod
    def setUpClass(cls):
        cls.M = Manin(389, 11)
        cls.data = build_eigenline(cls.M)
        cls.K = np.array(cls.data["manin_kernel"], dtype=np.int64)
        cls.vector = np.array(cls.data["eigenvector"], dtype=np.int64)

    def test_manin_rank_and_kernel(self):
        self.assertEqual(self.data["relation_rank"], 325)
        self.assertEqual(self.K.shape, (390,65))
        self.assertTrue(np.all(self.M.relations() @ self.K % 11 == 0))
        self.assertEqual(self.data["hecke_eigenspace_dimension"], 2)
        self.assertEqual(self.data["plus_eigenspace_dimension"], 1)

    def test_paths_recover_every_generator_on_whole_quotient(self):
        for i in range(390):
            start, end = self.M.generator_endpoints(i)
            path = self.M.path(start, end)
            self.assertTrue(np.all((path @ self.K - self.K[i]) % 11 == 0), i)

    def test_hecke_descends_and_commutes(self):
        R = self.M.relations()
        H2, H3 = self.M.hecke(2), self.M.hecke(3)
        self.assertTrue(np.all((R @ H2 % 11) @ self.K % 11 == 0))
        self.assertTrue(np.all((R @ H3 % 11) @ self.K % 11 == 0))
        self.assertTrue(np.all(((H2 @ H3-H3 @ H2) % 11) @ self.K % 11 == 0))

    def test_plus_and_out_of_sample_hecke(self):
        self.assertTrue(np.all(self.M.plus_relations() @ self.vector % 11 == 0))
        for ell in [7,13,17,19]:
            self.assertTrue(np.all((self.M.hecke(ell) @ self.vector-a_ell(ell)*self.vector) % 11 == 0))

    def test_point_counts(self):
        self.assertEqual([a_ell(q) for q in [2,3,5,7,11,13,17,19,397,991]],
                         [-2,-2,-3,-5,-4,-3,-6,5,24,-53])

    def test_generator_scalar_is_fixed_before_sum(self):
        nz = np.flatnonzero(self.vector)
        self.assertEqual(self.vector[nz[0]], 1)
        self.assertEqual(self.data["normalization_index"], int(nz[0]))

    def test_logs_are_complete_and_obey_group_law(self):
        for ell, root in [(397,5),(991,6)]:
            logs = discrete_logs(ell,root,11)
            self.assertEqual(logs[1],0)
            for a in range(1,ell):
                self.assertEqual(logs[a*root % ell], (logs[a]+1) % 11)
            with self.assertRaises(ValueError):
                discrete_logs(ell,1,11)

    def test_altered_symbol_is_rejected(self):
        wrong = self.vector.copy()
        wrong[0] = (wrong[0]+1) % 11
        self.assertTrue(np.any(self.M.relations() @ wrong % 11))

    def test_nullspace_small_known_counterexample(self):
        _, pivots, K = nullspace(np.array([[1,2],[1,2]],dtype=np.int64),11)
        self.assertEqual(len(pivots),1)
        self.assertTrue(np.all(np.array([[1,2],[1,2]]) @ K % 11 == 0))


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