#!/usr/bin/env python3
"""Exact certificates for periods, Pell signs and finite norm orbits.
Python 3 standard library only. Run with --output PATH to choose the JSON result.
The algorithms use integer arithmetic (Fraction only for the full-ring example),
not floating square roots, numerical logarithms or an arbitrary search cutoff.
"""
from math import isqrt
from fractions import Fraction
from pathlib import Path
from itertools import product
import argparse,json

def verify(condition,message):
    """Execute mathematical verification in normal and optimized Python."""
    if not condition:raise AssertionError(message)

def valid_D(D):
 if not isinstance(D,int) or isinstance(D,bool) or D<=0 or isqrt(D)**2==D:
  raise ValueError('D must be a positive nonsquare integer')
 return isqrt(D)

def sqrt_period(D):
 a0=valid_D(D);P,Q,a=0,1,a0;out=[]
 while True:
  Pnew=a*Q-P;num=D-Pnew*Pnew
  if num%Q:raise ArithmeticError('nonintegral complete quotient')
  P,Q=Pnew,num//Q;a=(a0+P)//Q
  out.append((P,Q,a))
  if (P,Q)==(a0,1):return out

def verify_period(D,states):
 """Reject missing/extra/repeated/tampered states, including an unclosed list."""
 try:
  a0=valid_D(D);P,Q,a=0,1,a0;seen=set()
  if not states:return False
  for i,state in enumerate(states):
   if len(state)!=3 or any(type(x) is not int for x in state):return False
   pn=a*Q-P;num=D-pn*pn
   if num<=0 or num%Q:return False
   qn=num//Q;an=(a0+pn)//qn
   if tuple(state)!=(pn,qn,an) or (pn,qn) in seen:return False
   seen.add((pn,qn));P,Q,a=pn,qn,an
   if (P,Q)==(a0,1) and i!=len(states)-1:return False
  if (P,Q)!=(a0,1):return False
  pn=a*Q-P;qn=(D-pn*pn)//Q;an=(a0+pn)//qn
  return tuple(states[0])==(pn,qn,an)
 except (ValueError,TypeError,ZeroDivisionError):return False

def convergents(D,count,states=None):
 a0=valid_D(D);states=sqrt_period(D) if states is None else states
 pp,p,qq,q=0,1,1,0;out=[]
 for n in range(count):
  a=a0 if n==0 else states[(n-1)%len(states)][2]
  pp,p=p,a*p+pp;qq,q=q,a*q+qq
  out.append((p,q))
 return out

def norm(z,D):return z[0]*z[0]-D*z[1]*z[1]
def multiply(z,w,D):return z[0]*w[0]+D*z[1]*w[1],z[0]*w[1]+z[1]*w[0]
def unit_power(e,k,D):
 if norm(e,D)!=1:raise ValueError('power base must have norm one')
 if k<0:e=(e[0],-e[1]);k=-k
 z=(1,0)
 while k:
  if k&1:z=multiply(z,e,D)
  e=multiply(e,e,D);k//=2
 return z

def fundamental(D):
 states=sqrt_period(D);L=len(states)
 cs=convergents(D,2*L,states)
 positive=cs[(L if L%2==0 else 2*L)-1]
 negative=cs[L-1] if L%2 else None
 verify(norm(positive,D)==1 and (negative is None or norm(negative,D)==-1), 'fundamental certificate or invariant failed: norm(positive,D)==1 and (negative is None or norm(negative,D)==-1)')
 return positive,negative

def norm_seeds(D,N):
 """All representatives in [sqrt(abs(N)), epsilon*sqrt(abs(N)))."""
 valid_D(D)
 if not isinstance(N,int) or isinstance(N,bool) or not N:raise ValueError('N must be a nonzero integer')
 (u,v),_=fundamental(D)
 last=isqrt(v*v*N-1) if N>0 else isqrt((u*u*(-N)-1)//D)
 out=[]
 for y in range(0 if N>0 else 1,last+1):
  xx=D*y*y+N
  if xx<0:continue
  x=isqrt(xx)
  if x*x==xx:out.append((x,y))
 return (u,v),last,out

def reduce_solution(D,N,z):
 """Return sign, canonical seed, k with z = sign*seed*epsilon**k."""
 valid_D(D)
 if type(N) is not int or not N or len(z)!=2 or any(type(t) is not int for t in z) or norm(z,D)!=N:raise ValueError('input must be an integer solution of the specified nonzero norm equation')
 e,_=fundamental(D);u,v=e
 sign=1 if (z[0]>0 if N>0 else z[1]>0) else -1
 x,y=sign*z[0],sign*z[1];k=0
 # Below the lower endpoint: multiply until first entering the interval.
 while (y<0 if N>0 else x<0):
  x,y=multiply((x,y),e,D);k-=1
 # Above or exactly at the excluded upper endpoint: divide by epsilon.
 while (y*y>=v*v*N if N>0 else D*y*y>=u*u*(-N)):
  x,y=multiply((x,y),(u,-v),D);k+=1
 seed=(x,y)
 verify(norm(seed,D)==N and x>=0 and y>=0, 'reduce_solution certificate or invariant failed: norm(seed,D)==N and x>=0 and y>=0')
 verify(multiply(seed,unit_power(e,k,D),D)==(sign*z[0],sign*z[1]), 'reduce_solution certificate or invariant failed: multiply(seed,unit_power(e,k,D),D)==(sign*z[0],sign*z[1])')
 return sign,seed,k

def main():
 parser=argparse.ArgumentParser(description=__doc__)
 parser.add_argument('--output',type=Path,default=Path(__file__).with_name('foundations-pell-orbit-results.json'))
 args=parser.parse_args();res={};counts={}
 cases=0;identities=0;minimality=0
 for D in range(2,501):
  if isqrt(D)**2==D:continue
  st=sqrt_period(D);verify(verify_period(D,st), 'main certificate or invariant failed: verify_period(D,st)');L=len(st)
  cs=convergents(D,2*L,st)
  for n,z in enumerate(cs):
   verify(norm(z,D)==(-1)**(n+1)*st[n%L][1], 'main certificate or invariant failed: norm(z,D)==(-1)**(n+1)*st[n%L][1]');identities+=1
  e,negative=fundamental(D)
  # Exhaustive bounded comparison is a separate check, not the existence proof.
  firstpos=firstneg=None
  for y in range(1,min(e[1],2000)+1):
   xx=D*y*y+1;x=isqrt(xx)
   if x*x==xx and firstpos is None:firstpos=(x,y)
   xx=D*y*y-1;x=isqrt(xx)
   if x*x==xx and firstneg is None:firstneg=(x,y)
  if e[1]<=2000:verify(firstpos==e, 'main certificate or invariant failed: firstpos==e');minimality+=1
  else:verify(firstpos is None, 'main certificate or invariant failed: firstpos is None')
  if negative is None:verify(firstneg is None, 'main certificate or invariant failed: firstneg is None')
  elif negative[1]<=min(e[1],2000):verify(firstneg==negative, 'main certificate or invariant failed: firstneg==negative')
  for k in range(-5,6):verify(norm(unit_power(e,k,D),D)==1, 'main certificate or invariant failed: norm(unit_power(e,k,D),D)==1')
  if negative:verify(multiply(negative,negative,D)==e, 'main certificate or invariant failed: multiply(negative,negative,D)==e')
  cases+=1
 counts.update(period_certificates_D_through_500=cases,convergent_norm_identities=identities,positive_minimality_exhaustive_to_fundamental=minimality)
 for D in (2,3,7,13,34,61):
  st=sqrt_period(D);e,neg=fundamental(D)
  res[str(D)]={'period_states':st,'length':len(st),'positive_fundamental':e,'negative_fundamental':neg,'convergents':[{'n':n,'x':z[0],'y':z[1],'norm':norm(z,D)} for n,z in enumerate(convergents(D,2*len(st),st))]}
 verify(fundamental(13)==((649,180),(18,5)), 'main certificate or invariant failed: fundamental(13)==((649,180),(18,5))')
 verify(unit_power((649,180),2,13)==(842401,233640), 'main certificate or invariant failed: unit_power((649,180),2,13)==(842401,233640)')
 verify(multiply((18,5),(649,180),13)==(23382,6485), 'main certificate or invariant failed: multiply((18,5),(649,180),13)==(23382,6485)')
 verify(fundamental(34)==((35,6),None) and (13*13+1)%34==0, 'main certificate or invariant failed: fundamental(34)==((35,6),None) and (13*13+1)%34==0')

 # Reject corrupt periods, not merely check the final polynomial identity.
 st=sqrt_period(13);bad=[[],st[:-1],st+st,[(*st[0][:2],st[0][2]+1),*st[1:]],[(st[0][0]+1,*st[0][1:]),*st[1:]],list(reversed(st))]
 for a in bad:verify(not verify_period(13,a), 'main certificate or invariant failed: not verify_period(13,a)')
 counts['corrupt_period_certificates_rejected']=len(bad)
 invalid=0
 for D in (0,-1,1,4,9,True,2.5):
  try:sqrt_period(D)
  except ValueError:invalid+=1
  else:raise AssertionError('invalid radicand accepted')
 counts['invalid_radicands_rejected']=invalid

 # Full integer-ring versus order example, exact rational coefficients.
 theta=(Fraction(3,2),Fraction(1,2));z=(Fraction(1),Fraction(0));powers=[]
 for k in range(1,7):
  z=multiply(z,theta,13)
  powers.append({'k':k,'x':str(z[0]),'y':str(z[1]),'norm':str(norm(z,13))})
 verify((powers[2]['x'],powers[2]['y'])==('18','5'), "main certificate or invariant failed: (powers[2]['x'],powers[2]['y'])==('18','5')")
 verify((powers[5]['x'],powers[5]['y'])==('649','180'), "main certificate or invariant failed: (powers[5]['x'],powers[5]['y'])==('649','180')")
 res['full_ring_theta_powers']=powers

 # Every reduced seed for D<=30, 1<=abs(N)<=12; generated orbit elements
 # through three positive/negative powers and both signs must reduce uniquely.
 models=0;seedcount=0;reductions=0;bounded=0
 for D in range(2,31):
  if isqrt(D)**2==D:continue
  for N in (*range(-12,0),*range(1,13)):
   e,last,seeds=norm_seeds(D,N);models+=1;seedcount+=len(seeds)
   verify(len(set(seeds))==len(seeds), 'main certificate or invariant failed: len(set(seeds))==len(seeds)')
   for seed in seeds:
    verify(reduce_solution(D,N,seed)==(1,seed,0), 'main certificate or invariant failed: reduce_solution(D,N,seed)==(1,seed,0)')
    for k in range(-3,4):
     for s in (-1,1):
      w=multiply(seed,unit_power(e,k,D),D);w=s*w[0],s*w[1]
      verify(reduce_solution(D,N,w)==(s,seed,k), 'main certificate or invariant failed: reduce_solution(D,N,w)==(s,seed,k)');reductions+=1
   # Independent small-box solution discovery must land in the full seed list.
   for y in range(-60,61):
    xx=D*y*y+N
    if xx<0:continue
    x=isqrt(xx)
    if x*x!=xx:continue
    for sign in ((1,) if x==0 else (-1,1)):
     z=(sign*x,y);_,seed,_=reduce_solution(D,N,z);verify(seed in seeds, 'main certificate or invariant failed: seed in seeds');bounded+=1
 counts.update(generalized_norm_models=models,canonical_seeds=seedcount,generated_orbit_reductions=reductions,independently_found_small_box_solutions=bounded)
 e,last,seeds=norm_seeds(13,4)
 verify(last==359 and seeds==[(2,0),(11,3),(119,33)], 'main certificate or invariant failed: last==359 and seeds==[(2,0),(11,3),(119,33)]')
 verify(reduce_solution(13,4,(1298,360))==(1,(2,0),1), 'main certificate or invariant failed: reduce_solution(13,4,(1298,360))==(1,(2,0),1)')
 verify(multiply((11,-3),e,13)==(119,33), 'main certificate or invariant failed: multiply((11,-3),e,13)==(119,33)')
 res['D13_N4']={'positive_unit':e,'last_y_inclusive':last,'seeds':seeds,'excluded_upper_endpoint':{'solution':[1298,360],'reduction':[1,[2,0],1]},'conjugate_seed_relation':'(11-3 sqrt(13))*epsilon = 119+33 sqrt(13)'}
 for D,N in ((13,-4),(2,-2),(2,3),(7,2)):
  e,last,seeds=norm_seeds(D,N);res[f'D{D}_N{N}']={'positive_unit':e,'last_y_inclusive':last,'seeds':seeds}
 verify(norm_seeds(2,-2)[2]==[(0,1)] and norm_seeds(2,3)[2]==[], 'main certificate or invariant failed: norm_seeds(2,-2)[2]==[(0,1)] and norm_seeds(2,3)[2]==[]')
 verify(reduce_solution(2,-2,(140,99))==(1,(0,1),3), 'main certificate or invariant failed: reduce_solution(2,-2,(140,99))==(1,(0,1),3)')
 res['transfer_reduction']={'D':2,'N':-2,'input':[140,99],'sign':1,'seed':[0,1],'exponent':3}
 res['status']='PASS';res['checks']=counts
 args.output.parent.mkdir(parents=True,exist_ok=True)
 args.output.write_text(json.dumps(res,ensure_ascii=False,indent=2)+'\n')
 print(json.dumps({'status':'PASS','checks':counts},ensure_ascii=False,indent=2))

if __name__=='__main__':main()
