"""D1258 C02.2. Peirce multiplication and the declared Roman map."""

from itertools import product, combinations, permutations
from collections import Counter
import json
import platform
import sympy as S
if not __debug__:
    raise SystemExit("Do not use -O: assertions are required.")
def report(key, value):
    print(key + " = " + str(value))
def va(x,y): return tuple(a+b for a,b in zip(x,y))
def vn(x): return tuple(-a for a in x)
def vs(x,y): return va(x,vn(y))
def sc(k,x): return tuple(k*a for a in x)
def dot(x,y): return sum(a*b for a,b in zip(x,y))
def cross(x,y):
    return (x[1]*y[2]-x[2]*y[1],x[2]*y[0]-x[0]*y[2],x[0]*y[1]-x[1]*y[0])
def qm(p,r):
    a,b,c,d=p; e,f,g,h=r
    return (a*e-b*f-c*g-d*h,a*f+b*e+c*h-d*g,
            a*g-b*h+c*e+d*f,a*h+b*g-c*f+d*e)
def qb(p): return (p[0],-p[1],-p[2],-p[3])
O_NAMES=("1","e1","e2","e3","f1","f2","f3","l")
OB=[tuple(S.Integer(i==j) for i in range(8)) for j in range(8)]
OZ=(S.Integer(0),)*8
def obar(x): return (x[0],)+vn(x[1:])
def omul(x,y,epsilon=1):
    # Same basis as D1256: f_i=-e_i*l, not +e_i*l.
    p=x[:4]; q=(x[7],)+vn(x[4:7])
    r=y[:4]; s=(y[7],)+vn(y[4:7])
    first=va(qm(p,r),sc(epsilon,qm(qb(s),q)))
    second=va(qm(s,p),qm(q,qb(r)))
    return first+vn(second[1:])+(second[0],)
def onorm(x,epsilon=1):
    return sum(a*a for a in x[:4])-epsilon*sum(a*a for a in x[4:])
def to_zorn(x):
    return (x[0]+x[7],)+va(x[1:4],x[4:7])+vs(x[4:7],x[1:4])+(x[0]-x[7],)
def from_zorn(z):
    a=z[0]; u=z[1:4]; v=z[4:7]; b=z[7]
    return ((a+b)/2,)+sc(S.Rational(1,2),vs(u,v))+sc(S.Rational(1,2),va(u,v))+((a-b)/2,)
def zmul(z,w):
    a=z[0]; u=z[1:4]; v=z[4:7]; b=z[7]
    c=w[0]; U=w[1:4]; V=w[4:7]; d=w[7]
    return (a*c+dot(u,V),)+va(va(sc(a,U),sc(d,u)),cross(v,V))+vs(va(sc(c,v),sc(b,V)),cross(u,U))+(dot(v,U)+b*d,)
def zero_vector(v): return all(S.expand(a)==0 for a in v)
def inertia(M):
    """Exact rational symmetric congruence elimination; no eigenvalue tolerance."""
    A=S.Matrix(M); assert A==A.T
    positive=negative=null=0
    while A.rows:
        k=next((i for i in range(A.rows) if A[i,i]!=0),None)
        if k is None:
            pair=next(((i,j) for i in range(A.rows) for j in range(i+1,A.rows) if A[i,j]!=0),None)
            if pair is None:
                null+=A.rows; break
            i,j=pair
            P=S.eye(A.rows); P[j,i]=1
            A=P.T*A*P
            k=i
        inds=[k]+[i for i in range(A.rows) if i!=k]
        A=A.extract(inds,inds); d=A[0,0]
        assert d.is_positive or d.is_negative
        positive+=int(bool(d>0)); negative+=int(bool(d<0))
        v=A[1:,0]; A=A[1:,1:]-(v*v.T)/d
    return (positive,negative,null)

def jmat(A):
    a,b,c=A[:3]; z=A[3:11]; y=A[11:19]; x=A[19:27]
    return [[sc(a,OB[0]),z,obar(y)],[obar(z),sc(b,OB[0]),x],[y,obar(x),sc(c,OB[0])]]
def jvec(M):
    assert all(zero_vector(M[i][i][1:]) for i in range(3))
    assert all(zero_vector(vs(M[i][j],obar(M[j][i]))) for i in range(3) for j in range(i+1,3))
    return tuple(M[i][i][0] for i in range(3))+tuple(M[0][1])+tuple(M[2][0])+tuple(M[1][2])
def mm(A,B):
    return [[va(va(omul(A[i][0],B[0][j]),omul(A[i][1],B[1][j])),omul(A[i][2],B[2][j]))
             for j in range(3)] for i in range(3)]
def jp(A,B):
    X=jmat(A); Y=jmat(B); XY=mm(X,Y); YX=mm(Y,X)
    return jvec([[sc(S.Rational(1,2),va(XY[i][j],YX[i][j])) for j in range(3)] for i in range(3)])
JB=[tuple(S.Integer(i==j) for i in range(27)) for j in range(27)]
JI=va(va(JB[0],JB[1]),JB[2]); JZ=(S.Integer(0),)*27
def tr(A): return sum(A[:3])
def sig(A):
    a,b,c=A[:3];z=A[3:11];y=A[11:19];x=A[19:27]
    return a*b+a*c+b*c-onorm(x)-onorm(y)-onorm(z)
def jnorm(A):
    a,b,c=A[:3];z=A[3:11];y=A[11:19];x=A[19:27]
    return a*b*c-a*onorm(x)-b*onorm(y)-c*onorm(z)+2*omul(omul(z,x),y)[0]
def sharp(A): return va(vs(jp(A,A),sc(tr(A),A)),sc(sig(A),JI))
def L(A): return S.Matrix.hstack(*[S.Matrix(jp(A,e)) for e in JB])
def brief(A): return {i:S.simplify(a) for i,a in enumerate(A) if S.simplify(a)!=0}

report("environment", {"python":platform.python_version(),"sympy":S.__version__})
report("data_basis","Declared H3(CD-split), diagonal adjoint and unit sphere.")
a,b,c=S.symbols("a b c")
D=(a,b,c)+(S.Integer(0),)*24
M=L(D)
expect=[a,b,c]+[(a+b)/2]*8+[(a+c)/2]*8+[(b+c)/2]*8
assert M==S.diag(*expect)
report("Peirce.dimensions",(3,8,8,8))
report("Peirce.diagonal_L_eigenvalues_multiplicities",[(a,1),(b,1),(c,1),((a+b)/2,8),((a+c)/2,8),((b+c)/2,8)])
assert all(jp(JB[i],JB[i])==JB[i] for i in range(3))
assert all(jp(JB[i],JB[j])==JZ for i in range(3) for j in range(i+1,3))
blocks=[range(3,11),range(11,19),range(19,27)]
ends=[(0,1),(0,2),(1,2)]
same=cross_checks=0
for block,(i,j) in zip(blocks,ends):
    for ix,iy in product(block,repeat=2):
        x=OB[ix-block.start]; y=OB[iy-block.start]
        inner=(onorm(va(x,y))-onorm(x)-onorm(y))/2
        assert jp(JB[ix],JB[iy])==sc(inner,va(JB[i],JB[j]));same+=1
for p,q in combinations(range(3),2):
    target=set(blocks[3-p-q])
    for ix,iy in product(blocks[p],blocks[q]):
        v=jp(JB[ix],JB[iy])
        assert all(v[k]==0 for k in range(27) if k not in target);cross_checks+=1
report("Peirce.same_block_product_checks",same)
report("Peirce.cross_block_support_checks",cross_checks)
report("diagonal.adjoint",tuple(S.expand(v) for v in sharp(D)[:3]))
assert all(S.expand(x-y)==0 for x,y in zip(sharp(sharp(D)),sc(jnorm(D),D)))
report("diagonal.double_adjoint_identity",True)
x,y,z=S.symbols("x y z",real=True)
U=S.Matrix([x,y,z]);T=S.Matrix([y*z,x*z,x*y]);dt=T.jacobian(U)
sphere=x*x+y*y+z*z-1
X,Y,Z=S.symbols("X Y Z")
F=X*X*Y*Y+Y*Y*Z*Z+Z*Z*X*X-X*Y*Z
pull=S.expand(F.subs({X:T[0],Y:T[1],Z:T[2]}))
assert S.expand(pull-x*x*y*y*z*z*sphere)==0
report("Roman.quartic_pullback",S.factor(pull))
assert T.subs({x:-x,y:-y,z:-z},simultaneous=True)==T
perm_count=0
for p in permutations(range(3)):
    assert T.subs(dict(zip(U,[U[i] for i in p])),simultaneous=True)==T.extract(p,[0])
    perm_count+=1
report("Roman.antipodal_identity",True)
report("Roman.quartic_axis_counterexample",F.subs({X:1,Y:0,Z:0}))
axis_quadratic=S.hessian(y*z,(y,z))/2
report("Roman.unit_sphere_axis_endpoint_bound",max(axis_quadratic.eigenvals()))
report("Roman.permutation_equivariance_checks",perm_count)
# Restricted differential: augmented rank<3 iff tangent differential rank<2.
pinch=[]
for k in range(3):
    j=[i for i in range(3) if i!=k]
    for s,t in product([-1,1],repeat=2):
        v=[S.Integer(0)]*3;v[j[0]]=s/S.sqrt(2);v[j[1]]=t/S.sqrt(2)
        subs=dict(zip(U,v))
        aug=dt.col_join(U.T).subs(subs)
        assert aug.rank()==2
        pinch.append(tuple(T.subs(subs)))
report("Roman.pinch_preimages",len(pinch))
report("Roman.pinch_images",sorted(set(pinch),key=str))
# A non-pinch point on a double-line axis has two RP2 preimages.
p=(0,S.Rational(3,5),S.Rational(4,5));q=(0,S.Rational(4,5),S.Rational(3,5))
assert T.subs(dict(zip(U,p)))==T.subs(dict(zip(U,q)))
assert p!=q and p!=tuple(-v for v in q)
report("Roman.distinct_projective_fibre_witness",(p,q,tuple(T.subs(dict(zip(U,p))))))
report("Roman.witness_tangent_rank",dt.col_join(U.T).subs(dict(zip(U,p))).rank()-1)
phi=(1+S.sqrt(5))/2;J=S.Matrix([phi,1,phi-1])
Js=S.simplify(T.subs(dict(zip(U,J))))
unitmap=S.simplify(T.subs(dict(zip(U,J/2))))
report("golden.norm_squared",S.simplify(J.dot(J)))
report("golden.adjoint",list(Js))
report("golden.unit_sphere_image",list(unitmap))
assert S.simplify(unitmap-Js/4)==S.zeros(3,1)
report("golden.unit_sphere_image_equals_adjoint_over4",True)
report("spinor_bundle_identification","NOT VERIFIED; no such map is constructed.")
