#!/usr/bin/env python3
"""Exact finite root/word certificates; Python standard library only.
Usage: python foundations-root-word-check.py --output results.json
Counts finite model checks; the general theorems are proved in the pages.
"""
from fractions import Fraction as Q
from itertools import product,permutations
from collections import deque,Counter
from pathlib import Path
import argparse,json
checks=Counter()
def check(test,group):
 checks[group]+=1
 if not test:raise ArithmeticError(group)
def vec(x):return tuple(Q(a) for a in x)
def plus(x,y):return tuple(a+b for a,b in zip(x,y))
def neg(x):return tuple(-a for a in x)
def times(c,x):return tuple(c*a for a in x)
def dot(x,y):return sum((a*b for a,b in zip(x,y)),Q(0))
def eye(n):return tuple(tuple(Q(i==j) for j in range(n)) for i in range(n))
def transpose(A):return tuple(zip(*A))
def mv(A,x):return tuple(dot(row,x) for row in A)
def mm(A,B):return tuple(tuple(dot(row,col) for col in transpose(B)) for row in A)
def inner(x,y,G):return dot(x,mv(G,y))
def inverse(A):
 n=len(A);M=[list(row)+list(e) for row,e in zip(A,eye(n))]
 for j in range(n):
  p=next(i for i in range(j,n) if M[i][j]);M[j],M[p]=M[p],M[j];d=M[j][j];M[j]=[z/d for z in M[j]]
  for i in range(n):
   if i!=j:
    d=M[i][j];M[i]=[a-d*b for a,b in zip(M[i],M[j])]
 return tuple(tuple(row[n:]) for row in M)
def reflection(a,G):
 n=len(a);den=inner(a,a,G);dual=mv(G,a)
 return tuple(tuple(Q(i==j)-2*a[i]*dual[j]/den for j in range(n)) for i in range(n))
def cartan(simple,G):return tuple(tuple(2*inner(a,b,G)/inner(a,a,G) for b in simple) for a in simple)
def positive_definite(G):
 M=[list(row) for row in G]
 for j in range(len(G)):
  if M[j][j]<=0:return False
  for i in range(j+1,len(G)):
   for k in range(j+1,len(G)):M[i][k]-=M[i][j]*M[j][k]/M[j][j]
 return True

def certify_roots(roots,G,simple,h):
 roots=set(roots);n=len(G);zero=vec([0]*n);check(positive_definite(G),'positive_Gram');check(zero not in roots,'nonzero_roots')
 B=transpose(simple);Binv=inverse(B);check(mm(B,Binv)==eye(n),'simple_basis_spans')
 for a in roots:
  R=reflection(a,G);check(mm(R,R)==eye(n),'reflection_involution');check(mm(mm(transpose(R),G),R)==G,'reflection_isometry')
  for b in roots:
   pairing=2*inner(a,b,G)/inner(a,a,G);check(pairing.denominator==1,'Cartan_integrality');check(mv(R,b) in roots,'all_reflections_preserve_roots')
   j=next(i for i in range(n) if a[i]);ratio=b[j]/a[j]
   if b==times(ratio,a):check(ratio in [-1,1],'reduced_lines')
  check(inner(a,h,G)!=0,'regular_polarization')
 positive={a for a in roots if inner(a,h,G)>0};check(len(positive)*2==len(roots),'positive_half')
 indecomp={a for a in positive if not any(plus(b,c)==a for b in positive for c in positive)};check(indecomp==set(simple),'indecomposable_simple_basis')
 for a in roots:
  coords=mv(Binv,a);check(all(t.denominator==1 for t in coords),'integral_simple_coordinates');check(all(t>=0 for t in coords) or all(t<=0 for t in coords),'same_sign_simple_coordinates')
 gens=[reflection(a,G) for a in simple]
 for i,S in enumerate(gens):check({mv(S,a) for a in positive if a!=simple[i]}==positive-{simple[i]},'simple_flips_only_one')
 return positive,gens

def all_group(gens):
 I=eye(len(gens[0]));words={I:()};todo=deque([I])
 while todo:
  A=todo.popleft()
  for i,S in enumerate(gens):
   C=mm(A,S)
   if C not in words:words[C]=words[A]+(i,);todo.append(C)
 return words

def evaluate(word,gens):
 A=eye(len(gens[0]))
 for i in word:A=mm(A,gens[i])
 return A

def inversions(W,positive):return {a for a in positive if mv(W,a) not in positive}
def exchange(word,j,gens,simple,positive):
 # Follow the root from the right; return the precise deleted position.
 b=simple[j]
 for t in range(len(word)-1,-1,-1):
  nxt=mv(gens[word[t]],b)
  if b in positive and nxt not in positive:
   check(b==simple[word[t]],'first_sign_crossing_simple_root')
   return word[:t]+word[t+1:]
  b=nxt
 raise ArithmeticError('missing first positive-to-negative crossing')
def certify_words(roots,positive,gens,simple):
 words=all_group(gens);I=eye(len(gens[0]));N=len(positive);lengthdist=Counter()
 for W,word in words.items():
  inv=inversions(W,positive);lengthdist[len(inv)]+=1
  check({mv(W,a) for a in roots}==set(roots),'group_preserves_roots');check(len(word)==len(inv),'BFS_shortest_equals_root_inversions')
  for i,S in enumerate(gens):
   descending=mv(W,simple[i]) not in positive
   check(len(inversions(mm(W,S),positive))==len(inv)+(-1 if descending else 1),'right_multiplication_length_change')
   if descending:check(evaluate(exchange(word,i,gens,simple,positive),gens)==mm(W,S),'exchange_deleted_word_product')
  A=W;record=[]
  while A!=I:
   i=next(i for i,a in enumerate(simple) if mv(A,a) not in positive);record.append(i);A=mm(A,gens[i])
  check(len(record)==len(inv),'descent_termination_length');check(evaluate(record[::-1],gens)==W,'descent_reverse_word')
  bad=word+(0,0);check(evaluate(bad,gens)==W and len(bad)>len(inv),'reject_correct_but_nonminimal_word')
 longs=[W for W in words if len(inversions(W,positive))==N];check(len(longs)==1,'unique_longest');w0=longs[0];check(mm(w0,w0)==I,'longest_involution')
 for W in words:check(len(inversions(mm(w0,W),positive))==N-len(inversions(W,positive)),'longest_complement_length')
 return words,dict(sorted(lengthdist.items())),w0

rankrows=[]
for m,n,name,N in [(0,0,'A1xA1',2),(1,1,'A2',3),(2,1,'B2',4),(3,1,'G2',6)]:
 positives=[(1,0),(0,1)]
 if m>=1:positives.append((1,1))
 if m>=2:positives.append((2,1))
 if m>=3:positives.extend([(3,1),(3,2)])
 roots={vec(a) for a in positives}|{neg(vec(a)) for a in positives};G=eye(2) if m==0 else (vec([2*n,-m*n]),vec([-m*n,2*m]));simple=[vec([1,0]),vec([0,1])];h=mv(inverse(G),vec([1,1]));pos,gens=certify_roots(roots,G,simple,h)
 check(cartan(simple,G)==(vec([2,-m]),vec([-n,2])),'rank_two_Cartan')
 words,dist,w0=certify_words(roots,pos,gens,simple);check(len(words)==2*N,'rank_two_group_order');R=mm(*gens);P=eye(2)
 for k in range(1,N+1):P=mm(P,R);check((P==eye(2))==(k==N),'exact_product_order')
 # All expressions up to length seven, including nonreduced words, test deletion.
 for length in range(8):
  for word in product(range(2),repeat=length):
   W=evaluate(word,gens)
   for j in range(2):
    if mv(W,simple[j]) not in pos:check(evaluate(exchange(word,j,gens,simple,pos),gens)==mm(W,gens[j]),'nonreduced_exchange_identity')
 rankrows.append({'type':name,'roots':len(roots),'group_order':len(words),'product_order':N,'length_distribution':dist})

classical=[];B3=None
for kind in ['B','C']:
 for n in range(1,5):
  I=eye(n);roots=set();c=1 if kind=='B' else 2
  for i in range(n):roots.update([times(c,I[i]),times(-c,I[i])])
  for i in range(n):
   for j in range(i+1,n):
    for a,b in product([-1,1],repeat=2):roots.add(plus(times(a,I[i]),times(b,I[j])))
  simple=[plus(I[i],neg(I[i+1])) for i in range(n-1)]+[times(c,I[-1])];h=vec(range(n,0,-1));positive,gens=certify_roots(roots,I,simple,h);words,dist,w0=certify_words(roots,positive,gens,simple)
  signed={transpose([times(sign[j],I[p[j]]) for j in range(n)]) for p in permutations(range(n)) for sign in product([-1,1],repeat=n)}
  check(set(words)==signed,'all_and_only_signed_permutations');check(w0==tuple(neg(row) for row in I),'BC_longest_minus_identity');check(len(roots)==2*n*n,'BC_root_count')
  classical.append({'type':kind+str(n),'roots':len(roots),'group_order':len(words),'length_distribution':dist})
  if n==3:
   A=cartan(simple,I);expected=(vec([2,-1,0]),vec([-1,2,-1]),vec([0,-2,2]));check(A==(expected if kind=='B' else transpose(expected)),'BC_Cartan_transpose');check(sum(inner(a,a,I)==(2 if kind=='B' else 4) for a in roots)==(12 if kind=='B' else 6),'BC_long_root_count')
   if kind=='B':B3=(roots,positive,gens,simple,words)

roots,positive,gens,simple,words=B3;w=evaluate([2,1,2,0],gens);expected=transpose([vec([0,0,-1]),vec([1,0,0]),vec([0,-1,0])]);check(w==expected,'B3_target_column_images')
check(inversions(w,positive)=={vec([1,0,0]),vec([0,0,1]),vec([1,-1,0]),vec([1,0,1])},'B3_target_exact_inversions')
longword=[0,1,2,1,0,1,2,1,2];check(evaluate(longword,gens)==tuple(neg(row) for row in eye(3)) and len(longword)==9,'B3_longest_nine_word')
# Every regular point in a finite integer box; check all possible first descents.
point_models=0
for raw in product(range(-4,5),repeat=3):
 x=vec(raw)
 if any(dot(x,a)==0 for a in roots):continue
 point_models+=1;state=x;record=[];q=lambda z:sum(dot(z,a)<0 for a in positive)
 while any(dot(state,a)<0 for a in simple):
  for i,a in enumerate(simple):
   if dot(state,a)<0:check(q(mv(gens[i],state))==q(state)-1,'each_legal_chamber_step')
  i=next(i for i,a in enumerate(simple) if dot(state,a)<0);state=mv(gens[i],state);record.append(i)
 check(state==tuple(sorted(map(abs,x),reverse=True)),'B3_sorted_absolute_representative');check(len(record)==q(x),'chamber_steps_equal_initial_negative_count')
 transporters=[W for W in words if mv(W,x)==state];check(len(transporters)==1,'regular_transporter_unique');check(evaluate(record[::-1],gens)==transporters[0],'chamber_composition_order')
x=vec([-3,1,-2]);sequence=[0,1,2,1,0,2,1];states=[x]
for i in sequence:x=mv(gens[i],x);states.append(x)
check(states==list(map(vec,[(-3,1,-2),(1,-3,-2),(1,-2,-3),(1,-2,3),(1,3,-2),(3,1,-2),(3,1,2),(3,2,1)])),'capstone_seven_states')
u=evaluate(sequence[::-1],gens);check(u==(vec([-1,0,0]),vec([0,0,-1]),vec([0,1,0])),'capstone_transporter_matrix');check(len(inversions(u,positive))==7,'transporter_word_minimal')
wall=vec([2,2,1]);check(mv(gens[0],wall)==wall and gens[0]!=eye(3),'wall_unique_transporter_rejected');check(any(dot(wall,a)==0 for a in roots),'wall_input_not_regular')
AB=cartan(simple,eye(3));check(AB!=transpose(AB),'reject_transposed_Cartan_for_fixed_simple_basis')
# Incorrect column order has the same length but must fail the product interface.
wrong=evaluate([0,2,1,2],gens);check(wrong!=w and len(inversions(wrong,positive))==4,'reject_inverse_order_even_with_same_length')
# Rank-two finite and noncrystallographic boundaries are kept distinct.
R=(vec([3,-2]),vec([2,-1]));N=(vec([2,-2]),vec([2,-2]));check(mm(N,N)==(vec([0,0]),vec([0,0])),'unipotent_square_zero')
powers=[];P=eye(2)
for k in range(21):
 target=tuple(tuple(Q(i==j)+k*N[i][j] for j in range(2)) for i in range(2));check(P==target,'unipotent_power_formula');powers.append(P);P=mm(P,R)
check(len(set(powers))==21,'finite_contract_rejected_by_unipotent_certificate')
# Counterexample: reflection length 1 but positive-root count 2 in a nonreduced set.
check({Q(-1),Q(1),Q(-2),Q(2)}=={-z for z in [-1,1,-2,2]},'nonreduced_reflection_stable');check(sum(z>0 for z in [-1,1,-2,2])==2,'nonreduced_length_mismatch')

def clean(x):
 if isinstance(x,Q):return int(x) if x.denominator==1 else str(x)
 if isinstance(x,dict):return {str(k):clean(v) for k,v in x.items()}
 if isinstance(x,(list,tuple)):return [clean(v) for v in x]
 return x
out={'status':'PASS','checks':sum(checks.values()),'groups':dict(checks),'rank_two':rankrows,'classical_models':classical,'B3_target_word':[3,2,3,1],'B3_longest_word':[i+1 for i in longword],'B3_point_execution':[i+1 for i in sequence],'B3_point_states':states,'B3_point_transporter':u,'regular_integer_points':point_models,'limits':['Finite exact models test implementations and displayed certificates; general classification and length theorems use the written proofs.','Product order uses rightmost-first composition; column images are not a point trajectory.','Word length uses simple reflections, not all root reflections.','No claim that an arbitrary reflection closure terminates.']}
ap=argparse.ArgumentParser();ap.add_argument('--output',type=Path);args=ap.parse_args();txt=json.dumps(clean(out),ensure_ascii=False,indent=2)+'\n'
if args.output:args.output.parent.mkdir(parents=True,exist_ok=True);args.output.write_text(txt)
print(txt,end='')
