#!/usr/bin/env python3
"""Exact finite checks for the encrypted-computation unit; no security claims.
Run with Python 3. Uses only the standard library and writes JSON to stdout.
"""
from fractions import Fraction as F
from itertools import product
from collections import Counter
import json,math,random
checks=[]
def eq(name,a,b):
    assert a==b,(name,a,b)
    checks.append({'name':name,'value':str(a),'expected':str(b),'status':'pass'})
def fact(name,**kw):checks.append(dict(name=name,status='pass',**kw))
def center(x,q):return (x+q//2)%q-q//2
def nearest(x):return (x.numerator*2+x.denominator)//(2*x.denominator)
def distance(a,b,q):return min((a-b)%q,(b-a)%q)
def decode_high(w,q):return int(distance(w,q//2,q)<=distance(w,0,q))
# Two-gate clear function.
eq('AND then XOR all inputs',[(a*b)^c for a,b,c in product(range(2),repeat=3)],[0,1,0,1,0,1,1,0])
# Paillier: all messages/randomizers and all ciphertext products at n=15.
n=15;q2=225;units=[r for r in range(1,n) if math.gcd(r,n)==1]
def enc(m,r):return pow(1+n,m,q2)*pow(r,n,q2)%q2
def dec(c):return ((pow(c,4,q2)-1)//n*4)%n
cts=[(m,r,enc(m,r)) for m in range(n) for r in units]
assert all(dec(c)==m for m,r,c in cts);fact('Paillier exhaustive decryption',cases=len(cts))
for m,r,c in cts:
    for mm,rr,cc in cts:assert dec(c*cc%q2)==(m+mm)%n
fact('Paillier exhaustive addition',cases=len(cts)**2)
eq('Paillier worked ciphers',[enc(2,2),enc(4,7),enc(2,2)*enc(4,7)%225],[158,223,134])
eq('Paillier product decrypt',dec(134),6)
for m,r,c in cts:
    assert Counter(c*pow(z,n,q2)%q2 for z in units)==Counter(enc(m,z) for z in units)
fact('Paillier exact rerandomization law',cases=len(cts))
# Regev example and all 32 choices of message/subset.
q=97;s=[3,5];A=[[2,7],[8,1],[4,6],[9,3]];err=[1,-1,2,-2];bs=[(sum(a*x for a,x in zip(row,s))+e)%q for row,e in zip(A,err)]
eq('Regev public b',bs,[42,28,44,40])
for r in product(range(2),repeat=4):
    u=[sum(r[i]*A[i][j] for i in range(4))%q for j in range(2)]
    for mu in range(2):
        v=(sum(r[i]*bs[i] for i in range(4))+mu*48)%q;w=(v-sum(x*y for x,y in zip(s,u)))%q
        assert w==(mu*48+sum(x*y for x,y in zip(r,err)))%q
        assert decode_high(w,q)==mu
fact('Regev all message/subset choices',cases=32)
eq('Regev failing residual',(48+26)%97,74);eq('Regev failing decode',decode_high(74,97),0)
# GSW exact bit decomposition and actual noisy NAND operations.
eq('gadget inner product modulo7',sum(x*y for x,y in zip([1,0,1,1,1,0],[2,4,1,1,2,4]))%7,6)
q=1009;ell=q.bit_length();N=2*ell;t=41;saug=[1,-t];v=[(s0*(1<<j))%q for s0 in saug for j in range(ell)];rng=random.Random(300105)
def bd(row):return [(a%q>>j)&1 for a in row for j in range(ell)]
def invbd(row):return [sum(row[i*ell+j]*(1<<j) for j in range(ell))%q for i in range(2)]
def flat(M):return [bd(invbd(row)) for row in M]
def mv(M,v):return [sum(a*b for a,b in zip(row,v))%q for row in M]
def mm(A,B):return [[sum(a*b for a,b in zip(row,col))%q for col in zip(*B)] for row in A]
B=[3,17,52,99];e0=[1,-1,0,1];pk=[[(a*t+e)%q,a] for a,e in zip(B,e0)]
def genc(mu):
    RR=[[rng.randrange(2) for _ in range(4)] for _ in range(N)];RA=mm(RR,pk);C0=[bd(row) for row in RA]
    for i in range(N):C0[i][i]+=mu
    C=flat(C0);er=[sum(a*b for a,b in zip(row,e0)) for row in RR]
    assert mv(C,v)==[(mu*x+e)%q for x,e in zip(v,er)]
    return C,er
def gdec(C):
    j=next(j for j in range(ell) if q/4<(1<<j)<=q/2);w=mv(C,v)[j];vv=1<<j
    return int(distance(w,vv,q)<distance(w,0,q))
for mu1,mu2 in product(range(2),repeat=2):
    for _ in range(8):
        C1,e1=genc(mu1);C2,e2=genc(mu2);M=mm(C1,C2);raw=[[(int(i==j)-M[i][j])%q for j in range(N)] for i in range(N)];C=flat(raw)
        out=1-mu1*mu2;expected=[-(mu2*e1[i]+sum(a*b for a,b in zip(C1[i],e2))) for i in range(N)]
        assert mv(C,v)==[(out*x+ee)%q for x,ee in zip(v,expected)]
        assert max(abs(x) for x in expected)<=3*(N+1)<q/8
        assert gdec(C)==out
fact('GSW noisy NAND exact relation and decoding',cases=32,N=N,q=q)
eq('GSW article depth bounds',[35**i for i in range(4)],[1,35,1225,42875])
# Relinearization: exact example and all coefficients w2.
eq('relinear product coefficients',[(252*249)%257,(252*3+2*249)%257,2*3],[40,226,6])
eq('relinear output',[(40+8+17)%257,(226+4+7)%257],[65,237]);eq('relinear new phase',(65+237*3)%257,5);eq('capstone base4 relinear output',[(40+2*5+17)%257,(226+2*2+7)%257],[67,237]);eq('capstone base4 phase',(67+237*3)%257,7)
q=257;s=3;ell=q.bit_length();keys=[];eps=[]
for j in range(ell):
    e=(-1)**j;rr=7*j+4;keys.append((((1<<j)*s*s+2*e-rr*s)%q,rr));eps.append(e)
for w2 in range(q):
    digs=[(w2>>j)&1 for j in range(ell)];c0=sum(d*k[0] for d,k in zip(digs,keys))%q;c1=sum(d*k[1] for d,k in zip(digs,keys))%q
    assert (c0+c1*s)%q==(w2*s*s+2*sum(d*e for d,e in zip(digs,eps)))%q
fact('relinear all public quadratic coefficients',cases=q)
# Modulus switching: exact fractions, all 257^2 ciphertext-coordinate pairs at s=2.
q=257;p=67;t=[1,-2]
def switch(c):return [nearest(F(p*x,q)) for x in c]
eq('modswitch worked coordinates',switch([155,10]),[40,3]);eq('modswitch decomposition',F(469-307+95,257),F(1))
eq('modswitch large-secret coordinates',switch([7,140]),[2,36]);eq('modswitch large-secret new phase',(2-34*36)%67,51);eq('modswitch large-secret failure',decode_high(51,67),0)
certified=0
for vv,uu in product(range(q),repeat=2):
    c=[vv,uu];cp=switch(c);phase=vv-2*uu;mu=decode_high(phase%q,q);e=center(phase-mu*(q//2),q);k=(phase-mu*(q//2)-e)//q
    rho=[F(y)-F(p*x,q) for x,y in zip(c,cp)]
    newe=F(p,q)*e+sum(a*b for a,b in zip(rho,t))+mu*(F(p*(q//2),q)-p//2)
    assert F(cp[0]-2*cp[1])==mu*(p//2)+newe+k*p
    bound=F(p,q)*abs(e)+F(3,2)+abs(F(p*(q//2),q)-p//2)
    assert abs(newe)<=bound
    if bound<F(p//2,2):assert decode_high((cp[0]-2*cp[1])%p,p)==mu;certified+=1
fact('modswitch exact identity for all coordinate pairs',cases=q*q,certified_correct_cases=certified)
eq('BGV ordinary-rounding failure',switch([29,10]),[8,3]);eq('BGV parity-preserving phase',7-2*2,3)
# Flooding exact TV and circuit budget.
def tv_shift(M,e):
    a=Counter(range(-M,M+1));b=Counter(range(-M+e,M+1+e));return sum(F(abs(a[x]-b[x]),2*(2*M+1)) for x in a.keys()|b.keys())
for M in range(1,16):
    for e in range(-2*M-1,2*M+2):assert tv_shift(M,e)==min(F(1),F(abs(e),2*M+1))
fact('uniform flooding exact TV formula',cases=sum(4*M+3 for M in range(1,16)))
eq('flooding article bound',tv_shift(100,3),F(1,67))
# Capstone query: Paillier one-hot versus index lookup, all 16 databases x4indices.
for D in product(range(2),repeat=4):
    for i in range(4):
        b1,b0=i//2,i%2;sel=[(1-b1)*(1-b0),(1-b1)*b0,b1*(1-b0),b1*b0];assert sel[i]==1 and sum(sel)==1
        assert sum(a*b for a,b in zip(D,sel))==D[i]
fact('four-record lookup all databases and indices',cases=64)
D=[1,0,1,1];r=[2,4,7,11];query=[enc(int(j==2),r[j]) for j in range(4)];answer=math.prod(c**d for c,d in zip(query,D))%225
eq('capstone Paillier one-hot query',query,[143,199,88,26]);eq('capstone Paillier response',answer,34);eq('capstone Paillier selected answer',dec(answer),1)
bad=[enc(1,rj) for rj in r];eq('capstone malicious all-ones sum',dec(math.prod(c**d for c,d in zip(bad,D))%225),3)
q=65537;threshold=F(q,8);E2=35**2
eq('capstone depth3 naive bound',35*E2,42875);assert 35*E2>=threshold
eq('capstone refresh both',4+34*4,140);eq('capstone refresh right',E2+34*4,1361);eq('capstone refresh left',4+34*E2,41654)
assert 140<threshold and 1361<threshold and 41654>=threshold
inputs=[1,1,1,0,1,0,1,1];layers=[inputs]
while len(layers[-1])>1:layers.append([1-a*b for a,b in zip(layers[-1][::2],layers[-1][1::2])])
eq('capstone balanced NAND clear layers',layers[1:],[[0,1,1,0],[1,1],[0]])
eq('capstone flooding at B=4 M=1000',tv_shift(1000,4),F(4,2001));assert 1000+4<threshold
result={'status':'pass','checks':checks,'check_count':len(checks),'scope':'Exact toy arithmetic, algebraic identities and finite enumeration only. No cryptographic hardness, real parameter security, or real bootstrapping construction is asserted.'}
print(json.dumps(result,ensure_ascii=False,indent=2))
