#!/usr/bin/env python3
"""Recompute the finite examples in the fault-tolerant quantum unit.
Python 3 standard library only. No simulation is a proof of a general threshold.
"""
import cmath, itertools, json, math
from collections import Counter
from fractions import Fraction
from pathlib import Path
CHECKS=[]
def check(name,condition,details=None):
 if not condition: raise AssertionError(name)
 CHECKS.append({'check':name,'passed':True,'details':details})
def parity(x):return x.bit_count()&1
def span(rows):
 out={0}
 for r in rows:out|={v^r for v in list(out)}
 return out
def rank(rows):return (len(span(rows))).bit_length()-1
def comm(p,q):return parity((p[0]&q[1])^(p[1]&q[0]))
def cx(p,a,b):
 x,z=p
 if x>>a&1:x^=1<<b
 if z>>b&1:z^=1<<a
 return x,z
def hgate(p,a):
 x,z=p
 if ((x^z)>>a)&1:x^=1<<a;z^=1<<a
 return x,z
def cz(p,a,b):
 x,z=p
 z^=((x>>a)&1)<<b;z^=((x>>b)&1)<<a
 return x,z
def fault_paulis(support):
 for digits in itertools.product(range(4),repeat=len(support)):
  if not any(digits):continue
  x=z=0
  for q,d in zip(support,digits):
   if d&1:x|=1<<q
   if d&2:z|=1<<q
  yield x,z

def matmul(a,b):return [[sum(x*y for x,y in zip(row,col)) for col in zip(*b)]for row in a]
def adj(a):return [[v.conjugate() for v in col]for col in zip(*a)]
def kron(a,b):return [[x*y for x in ar for y in br]for ar in a for br in b]
def scale(c,a):return [[c*v for v in r]for r in a]
def close(a,b,tol=1e-12):return all(abs(x-y)<tol for ar,br in zip(a,b)for x,y in zip(ar,br))
I=[[1+0j,0j],[0j,1+0j]];X=[[0j,1+0j],[1+0j,0j]];Z=[[1+0j,0j],[0j,-1+0j]];Y=[[0j,-1j],[1j,0j]]
H=scale(1/math.sqrt(2),[[1+0j,1+0j],[1+0j,-1+0j]]);S=[[1+0j,0j],[0j,1j]]
CNOT=[[1+0j,0j,0j,0j],[0j,1+0j,0j,0j],[0j,0j,0j,1+0j],[0j,0j,1+0j,0j]]
def pauli_mat(x,z,r,n):
 a=[[complex((-1)**r)]]
 for j in range(n):a=kron(a,{(0,0):I,(1,0):X,(0,1):Z,(1,1):Y}[x>>j&1,z>>j&1])
 return a
# Signed tableau updates: qubit 0 is the left tensor factor in these matrices.
for gate,U in [('H',kron(H,I)),('S',kron(S,I)),('CNOT',CNOT)]:
 for x,z,r in itertools.product(range(4),range(4),range(2)):
  xa=x&1;za=z&1;xb=x>>1&1;zb=z>>1&1
  if gate=='H':xn,zn=hgate((x,z),0);rn=r^(xa&za)
  elif gate=='S':xn,zn=x,z^(xa<<0);rn=r^(xa&za)
  else:xn,zn=cx((x,z),0,1);rn=r^(xa&zb&(xb^za^1))
  assert close(matmul(matmul(U,pauli_mat(x,z,r,2)),adj(U)),pauli_mat(xn,zn,rn,2))
check('All signed H/S/CNOT tableau updates',True,{'conjugations':96})
check('Steane transversal S has the inverse logical phase',1j**7==-1j and 1j**4==1)
# Steane code, columns are 1,...,7, low bit first.
Hr=[sum(((j>>k)&1)<<(j-1) for j in range(1,8))for k in range(3)]
D=span(Hr);ker={e for e in range(128) if all(not parity(e&r) for r in Hr)}
check('Steane orthogonality and ranks',all(not parity(a&b)for a in Hr for b in Hr) and rank(Hr)==3,{'n':7,'k':1})
check('Steane exact X/Z logical distance',min(e.bit_count()for e in ker-D)==3,{'row_space_size':len(D),'kernel_size':len(ker)})
check('Steane all-ones logical representatives',127 in ker-D and parity(127)==1 and (127^7) in D)
def syn(p):
 x,z=p
 return tuple(parity(z&r)for r in Hr)+tuple(parity(x&r)for r in Hr)
def recovery(s):
 iz=sum(s[k]<<k for k in range(3));ix=sum(s[k+3]<<k for k in range(3))
 return (1<<(ix-1) if ix else 0,1<<(iz-1) if iz else 0)
check('Recovery table covers all 64 syndromes',all(syn(recovery(s))==s for s in itertools.product([0,1],repeat=6)))
singles=[(0,0)]+[(1<<j,0)for j in range(7)]+[(0,1<<j)for j in range(7)]+[(1<<j,1<<j)for j in range(7)]
check('All 21 Steane single Paulis corrected',all(recovery(syn(e))==e for e in singles))
# Every single location fault in a concrete cat schedule.
# Five initializations, H, three chain CNOTs, two verifier CNOTs, measurement.
cat_ops=[('prep',(q,))for q in range(5)]+[('H',(0,)),('cx',(0,1)),('cx',(1,2)),('cx',(2,3)),('cx',(0,4)),('cx',(3,4)),('measure',(4,))]
def propagate(p,ops):
 for op,s in ops:
  if op=='cx':p=cx(p,*s)
  elif op=='H':p=hgate(p,*s)
  elif op=='cz':p=cz(p,*s)
 return p
cat_cases=[]
for t,(op,support) in enumerate(cat_ops):
 # At measurement, an X on the measured qubit represents a flipped result.
 # Each initialized inactive wire has an idle, including during measurement.
 candidates=[('operation',support)]
 initialized=set(range(min(t+1,5)))
 candidates += [('idle',(q,))for q in initialized if q not in support]
 for kind,sup in candidates:
  for p in fault_paulis(sup):
   x,z=propagate(p,cat_ops[t+1:]);accepted=not(x>>4&1)
   catx=x&15;effective_x_weight=min(catx.bit_count(),4-catx.bit_count())
   assert not accepted or effective_x_weight<=1
   cat_cases.append((t,kind,p,accepted,effective_x_weight))
check('Every scheduled single cat fault passes safe or is rejected',True,{'cases':len(cat_cases),'accepted':sum(c[3]for c in cat_cases),'rejected':sum(not c[3]for c in cat_cases)})
# Dangerous pair explicitly arises after the middle CNOT and is caught.
x,z=propagate((1<<2,0),cat_ops[8:]) # after op 7: CNOT a2 -> a3
check('Cat X3X4 is detected by the endpoint verifier',(x&15)==12 and bool(x>>4&1),{'final_x_mask':x})
# Preparation faults accepted by verification, then all four disjoint CZs.
for t,kind,p,accepted,w in cat_cases:
 if not accepted:continue
 x,z=propagate(p,cat_ops[t+1:]);x&=15;z&=15
 for j in range(4):x,z=cz((x,z),j,5+j)
 zd=(z>>5)&15
 assert min(zd.bit_count(),4-zd.bit_count())<=1
coupling_cases=0
coupling=[('cz',(j,5+j))for j in range(4)]
for t,(_,s) in enumerate(coupling):
 for p in fault_paulis(s):
  x,z=propagate(p,coupling[t+1:]);support=((x|z)>>5)&15
  assert support.bit_count()<=1
  coupling_cases+=1
check('Accepted cat propagation and every faulty cat-data CZ are safe',True,{'CZ_fault_cases':coupling_cases})
# Four-round full-syndrome rule. A single faulty check can change one data
# Pauli AND flip its reported bit, counted together as ONE fault.
def select_rounds(records):
 rounds=[tuple(records[6*r:6*r+6])for r in range(4)]
 for a,b in zip(rounds,rounds[1:]):
  if a==b:return a
 raise AssertionError('No consecutive agreement under the one-fault promise')
ec_cases=0
for t in range(24):
 for e in singles:
  for flip in range(2):
   data=(0,0);records=[]
   for k in range(24):
    bit=syn(data)[k%6]
    if k==t:bit^=flip;data=e
    records.append(bit)
   chosen=select_rounds(records);r=recovery(chosen)
   residual=(data[0]^r[0],data[1]^r[1])
   assert (residual[0]|residual[1]).bit_count()<=1
   ec_cases+=1
check('Four complete syndrome rounds: one correlated data/readout fault',True,{'cases':ec_cases,'rounds':4,'checks_per_round':6})
# Arbitrary input Pauli sector: the chosen syndrome occurred either before or
# after the one fault. This certifies distance from SOME codeword, not recovery
# of the initially intended logical state for arbitrary multi-errors.
arbitrary_cases=0
all_syn=list(itertools.product([0,1],repeat=6))
for s0 in all_syn:
 for t in range(24):
  for e in singles:
   s1=tuple(a^b for a,b in zip(s0,syn(e)))
   for flip in range(2):
    records=[s0[k%6] if k<=t else s1[k%6]for k in range(24)];records[t]^=flip
    chosen=select_rounds(records)
    assert chosen==s0 or chosen==s1
    arbitrary_cases+=1
check('Arbitrary input syndrome sectors choose an actually occurring syndrome',True,{'cases':arbitrary_cases})
# Toric L=3. Original and dual incidence ranks; explicit logicals.
L=3
def hh(r,c):return (r%L)*L+c%L
def vv(r,c):return L*L+(r%L)*L+c%L
def bits(indices):
 out=0
 for j in indices:out^=1<<j
 return out
stars=[bits([hh(r,c),hh(r,c-1),vv(r,c),vv(r-1,c)])for r in range(L)for c in range(L)]
faces=[bits([hh(r,c),vv(r,c+1),hh(r+1,c),vv(r,c)])for r in range(L)for c in range(L)]
check('Toric L3 commutation and independent checks',all(not parity(a&b)for a in stars for b in faces) and rank(stars)==rank(faces)==8,{'n':18,'k':2})
zl=[bits(hh(0,c)for c in range(L)),bits(vv(r,0)for r in range(L))];xl=[bits(hh(r,0)for r in range(L)),bits(vv(0,c)for c in range(L))]
check('Toric two pairs of logical Paulis',all(parity(xl[i]&zl[j])==(i==j)for i in range(2)for j in range(2)))
E=1<<hh(0,0);R=bits([hh(0,1),hh(0,2)])
check('Same toric syndrome, different homology',all(parity(s&E)==parity(s&R)for s in stars) and E^R==zl[0] and E^R not in span(faces))
for w in [1,2]:
 for support in itertools.combinations(range(18),w):
  e=bits(support)
  assert any(parity(e&s)for s in stars) or e in span(faces)
  assert any(parity(e&f)for f in faces) or e in span(stars)
check('Toric L3 has no weight1/2 nontrivial zero-syndrome X/Z',True)
# Bacon-Shor, horizontally XX and vertically ZZ.
def pos(r,c):return r*3+c
gx=[(bits([pos(r,c),pos(r,c+1)]),0)for r in range(3)for c in range(2)]
gz=[(0,bits([pos(r,c),pos(r+1,c)]))for r in range(2)for c in range(3)]
sx=[(bits(pos(r,j)for r in range(3)for j in [c,c+1]),0)for c in range(2)]
sz=[(0,bits(pos(j,c)for c in range(3)for j in [r,r+1]))for r in range(2)]
gauge=gx+gz;st=sx+sz
gr=rank([x|(z<<9)for x,z in gauge]);sr=rank([x|(z<<9)for x,z in st])
check('Bacon-Shor gauge and stabilizer ranks',gr==12 and sr==4 and all(not comm(s,g)for s in st for g in gauge),{'physical':9,'logical':9-(gr+sr)//2,'gauge_qubits':(gr-sr)//2})
lx=(bits(pos(r,0)for r in range(3)),0);lz=(0,bits(pos(0,c)for c in range(3)))
check('Bacon-Shor bare logicals commute with every gauge generator',all(not comm(l,g)for l in [lx,lz]for g in gauge) and comm(lx,lz)==1)
check('Bacon-Shor individual gauges need not commute',comm(gx[0],gz[0])==1)
# Injection branch matrices, including a Z-corrupted resource.
omega=cmath.exp(1j*math.pi/4);T=[[1+0j,0j],[0j,omega]]
for bad in [0,1]:
 a=1/math.sqrt(2);b=(-1)**bad*omega/math.sqrt(2)
 K0=[[a,0j],[0j,b]];K1=[[b,0j],[0j,a]]
 target=matmul(Z,T) if bad else T
 assert close(K0,scale(1/math.sqrt(2),target))
 assert close(matmul(S,K1),scale(((-1)**bad)*omega/math.sqrt(2),target))
 assert close(matmul(adj(K0),K0),scale(.5,I)) and close(matmul(adj(K1),K1),scale(.5,I))
check('Injection both branches and Z-resource error channel',True,{'branches_checked':4})
# RM15 construction and exhaustive all32768 phase-error patterns.
H15=[sum(((j>>k)&1)<<(j-1)for j in range(1,16))for k in range(4)]
D15=span(H15);G15=span(H15+[32767])
check('RM15 row spaces and transverse T phases',len(D15)==16 and len(G15)==32 and {e.bit_count()for e in D15}=={0,8} and {(32767^e).bit_count()for e in D15}=={7,15})
accepted=Counter();odd=Counter();triples=[]
for e in range(32768):
 if all(not parity(e&r)for r in H15):
  w=e.bit_count();accepted[w]+=1
  if w&1:odd[w]+=1
  if w==3:triples.append(e)
check('RM15 exhaustive acceptance and logical-error classification',sum(accepted.values())==2048 and sum(odd.values())==1024 and accepted[1]==accepted[2]==0 and odd[3]==35,{'accepted_weight_enumerator':dict(sorted(accepted.items())),'odd_weight_enumerator':dict(sorted(odd.items()))})
check('All 35 correlated triples have equal marginals',all(sum(e>>j&1 for e in triples)==7 for j in range(15)))
numerics=[]
for p in [Fraction(1,100),Fraction(14,100),Fraction(15,100)]:
 t=1-2*p;acc=(1+15*t**8)/16;pout=(1-15*t**7+15*t**8-t**15)/(2*(1+15*t**8))
 brute_acc=sum(count*p**w*(1-p)**(15-w)for w,count in accepted.items())
 brute_odd=sum(count*p**w*(1-p)**(15-w)for w,count in odd.items())
 assert acc==brute_acc and pout==brute_odd/acc
 numerics.append({'p':float(p),'acceptance':float(acc),'p_out':float(pout),'expected_inputs':float(15/acc)})
check('Exact rational distillation formulas equal exhaustive enumerators',True,numerics)
A=100;ps=[Fraction(1,1000)]
for _ in range(3):ps.append(A*ps[-1]**2)
check('Threshold teaching recursion',ps==[Fraction(1,1000),Fraction(1,10000),Fraction(1,10**6),Fraction(1,10**10)] and 10**6*ps[3]<Fraction(1,100),{'p_bounds':[float(x)for x in ps],'million_location_union_bound':float(10**6*ps[3]),'scope':'Assumes the proved recurrence; not a hardware threshold or a proof by enumeration.'})
# Fixed retry cap integrates resources with the location-bound promise.
p=Fraction(1,100);t=1-2*p;acc=(1+15*t**8)/16
out=(1-15*t**7+15*t**8-t**15)/(2*(1+15*t**8))
exhaust=(1-acc)**10;total=Fraction(1,10000)+100*exhaust+100*out
check('Capstone bounded factory and total failure budget',total<Fraction(1,100),{'attempts_per_output':10,'outputs':100,'max_raw_states':15000,'exhaustion_union_bound':float(100*exhaust),'accepted_resource_error_union_bound':float(100*out),'device_union_bound':.0001,'total_failure_bound':float(total)})
RESULT={'unit':'e-fault-tolerant-quantum','passed':True,'check_count':len(CHECKS),'checks':CHECKS}
if __name__=='__main__':
 print(json.dumps(RESULT,ensure_ascii=False,indent=2))
