#!/usr/bin/env python3
"""Exact rank-metric and interpolation-decoder certificates. Standard library only.
Integers label base-p coefficient vectors, not arithmetic modulo the field size.
q-polynomial coefficient i belongs to X^(q^i); compositions remain FORMAL.
Explicit checks run unchanged under python -O. Small-field tables are pedagogical.
"""
import argparse
import itertools
import json
from collections import Counter
from pathlib import Path
CHECKS = 0

def check(ok, why):
    global CHECKS
    CHECKS += 1
    if not ok: raise ValueError(why)

def trim(a):
    a = list(a)
    while len(a)>1 and a[-1]==0: a.pop()
    return a or [0]

class Field:
    def __init__(self, p, modulus, d=1):
        self.p,self.mod = p,list(modulus); self.absolute_degree = len(modulus)-1
        if p<2 or any(p%j==0 for j in range(2,int(p**0.5)+1)): raise ValueError('p must be prime')
        if d<1 or self.absolute_degree<1 or self.absolute_degree%d or modulus[-1]!=1 or any(not 0<=x<p for x in modulus): raise ValueError('bad field presentation')
        self.N,self.q,self.m = p**self.absolute_degree,p**d,self.absolute_degree//d
        self.digits = [tuple((x//p**i)%p for i in range(self.absolute_degree)) for x in range(self.N)]
        self.add = [[self.encode([(u+v)%p for u,v in zip(a,b)]) for b in self.digits] for a in self.digits]
        self.neg = [self.encode([(-u)%p for u in a]) for a in self.digits]
        self.mul = [[self.rawmul(a,b) for b in self.digits] for a in self.digits]
        self.inv = [0]+[next((y for y in range(1,self.N) if self.mul[x][y]==1),None) for x in range(1,self.N)]
        if None in self.inv: raise ValueError('reducible modulus')
        self.K = [x for x in range(self.N) if self.power(x,self.q)==x]
        check(len(self.K)==self.q,'fixed field size')
        self.basis = [self.power(p,j) for j in range(self.m)]
        self.coords = {self.dot(c,self.basis):c for c in itertools.product(self.K,repeat=self.m)}
        check(len(self.coords)==self.N,'relative power basis')
    def encode(self,a): return sum(v*self.p**i for i,v in enumerate(a))
    def rawmul(self,a,b):
        m=self.absolute_degree; c=[0]*(2*m-1)
        for i,u in enumerate(a):
            for j,v in enumerate(b): c[i+j]=(c[i+j]+u*v)%self.p
        for i in range(len(c)-1,m-1,-1):
            for j in range(m): c[i-m+j]=(c[i-m+j]-c[i]*self.mod[j])%self.p
        return self.encode(c[:m])
    def sub(self,a,b): return self.add[a][self.neg[b]]
    def power(self,a,k):
        if k<0: raise ValueError('negative exponent')
        z=1
        while k:
            if k&1: z=self.mul[z][a]
            a=self.mul[a][a]; k//=2
        return z
    def total(self,seq):
        z=0
        for x in seq: z=self.add[z][x]
        return z
    def dot(self,a,b): return self.total(self.mul[x][y] for x,y in zip(a,b))
    def qeval(self,a,x): return self.total(self.mul[c][self.power(x,self.q**i)] for i,c in enumerate(a))
    def compose(self,a,b):
        out=[0]*(len(a)+len(b)-1)
        for i,x in enumerate(a):
            for j,y in enumerate(b): out[i+j]=self.add[out[i+j]][self.mul[x][self.power(y,self.q**i)]]
        return trim(out)
    def psub(self,a,b): return trim([self.sub(a[i] if i<len(a) else 0,b[i] if i<len(b) else 0) for i in range(max(len(a),len(b)))])
    def divide_left_factor(self,N,lam):
        """Return (f,R) with N=lam compose f+R and qdeg(R)<qdeg(lam)."""
        N,lam=trim(N),trim(lam)
        if lam==[0]: raise ValueError('zero composition divisor')
        t=len(lam)-1; out=[0]*max(1,len(N)-t)
        while N!=[0] and len(N)>=len(lam):
            j=len(N)-len(lam); u=self.mul[N[-1]][self.inv[lam[-1]]]
            c=self.power(u,self.q**((-t)%self.m)); out[j]=c
            N=self.psub(N,self.compose(lam,[0]*j+[c]))
        return trim(out),N
    def rref(self,A,b=None):
        width=len(A[0]) if A else 0
        if any(len(r)!=width for r in A) or (b is not None and len(b)!=len(A)): raise ValueError('matrix shape')
        B=[list(r)+([] if b is None else [b[i]]) for i,r in enumerate(A)]; pivots=[]; row=0
        for j in range(width):
            pivot=next((i for i in range(row,len(B)) if B[i][j]),None)
            if pivot is None: continue
            B[row],B[pivot]=B[pivot],B[row]
            z=self.inv[B[row][j]]; B[row]=[self.mul[z][x] for x in B[row]]
            for i in range(len(B)):
                if i!=row:
                    z=B[i][j]; B[i]=[self.sub(x,self.mul[z][y]) for x,y in zip(B[i],B[row])]
            pivots.append(j); row+=1
            if row==len(B): break
        if b is None: return len(pivots)
        if any(not any(r[:width]) and r[-1] for r in B): return None,[],B,pivots
        x=[0]*width
        for i,j in enumerate(pivots): x[j]=B[i][-1]
        null=[]
        for j in range(width):
            if j not in pivots:
                v=[0]*width; v[j]=1
                for i,h in enumerate(pivots): v[h]=self.neg[B[i][j]]
                null.append(v)
        return x,null,B,pivots
    def rankword(self,word): return self.rref([list(r) for r in zip(*(self.coords[x] for x in word))]) if word else 0
    def encode_message(self,f,g): return tuple(self.qeval(f,x) for x in g)

def validate(F,g,k,t,y):
    if type(k) is not int or type(t) is not int or not 1<=k<=len(g)<=F.m or not 0<=t<=(len(g)-k)//2: raise ValueError('invalid code parameters or radius')
    if len(y)!=len(g) or any(type(x) is not int or not 0<=x<F.N for x in list(g)+list(y)): raise ValueError('invalid field word')
    if F.rankword(g)!=len(g): raise ValueError('evaluation points must be K-independent')

def interpolate(F,g,k,t,y):
    A=[[F.power(x,F.q**j) for j in range(k+t)]+[F.neg[F.power(v,F.q**j)] for j in range(t)] for x,v in zip(g,y)]
    b=[F.power(v,F.q**t) for v in y]
    return A,b,F.rref(A,b)

def candidate(F,g,k,t,y,x):
    N,lam=trim(x[:k+t]),list(x[k+t:])+[1]; f,rem=F.divide_left_factor(N,lam)
    if rem!=[0]: return {'status':'NO_CODEWORD_WITHIN_RADIUS','reason':'nonzero_formal_remainder','N':N,'lambda':lam,'remainder':rem}
    if f!=[0] and len(f)>k: return {'status':'NO_CODEWORD_WITHIN_RADIUS','reason':'message_degree'}
    c=F.encode_message(f,g); e=tuple(F.sub(v,w) for v,w in zip(y,c)); r=F.rankword(e)
    if r>t: return {'status':'NO_CODEWORD_WITHIN_RADIUS','reason':'residual_rank','candidate':f,'residual_rank':r}
    return {'status':'DECODED','message':f,'codeword':list(c),'error':list(e),'error_rank':r,'N':N,'lambda':lam,'remainder':rem}

def decode(F,g,k,t,y):
    validate(F,g,k,t,y); A,b,(x,null,B,pivots)=interpolate(F,g,k,t,y)
    if x is None: return {'status':'NO_CODEWORD_WITHIN_RADIUS','reason':'inconsistent_system','rref':B}
    out=candidate(F,g,k,t,y,x); out['auxiliary_free_dimension']=len(null)
    return out

def rank_one_errors(F,n): return sorted({tuple(F.mul[u][v] for v in c) for u in range(F.N) for c in itertools.product(F.K,repeat=n)})

def code_and_balls(F,g,k):
    code={f:F.encode_message(f,g) for f in itertools.product(range(F.N),repeat=k)}
    check(len(set(code.values()))==F.N**k,'injective encoding')
    weights=Counter(F.rankword(c) for f,c in code.items() if any(f)); check(min(weights)==len(g)-k+1,'exact MRD distance')
    errors=rank_one_errors(F,len(g)); check(len(errors)==1+(F.N-1)*(F.q**len(g)-1)//(F.q-1),'rank-one ball size')
    balls={}
    for f,c in code.items():
        for e in errors:
            y=tuple(F.add[u][v] for u,v in zip(c,e)); check(y not in balls,'disjoint radius-one balls'); balls[y]=trim(f)
    return code,errors,balls,dict(sorted(weights.items()))

def test_received(F,g,k,t,words,balls):
    status=Counter(); reasons=Counter()
    for y in words:
        out=decode(F,g,k,t,y); expect=balls.get(tuple(y))
        check((out['status']=='DECODED')==(expect is not None),'decoder versus exhaustive ball dictionary')
        if expect is not None:
            check(out['message']==expect,'unique decoded message')
            check(F.compose(out['lambda'],out['message'])==out['N'],'formal division identity')
            check(all(F.qeval(out['lambda'],v)==F.qeval(out['N'],u) for u,v in zip(g,y)),'interpolation equations')
        else: reasons[out['reason']]+=1
        status[out['status']]+=1
    return {'statuses':dict(status),'failure_reasons':dict(reasons)}

def expect_reject(fn):
    try: fn()
    except ValueError: check(True,'invalid input rejected')
    else: check(False,'invalid input accepted')

def run():
    global CHECKS
    CHECKS = 0
    results=[]; main_case=None
    specs=[(2,[1,1,0,1],1,1,None,True),(2,[1,1,0,0,1],1,2,None,True),(2,[1,1,0,0,1],1,1,3,True),(3,[1,2,0,1],1,1,None,True),(2,[1,1,0,0,0,0,1],2,1,None,False)]
    for p,mod,d,k,length,exhaustive in specs:
        F=Field(p,mod,d); g=F.basis[:length]; n=len(g); code,errors,balls,weights=code_and_balls(F,g,k)
        words=itertools.product(range(F.N),repeat=n) if exhaustive else sorted(set(errors+list(balls)[::47]+[(i,(7*i+3)%F.N,(11*i+5)%F.N) for i in range(F.N)]))
        outcome=test_received(F,g,k,1,words,balls); erase_checks=0
        for f,c in code.items():
            for J in itertools.combinations(range(n),k):
                A=[[F.power(g[i],F.q**j) for j in range(k)] for i in J]
                recovered,free,_,_=F.rref(A,[c[i] for i in J])
                check(recovered==list(f) and free==[],'all k-survivor interpolations'); erase_checks+=1
        for h in range(F.N):
            lam=[h,1]; f=[F.add[h][1],h,0,1]+[0]*F.m+[1]
            N=F.compose(lam,f); quotient,rem=F.divide_left_factor(N,lam)
            check(quotient==trim(f) and rem==[0],'high formal-degree division')
        expect_reject(lambda:decode(F,[g[0]]*n,k,1,[0]*n))
        expect_reject(lambda:decode(F,g,k,(n-k)//2+1,[0]*n))
        expect_reject(lambda:decode(F,g,k,1,[0]*(n-1)))
        expect_reject(lambda:decode(F,g,k,1,[F.N]*n))
        for f,c in list(code.items())[:F.N]:
            zero=decode(F,g,k,0,c); check(zero['status']=='DECODED' and zero['message']==trim(f),'zero-radius codeword')
        for h in range(F.N):
            y=tuple((h+3*i)%F.N for i in range(n)); out=decode(F,g,n,0,y)
            check(out['status']=='DECODED' and out['codeword']==list(y),'full-space radius zero')
        results.append({'p':p,'modulus':mod,'q':F.q,'extension_degree':F.m,'n':n,'k':k,'g':g,'nonzero_codeword_rank_counts':weights,'rank_one_ball_size':len(errors),'covered_received_words':len(balls),'ambient_words':F.N**n,'received_scope':'all ambient words' if exhaustive else 'deterministic sample: all rank-one errors of zero, every 47th covered word, 64 ambient probes','tested_received_words':sum(outcome['statuses'].values()),'erasure_interpolations':erase_checks,**outcome})
        if F.N==16 and k==2:
            f=[2,1]; c=F.encode_message(f,g); received=tuple(F.add[v][8] for v in c); out=decode(F,g,k,1,received)
            A,b,(x,null,B,pivots)=interpolate(F,g,k,1,received)
            check(out['message']==f and out['error']==[8]*4,'main four-position rank-one correction')
            check(out['lambda']==[8,1] and out['N']==[3,12,1],'main exact auxiliary polynomials')
            A0,b0,(x0,null0,B0,pivots0)=interpolate(F,g,k,1,c); check(len(null0)==1,'one free auxiliary parameter at zero error')
            for u in range(F.N):
                v=[F.add[a][F.mul[u][b]] for a,b in zip(x0,null0[0])]; alt=candidate(F,g,k,1,c,v)
                check(alt['status']=='DECODED' and alt['message']==f,'all zero-error auxiliary choices give same message')
            fail_y=next(v for v in itertools.product(range(F.N),repeat=n) if v not in balls); failure=decode(F,g,k,1,fail_y)
            mis=decode(F,g,k,1,c); check(F.rankword(c)>1 and mis['message']==f,'successful radius certificate does not identify actual sender')
            check(F.rankword([1,1])==1 and F.rankword([1,2])==2,'coordinatewise L-scalings need not preserve rank')
            main_case={'field':'F2[a]/(a^4+a+1)','g':g,'message':f,'codeword':list(c),'received':list(received),'interpolation_matrix':A,'right_hand_side':b,'rref':B,'pivots':pivots,'decoded':out,'zero_error_auxiliary_choices':F.N,'zero_error_particular':x0,'zero_error_null_basis':null0,'failure_received':list(fail_y),'failure':failure,'outside_promise_received_codeword':list(c),'outside_promise_actual_error_rank':F.rankword(c)}
    # Radius two, message dimension two, full six-dimensional binary extension.
    F=Field(2,[1,1,0,0,0,0,1]); g=F.basis; radius_two=0
    for u in range(F.N):
        for v in range(F.N):
            f=[u,(5*v+3)%F.N]; c=F.encode_message(f,g)
            e=(u,v,F.add[u][v],0,u,v); y=tuple(F.add[a][b] for a,b in zip(c,e))
            out=decode(F,g,2,2,y)
            check(out['status']=='DECODED' and out['message']==trim(f),'radius-two recovery')
            check(F.compose(out['lambda'],out['message'])==out['N'],'radius-two formal identity')
            check(out['error_rank']==F.rankword(e)<=2,'radius-two residual certificate'); radius_two+=1
    # Left-factor division with nonmonic, high-degree divisors and nonzero remainders.
    for t in [0,1,2,7]:
        for h in range(1,F.N):
            lam=[h]*(t+1); f=[h,1,h]; N=F.compose(lam,f)
            quo,rem=F.divide_left_factor(N,lam)
            check(quo==f and rem==[0],'nonmonic arbitrary-degree formal division')
            if t:
                noisy=F.psub(N,[1]); quo,rem=F.divide_left_factor(noisy,lam)
                check(quo==f and rem==[1],'formal nonzero remainder retained')
    F4=Field(2,[1,1,0,0,0,0,1],2); transfer=decode(F4,F4.basis,1,1,[10,12,0])
    check(transfer['message']==[2] and transfer['lambda']==[24,1] and transfer['N']==[48,16],'nonprime-base terminal witness')
    check(transfer['error']==[8,8,8] and transfer['error_rank']==1,'nonprime-base terminal residual')
    F2=Field(2,[1,1,0,0,1]); F4small=Field(2,[1,1,0,0,1],2)
    check(F2.rankword([1,6])==2 and F4small.rankword([1,6])==1,'same vector different base-field rank')
    return {'schema':1,'checks':CHECKS,'field_label_convention':'base-p coefficient vectors; q-polynomial index i means X^(q^i)','field_cases':results,'radius_two_promised_cases':radius_two,'nonprime_base_terminal':transfer,'main':main_case}

def main():
    parser=argparse.ArgumentParser(); parser.add_argument('--output',type=Path,default=Path(__file__).with_name('foundations-rank-metric-results.json')); args=parser.parse_args()
    result=run(); args.output.parent.mkdir(parents=True,exist_ok=True); args.output.write_text(json.dumps(result,ensure_ascii=False,indent=2)+'\n')
    print(json.dumps({'checks':result['checks'],'field_cases':len(result['field_cases']),'output':str(args.output)}))
if __name__=='__main__': main()
