"""Independent exact verifier A (SymPy) for the power-weighted Keller family.
No numerical approximations are used.
"""
import hashlib
import json
from pathlib import Path
import sympy as sp

x, y, z, s = sp.symbols("x y z s")
ROOT = Path(__file__).resolve().parent


def construct(k: int, d: int):
    assert 1 <= k < d
    v = x**k * y
    t = x**(k + 1) * z
    u = 1 + v
    gamma = 1 - sp.Rational(d + k, d) * v - t
    w = u * gamma
    q_s = sp.Rational(k + 1, d - k) * s**k - sp.Rational(d + 1, d - k) * s**d
    q = q_s.subs(s, w)
    Q = sp.integrate(q_s, (s, 0, w))
    p = sp.Rational(1, k + 1) * (w * q - Q)
    alpha = sp.cancel(p / gamma ** (k + 1) + sp.Rational(1, k + 1) * u)
    beta = sp.cancel(q / gamma**k + 1)
    assert sp.denom(alpha) == 1
    assert sp.denom(beta) == 1
    F = [
        sp.cancel(alpha / x ** (k + 1)),
        sp.cancel(beta / x**k),
        sp.expand(x * gamma),
    ]
    assert all(sp.denom(component) == 1 for component in F)
    F = [sp.expand(component) for component in F]
    return F


def verify_family_member(k: int, d: int):
    F = construct(k, d)
    determinant = sp.factor(sp.Matrix(F).jacobian((x, y, z)).det())
    assert determinant == -sp.Rational(k, k + 1)
    return {
        "k": k,
        "d": d,
        "jacobian": str(determinant),
        "degrees": [sp.Poly(component, x, y, z).total_degree() for component in F],
    }


def verify_uniform_odd_collision(k: int, d: int):
    assert k % 2 == 1 and d % 2 == 1 and k < d
    F = construct(k, d)
    left = {x: -1, y: 0, z: 2}
    right = {x: 1, y: 0, z: 0}
    assert any(left[var] != right[var] for var in (x, y, z))
    image_left = [sp.factor(component.subs(left)) for component in F]
    image_right = [sp.factor(component.subs(right)) for component in F]
    assert image_left == image_right == [0, 0, 1]
    return {"k": k, "d": d, "left": [-1, 0, 2], "right": [1, 0, 0], "image": [0, 0, 1]}


def verify_expanded_degree_eight_specimen():
    F = construct(2, 3)
    a = x**2 * y
    expected = [
        sp.expand(z * (1 + a) ** 4 + x * y**2 * (5 * a**3 + 17 * a**2 + 20 * a + 8) / 3),
        sp.expand(4 * x * z * (1 + a) ** 3 + y * (20 * a**3 + 48 * a**2 + 33 * a + 2) / 3),
        sp.expand(-x**4 * z + x * (3 - 5 * a) / 3),
    ]
    assert all(sp.expand(lhs - rhs) == 0 for lhs, rhs in zip(F, expected))
    determinant = sp.factor(sp.Matrix(expected).jacobian((x, y, z)).det())
    assert determinant == -sp.Rational(2, 3)
    point_1 = {x: sp.Rational(8, 3), y: sp.Rational(-9, 64), z: sp.Rational(495, 4096)}
    point_2 = {x: sp.Rational(8, 3), y: sp.Rational(9, 64), z: sp.Rational(-225, 4096)}
    assert any(point_1[var] != point_2[var] for var in (x, y, z))
    image_1 = [sp.factor(component.subs(point_1)) for component in expected]
    image_2 = [sp.factor(component.subs(point_2)) for component in expected]
    assert image_1 == image_2 == [0, sp.Rational(9, 64), 1]
    expanded = [str(component) for component in expected]
    return {
        "map": expanded,
        "jacobian": str(determinant),
        "point_1": [str(point_1[var]) for var in (x, y, z)],
        "point_2": [str(point_2[var]) for var in (x, y, z)],
        "image": [str(value) for value in image_1],
    }


def main():
    family = [verify_family_member(k, d) for k in range(1, 7) for d in range(k + 1, 9)]
    odd_collisions = [verify_uniform_odd_collision(k, d) for k in range(1, 8, 2) for d in range(k + 2, 10, 2)]
    specimen = verify_expanded_degree_eight_specimen()
    certificate = {
        "engine": "SymPy",
        "sympy_version": sp.__version__,
        "exact_arithmetic": True,
        "family_grid_count": len(family),
        "family_grid": family,
        "odd_collision_count": len(odd_collisions),
        "odd_collisions": odd_collisions,
        "degree_eight_specimen": specimen,
    }
    payload = json.dumps(certificate, sort_keys=True, indent=2)
    (ROOT / "certificate_sympy.json").write_text(payload, encoding="utf-8")
    digest = hashlib.sha256(payload.encode("utf-8")).hexdigest()
    print(json.dumps({"status": "PASS", "certificate": "certificate_sympy.json", "sha256": digest, "family_grid_count": len(family), "odd_collision_count": len(odd_collisions)}, indent=2))


if __name__ == "__main__":
    main()
