"""D1259 C03.1. Declared split-Albert FTS boundary and transverse-velocity tests."""

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("object", "standard split-Albert Freudenthal coordinates; NOT frozen E7 arrays")
a,b,t,r=S.symbols("a b t r", real=True, nonzero=True)
xx=S.symbols("x0:27", real=True); yy=S.symbols("y0:27", real=True)
metric=S.diag(1,1,1,*([2]*4+[-2]*4)*3)
Nx=S.expand(jnorm(xx)); Ny=S.expand(jnorm(yy))
gx=S.Matrix([S.diff(Nx,v) for v in xx])
gy=S.Matrix([S.diff(Ny,v) for v in yy])
sx=S.Matrix(sharp(xx))
assert all(S.expand(v)==0 for v in metric*sx-gx)
sy=metric.inv()*gy
Txy=(S.Matrix(xx).T*metric*S.Matrix(yy))[0]
Tsharp=(sx.T*metric*sy)[0]
I4=-(a*b-Txy)**2-4*a*Nx-4*b*Ny+4*Tsharp
# Simultaneously scale all conjugate coordinates by t. This is an exact
# polynomial coefficient test in the complete transverse direction (b,y).
transverse_map={b:t*b, **{y:t*y for y in yy}}
It=I4.subs(transverse_map, simultaneous=True).expand()
Ip=S.Poly(It,t)
assert S.expand(Ip.coeff_monomial(1)+4*a*Nx)==0
assert Ip.coeff_monomial(t)==0
assert S.expand(Ip.coeff_monomial(t**2)+(a*b-Txy)**2-4*Tsharp)==0
assert S.expand(Ip.coeff_monomial(t**4)+4*b*Ny)==0
assert set(k[0] for k in Ip.monoms()) <= {0,2,4}
report("jordan_coordinate_count",len(xx))
report("norm_polynomial_monomials",len(S.Poly(Nx,*xx).terms()))
report("gradient_equals_trace_metric_times_adjoint",True)
report("transverse_polynomial_degrees",sorted(k[0] for k in Ip.monoms()))
report("I4_on_mass_boundary","-4*a*N(x)")
report("all_first_conjugate_derivatives_on_boundary",0)
# q=(a,x), canonical covector p=(b, metric*y), not raw y coordinates.
I=S.eye(28); Z=S.zeros(28); Omega=Z.row_join(I).col_join((-I).row_join(Z))
E=I.col_join(Z)
assert E.T*Omega*E==S.zeros(28)
K=(E.T*Omega).nullspace()
Kmat=S.Matrix.hstack(*K)
assert Kmat[28:,:]==S.zeros(28)
assert Kmat.rank()==E.rank()==28 and Omega.det()==1
report("FTS_dimension",Omega.rows)
report("boundary_rank_and_symplectic_perp_dimension",(E.rank(),len(K)))
report("boundary_pullback_Omega_zero",True)
# The quartic has weights zero for a:-6,x:+2,b:+6,y:-2.
for monom,_ in S.Poly(Nx,*xx).terms(): assert sum(monom)==3
for v in sx:
    assert all(sum(monom)==2 for monom,_ in S.Poly(S.expand(v),*xx).terms())
weights=[-6]+[2]*27+[6]+[-2]*27
assert all(weights[i]+weights[28+i]==0 for i in range(28))
assert -6+3*2==6+3*(-2)==2*2+2*(-2)==0
report("grading_weights_a_x_b_y",(-6,2,6,-2))
report("symplectic_and_quartic_grading_test",True)
# A nonzero boundary Hamiltonian vector, despite vanishing q velocity.
# Convention: qdot=partial_p H, pdot=-partial_q H.
witness={a:S.Integer(1), **{xx[i]:JI[i] for i in range(27)}}
qdot=S.zeros(28,1)
pdot=S.Matrix([4*Nx]+[4*a*v for v in gx]).subs(witness)
assert list(pdot[:4])==[4,4,4,4]
assert all(v==0 for v in pdot[4:])
report("witness_boundary_H",S.simplify((-4*a*Nx).subs(witness)))
report("witness_qdot_nonzero_entries",{})
report("witness_pdot_nonzero_entries",brief(pdot))
report("full_Hamiltonian_vector_is_zero",False)
report("registered_E7_solder_and_local_polarization","NOT VERIFIED: original arrays absent")
report("result","PASS: stated exact coordinate assertions only")
