"""D1258 C01.2. Declarative object: derivations of the stated CD algebra."""

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 derivation_system(epsilon):
    table=[[omul(x,y,epsilon) for y in OB] for x in OB]
    # D_{row,column}: D acts on column vectors.
    rows=[]
    for i,j,k in product(range(8), repeat=3):
        eq=[0]*64
        for m in range(8):
            eq[k*8+m]+=table[i][j][m]
            eq[m*8+i]-=table[m][j][k]
            eq[m*8+j]-=table[i][m][k]
        rows.append(eq)
    A=S.Matrix(rows)
    ns=A.nullspace()
    ds=[S.Matrix(8,8,list(v)) for v in ns]
    for D in ds:
        assert A*S.Matrix(list(D))==S.zeros(A.rows,1)
    return A,ds
def flat(D): return S.Matrix(list(D))
def coordinates(ds):
    B=S.Matrix.hstack(*[flat(D) for D in ds])
    ix=list(B.T.rref()[1]); C=B.extract(ix,range(B.cols)).inv()
    def coord(M):
        z=C*flat(M).extract(ix,[0])
        assert B*z==flat(M)
        return z
    return coord
report("environment", {"python":platform.python_version(),"sympy":S.__version__})
report("data_basis","Declared CD definition; no banked arrays; all arithmetic exact.")
for eps,label in [(1,"split"),(-1,"compact")]:
    A,ds=derivation_system(eps); d=len(ds)
    coord=coordinates(ds)
    norm=S.diag(*[onorm(e,eps) for e in OB])
    assert all(D[:,0]==S.zeros(8,1) for D in ds)
    assert all(D.T*norm+norm*D==S.zeros(8) for D in ds)
    ads=[S.Matrix.hstack(*[coord(D*E-E*D) for E in ds]) for D in ds]
    K=S.Matrix(d,d,lambda i,j:S.trace(ads[i]*ads[j]))
    rep=S.Matrix(d,d,lambda i,j:S.trace(ds[i]*ds[j]))
    assert K==4*rep
    theta=S.Matrix.hstack(*[coord(-D.T) for D in ds])
    assert theta**2==S.eye(d)
    kdim=len((theta-S.eye(d)).nullspace())
    pdim=len((theta+S.eye(d)).nullspace())
    report(label+".derivation_system_shape",A.shape)
    report(label+".derivation_rank_nullity",(64-d,d))
    report(label+".unit_fixed_and_norm_skew",True)
    report(label+".closed_commutator_pairs",d*d)
    report(label+".Killing_equals_4_trace8",True)
    report(label+".Killing_inertia(+,-,0)",inertia(K))
    report(label+".Cartan_dimensions(k,p)",(kdim,pdim))
    report(label+".norm_inertia(+,-,0)",inertia(norm))
    if eps==1:
        split_ds=ds
        # A split Cartan directly from native Zorn scaling u->H u,v->-H v.
        C=S.Matrix.hstack(*[S.Matrix(to_zorn(e)) for e in OB])
        Hs=[]
        for h in [(1,-1,0),(0,1,-1)]:
            Hs.append(C.inv()*S.diag(0,*h,*[-a for a in h],0)*C)
        adsH=[S.Matrix.hstack(*[coord(H*D-D*H) for D in ds]) for H in Hs]
        weights=[]
        for v1 in range(-3,4):
            for v2 in range(-3,4):
                m=len((adsH[0]-v1*S.eye(d)).col_join(adsH[1]-v2*S.eye(d)).nullspace())
                if m: weights.append(((v1,v2),m))
        assert sum(m for _,m in weights)==d
        report("split.Cartan_joint_weights",weights)
        report("split.Cartan_centralizer_dimension",dict(weights)[(0,0)])
        KC=S.Matrix(2,2,lambda i,j:4*S.trace(Hs[i]*Hs[j]))
        lens=Counter()
        for w,m in weights:
            if w!=(0,0):lens[(S.Matrix([w])*KC.inv()*S.Matrix(w))[0]]+=m
        report("split.Cartan_Killing_matrix",KC.tolist())
        report("split.root_squared_lengths_multiplicities",sorted(lens.items()))
# Negative control: norm-preserving linear map need not be a derivation.
D=S.diag(0,0,0,0,0,0,0,0)
D[1,2]=1; D[2,1]=-1
assert D.T*S.diag(1,1,1,1,-1,-1,-1,-1)+S.diag(1,1,1,1,-1,-1,-1,-1)*D==S.zeros(8)
def leibniz(D,x,y):
    return vs(tuple(D*S.Matrix(omul(x,y))),
              va(omul(tuple(D*S.Matrix(x)),y),omul(x,tuple(D*S.Matrix(y)))))
w=next((i,j,leibniz(D,OB[i],OB[j])) for i,j in product(range(8),repeat=2)
       if not zero_vector(leibniz(D,OB[i],OB[j])))
report("negative_control.norm_skew_not_derivation",w)
# Quaternion automorphism family F_{a,b}(p,q)=(a p a^-1,b q a^-1).
units=[(S.Integer(1),0,0,0),(0,S.Integer(1),0,0),(0,0,S.Integer(1),0),(0,0,0,S.Integer(1))]
units+=[tuple(-t for t in u) for u in units.copy()]
def F(x,a,b):
    p=x[:4]; q=(x[7],)+vn(x[4:7])
    P=qm(qm(a,p),qb(a)); Q=qm(qm(b,q),qb(a))
    return P+vn(Q[1:])+(Q[0],)
tests=0; kernels=[]
for a,b in product(units,repeat=2):
    if all(F(e,a,b)==e for e in OB):kernels.append((a,b))
    for x,y in product(OB,repeat=2):
        assert F(omul(x,y),a,b)==omul(F(x,a,b),F(y,a,b)); tests+=1
report("quaternion_action.finite_basis_product_checks",tests)
report("quaternion_action.finite_sample_kernel",kernels)
report("dimensions.half_pair",(4,8))
count=0
for x,y in product(OB,repeat=2):
    assert zero_vector(vs(from_zorn(zmul(to_zorn(x),to_zorn(y))),omul(x,y)))
    count+=1
report("CD_Zorn.all_basis_pairs",count)
report("source_status","No labels changed; full banked embedding alignment NOT VERIFIED.")
