#!/usr/bin/env python3
"""Exact S3 convolution certificates over Q(sqrt(3),i), plus modular rank and erasure checks.
Standard library only. No floating point; this is not a generic irrep-discovery or FFT package.
"""
from fractions import Fraction as Q
from itertools import product
from collections import Counter
from pathlib import Path
import argparse,json
from math import isqrt

class K:
 """a+b sqrt(3)+i(c+d sqrt(3)); exact rational coefficients."""
 __slots__=('v',)
 def __init__(self,a=0,b=0,c=0,d=0):
  if isinstance(a,K):self.v=a.v;return
  if any(not isinstance(x,(int,Q))for x in(a,b,c,d)):raise ValueError('exact rational coefficients required')
  self.v=tuple(Q(x)for x in(a,b,c,d))
 def __add__(self,z):z=K(z);return K(*(x+y for x,y in zip(self.v,z.v)))
 __radd__=__add__
 def __neg__(self):return K(*(-x for x in self.v))
 def __sub__(self,z):return self+-K(z)
 def __rsub__(self,z):return K(z)+-self
 def __mul__(self,z):
  z=K(z);a,b,c,d=self.v;e,f,g,h=z.v
  return K(a*e+3*b*f-c*g-3*d*h,a*f+b*e-c*h-d*g,a*g+3*b*h+c*e+3*d*f,a*h+b*g+c*f+d*e)
 __rmul__=__mul__
 def conj(self):a,b,c,d=self.v;return K(a,b,-c,-d)
 def inv(self):
  if not self:raise ValueError('zero inverse')
  n=self*self.conj();a,b,c,d=n.v
  if c or d:raise ArithmeticError('norm not real')
  return self.conj()*K(a/(a*a-3*b*b),-b/(a*a-3*b*b))
 def __truediv__(self,z):return self*K(z).inv()
 def __bool__(self):return any(self.v)
 def __eq__(self,z):return self.v==K(z).v
 def __repr__(self):
  names=['','*sqrt(3)','*i','*sqrt(3)*i'];return '+'.join(str(v)+n for v,n in zip(self.v,names)if v)or'0'

class Field:
 def __init__(self,p=0):
  if not isinstance(p,int)or p<0 or p==1 or p and any(p%d==0 for d in range(2,isqrt(p)+1)):raise ValueError('zero or prime characteristic')
  self.p=p;self.z=self.c(0);self.o=self.c(1)
 def c(self,x):
  if not self.p:return K(x)
  if not isinstance(x,(int,Q)):raise ValueError('exact prime-field input')
  x=Q(x)
  if x.denominator%self.p==0:raise ValueError('denominator vanishes')
  return x.numerator*pow(x.denominator,-1,self.p)%self.p
 def add(self,a,b):return(a+b)%self.p if self.p else a+b
 def neg(self,a):return(-a)%self.p if self.p else-a
 def mul(self,a,b):return a*b%self.p if self.p else a*b
 def inv(self,a):
  if not a:raise ValueError('zero inverse')
  return pow(a,-1,self.p)if self.p else a.inv()
 def star(self,a):return a if self.p else a.conj()
 def total(self,xs):
  s=self.z
  for x in xs:s=self.add(s,x)
  return s

def zeros(F,m,n):return[[F.z for _ in range(n)]for _ in range(m)]
def ident(F,n):return[[F.o if i==j else F.z for j in range(n)]for i in range(n)]
def trans(A):return[list(x)for x in zip(*A)]
def adj(F,A):return[[F.star(x)for x in row]for row in trans(A)]
def mm(F,A,B):return[[F.total(F.mul(a,b)for a,b in zip(row,col))for col in zip(*B)]for row in A]
def scale(F,c,A):return[[F.mul(c,x)for x in row]for row in A]
def ma(F,A,B):return[[F.add(a,b)for a,b in zip(x,y)]for x,y in zip(A,B)]
def rref(F,A,n=None):
 B=[r[:]for r in A];n=len(B[0])if n is None and B else(n or 0);r=0;piv=[]
 for c in range(n):
  j=next((j for j in range(r,len(B))if B[j][c]),None)
  if j is None:continue
  B[r],B[j]=B[j],B[r];q=F.inv(B[r][c]);B[r]=[F.mul(q,x)for x in B[r]]
  for j in range(len(B)):
   if j!=r and B[j][c]:
    q=B[j][c];B[j]=[F.add(x,F.neg(F.mul(q,y)))for x,y in zip(B[j],B[r])]
  piv.append(c);r+=1
  if r==len(B):break
 return B,piv

def rank(F,A,n=None):return len(rref(F,A,n)[1])
def solve(F,A,b,n=None):
 if len(A)!=len(b)or A and any(len(row)!=len(A[0])for row in A):raise ValueError('system dimensions')
 n=len(A[0])if A else(n or 0);B,p=rref(F,[r+[v]for r,v in zip(A,b)],n)
 if any(not any(r[:n])and r[n]for r in B):return None
 x=[F.z]*n
 for i,j in enumerate(p):x[j]=B[i][n]
 basis=[]
 for j in range(n):
  if j in p:continue
  v=[F.z]*n;v[j]=F.o
  for i,c in enumerate(p):v[c]=F.neg(B[i][j])
  basis.append(v)
 return x,basis

def inverse(F,A):
 n=len(A);B,p=rref(F,[r+s for r,s in zip(A,ident(F,n))],n)
 if len(p)!=n:return None
 return[r[n:]for r in B]

def dihedral(n):
 G=[(a,b)for b in range(2)for a in range(n)]
 return[[G.index(((a+(-1)**b*c)%n,(b+d)%2))for c,d in G]for a,b in G]
def validate_group(T):
 n=len(T)
 if not n or any(len(r)!=n or any(not isinstance(v,int)or not 0<=v<n for v in r)for r in T):raise ValueError('bad multiplication table')
 if any(T[0][g]!=g or T[g][0]!=g for g in range(n)):raise ValueError('identity must be first')
 inv=[]
 for g in range(n):
  h=next((h for h in range(n)if T[g][h]==T[h][g]==0),None)
  if h is None:raise ValueError('inverse missing')
  inv.append(h)
 if any(T[T[a][b]][c]!=T[a][T[b][c]]for a,b,c in product(range(n),repeat=3)):raise ValueError('nonassociative table')
 return inv

def conv(F,T,a,b):
 if len(a)!=len(T)or len(b)!=len(T):raise ValueError('function length')
 out=[F.z]*len(T)
 for x,y in product(range(len(T)),repeat=2):out[T[x][y]]=F.add(out[T[x][y]],F.mul(a[x],b[y]))
 return out

def left_matrix(F,T,iv,a):return[[a[T[x][iv[t]]]for t in range(len(T))]for x in range(len(T))]
def delta(F,n,g):return[F.o if i==g else F.z for i in range(n)]
def reps(F):
 if F.p:raise ValueError('complex S3 representation contract')
 R=[[K(Q(-1,2)),K(0,Q(-1,2))],[K(0,Q(1,2)),K(Q(-1,2))]];S=[[K(1),K(0)],[K(0),K(-1)]];I=ident(F,2);RR=mm(F,R,R)
 return[[[[K(1)]]for _ in range(6)],[[[K(1 if i<3 else-1)]]for i in range(6)],[I,R,RR,S,mm(F,R,S),mm(F,RR,S)]]
def ft(F,reps,a):
 if any(len(rep)!=len(a)for rep in reps):raise ValueError('transform function length')
 out=[]
 for rep in reps:
  M=zeros(F,len(rep[0]),len(rep[0]))
  for c,A in zip(a,rep):M=ma(F,M,scale(F,c,A))
  out.append(M)
 return out

def ift(F,reps,iv,blocks):
 if len(reps)!=len(blocks)or any(len(B)!=len(rep[0])or any(len(row)!=len(B)for row in B)for rep,B in zip(reps,blocks)):raise ValueError('transform block dimensions')
 n=len(iv);out=[]
 for g in range(n):
  vals=[]
  for rep,B in zip(reps,blocks):
   A=mm(F,B,rep[iv[g]]);vals.append(F.mul(F.c(len(B)),F.total(A[i][i]for i in range(len(B)))))
  out.append(F.mul(F.inv(F.c(n)),F.total(vals)))
 return out

def block_solve(F,reps,iv,A,b):
 X=[];freedom=[]
 for j,(a,y)in enumerate(zip(A,b)):
  d=len(a);x=zeros(F,d,d)
  for col in range(d):
   ans=solve(F,a,[r[col]for r in y])
   if ans is None:return None
   v,ns=ans
   for row in range(d):x[row][col]=v[row]
   for w in ns:
    Z=[zeros(F,len(vv),len(vv))for vv in A]
    for row in range(d):Z[j][row][col]=w[row]
    freedom.append(ift(F,reps,iv,Z))
  X.append(x)
 return ift(F,reps,iv,X),freedom

COUNT=Counter()
def check(ok,key):
 if not ok:raise AssertionError(key)
 COUNT[key]+=1

def enc(x):
 if isinstance(x,K):return repr(x)
 if isinstance(x,(list,tuple)):return[enc(v)for v in x]
 if isinstance(x,dict):return{k:enc(v)for k,v in x.items()}
 return x

def rejected(fn):
 try:fn()
 except(ValueError,TypeError,IndexError):return True
 return False

def run():
 COUNT.clear();F=Field();T=dihedral(3);iv=validate_group(T);Rs=reps(F);N=6;I=delta(F,N,0);records={'group_order':['e','r','r^2','s','rs','r^2s']}
 check(sum(len(r[0])**2 for r in Rs)==N,'complete representation dimensions')
 for rep in Rs:
  d=len(rep[0])
  for g,h in product(range(N),repeat=2):check(mm(F,rep[g],rep[h])==rep[T[g][h]],'representation multiplication')
  for A in rep:check(mm(F,adj(F,A),A)==ident(F,d),'unitary matrices')
 entries=[(len(rep[0]),[rep[g][i][j]for g in range(N)])for rep in Rs for i,j in product(range(len(rep[0])),repeat=2)]
 for i,(d,a)in enumerate(entries):
  for j,(_,b)in enumerate(entries):check(F.total(F.mul(F.star(x),y)for x,y in zip(a,b))==F.c(Q(N,d)if i==j else 0),'matrix coefficient orthogonality')
 for g,h in product(range(N),repeat=2):
  a=delta(F,N,g);b=delta(F,N,h);check(conv(F,T,a,b)==delta(F,N,T[g][h]),'delta convolution direction')
  check(ft(F,Rs,conv(F,T,a,b))==[mm(F,x,y)for x,y in zip(ft(F,Rs,a),ft(F,Rs,b))],'full bilinear convolution on basis')
 for raw in product([-1,0,1],repeat=N):
  a=list(map(F.c,raw));A=ft(F,Rs,a);L=left_matrix(F,T,iv,a);rr=rank(F,L)
  check(ift(F,Rs,iv,A)==a,'all ternary S3 inversion')
  check(rr==sum(len(x)*rank(F,x)for x in A),'direct versus weighted block rank')
  check(not any(a)or sum(bool(x)for x in a)*rr>=N,'complex support rank inequality')
  check(F.total(F.mul(x,x)for x in a)==F.mul(F.c(Q(1,N)),F.total(F.mul(F.c(len(B)),F.total(F.mul(x,x)for row in B for x in row))for B in A)),'real Plancherel')
  central=all(a[T[T[g][x]][iv[g]]]==a[x]for g,x in product(range(N),repeat=2))
  scalar=all(B==scale(F,B[0][0],ident(F,len(B)))for B in A)
  check(central==scalar,'class function iff scalar blocks')
 # Complex inputs and both orders; includes nonzero imaginary and irrational parts.
 for k in range(36):
  a=[K((k+2*j)%5-2,Q((k+j)%3-1,2),((k+1)*(j+1))%3-1,Q((k+3*j)%3-1,3))for j in range(N)]
  b=[K((k+3*j)%3-1,0,(k*j+2)%5-2)for j in range(N)];A=ft(F,Rs,a);B=ft(F,Rs,b)
  check(ift(F,Rs,iv,A)==a,'complex irrational inversion')
  check(ft(F,Rs,conv(F,T,a,b))==[mm(F,x,y)for x,y in zip(A,B)],'complex ordered convolution')
  check(F.total(F.mul(F.star(x),y)for x,y in zip(a,b))==F.mul(F.c(Q(1,N)),F.total(F.mul(F.c(len(X)),F.total(F.mul(F.star(x),y)for r,s in zip(X,Y)for x,y in zip(r,s)))for X,Y in zip(A,B))),'complex Plancherel')
  star=[F.star(a[iv[g]])for g in range(N)];check(ft(F,Rs,star)==[adj(F,x)for x in A],'group involution adjoint')
 u=list(map(K,[0,-1,1,0,1,-1]));A=ft(F,Rs,u);q=[F.add(I[j],F.neg(delta(F,N,3)[j]))for j in range(N)];a=list(map(K,[2,1,0,1,0,0]));H=list(map(K,[1,0,0,1,0,0]));aa=[F.add(I[j],u[j])for j in range(N)];ainv=[F.add(I[j],F.neg(u[j]))for j in range(N)]
 check(A==[[[K(0)]],[[K(0)]],[[K(0),K(0,2)],[K(0),K(0)]]],'noncentral nilpotent block')
 check(conv(F,T,u,u)==[K(0)]*N and any(u),'nonzero nilpotent convolution')
 check(conv(F,T,aa,ainv)==I and conv(F,T,ainv,aa)==I,'nilpotent unit two sided inverse')
 B=ft(F,Rs,a);inva=ift(F,Rs,iv,[inverse(F,x)for x in B]);check(inva==list(map(K,[Q(5,8),Q(-3,8),Q(1,8),Q(-3,8),Q(1,8),Q(1,8)])),'explicit general kernel inverse')
 check(conv(F,T,a,inva)==I and conv(F,T,inva,a)==I,'original group inverse verification')
 C=mm(F,adj(F,ft(F,Rs,aa)[2]),ft(F,Rs,aa)[2]);check(C[0][0]+C[1][1]==K(14)and C[0][0]*C[1][1]-C[0][1]*C[1][0]==K(1),'shear singular square polynomial')
 records['noncentral_nilpotent']={'values':u,'blocks':A,'rank':2,'support':4};records['general_inverse']={'kernel':a,'inverse':inva,'blocks':B};records['singular_kernel']={'values':q,'blocks':ft(F,Rs,q),'rank':3}
 for kernel in [u,q,a,aa,H,I,[K(0)]*N]:
  L=left_matrix(F,T,iv,kernel);A=ft(F,Rs,kernel)
  rights=[delta(F,N,j)for j in range(N)]+[conv(F,T,kernel,[K(j+1)for j in range(N)]),kernel]
  for b in rights:
   ans=block_solve(F,Rs,iv,A,ft(F,Rs,b));direct=solve(F,L,b)
   check((ans is None)==(direct is None),'block and direct compatibility')
   if ans is None:continue
   x,ns=ans;check(conv(F,T,kernel,x)==b,'block solution original residual')
   check(len(ns)==N-rank(F,L)and rank(F,trans(ns),len(ns))==len(ns),'all solution freedom complete')
   for v in ns:check(conv(F,T,kernel,v)==[K(0)]*N,'kernel freedom original residual')
 # Changes of irrep basis preserve algebra/rank; a nonunitary change need not preserve raw Frobenius norm.
 for P in [[[K(0),K(-1)],[K(1),K(0)]],[[K(Q(3,5)),K(Q(-4,5))],[K(Q(4,5)),K(Q(3,5))]],[[K(1),K(1)],[K(0),K(1)]]]:
  Pi=inverse(F,P);new=[Rs[0],Rs[1],[mm(F,mm(F,Pi,M),P)for M in Rs[2]]]
  for v in [u,q,a,aa]:
   blocks=ft(F,new,v);check(ift(F,new,iv,blocks)==v,'changed basis algebraic inversion')
   check(sum(len(B)*rank(F,B)for B in blocks)==rank(F,left_matrix(F,T,iv,v)),'changed basis ranks')
 original=ft(F,Rs,q)[2];changed=ft(F,new,q)[2]
 check(F.total(x*x for row in original for x in row)!=F.total(x*x for row in changed for x in row),'nonunitary basis changes raw Frobenius norm')
 # Any-field inequality checked from raw translation matrices, with actual pivot-column cover certificates.
 for p,table in [(2,T),(3,T),(2,dihedral(4)),(5,[[ (i+j)%5 for j in range(5)]for i in range(5)])]:
  E=Field(p);inv=validate_group(table);n=len(table)
  for raw in product(range(p),repeat=n):
   L=left_matrix(E,table,inv,raw);_,piv=rref(E,L);r=len(piv)
   check(not any(raw)or sum(bool(x)for x in raw)*r>=n,'arbitrary field support rank')
   if not any(raw):continue
   supports=[{x for x in range(n)if L[x][j]}for j in piv]
   check(set().union(*supports)==set(range(n))and all(len(s)==sum(bool(x)for x in raw)for s in supports),'actual translate basis support cover')
 # Every subgroup of S3, every erasure set, in three characteristics.
 subgroups=[]
 for mask in range(1<<N):
  Hset={j for j in range(N)if mask>>j&1}
  if 0 in Hset and all(T[g][h]in Hset for g,h in product(Hset,repeat=2)):subgroups.append(Hset)
 check(len(subgroups)==6,'all S3 subgroups')
 for p in [2,3,5]:
  E=Field(p)
  for Hset in subgroups:
   cosets=[]
   for t in range(N):
    Sset={T[h][t]for h in Hset}
    if Sset not in cosets:cosets.append(Sset)
   r=len(cosets);B=[[E.c(x in C)for C in cosets]for x in range(N)];c=[E.c(j+1)for j in range(r)];v=[E.total(E.mul(x,y)for x,y in zip(row,c))for row in B]
   check(rank(E,B)==r and len(Hset)*r==N,'subgroup equality all fields')
   for mask in range(1<<N):
    known=[x for x in range(N)if not mask>>x&1];BK=[B[x]for x in known];rr=rank(E,BK,r);ans=solve(E,BK,[v[x]for x in known],r)
    check((rr==r)==all(set(known)&C for C in cosets),'exact retained row coset criterion')
    check(not ((N-len(known))*r<N)or rr==r,'uniform erasure guarantee')
    check(ans is not None and len(ans[1])==r-rr,'erasure ambiguity dimension')
    if rr==r:check(ans[0]==c,'erasure original message recovery')
    C=next((C for C in cosets if len(set(known)&C)>=2),None)
    if C:
     bad=[v[x]for x in known];bad[known.index(next(iter(set(known)&C)))]=E.add(bad[known.index(next(iter(set(known)&C)))],1)
     check(solve(E,BK,bad,r)is None,'incompatible retained data rejected')
 E=Field(2);C2=[[0,1],[1,0]];f=[1,1];check(rank(E,left_matrix(E,C2,[0,1],f))==1 and conv(E,C2,f,f)==[0,0]and sum(f)%2==0,'modular simple block loses nonzero rank')
 for fn in [lambda:K(.5),lambda:Field(4),lambda:validate_group([[0,1],[1,1]]),lambda:conv(F,T,[K(1)],I),lambda:K(0).inv(),lambda:reps(Field(2))]:check(rejected(fn),'invalid contract rejected')
 return enc({'status':'PASS','checks':sum(COUNT.values()),'groups':dict(COUNT),'records':records,'scope':'Exact S3 arithmetic over Q(sqrt(3),i); original-matrix rank over F2/F3/F5 and finite erasure tests. General proofs, arbitrary complex input and arbitrary groups are justified by the articles, not by finite enumeration.'})
if __name__=='__main__':
 p=argparse.ArgumentParser(description=__doc__);p.add_argument('--output',type=Path);args=p.parse_args();out=json.dumps(run(),ensure_ascii=False,indent=2)+'\n'
 if args.output:args.output.parent.mkdir(parents=True,exist_ok=True);args.output.write_text(out);print('PASS',sum(COUNT.values()),'checks',args.output)
 else:print(out,end='')
