#!/usr/bin/env python3
"""Global geodesic examples: exact rational certificates and floating diagnostics.

Python standard library only. Writes JSON only with an explicit --output.
The checks do not replace the accompanying global existence and curvature proofs.
"""
from __future__ import annotations
import argparse
from fractions import Fraction as F
import json
import math
from pathlib import Path


class Checks:
    def __init__(self):
        self.counts = {}
    def check(self, condition, group, explanation):
        self.counts[group] = self.counts.get(group, 0) + 1
        if not condition:
            raise ArithmeticError(f'{group}: {explanation}')
    def close(self, x, y, group, explanation, tol=2e-11):
        self.check(abs(x-y) <= tol * max(1.0, abs(x), abs(y)), group, explanation)


def cylinder_minimizers(circumference, dx, dy):
    """Return exact squared distance and ALL nearest endpoint lift indices."""
    circumference, dx, dy = map(F, (circumference, dx, dy))
    if circumference <= 0:
        raise ValueError('circumference must be strictly positive')
    lower = (-dx) // circumference
    candidates = (lower, lower + 1)
    squares = {m: (dx + circumference*m)**2 + dy**2 for m in candidates}
    minimum = min(squares.values())
    return minimum, tuple(m for m in candidates if squares[m] == minimum)


def index_pi_coefficient(dimension, k, length_over_pi):
    """Coefficient of pi in the summed index upper bound for L=alpha*pi."""
    k, alpha = F(k), F(length_over_pi)
    if dimension < 2 or not isinstance(dimension, int) or k <= 0 or alpha <= 0:
        raise ValueError('need integer dimension >= 2, k > 0 and alpha > 0')
    return F(dimension-1, 2) * (1/alpha - k*alpha)


def paraboloid_curvature(radius):
    radius = F(radius)
    if radius < 0:
        raise ValueError('radius must be nonnegative')
    return 4 / (1 + 4*radius**2)**2


def refute_paraboloid_lower_bound(k):
    k = F(k)
    if k <= 0:
        raise ValueError('the proposed lower bound must be positive')
    radius = 1
    while 4*k*radius**4 <= 1:
        radius *= 2
    return radius, paraboloid_curvature(radius)


def sphere_exp(radius, v):
    """Exp at (0,0,radius), with Cartesian tangent vector v=(vx,vy)."""
    if radius <= 0:
        raise ValueError('radius must be positive')
    r = math.hypot(*v)
    if r == 0:
        return (0.0, 0.0, float(radius))
    factor = radius * math.sin(r/radius) / r
    return (factor*v[0], factor*v[1], radius*math.cos(r/radius))


def sphere_dexp(radius, v, w):
    if radius <= 0:
        raise ValueError('radius must be positive')
    r = math.hypot(*v)
    if r == 0:
        return (float(w[0]), float(w[1]), 0.0)
    u = (v[0]/r, v[1]/r)
    radial = u[0]*w[0] + u[1]*w[1]
    angular = (w[0]-radial*u[0], w[1]-radial*u[1])
    c, s = math.cos(r/radius), math.sin(r/radius)
    return (c*radial*u[0] + radius*s/r*angular[0],
            c*radial*u[1] + radius*s/r*angular[1], -s*radial)


def dot(x, y):
    return sum(a*b for a, b in zip(x, y))


def simpson(f, end, intervals=2048):
    if intervals <= 0 or intervals % 2:
        raise ValueError('positive even number of intervals required')
    step = end/intervals
    total = f(0.0)+f(end)
    total += sum((4 if j % 2 else 2)*f(j*step) for j in range(1, intervals))
    return step*total/3


def run():
    c = Checks()
    # Exhaustive nearest-integer tests, using a wider brute-force window as oracle.
    cylinder_cases = 0
    for ell in (F(1,2), F(1), F(2), F(3), F(6)):
        for denominator in (1,2,3,7):
            for numerator in range(-20,21):
                dx = F(numerator, denominator)
                for dy in (F(0), F(2,3), F(4)):
                    minimum, indices = cylinder_minimizers(ell, dx, dy)
                    center = math.floor(-dx/ell)
                    values = {m: (dx+ell*m)**2+dy**2 for m in range(center-5, center+7)}
                    brute = min(values.values())
                    brute_indices = tuple(m for m in values if values[m] == brute)
                    c.check((minimum, indices) == (brute, brute_indices), 'exact_cylinder', 'all minimizing lifts')
                    c.check(len(indices) in (1,2), 'exact_cylinder', 'one or two minimizers')
                    # Moving the representative by two turns changes indices by -2.
                    shifted, shifted_indices = cylinder_minimizers(ell, dx+2*ell, dy)
                    c.check(shifted == minimum and shifted_indices == tuple(m-2 for m in indices),
                            'exact_cylinder', 'representative invariance')
                    cylinder_cases += 1
    for delta in (F(j,7) for j in range(-20,21)):
        squared, indices = cylinder_minimizers(6, 3+delta, 4)
        expected = (0,) if delta < 0 else ((-1,) if delta > 0 else (-1,0))
        c.check(squared == (3-abs(delta))**2+16 and indices == expected,
                'exact_cylinder', 'cut branch on -3 < delta < 3')
    c.check(cylinder_minimizers(6,3,4) == (F(25),(-1,0)), 'exact_capstone', 'two 3-4-5 lifts')
    c.check(cylinder_minimizers(6,2,4) == (F(20),(0,)), 'exact_capstone', 'unique perturbed lift')
    c.check((3+6)**2+16 == 97 and (2-6)**2+16 == 32, 'exact_capstone', 'next length squares')
    # Exact inequalities behind the punctured-plane sequence, no rounded sqrt tests.
    for j in range(1,401):
        square = F(4)+F(4,j*j)
        c.check(square > 4 and square < (2+F(1,j*j))**2,
                'exact_missing_point', '0 < L_j-2 < 1/j^2 by squaring positive terms')
        next_square = F(4)+F(4,(j+1)**2)
        c.check(4 < next_square < square, 'exact_missing_point', 'decreasing above infimum')
        # Open disk radii and mutual distances are exact Euclidean distances.
        z_j, z_next = 1-F(1,j+1), 1-F(1,j+2)
        c.check(0 < z_j < z_next < 1 and z_next-z_j == F(1,(j+1)*(j+2)),
                'exact_missing_point', 'open-disk escape sequence')
    # Sign is controlled solely by L^2*k versus pi^2.
    index_cases = 0
    for n in (2,3,4,8):
        for root_k in (F(1,3),F(1,2),F(1),F(3,2),F(3)):
            for s in (F(j,12) for j in range(1,37)):
                alpha = s/root_k
                coefficient = index_pi_coefficient(n, root_k**2, alpha)
                normalized = F(n-1,2)*root_k*(1/s-s)
                c.check(coefficient == normalized, 'exact_index', 'dimensional and dimensionless formulas')
                c.check((coefficient > 0) == (s < 1) and (coefficient == 0) == (s == 1)
                        and (coefficient < 0) == (s > 1), 'exact_index', 'strict sign threshold')
                # g -> 9g gives L -> 3L and k -> k/9; the index scales by 1/3
                # for the normalized unit parallel test vectors used in this script.
                scaled = index_pi_coefficient(n, root_k**2/9, 3*alpha)
                c.check(scaled == coefficient/3, 'exact_index', 'length-curvature scaling')
                index_cases += 1
    c.check(index_pi_coefficient(4,F(9,4),1) == F(-15,8), 'exact_capstone', 'four-dimensional negative certificate')
    c.check(index_pi_coefficient(3,F(1,4),3) == F(-5,12), 'exact_capstone', 'formal-page three-dimensional example')
    c.check(index_pi_coefficient(4,F(9,4),F(2,3)) == 0, 'exact_capstone', 'sharp sphere threshold')
    # Uniform positivity really fails on the complete paraboloid.
    curvature_witnesses = []
    for k in [F(j,11) for j in range(1,34)] + [F(1,10**j) for j in range(1,25)]:
        r, curvature = refute_paraboloid_lower_bound(k)
        c.check(0 < curvature < F(1,4*r**4) < k,
                'exact_paraboloid', 'finite witness below each proposed positive bound')
        c.check(curvature == F(4,(1+4*r*r)**2), 'exact_paraboloid', 'exact Gaussian curvature')
        curvature_witnesses.append({'k':str(k), 'radius':r, 'curvature':str(curvature)})
    for r, expected in [(0,F(4)),(1,F(4,25)),(2,F(4,289)),(10,F(4,160801))]:
        c.check(paraboloid_curvature(r) == expected, 'exact_capstone', 'stated paraboloid sample')
    for radius in (F(j,9) for j in range(1,91)):
        E, G = 1+4*radius**2, radius**2
        # e*g_second = 4*r^2/(1+4*r^2), with no square roots needed.
        numerator = 4*radius**2/(1+4*radius**2)
        c.check(numerator/(E*G) == paraboloid_curvature(radius),
                'exact_paraboloid', 'first and second fundamental forms')
    # Rational representatives for projective points: exact norm and sign invariance.
    points = []
    for t in (F(j,5) for j in range(-10,11)):
        points.append(((1-t*t)/(1+t*t),2*t/(1+t*t),F(0)))
    points.extend([(F(0),F(0),F(1)),(F(3,5),F(4,5),F(0))])
    for p in points:
        c.check(dot(p,p) == 1, 'exact_projective', 'unit representative')
        for q in points:
            inner = dot(p,q)
            c.check(-1 <= inner <= 1 and abs(inner) == abs(dot(p,tuple(-x for x in q))),
                    'exact_projective', 'absolute inner product under antipodal identification')
    # Analytic sphere differential versus central differences and Gauss identity.
    max_gauss_error, max_fd_error = 0.0, 0.0
    sphere_cases = 0
    directions = ((1.0,0.0),(0.0,1.0),(0.6,0.8),(-0.8,0.6))
    test_vectors = ((1.,0.),(0.,1.),(2.,-3.),(-0.5,0.25))
    for a in (0.5,1.,2.,3.):
        for multiple in (0.,0.1,0.25,0.5,0.75,1.,1.5,2.):
            r = multiple*math.pi*a
            for u in directions:
                v = (r*u[0],r*u[1])
                point = sphere_exp(a,v)
                radial = sphere_dexp(a,v,v)
                c.close(dot(point,point),a*a,'float_sphere','exponential remains on sphere')
                c.close(dot(radial,radial),dot(v,v),'float_sphere','radial norm')
                for w in test_vectors:
                    dw = sphere_dexp(a,v,w)
                    lhs, rhs = dot(radial,dw), dot(v,w)
                    max_gauss_error = max(max_gauss_error,abs(lhs-rhs))
                    c.close(lhs,rhs,'float_sphere','full radial inner-product identity')
                    c.close(dot(point,dw),0.,'float_sphere','differential is tangent')
                    h = 1e-6
                    plus = sphere_exp(a,tuple(x+h*y for x,y in zip(v,w)))
                    minus = sphere_exp(a,tuple(x-h*y for x,y in zip(v,w)))
                    finite = tuple((x-y)/(2*h) for x,y in zip(plus,minus))
                    max_fd_error = max(max_fd_error,max(abs(x-y) for x,y in zip(finite,dw)))
                    for x,y in zip(finite,dw):
                        c.close(x,y,'float_sphere','central-difference comparison',tol=5e-8)
                    sphere_cases += 1
    c.close(sphere_exp(2.,(math.pi,0.))[0],2.,'float_capstone','north to equator')
    c.close(sphere_dexp(2.,(math.pi,0.),(0.,1.))[1],2/math.pi,'float_capstone','angular contraction')
    # Independently integrate the trial functions rather than substitute L/2.
    max_integral_error = 0.0
    for n in (2,4,7):
        for k in (F(1,4),F(1),F(9,4)):
            for alpha in (F(1,3),F(2,3),F(1),F(3,2),F(3)):
                L = float(alpha)*math.pi
                first = simpson(lambda t: (math.pi/L*math.cos(math.pi*t/L))**2,L)
                second = simpson(lambda t: math.sin(math.pi*t/L)**2,L)
                actual = (n-1)*(first-float(k)*second)
                expected = float(index_pi_coefficient(n,k,alpha))*math.pi
                max_integral_error = max(max_integral_error,abs(actual-expected))
                c.close(second,L/2,'float_integrals','sin squared integral')
                c.close(first,math.pi**2/(2*L),'float_integrals','derivative squared integral')
                c.close(actual,expected,'float_integrals','summed index integral')
    # Boundary contract: invalid input must be rejected, also under python -O.
    invalid = [lambda:cylinder_minimizers(0,1,2),lambda:cylinder_minimizers(-1,0,0),
               lambda:index_pi_coefficient(1,1,1),lambda:index_pi_coefficient(2,0,1),
               lambda:index_pi_coefficient(2,1,0),lambda:refute_paraboloid_lower_bound(0),
               lambda:paraboloid_curvature(-1),lambda:sphere_exp(0,(0,0))]
    for f in invalid:
        rejected = False
        try:
            f()
        except ValueError:
            rejected = True
        c.check(rejected,'input_validation','invalid model parameter rejected')
    return {
        'status':'PASS',
        'scope':'Finite checks of explicit models; not a replacement for the proofs.',
        'assertions':sum(c.counts.values()),'groups':c.counts,
        'exact_models':{'cylinder_cases':cylinder_cases,'index_cases':index_cases,
            'cylinder_main':{'distance_squared':'25','lift_indices':[-1,0],'next_squared':'97'},
            'cylinder_perturbed':{'distance_squared':'20','lift_indices':[0],'next_squared':'32'},
            'negative_index_coefficient_of_pi':'-15/8','sharp_length_coefficient_of_pi':'2/3',
            'punctured_plane_infimum':'2; not attained (geometric equality proof in text)',
            'projective_fiber_size':2,'paraboloid_witnesses':curvature_witnesses},
        'floating_diagnostics':{'sphere_vector_cases':sphere_cases,
            'max_gauss_absolute_error':format(max_gauss_error,'.3e'),
            'max_central_difference_absolute_error':format(max_fd_error,'.3e'),
            'max_integral_absolute_error':format(max_integral_error,'.3e'),
            'projective_sample_distance':format(math.acos(3/5),'.12f'),
            'note':'Tolerance diagnostics only; exact identities and global conclusions are proved in the text.'}}


def main():
    parser=argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--output',type=Path,help='write results only to this explicit path')
    args=parser.parse_args()
    result=run()
    output=json.dumps(result,ensure_ascii=False,indent=2)+'\n'
    if args.output:
        args.output.parent.mkdir(parents=True,exist_ok=True)
        args.output.write_text(output,encoding='utf-8')
    print(output,end='')

if __name__ == '__main__':
    main()
