#!/usr/bin/env python3
"""Exact audit for Chapter 7's two-singular-loci benchmark.

Dependencies:
    Python 3.10+
    sympy 1.13+
    mpmath 1.3+

The script verifies:

1. the Frobenius pole and obstruction at theta = 2;
2. an exact resonant but log-free ODE pair;
3. the c = 28 level-two Kac determinant;
4. two nonzero Virasoro-block residues; and
5. their high-precision numerical limits.
"""

from __future__ import annotations

import mpmath as mp
import sympy as sp


def require_zero(expression: sp.Expr, label: str) -> None:
    """Raise a visible error if an exact symbolic residual is nonzero."""
    residual = sp.simplify(expression)
    if residual != 0:
        raise RuntimeError(f"{label} failed: residual = {residual}")


def symbolic_audit() -> dict[str, sp.Expr]:
    """Return the exact identities after raising on any failed check."""
    theta, lam, mu, x = sp.symbols("theta lambda mu x")

    a_1 = lam / (theta - 1)
    a_2 = sp.factor((lam * a_1 + mu) / (2 * (theta - 2)))
    obstruction = lam**2 + mu
    frobenius_residue = sp.limit((theta - 2) * a_2, theta, 2)
    require_zero(
        frobenius_residue - obstruction / 2,
        "Frobenius residue",
    )
    larger_b_1 = -lam / 3
    larger_b_2 = (lam**2 - 3 * mu) / 24
    logarithmic_coefficient = sp.factor(
        -2 * (3 * larger_b_1**2 - 2 * larger_b_2)
    )
    require_zero(
        logarithmic_coefficient + obstruction / 2,
        "logarithmic coefficient",
    )

    operator = lambda y: sp.simplify(
        sp.diff(y, x, 2)
        - sp.diff(y, x) / x
        + (lam / x - lam**2) * y
    )
    y_0 = sp.exp(lam * x)
    y_2 = (
        sp.exp(lam * x)
        - (1 + 2 * lam * x) * sp.exp(-lam * x)
    ) / (2 * lam**2)
    require_zero(operator(y_0), "first apparent solution")
    require_zero(operator(y_2), "second apparent solution")
    require_zero(
        sp.limit(y_2 / x**2, x, 0) - 1,
        "second-solution normalization",
    )
    require_zero(
        sp.limit(y_2, lam, 0) - x**2,
        "continuous lambda-zero extension",
    )

    delta = sp.symbols("Delta")
    gram = sp.Matrix(
        [
            [4 * delta + 14, 6 * delta],
            [6 * delta, 4 * delta * (2 * delta + 1)],
        ]
    )
    gram_determinant = sp.factor(gram.det())
    expected_determinant = 4 * delta * (delta + 2) * (8 * delta + 7)
    require_zero(
        gram_determinant - expected_determinant,
        "level-two Gram determinant",
    )

    a = delta + 1
    b = delta - 1
    p = delta + 3
    q = delta + 2
    gamma_left = sp.Matrix([p, a * (a + 1)])
    gamma_right = sp.Matrix([[q, b * (b + 1)]])
    block_level_two = sp.factor(
        (gamma_right * gram.inv() * gamma_left)[0]
    )

    residue_21 = sp.limit(
        (delta + 2) * block_level_two,
        delta,
        -2,
    )
    residue_12 = sp.limit(
        (delta + sp.Rational(7, 8)) * block_level_two,
        delta,
        -sp.Rational(7, 8),
    )
    require_zero(residue_21 - 1, "(2,1) Kac residue")
    require_zero(
        residue_12 + sp.Rational(3619, 4096),
        "(1,2) Kac residue",
    )

    return {
        "frobenius_residue": frobenius_residue,
        "logarithmic_coefficient": logarithmic_coefficient,
        "gram_determinant": gram_determinant,
        "block_level_two": block_level_two,
        "residue_21": residue_21,
        "residue_12": residue_12,
    }


def numerical_residue_table() -> list[tuple[int, mp.mpf, mp.mpf, mp.mpf]]:
    """Return convergent residue estimates at increasing precision scales."""
    mp.mp.dps = 70

    def block_level_two(z: mp.mpf) -> mp.mpf:
        return (
            2 * z**4
            + 9 * z**3
            + 13 * z**2
            + 8 * z
            - 14
        ) / (2 * (z + 2) * (8 * z + 7))

    rows = []
    for exponent in (4, 8, 16, 30):
        epsilon = mp.mpf(10) ** (-exponent)
        frobenius = 1 / (2 * (1 + epsilon))
        kac_21 = epsilon * block_level_two(-2 + epsilon)
        kac_12 = epsilon * block_level_two(
            -mp.mpf(7) / 8 + epsilon
        )
        rows.append((exponent, frobenius, kac_21, kac_12))
    return rows


def verify_numerical_residues(
    rows: list[tuple[int, mp.mpf, mp.mpf, mp.mpf]],
) -> None:
    """Raise unless every error decreases and the last row is accurate."""
    targets = (mp.mpf("0.5"), mp.mpf(1), -mp.mpf(3619) / 4096)
    columns = tuple(zip(*(row[1:] for row in rows)))
    for label, values, target in zip(
        ("Frobenius", "Kac (2,1)", "Kac (1,2)"),
        columns,
        targets,
    ):
        errors = [abs(value - target) for value in values]
        if not all(right < left for left, right in zip(errors, errors[1:])):
            raise RuntimeError(f"{label} errors do not decrease: {errors}")
        if errors[-1] >= mp.mpf("1e-25"):
            raise RuntimeError(
                f"{label} final residue error is too large: {errors[-1]}"
            )


def main() -> None:
    exact = symbolic_audit()
    rows = numerical_residue_table()
    verify_numerical_residues(rows)

    print("exact symbolic audit")
    for name, value in exact.items():
        print(f"  {name} = {value}")

    print("\nnumerical limits")
    print("  k       Frobenius residue       Kac (2,1)       Kac (1,2)")
    for exponent, frobenius, kac_21, kac_12 in rows:
        print(
            f"  {exponent:<2d}  "
            f"{mp.nstr(frobenius, 16):>23}  "
            f"{mp.nstr(kac_21, 16):>14}  "
            f"{mp.nstr(kac_12, 16):>14}"
        )


if __name__ == "__main__":
    main()
