"""
Tests numériques pour la voie K^p.
Examen des conventions pour T_p sous le candidat (a) Cayley local.

Conventions testées :
  V1 = QR        : {1, 2, ν_p} avec ν_p = plus petit non-résidu quadratique > 1
  V2 = GEN       : {1, g, g^2 mod p} avec g = plus petit générateur de (Z/pZ)*
  V3 = COURET    : {1, 11 mod p, 29 mod p} (transposition naïve)
  V4 = SYM       : {1, 2, p-2} (triplet symétrique autour de p/2)

Mesures pour chaque (p, convention) :
  - T_p effectif (cardinalité, distinctness)
  - spectre σ(B̃_p)
  - ||B̃_p||_op et ||B̃_p||_HS^2
  - Tr(B̃_p^m) pour m = 1, 2, 3, 4
"""
import numpy as np
from sympy import primerange


def is_quadratic_residue(a, p):
    """Symbole de Legendre par puissance modulaire."""
    if a % p == 0:
        return None  # ni résidu ni non-résidu
    return pow(a, (p - 1) // 2, p) == 1


def smallest_non_residue_above(p, threshold=1):
    """Plus petit non-résidu quadratique > threshold."""
    for n in range(threshold + 1, p):
        if not is_quadratic_residue(n, p):
            return n
    return None


def smallest_generator(p):
    """Plus petit générateur de (Z/pZ)*."""
    if p == 2:
        return 1
    phi = p - 1
    # Diviseurs premiers de phi
    prime_divisors = []
    n = phi
    d = 2
    while d * d <= n:
        if n % d == 0:
            prime_divisors.append(d)
            while n % d == 0:
                n //= d
        d += 1
    if n > 1:
        prime_divisors.append(n)
    
    for g in range(2, p):
        # g est générateur ssi g^(phi/q) != 1 pour tout diviseur premier q de phi
        if all(pow(g, phi // q, p) != 1 for q in prime_divisors):
            return g
    return None


def get_T_p(p, convention):
    """Retourne T_p selon la convention, ou None si la convention ne s'applique pas."""
    if convention == "QR":
        # {1, 2, ν_p}
        if p < 7:
            return None
        nu = smallest_non_residue_above(p, threshold=1)
        if nu is None or nu == 2:
            # 2 peut être lui-même non-résidu : essayer {1, 2, ν > 2}
            nu = smallest_non_residue_above(p, threshold=2)
            if nu is None:
                return None
            T = sorted({1, 2, nu})
        else:
            T = sorted({1, 2, nu})
        if len(T) < 3:
            return None
        return tuple(T)
    
    elif convention == "GEN":
        # {1, g, g^2}
        g = smallest_generator(p)
        if g is None:
            return None
        g2 = pow(g, 2, p)
        T = sorted({1, g, g2})
        if len(T) < 3:
            return None
        return tuple(T)
    
    elif convention == "COURET":
        # {1, 11 mod p, 29 mod p}
        if p in (2, 3, 5, 7):
            # 11 mod 7 = 4, 29 mod 7 = 1 → collision avec 1
            return None
        a = 11 % p
        b = 29 % p
        T = sorted({1, a, b})
        if len(T) < 3:
            return None
        return tuple(T)
    
    elif convention == "SYM":
        # {1, 2, p-2}
        if p < 7:
            return None
        T = sorted({1, 2, p - 2})
        if len(T) < 3:
            return None
        return tuple(T)
    
    return None


def cayley_matrix(p, T):
    """Matrice de Cayley A[g,h] = 1 si g^{-1} h ∈ T mod p, 0 sinon.
    Sur les éléments inversibles G_p = {1, 2, ..., p-1}."""
    G = list(range(1, p))
    n = len(G)
    A = np.zeros((n, n), dtype=float)
    T_set = set(T)
    for i, g in enumerate(G):
        g_inv = pow(g, -1, p)
        for j, h in enumerate(G):
            if (g_inv * h) % p in T_set:
                A[i, j] = 1.0
    return A


def centered_projector(p):
    """Projecteur P_p^0 sur l'espace orthogonal au mode constant, dim G_p = p-1."""
    n = p - 1
    P = np.eye(n) - np.ones((n, n)) / n
    return P


def compute_B_p(p, T):
    """Calcule B̃_p = (1/|T|) P_p^0 A_T P_p^0."""
    A = cayley_matrix(p, T)
    P = centered_projector(p)
    B = (1.0 / len(T)) * P @ A @ P
    return B


def measure(p, T):
    """Mesure spectre, normes, traces de puissances pour (p, T)."""
    B = compute_B_p(p, T)
    eigvals = np.linalg.eigvalsh((B + B.T) / 2)  # B est-il symétrique ? Pas nécessairement
    # En général B n'est pas symétrique : utiliser valeurs propres complexes
    eigvals_full = np.linalg.eigvals(B)
    
    op_norm = np.max(np.abs(eigvals_full))
    hs_norm_sq = np.sum(np.abs(eigvals_full) ** 2)  # = ||B||_HS^2 pour B normal
    
    # Calcul direct ||B||_HS^2 = trace(B^* B) (plus robuste)
    hs_norm_sq_direct = np.real(np.trace(B.conj().T @ B))
    
    # Traces de puissances Tr(B^m) pour m = 1, 2, 3, 4
    traces = {}
    Bm = np.eye(p - 1)
    for m in range(1, 5):
        Bm = Bm @ B
        traces[m] = np.real(np.trace(Bm))
    
    return {
        "p": p,
        "T": T,
        "dim": p - 1,
        "op_norm": float(op_norm),
        "hs_norm_sq": float(hs_norm_sq_direct),
        "hs_norm_sq_eig": float(hs_norm_sq),
        "tr_1": float(traces[1]),
        "tr_2": float(traces[2]),
        "tr_3": float(traces[3]),
        "tr_4": float(traces[4]),
        "eigvals_abs_max": float(np.max(np.abs(eigvals_full))),
        "is_real_spectrum": float(np.max(np.abs(np.imag(eigvals_full)))) < 1e-10,
    }


if __name__ == "__main__":
    print("Setup OK.")
    print("Test rapide p=7 :")
    for conv in ["QR", "GEN", "COURET", "SYM"]:
        T = get_T_p(7, conv)
        if T is None:
            print(f"  {conv:8s} : non défini pour p=7")
        else:
            m = measure(7, T)
            print(f"  {conv:8s} : T={T}, ||B||_op={m['op_norm']:.4f}, ||B||_HS^2={m['hs_norm_sq']:.4f}")
