"""Exact weight-two Manin-symbol calculation over a prime finite field.

Rows act on column vectors of functional values.  The geometric plus
condition is lambda(c,d)=lambda(-c,d); see notes/twin_conventions.md.
This module has no Kurihara target value or post-sum normalization.
"""
from functools import lru_cache
from math import gcd
import numpy as np


def nullspace(matrix, p):
    """Return RREF, pivot columns, and a canonical column kernel over F_p."""
    A = np.array(matrix, dtype=np.int64, copy=True) % p
    if A.ndim != 2:
        raise ValueError("matrix must have two dimensions")
    m, n = A.shape
    pivots = []
    r = 0
    for c in range(n):
        candidates = np.flatnonzero(A[r:, c])
        if not len(candidates):
            continue
        j = r + int(candidates[0])
        A[[r, j]] = A[[j, r]]
        A[r] = A[r] * pow(int(A[r, c]), -1, p) % p
        factors = A[:, c].copy()
        factors[r] = 0
        A = (A - factors[:, None] * A[r][None, :]) % p
        pivots.append(c)
        r += 1
        if r == m:
            break
    free = [c for c in range(n) if c not in set(pivots)]
    K = np.zeros((n, len(free)), dtype=np.int64)
    for j, c in enumerate(free):
        K[c, j] = 1
        for i, pivot in enumerate(pivots):
            K[pivot, j] = -A[i, c] % p
    return A, pivots, K


def a_ell(ell):
    """Exact trace for y^2+y=x^3+x^2-2x, including ell=2."""
    if ell == 2:
        points = 1 + sum((y*y+y-x*x*x-x*x+2*x) % 2 == 0
                         for x in range(2) for y in range(2))
        return 3 - points
    total = 0
    for x in range(ell):
        d = (4*x*x*x+4*x*x-8*x+1) % ell
        if d:
            s = pow(d, (ell-1)//2, ell)
            if s not in (1, ell-1):
                raise ValueError("ell is not an odd prime")
            total += 1 if s == 1 else -1
    return -total


def discrete_logs(ell, root, p):
    """Full primitive-root walk; entry zero is an explicit invalid sentinel."""
    if (ell-1) % p:
        raise ValueError("p must divide ell-1")
    logs = [-1] * ell
    x = 1
    for exponent in range(ell-1):
        if x == 0 or logs[x] != -1:
            raise ValueError("root does not generate the full unit group")
        logs[x] = exponent % p
        x = x * root % ell
    if x != 1 or any(v == -1 for v in logs[1:]):
        raise ValueError("incomplete discrete-log table")
    return logs


class Manin:
    def __init__(self, N, p):
        self.N, self.p, self.size = N, p, N+1
        self.inverses = [0] + [pow(c, -1, N) for c in range(1, N)]

    def index(self, c, d):
        c, d = c % self.N, d % self.N
        if c:
            return d * self.inverses[c] % self.N
        if d:
            return self.N
        raise ValueError("zero bottom row does not define a projective point")

    def pair(self, i):
        if not 0 <= i <= self.N:
            raise ValueError("generator index out of range")
        return (1, i) if i < self.N else (0, 1)

    def generator_endpoints(self, i):
        # For i<N, use [[0,-1],[1,i]]. For infinity, use the identity.
        return ((-1, i), (0, 1)) if i < self.N else ((0, 1), (1, 0))

    @lru_cache(maxsize=None)
    def relations(self):
        rows = np.zeros((2*self.size, self.size), dtype=np.int64)
        for i in range(self.size):
            c, d = self.pair(i)
            for j in [i, self.index(d, -c)]:
                rows[i, j] += 1
            for j in [i, self.index(d, -c-d), self.index(-c-d, c)]:
                rows[self.size+i, j] += 1
        return rows % self.p

    @lru_cache(maxsize=None)
    def plus_relations(self):
        rows = np.zeros((self.size, self.size), dtype=np.int64)
        for i in range(self.size):
            c, d = self.pair(i)
            rows[i, i] += 1
            rows[i, self.index(-c, d)] -= 1
        return rows % self.p

    @lru_cache(maxsize=100000)
    def cusp_word(self, a, b):
        """Ordinary floor CF decomposition of {infinity,a/b} into symbols."""
        if b == 0:
            return ()
        if b < 0:
            a, b = -a, -b
        common = gcd(a, b)
        a, b = a//common, b//common
        p_old, q_old, p_prev, q_prev = 0, 1, 1, 0
        word = []
        while b:
            digit, remainder = divmod(a, b)
            p_next = digit*p_prev+p_old
            q_next = digit*q_prev+q_old
            determinant = p_next*q_prev-p_prev*q_next
            if determinant not in (-1, 1):
                raise ArithmeticError("non-unimodular CF edge")
            word.append(self.index(determinant*q_next, q_prev))
            p_old, q_old, p_prev, q_prev = p_prev, q_prev, p_next, q_next
            a, b = b, remainder
        return tuple(word)

    def path(self, start, end):
        row = np.zeros(self.size, dtype=np.int64)
        for i in self.cusp_word(*end):
            row[i] += 1
        for i in self.cusp_word(*start):
            row[i] -= 1
        return row % self.p

    @lru_cache(maxsize=None)
    def hecke(self, q):
        if self.N % q == 0:
            raise ValueError("this formula requires q not dividing N")
        H = np.zeros((self.size, self.size), dtype=np.int64)
        for i in range(self.size):
            start, end = self.generator_endpoints(i)
            H[i] += self.path((q*start[0], start[1]), (q*end[0], end[1]))
            for t in range(q):
                H[i] += self.path((start[0]+t*start[1], q*start[1]),
                                  (end[0]+t*end[1], q*end[1]))
        return H % self.p

    def scalar_cusp(self, a, b, vector):
        """Uncached streaming CF evaluator, used on the large finite sum."""
        if b == 0:
            return 0
        if b < 0:
            a, b = -a, -b
        q_old, q_prev, determinant, value = 1, 0, -1, 0
        while b:
            digit, remainder = divmod(a, b)
            q_next = digit*q_prev+q_old
            value += int(vector[self.index(determinant*q_next, q_prev)])
            q_old, q_prev = q_prev, q_next
            determinant = -determinant
            a, b = b, remainder
        return value % self.p


def build_eigenline(M):
    """Construct and normalize the line before any Kurihara calculation."""
    p = M.p
    R = M.relations()
    rref, pivots, K = nullspace(R, p)
    eigen_primes = (2, 3, 5)
    A = np.vstack([(M.hecke(q)-a_ell(q)*np.eye(M.size, dtype=np.int64)) % p
                   for q in eigen_primes])
    A_on_K = A @ K % p
    erref, epivots, L = nullspace(A_on_K, p)
    B = K @ L % p
    plus_on_B = M.plus_relations() @ B % p
    prref, ppivots, C = nullspace(plus_on_B, p)
    line = B @ C % p
    if line.shape[1] != 1:
        raise ArithmeticError(f"geometric plus eigenspace dimension is {line.shape[1]}, expected 1")
    vector = line[:, 0]
    first = int(np.flatnonzero(vector)[0])
    vector = vector * pow(int(vector[first]), -1, p) % p
    return {
        "relation_rank": len(pivots),
        "relation_pivots": pivots,
        "relation_rref_nonzero": rref[:len(pivots)].tolist(),
        "manin_kernel": K.tolist(),
        "hecke_primes": list(eigen_primes),
        "hecke_constraints_on_kernel": A_on_K.tolist(),
        "hecke_constraint_pivots": epivots,
        "hecke_constraint_rref_nonzero": erref[:len(epivots)].tolist(),
        "hecke_kernel_coordinates": L.tolist(),
        "hecke_eigenbasis": B.tolist(),
        "hecke_eigenspace_dimension": B.shape[1],
        "plus_constraints_on_eigenbasis": plus_on_B.tolist(),
        "plus_constraint_pivots": ppivots,
        "plus_constraint_rref_nonzero": prref[:len(ppivots)].tolist(),
        "plus_kernel_coordinates": C.tolist(),
        "plus_eigenspace_dimension": line.shape[1],
        "normalization_rule": "first nonzero coordinate in the fixed generator order equals 1",
        "normalization_index": first,
        "eigenvector": vector.tolist(),
    }
