"""
Primality Testing — MAT5EJ302(1) Module I
Trial division and Fermat's primality test.
"""

import random
import math


def mod_exp(base, exp, mod):
    """Fast modular exponentiation. Time: O(log exp)."""
    result = 1
    base %= mod
    while exp > 0:
        if exp % 2 == 1:
            result = result * base % mod
        exp //= 2
        base = base * base % mod
    return result


def is_prime_trial(n):
    """
    Trial division: check divisibility up to sqrt(n).
    Time: O(sqrt(n)) — exponential in the number of bits of n.
    Only practical for small n (say n < 10^12).
    """
    if n < 2:
        return False
    if n == 2:
        return True
    if n % 2 == 0:
        return False
    d = 3
    while d * d <= n:
        if n % d == 0:
            return False
        d += 2
    return True


def fermat_test(n, k=20):
    """
    Fermat primality test. Performs k random trials.
    - Returns False: n is DEFINITELY composite.
    - Returns True:  n is PROBABLY prime (may be wrong for Carmichael numbers).
    Time: O(k * log^2(n)).
    """
    if n < 2:
        return False
    if n == 2 or n == 3:
        return True
    if n % 2 == 0:
        return False

    for _ in range(k):
        a = random.randint(2, n - 2)
        if mod_exp(a, n - 1, n) != 1:
            return False   # a is a Fermat witness — n is composite
    return True


def fermat_test_verbose(n, k=5):
    """Same as fermat_test but prints each trial."""
    print(f"\nFermat test for n = {n}:")
    if n < 2:
        print("  Not prime (n < 2)"); return False
    if n % 2 == 0 and n != 2:
        print(f"  {n} is even — composite"); return False

    for i in range(k):
        a = random.randint(2, n - 2)
        result = mod_exp(a, n - 1, n)
        status = "PASS" if result == 1 else "FAIL (composite!)"
        print(f"  Trial {i+1}: a={a}, a^(n-1) mod n = {result}  [{status}]")
        if result != 1:
            return False
    print(f"  All {k} trials passed → probably prime")
    return True


if __name__ == "__main__":
    print("=== Prime Classification: Trial Division vs Fermat ===")
    test_numbers = [2, 3, 4, 5, 11, 13, 15, 17, 91, 97, 101, 341, 561, 1009]

    print(f"{'n':>6}  {'Trial':>7}  {'Fermat(k=20)':>14}  {'Note'}")
    print("-" * 50)
    for n in test_numbers:
        trial = is_prime_trial(n)
        fermat = fermat_test(n, k=20)
        note = ""
        if n == 341:
            note = "341=11×31, Carmichael-like"
        elif n == 561:
            note = "561=3×11×17, Carmichael number"
        elif n == 91:
            note = "91=7×13"
        print(f"{n:>6}  {str(trial):>7}  {str(fermat):>14}  {note}")

    print()
    print("=== Fermat Test — Verbose ===")
    fermat_test_verbose(97, k=5)
    fermat_test_verbose(91, k=5)

    print()
    print("=== Primes up to 50 (trial division) ===")
    primes = [n for n in range(2, 51) if is_prime_trial(n)]
    print(primes)

    print()
    print("=== Generating a large probable prime ===")
    import time
    bits = 20
    count = 0
    t0 = time.time()
    while True:
        candidate = random.getrandbits(bits) | 1   # ensure odd
        count += 1
        if fermat_test(candidate, k=30):
            t1 = time.time()
            print(f"  Found probable prime: {candidate}")
            print(f"  Tested {count} candidates in {(t1-t0)*1000:.1f} ms")
            break
