#!/usr/bin/env python3
"""Static Foldy--Wouthuysen expansion: ordered algebra and error studies.

Run with --help. Standard library + NumPy only; offline and deterministic.
Let H=beta*m*c^2+c*D+V, D=alpha.pi, V=q*Phi, pi=-i*hbar*grad-q*A.
With psi_FW=U*psi, U=exp(G3)exp(G2)exp(G1), each G is anti-Hermitian.
The even Hamiltonian through c^-2 is
 beta*m*c^2+V+beta*D^2/(2m)-beta*D^4/(8m^3*c^2)
 -[D,[D,V]]/(8m^2*c^2).
Here D^2=F=pi^2-q*hbar*Sigma.B; the FULL ordered F^2 must be retained.

Three independent calculations make different claims:
1. Exact Fraction noncommutative BCH coefficients, with ONLY beta^2=I,
   beta*D=-D*beta, beta*V=V*beta. D and V never commute by assumption.
2. Exact rational-complex polynomial differential-operator fixtures for
   nonuniform static electric/magnetic fields, both charge signs and q=0.
3. Numerical exact matrix exponentials of generic finite even/odd fixtures,
   plus Decimal Taylor bounds for nonnegative pure-magnetic spectral F.
   The finite matrix fixture tests algebra; it is not a spatial EM solver.

G2 removes the c^-1 odd term but GENERICALLY leaves an odd c^-3 term.
G3 removes that term before claiming a formal full-operator O(c^-4) remainder.
At fixed smooth A,Phi and m>0 this is an asymptotic expansion on a common
operator core, not a global operator-norm expansion of unbounded momenta.
The pure-magnetic bound becomes a uniform operator bound only on a bounded
spectral interval. Time-dependent fields need the separate i*hbar*dot(U)*Udagger
term. This script does not derive a time-dependent FW transformation.

CSVs: formal coefficients, exact field checks, matrix refinement, and scalar
root remainders. JSON includes checks, dependency versions and SHA-256 hashes.
New output files only; output directory defaults to a fresh temporary folder.
Owner: /relativistic-qm/foldy-wouthuysen-expansion/.
References: Foldy & Wouthuysen, Phys. Rev. 78, 29 (1950);
Bjorken & Drell, Relativistic Quantum Mechanics (1964).
Tested environment: Python 3.12.14, NumPy 2.3.5.
"""
from __future__ import annotations

import argparse
from collections import defaultdict
import csv
from decimal import Decimal, localcontext
from fractions import Fraction
import hashlib
import json
import math
from pathlib import Path
import platform
import sys
import tempfile

import numpy as np


class Checks:
    def __init__(self):
        self.rows = []

    def condition(self, name, passed, **details):
        self.rows.append({"name": name, "passed": bool(passed), **details})

    def close(self, name, actual, expected, tolerance=2e-12):
        absolute = float(np.linalg.norm(np.asarray(actual) - np.asarray(expected)))
        scale = max(1.0, float(np.linalg.norm(expected)))
        self.condition(name, absolute <= tolerance * scale, absolute_error=absolute,
                       scaled_error=absolute / scale, tolerance=tolerance)

    @property
    def passed(self):
        return all(row["passed"] for row in self.rows)


# Laurent-series keys are (power of lambda=1/c, power of m, ordered word).
# Move beta left, acquiring a minus sign for every crossed D; V is even.
def normalize(word):
    sign = (-1) ** sum(word[:i].count("D") for i, symbol in enumerate(word) if symbol == "B")
    return sign, ("B" if word.count("B") % 2 else "") + word.replace("B", "")


def tidy(series):
    return {key: value for key, value in series.items() if value}


def term(degree, mass, word, coefficient=1):
    sign, canonical = normalize(word)
    return {(degree, mass, canonical): sign * Fraction(coefficient)}


def add(*series):
    result = defaultdict(Fraction)
    for item in series:
        for key, value in item.items():
            result[key] += value
    return tidy(result)


def scale(series, coefficient):
    return tidy({key: value * coefficient for key, value in series.items()})


def multiply(left, right, maxdegree=4):
    result = defaultdict(Fraction)
    for (p, m, w), a in left.items():
        for (q, n, v), b in right.items():
            if p + q <= maxdegree:
                sign, word = normalize(w + v)
                result[(p + q, m + n, word)] += sign * a * b
    return tidy(result)


def commutator(left, right):
    return add(multiply(left, right), scale(multiply(right, left), -1))


def adjoint(series):
    result = defaultdict(Fraction)
    for (p, m, w), a in series.items():
        sign, word = normalize(w[::-1])
        result[(p, m, word)] += sign * a
    return tidy(result)


def select(series, degree=None, parity=None, through=None):
    return {key: value for key, value in series.items()
            if (degree is None or key[0] == degree)
            and (parity is None or key[2].count("D") % 2 == parity)
            and (through is None or key[0] <= through)}


def bch(generator, hamiltonian):
    # H starts at lambda^-2; eight nested G1 commutators suffice through +4.
    result, nested = dict(hamiltonian), dict(hamiltonian)
    for n in range(1, 9):
        nested = commutator(generator, nested)
        if not nested:
            break
        result = add(result, scale(nested, Fraction(1, math.factorial(n))))
    return result


def formal_study(checks):
    def exact(name, actual, expected):
        checks.condition(name, not add(actual, scale(expected, -1)), arithmetic="exact Fraction")
    h = add(term(-2, 1, "B"), term(-1, 0, "D"), term(0, 0, "V"))
    g1 = term(1, -1, "BD", Fraction(1, 2))
    h1 = bch(g1, h)
    nested = add(term(2, -2, "DDV"), term(2, -2, "DVD", -2), term(2, -2, "VDD"))
    even = add(term(-2, 1, "B"), term(0, 0, "V"), term(0, -1, "BDD", Fraction(1, 2)),
               term(2, -3, "BDDDD", Fraction(-1, 8)), scale(nested, Fraction(-1, 8)))
    odd1 = add(term(1, -1, "BDV", Fraction(1, 2)), term(1, -1, "BVD", Fraction(-1, 2)),
               term(1, -2, "DDD", Fraction(-1, 3)))
    exact("G1 even terms", select(h1, parity=0, through=2), even)
    exact("G1 remaining odd c^-1", select(h1, degree=1), odd1)
    exact("G1 cancels odd c^1", select(h1, degree=-1), {})
    g2 = scale(multiply(term(2, -1, "B"), odd1), Fraction(1, 2))
    exact("G2 coefficients", g2, add(term(3, -2, "DV", Fraction(1, 4)),
          term(3, -2, "VD", Fraction(-1, 4)), term(3, -3, "BDDD", Fraction(-1, 6))))
    h2 = bch(g2, h1)
    exact("G2 cancels odd c^-1", select(h2, degree=1), {})
    exact("G2 preserves even terms", select(h2, parity=0, through=2), even)
    odd3 = select(h2, degree=3, parity=1)
    checks.condition("G2 has generic odd c^-3 remainder", bool(odd3), monomials=len(odd3))
    g3 = scale(multiply(term(2, -1, "B"), odd3, maxdegree=5), Fraction(1, 2))
    h3 = bch(g3, h2)
    exact("G3 cancels odd c^-3", select(h3, degree=3, parity=1), {})
    exact("full transformed H through c^-3", select(h3, through=3), even)
    for name, generator in (("G1", g1), ("G2", g2), ("G3", g3)):
        exact(name + " anti-Hermitian", adjoint(generator), scale(generator, -1))
    for name, ham in (("H", h), ("H1", h1), ("H2", h2), ("H3", h3)):
        exact(name + " Hermitian", adjoint(ham), ham)
    series = {"H": h, "G1": g1, "G2": g2, "G3": g3, "odd_after_G2": odd3, "even_truncation": even}
    rows = [{"expression": name, "lambda_power": p, "mass_power": m, "word": w or "I", "coefficient": str(a)}
            for name, expression in series.items() for (p, m, w), a in sorted(expression.items())]
    return series, rows


class RationalComplex:
    """Exact complex pair of Fractions, used only for polynomial coefficients."""
    def __init__(self, real=0, imag=0):
        if isinstance(real, RationalComplex):
            self.real, self.imag = real.real, real.imag
        else:
            self.real, self.imag = Fraction(real), Fraction(imag)

    def __add__(self, other):
        z = RationalComplex(other)
        return RationalComplex(self.real + z.real, self.imag + z.imag)

    __radd__ = __add__

    def __neg__(self):
        return RationalComplex(-self.real, -self.imag)

    def __sub__(self, other):
        return self + -RationalComplex(other)

    def __mul__(self, other):
        z = RationalComplex(other)
        return RationalComplex(self.real * z.real - self.imag * z.imag,
                               self.real * z.imag + self.imag * z.real)

    __rmul__ = __mul__

    def __eq__(self, other):
        z = RationalComplex(other)
        return self.real == z.real and self.imag == z.imag


# Sparse 3D polynomials map (x power,y power,z power) to exact coefficients.
def pclean(poly):
    return {power: value for power, value in poly.items() if value != 0}


def constant(value):
    return pclean({(0, 0, 0): RationalComplex(value)})


def padd(*polys):
    result = {}
    for poly in polys:
        for power, value in poly.items():
            result[power] = result.get(power, RationalComplex()) + value
    return pclean(result)


def pscale(poly, scalar):
    return pclean({power: value * scalar for power, value in poly.items()})


def pmul(left, right):
    result = {}
    for p, a in left.items():
        for q, b in right.items():
            power = tuple(x + y for x, y in zip(p, q))
            result[power] = result.get(power, RationalComplex()) + a * b
    return pclean(result)


def derivative(poly, axis):
    result = {}
    for power, value in poly.items():
        if power[axis]:
            reduced = list(power)
            reduced[axis] -= 1
            result[tuple(reduced)] = value * power[axis]
    return pclean(result)


def vadd(*vectors):
    return [padd(*(v[i] for v in vectors)) for i in range(4)]


def vscale(vector, scalar):
    return [pscale(p, scalar) for p in vector]


def vmul(poly, vector):
    return [pmul(poly, p) for p in vector]


def matrix_action(matrix, vector):
    return [padd(*(pscale(p, coefficient) for coefficient, p in zip(row, vector))) for row in matrix]


def epsilon(i, j, k):
    return (j - i) * (k - i) * (k - j) // 2


def field_study(checks):
    q = RationalComplex
    z2 = [[q(), q()], [q(), q()]]
    pauli = [[[q(), q(1)], [q(1), q()]], [[q(), q(0, -1)], [q(0, 1), q()]], [[q(1), q()], [q(), q(-1)]]]
    def block(a, b, c, d):
        return [ar + br for ar, br in zip(a, b)] + [cr + dr for cr, dr in zip(c, d)]
    alpha = [block(z2, s, s, z2) for s in pauli]
    sigma = [block(s, z2, z2, s) for s in pauli]
    x, y, z = ({power: q(1)} for power in ((1, 0, 0), (0, 1, 0), (0, 0, 1)))
    zero = constant(0)
    backgrounds = [
        ("zero", zero, [zero, zero, zero]),
        ("uniform_electric", pscale(z, -3), [zero, zero, zero]),
        ("uniform_magnetic", zero, [pscale(y, -1), x, zero]),
        ("nonuniform_static", padd(pmul(x, y), pmul(z, z), pscale(pmul(x, pmul(y, z)), 2)),
         [pmul(y, z), padd(pmul(x, x), pscale(z, -1)), pmul(x, y)]),
        ("nonuniform_magnetic_electric", padd(pmul(x, x), pmul(y, pmul(z, z))),
         [pmul(y, pmul(z, z)), pmul(x, pmul(x, z)), pmul(x, pmul(y, y))]),
    ]
    # Fixed polynomial jets exercise derivatives and all complex components.
    jets = [[padd(constant(q(i + 1, j - i)), pscale(x, i - j), pscale(pmul(y, z), q(j + 1, i)),
                  pscale(pmul(x, pmul(x, z)), i + j + 1)) for i in range(4)] for j in range(3)]
    rows, omissions = [], 0
    for label, phi, avec in backgrounds:
        electric = [pscale(derivative(phi, i), -1) for i in range(3)]
        magnetic = [padd(*(pscale(derivative(avec[k], j), epsilon(i, j, k))
                            for j in range(3) for k in range(3))) for i in range(3)]
        for charge in (-2, 0, 2):
            hbar = 2
            potential = pscale(phi, charge)
            def pi(i, v):
                return vadd([pscale(derivative(p, i), q(0, -hbar)) for p in v], vscale(vmul(avec[i], v), -charge))
            def d(v):
                return vadd(*(matrix_action(alpha[i], pi(i, v)) for i in range(3)))
            def pi2(v):
                return vadd(*(pi(i, pi(i, v)) for i in range(3)))
            def sb(v):
                return vadd(*(matrix_action(sigma[i], vmul(magnetic[i], v)) for i in range(3)))
            def field(v):
                return vadd(pi2(v), vscale(sb(v), -charge * hbar))
            def cross(v, reverse=False):
                return vadd(*(vscale(matrix_action(sigma[k], pi(i, vmul(electric[j], v)) if reverse
                                         else vmul(electric[i], pi(j, v))), epsilon(k, i, j))
                              for k in range(3) for i in range(3) for j in range(3)))
            divergence = padd(*(derivative(electric[i], i) for i in range(3)))
            b2 = padd(*(pmul(b, b) for b in magnetic))
            for number, v in enumerate(jets):
                double = vadd(d(d(vmul(potential, v))), vscale(d(vmul(potential, d(v))), -2), vmul(potential, d(d(v))))
                electric_rhs = vadd(vscale(vmul(divergence, v), charge * hbar ** 2),
                                    vscale(vadd(cross(v), vscale(cross(v, True), -1)), charge * hbar))
                expanded = vadd(pi2(pi2(v)), vscale(vadd(pi2(sb(v)), sb(pi2(v))), -charge * hbar),
                                vscale(vmul(b2, v), charge ** 2 * hbar ** 2))
                comparisons = [("electric_nested_commutator", double, electric_rhs),
                               ("curl_free_cross_order", cross(v, True), vscale(cross(v), -1)),
                               ("D_squared_equals_F", d(d(v)), field(v)),
                               ("F_squared_full_ordering", field(field(v)), expanded),
                               ("Sigma_B_squared", sb(sb(v)), vmul(b2, v))]
                for name, actual, expected in comparisons:
                    passed = actual == expected
                    checks.condition(f"{label} q={charge} jet={number} {name}", passed, arithmetic="exact rational complex")
                    rows.append({"background": label, "charge": charge, "hbar": hbar, "jet": number,
                                 "identity": name, "passed": passed})
                omissions += field(field(v)) != pi2(pi2(v))
    checks.condition("omitting magnetic terms from F^2 is detected", omissions > 0, detected_cases=omissions)
    return rows, omissions


def evaluate(series, mass, speed, beta, d, v):
    result = np.zeros_like(beta)
    letters = {"B": beta, "D": d, "V": v}
    for (power, mpower, word), coefficient in series.items():
        value = np.eye(len(beta), dtype=complex)
        for letter in word:
            value = value @ letters[letter]
        result += float(coefficient) * speed ** (-power) * mass ** mpower * value
    return result


def exponential_antihermitian(generator):
    # iG is Hermitian. eigh computes exp(G) independently of truncated BCH.
    eigenvalues, vectors = np.linalg.eigh(1j * generator)
    return (vectors * np.exp(-1j * eigenvalues)) @ vectors.conj().T


def matrix_study(series, checks, mass, first_speed):
    beta = np.diag([1, 1, -1, -1]).astype(complex)
    k = np.array([[0.8, 0.2 + 0.3j], [-0.1j, -0.6]], dtype=complex)
    zero = np.zeros((2, 2), dtype=complex)
    d = np.block([[zero, k], [k.conj().T, zero]])
    v = np.block([[np.array([[0.3, 0.12j], [-0.12j, -0.4]]), zero],
                  [zero, np.array([[-0.2, 0.1 + 0.05j], [0.1 - 0.05j, 0.45]])]])
    checks.close("matrix D odd", beta @ d + d @ beta, np.zeros_like(beta))
    checks.close("matrix V even", beta @ v - v @ beta, np.zeros_like(beta))
    checks.condition("matrix D,V do not commute", np.linalg.norm(d @ v - v @ d) > 0.1)
    rows = []
    for speed in first_speed * 2.0 ** np.arange(4):
        h = mass * speed ** 2 * beta + speed * d + v
        generators = [evaluate(series[name], mass, speed, beta, d, v) for name in ("G1", "G2", "G3")]
        transformed = h.copy()
        combined = np.eye(4, dtype=complex)
        odd_errors = []
        for stage, generator in enumerate(generators, 1):
            checks.close(f"matrix c={speed:g} G{stage} anti-Hermitian", generator.conj().T, -generator)
            u = exponential_antihermitian(generator)
            combined = u @ combined
            transformed = u @ transformed @ u.conj().T
            odd = (transformed - beta @ transformed @ beta) / 2
            odd_errors.append(float(np.linalg.norm(odd)))
        checks.close(f"matrix c={speed:g} total U unitary", combined @ combined.conj().T, np.eye(4))
        checks.close(f"matrix c={speed:g} spectrum unchanged", np.linalg.eigvalsh(transformed), np.linalg.eigvalsh(h))
        even_truncated = evaluate(series["even_truncation"], mass, speed, beta, d, v)
        even_part = (transformed + beta @ transformed @ beta) / 2
        row = {"mass": mass, "c": float(speed), "G1_odd_norm": odd_errors[0],
               "G2_odd_norm": odd_errors[1], "G3_odd_norm": odd_errors[2],
               "even_truncation_error": float(np.linalg.norm(even_part - even_truncated)),
               "full_truncation_error": float(np.linalg.norm(transformed - even_truncated))}
        rows.append(row)
    # A fixed fixture, moderate c, gives a robust independent asymptotic test.
    # Configurable matrix studies are diagnostics; only fixed m=1,c=4 is gated.
    if mass == 1.0 and first_speed == 4.0:
        for key, expected in (("G1_odd_norm", 1), ("G2_odd_norm", 3), ("G3_odd_norm", 5), ("even_truncation_error", 4)):
            order = math.log2(rows[-2][key] / rows[-1][key])
            checks.condition("matrix asymptotic order " + key, abs(order - expected) < 0.2,
                             observed_order=order, expected_order=expected, tolerance=0.2)
    for index, row in enumerate(rows):
        for key in ("G1_odd_norm", "G2_odd_norm", "G3_odd_norm", "even_truncation_error"):
            row[key + "_observed_order"] = None if index == 0 else math.log2(rows[index - 1][key] / row[key]) if row[key] else None
    return rows


def root_study(checks, precision):
    rows = []
    with localcontext() as context:
        context.prec = precision
        for text in ("0", "1e-30", "1e-12", "1e-6", ".001", ".03", ".1", "1", "3", "10", "1e4"):
            r = Decimal(text)
            root = (1 + r).sqrt()
            approximation = 1 + r / 2 - r ** 2 / 8
            remainder = root - approximation
            lower, upper = r ** 3 / (16 * (1 + r) ** 2 * root), r ** 3 / 16
            checks.condition("root Taylor bounds r=" + text, lower <= remainder <= upper, arithmetic=f"Decimal {precision} digits")
            rows.append({"r_F_over_m2c2": text, "sqrt_1_plus_r": str(root), "truncation": str(approximation),
                         "remainder": str(remainder), "lower_bound": str(lower), "upper_bound": str(upper),
                         "small_parameter": r <= Decimal(".1")})
    return rows


def write_csv(path, rows):
    fields = list(dict.fromkeys(key for row in rows for key in row))
    with path.open("x", encoding="utf8", newline="") as stream:
        writer = csv.DictWriter(stream, fieldnames=fields)
        writer.writeheader()
        writer.writerows(rows)


def parser():
    p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
    p.add_argument("--output-dir", type=Path, help="new CSV/JSON outputs; default fresh temporary directory")
    p.add_argument("--json-output", type=Path, help="override summary JSON destination")
    p.add_argument("--coefficients-csv", type=Path, help="override exact coefficient CSV destination")
    p.add_argument("--mass", type=float, default=1.0, help="finite-matrix study mass, 0.25..4 (default 1)")
    p.add_argument("--first-c", type=float, default=4.0, help="matrix-study first c, 1..16; then three doublings (default 4)")
    p.add_argument("--precision", type=int, default=160, help="Decimal precision, 120..500 (default 160)")
    return p


def main(argv=None):
    p = parser()
    args = p.parse_args(argv)
    if not math.isfinite(args.mass) or not 0.25 <= args.mass <= 4:
        p.error("mass must be finite and between 0.25 and 4")
    if not math.isfinite(args.first_c) or not 1 <= args.first_c <= 16:
        p.error("first-c must be finite and between 1 and 16")
    if not 120 <= args.precision <= 500:
        p.error("precision must be between 120 and 500 digits")
    output = args.output_dir or Path(tempfile.mkdtemp(prefix="fw-expansion-"))
    paths = {"coefficients": args.coefficients_csv or output / "coefficients.csv", "fields": output / "fields.csv",
             "matrices": output / "matrices.csv", "roots": output / "roots.csv", "summary": args.json_output or output / "summary.json"}
    if len({path.resolve() for path in paths.values()}) != len(paths):
        p.error("output paths must be distinct")
    for path in paths.values():
        if path.exists() or path.is_symlink():
            p.error(f"existing outputs are never overwritten: {path}")
        if ".git" in [part.lower() for part in path.resolve().parts]:
            p.error(f"output cannot be inside .git: {path}")
    for path in paths.values():
        path.parent.mkdir(parents=True, exist_ok=True)
    checks = Checks()
    series, coefficients = formal_study(checks)
    fields, omissions = field_study(checks)
    matrices = [{"study": "fixed_self_check", **row} for row in matrix_study(series, checks, 1.0, 4.0)]
    if (args.mass, args.first_c) != (1.0, 4.0):
        matrices += [{"study": "configured_diagnostic", **row} for row in matrix_study(series, checks, args.mass, args.first_c)]
    roots = root_study(checks, args.precision)
    for name, rows in (("coefficients", coefficients), ("fields", fields), ("matrices", matrices), ("roots", roots)):
        write_csv(paths[name], rows)
    report = {"program": Path(__file__).name, "passed": checks.passed, "check_count": len(checks.rows),
              "script_sha256": hashlib.sha256(Path(__file__).read_bytes()).hexdigest(),
              "environment": {"python": platform.python_version(), "numpy": np.__version__, "platform": platform.platform()},
              "randomness": "none; fixed rational polynomial and complex matrix fixtures",
              "parameters": {"mass": args.mass, "first_c": args.first_c, "precision": args.precision},
              "exact_field_identities": len(fields), "detected_magnetic_omission_cases": omissions,
              "checks": checks.rows,
              "interpretation": ["Static potentials held fixed as c tends to infinity; m>0.",
                  "Fraction BCH coefficients are exact in the declared free even/odd algebra.",
                  "Polynomial tests are exact finite examples, not an analytic proof for all domains.",
                  "G2 leaves odd c^-3 terms; G3 is needed for the stated full-operator formal remainder.",
                  "Finite matrix exponentials independently test generic noncommuting even/odd algebra, not EM discretization.",
                  "Observed-order checks use the fixed m=1,c=4 fixture; custom studies expose asymptotic failure or roundoff.",
                  "Root bound holds for all nonnegative r, but useful relative truncation requires r much smaller than one.",
                  "Unbounded operators need common cores and spectral restrictions; no global operator-norm convergence is claimed.",
                  "Time-dependent FW transformation has an extra i*hbar*dot(U)*Udagger term."],
              "output_files": {key: str(path) for key, path in paths.items()},
              "output_sha256": {key: hashlib.sha256(path.read_bytes()).hexdigest() for key, path in paths.items() if key != "summary"}}
    with paths["summary"].open("x", encoding="utf8") as stream:
        json.dump(report, stream, indent=2, allow_nan=False)
        stream.write("\n")
    print(json.dumps({"passed": checks.passed, "checks": len(checks.rows),
                      "failed": [row for row in checks.rows if not row["passed"]], "summary_json": str(paths["summary"])}, indent=2))
    return 0 if checks.passed else 1


if __name__ == "__main__":
    sys.exit(main())
