#!/usr/bin/env python3
# SPDX-License-Identifier: MIT
"""Reproduce the one-parameter effective-charge calculation for helium.

The clamped-nucleus nonrelativistic Hamiltonian is evaluated in the normalized
product of two hydrogenic 1s orbitals with exponent zeta. In Hartree atomic
units,

    E(zeta; Z) = zeta**2 - 2*Z*zeta + 5*zeta/8.

The program compares the analytic minimum with an independent golden-section
search, writes machine-readable tables, and runs validation checks. It uses
only the Python standard library and is deterministic.
"""

from __future__ import annotations

import argparse
import csv
import json
import math
import platform
import sys
from dataclasses import asdict, dataclass
from pathlib import Path


HELIUM_EXACT_NONREL = -2.9037243770341196
HELIUM_HARTREE_FOCK_LIMIT = -2.861679995612


@dataclass(frozen=True)
class Components:
    zeta: float
    kinetic: float
    nuclear_attraction: float
    electron_repulsion: float

    @property
    def potential(self) -> float:
        return self.nuclear_attraction + self.electron_repulsion

    @property
    def total(self) -> float:
        return self.kinetic + self.potential

    @property
    def virial_residual(self) -> float:
        return 2.0 * self.kinetic + self.potential


def energy_components(zeta: float, nuclear_charge: float = 2.0) -> Components:
    """Return analytic expectation-value components in Hartree."""
    if zeta <= 0.0:
        raise ValueError("zeta must be positive")
    if nuclear_charge <= 0.0:
        raise ValueError("nuclear charge must be positive")
    return Components(
        zeta=zeta,
        kinetic=zeta * zeta,
        nuclear_attraction=-2.0 * nuclear_charge * zeta,
        electron_repulsion=5.0 * zeta / 8.0,
    )


def energy(zeta: float, nuclear_charge: float = 2.0) -> float:
    return energy_components(zeta, nuclear_charge).total


def analytic_optimum(nuclear_charge: float = 2.0) -> tuple[float, float]:
    """Return zeta_star and E(zeta_star)."""
    zeta_star = nuclear_charge - 5.0 / 16.0
    if zeta_star <= 0.0:
        raise ValueError("this trial family has no positive interior minimum")
    return zeta_star, -(zeta_star * zeta_star)


def golden_section_minimum(
    function,
    lower: float,
    upper: float,
    *,
    relative_tolerance: float = 2.0e-14,
    max_iterations: int = 512,
) -> tuple[float, float, int]:
    """Minimize a unimodal scalar function while retaining a bracket."""
    if not 0.0 < lower < upper:
        raise ValueError("require 0 < lower < upper")

    invphi = (math.sqrt(5.0) - 1.0) / 2.0
    a, b = lower, upper
    c = b - invphi * (b - a)
    d = a + invphi * (b - a)
    fc, fd = function(c), function(d)

    for iteration in range(1, max_iterations + 1):
        midpoint = 0.5 * (a + b)
        if b - a <= relative_tolerance * max(1.0, abs(midpoint)):
            return midpoint, function(midpoint), iteration
        if fc <= fd:
            b, d, fd = d, c, fc
            c = b - invphi * (b - a)
            fc = function(c)
        else:
            a, c, fc = c, d, fd
            d = a + invphi * (b - a)
            fd = function(d)

    raise RuntimeError("golden-section search did not converge")


def write_curve(
    path: Path,
    *,
    nuclear_charge: float,
    lower: float,
    upper: float,
    points: int,
) -> None:
    if points < 3:
        raise ValueError("points must be at least three")
    step = (upper - lower) / (points - 1)
    with path.open("w", newline="", encoding="utf-8") as handle:
        writer = csv.writer(handle)
        writer.writerow(
            [
                "zeta",
                "kinetic_Eh",
                "nuclear_attraction_Eh",
                "electron_repulsion_Eh",
                "potential_Eh",
                "total_Eh",
                "virial_residual_Eh",
            ]
        )
        for index in range(points):
            zeta = lower + index * step
            result = energy_components(zeta, nuclear_charge)
            writer.writerow(
                [
                    f"{zeta:.16g}",
                    f"{result.kinetic:.16g}",
                    f"{result.nuclear_attraction:.16g}",
                    f"{result.electron_repulsion:.16g}",
                    f"{result.potential:.16g}",
                    f"{result.total:.16g}",
                    f"{result.virial_residual:.16g}",
                ]
            )


def write_summary(
    path: Path,
    *,
    nuclear_charge: float,
    analytic: Components,
    numeric: Components,
) -> None:
    fieldnames = [
        "label",
        "hamiltonian",
        "zeta",
        "kinetic_Eh",
        "nuclear_attraction_Eh",
        "electron_repulsion_Eh",
        "total_Eh",
        "variational_for_full_H",
        "note",
    ]
    no_repulsion_total = -(nuclear_charge * nuclear_charge)
    unoptimized = energy_components(nuclear_charge, nuclear_charge)
    rows = [
        {
            "label": "independent_no_repulsion",
            "hamiltonian": "H_without_1_over_r12",
            "zeta": f"{nuclear_charge:.16g}",
            "kinetic_Eh": f"{nuclear_charge**2:.16g}",
            "nuclear_attraction_Eh": f"{-2*nuclear_charge**2:.16g}",
            "electron_repulsion_Eh": "0",
            "total_Eh": f"{no_repulsion_total:.16g}",
            "variational_for_full_H": "false",
            "note": "different Hamiltonian; not an upper bound to full helium",
        },
        {
            "label": "unoptimized_full_H",
            "hamiltonian": "full_nonrelativistic_clamped_nucleus",
            "zeta": f"{unoptimized.zeta:.16g}",
            "kinetic_Eh": f"{unoptimized.kinetic:.16g}",
            "nuclear_attraction_Eh": f"{unoptimized.nuclear_attraction:.16g}",
            "electron_repulsion_Eh": f"{unoptimized.electron_repulsion:.16g}",
            "total_Eh": f"{unoptimized.total:.16g}",
            "variational_for_full_H": "true",
            "note": "bare Z exponent with repulsion evaluated",
        },
        {
            "label": "optimized_analytic",
            "hamiltonian": "full_nonrelativistic_clamped_nucleus",
            "zeta": f"{analytic.zeta:.16g}",
            "kinetic_Eh": f"{analytic.kinetic:.16g}",
            "nuclear_attraction_Eh": f"{analytic.nuclear_attraction:.16g}",
            "electron_repulsion_Eh": f"{analytic.electron_repulsion:.16g}",
            "total_Eh": f"{analytic.total:.16g}",
            "variational_for_full_H": "true",
            "note": "one-parameter minimum",
        },
        {
            "label": "optimized_numeric",
            "hamiltonian": "full_nonrelativistic_clamped_nucleus",
            "zeta": f"{numeric.zeta:.16g}",
            "kinetic_Eh": f"{numeric.kinetic:.16g}",
            "nuclear_attraction_Eh": f"{numeric.nuclear_attraction:.16g}",
            "electron_repulsion_Eh": f"{numeric.electron_repulsion:.16g}",
            "total_Eh": f"{numeric.total:.16g}",
            "variational_for_full_H": "true",
            "note": "golden-section result",
        },
    ]

    if math.isclose(nuclear_charge, 2.0, rel_tol=0.0, abs_tol=1.0e-15):
        rows.extend(
            [
                {
                    "label": "hartree_fock_limit_reference",
                    "hamiltonian": "full_nonrelativistic_clamped_nucleus",
                    "zeta": "",
                    "kinetic_Eh": "",
                    "nuclear_attraction_Eh": "",
                    "electron_repulsion_Eh": "",
                    "total_Eh": f"{HELIUM_HARTREE_FOCK_LIMIT:.16g}",
                    "variational_for_full_H": "true",
                    "note": "literature reference; not computed by this program",
                },
                {
                    "label": "exact_nonrelativistic_reference",
                    "hamiltonian": "full_nonrelativistic_clamped_nucleus",
                    "zeta": "",
                    "kinetic_Eh": "",
                    "nuclear_attraction_Eh": "",
                    "electron_repulsion_Eh": "",
                    "total_Eh": f"{HELIUM_EXACT_NONREL:.16g}",
                    "variational_for_full_H": "true",
                    "note": "literature reference; not computed by this program",
                },
            ]
        )

    with path.open("w", newline="", encoding="utf-8") as handle:
        writer = csv.DictWriter(handle, fieldnames=fieldnames)
        writer.writeheader()
        writer.writerows(rows)


def validate(
    *,
    nuclear_charge: float,
    analytic: Components,
    numeric: Components,
    numeric_iterations: int,
) -> dict[str, float | int | bool]:
    zeta_error = abs(numeric.zeta - analytic.zeta)
    energy_error = abs(numeric.total - analytic.total)
    derivative_step = 1.0e-5
    derivative = (
        energy(analytic.zeta + derivative_step, nuclear_charge)
        - energy(analytic.zeta - derivative_step, nuclear_charge)
    ) / (2.0 * derivative_step)
    unoptimized = energy_components(nuclear_charge, nuclear_charge)

    checks = {
        # Energy-only minimization loses parameter resolution quadratically
        # near a stationary point in binary64 arithmetic.
        "zeta_agreement": zeta_error < 5.0e-9,
        "energy_agreement": energy_error < 2.0e-13,
        "stationary_finite_difference": abs(derivative) < 2.0e-10,
        "virial_stationarity": abs(analytic.virial_residual) < 2.0e-14,
        "optimization_lowers_full_H": analytic.total <= unoptimized.total,
    }

    if math.isclose(nuclear_charge, 2.0, rel_tol=0.0, abs_tol=1.0e-15):
        checks["variational_upper_bound"] = analytic.total >= HELIUM_EXACT_NONREL
        checks["hartree_fock_improves_trial"] = (
            HELIUM_HARTREE_FOCK_LIMIT <= analytic.total
        )

    failed = [name for name, passed in checks.items() if not passed]
    if failed:
        raise AssertionError("validation failed: " + ", ".join(failed))

    return {
        **checks,
        "zeta_absolute_error": zeta_error,
        "energy_absolute_error_Eh": energy_error,
        "finite_difference_gradient_Eh": derivative,
        "analytic_virial_residual_Eh": analytic.virial_residual,
        "numeric_iterations": numeric_iterations,
    }


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--nuclear-charge", type=float, default=2.0)
    parser.add_argument("--zeta-min", type=float, default=0.5)
    parser.add_argument("--zeta-max", type=float, default=3.0)
    parser.add_argument("--points", type=int, default=501)
    parser.add_argument(
        "--output-dir",
        type=Path,
        default=Path("variational-helium-output"),
    )
    return parser.parse_args()


def main() -> None:
    args = parse_args()
    if not 0.0 < args.zeta_min < args.zeta_max:
        raise ValueError("require 0 < zeta-min < zeta-max")

    zeta_star, expected_energy = analytic_optimum(args.nuclear_charge)
    if not args.zeta_min < zeta_star < args.zeta_max:
        raise ValueError("the analytic optimum must lie inside the search interval")

    numeric_zeta, numeric_energy, iterations = golden_section_minimum(
        lambda value: energy(value, args.nuclear_charge),
        args.zeta_min,
        args.zeta_max,
    )
    analytic = energy_components(zeta_star, args.nuclear_charge)
    numeric = energy_components(numeric_zeta, args.nuclear_charge)
    if not math.isclose(
        analytic.total,
        expected_energy,
        rel_tol=0.0,
        abs_tol=2.0e-15,
    ):
        raise AssertionError("analytic energy formula is internally inconsistent")

    validation = validate(
        nuclear_charge=args.nuclear_charge,
        analytic=analytic,
        numeric=numeric,
        numeric_iterations=iterations,
    )

    args.output_dir.mkdir(parents=True, exist_ok=True)
    curve_path = args.output_dir / "variational-helium-energy-curve.csv"
    summary_path = args.output_dir / "variational-helium-summary.csv"
    metadata_path = args.output_dir / "variational-helium-metadata.json"

    write_curve(
        curve_path,
        nuclear_charge=args.nuclear_charge,
        lower=args.zeta_min,
        upper=args.zeta_max,
        points=args.points,
    )
    write_summary(
        summary_path,
        nuclear_charge=args.nuclear_charge,
        analytic=analytic,
        numeric=numeric,
    )

    metadata = {
        "model": "two hydrogenic 1s orbitals with common exponent zeta",
        "units": "Hartree atomic units",
        "nuclear_charge": args.nuclear_charge,
        "zeta_interval": [args.zeta_min, args.zeta_max],
        "curve_points": args.points,
        "analytic_components": asdict(analytic)
        | {
            "potential": analytic.potential,
            "total": analytic.total,
            "virial_residual": analytic.virial_residual,
        },
        "numeric_components": asdict(numeric)
        | {
            "potential": numeric.potential,
            "total": numeric.total,
            "virial_residual": numeric.virial_residual,
        },
        "validation": validation,
        "python": platform.python_version(),
        "implementation": platform.python_implementation(),
        "platform": platform.platform(),
        "random_seed": None,
        "dependencies": "Python standard library only",
        "license": "MIT",
        "outputs": [curve_path.name, summary_path.name, metadata_path.name],
    }
    metadata_path.write_text(
        json.dumps(metadata, indent=2, sort_keys=True) + "\n",
        encoding="utf-8",
    )

    print(f"analytic zeta       = {analytic.zeta:.15f}")
    print(f"numeric zeta        = {numeric.zeta:.15f}")
    print(f"optimized energy    = {analytic.total:.15f} Eh")
    print(
        "unoptimized energy  = "
        f"{energy(args.nuclear_charge, args.nuclear_charge):.15f} Eh"
    )
    print(
        "no-repulsion model  = "
        f"{-(args.nuclear_charge**2):.15f} Eh"
    )
    if math.isclose(args.nuclear_charge, 2.0, rel_tol=0.0, abs_tol=1.0e-15):
        print(
            "trial deficit       = "
            f"{analytic.total - HELIUM_EXACT_NONREL:.15f} Eh"
        )
        print(
            "HF correlation gap  = "
            f"{HELIUM_HARTREE_FOCK_LIMIT - HELIUM_EXACT_NONREL:.15f} Eh"
        )
    print(f"outputs             = {args.output_dir.resolve()}")
    print("validation          = all checks passed")


if __name__ == "__main__":
    main()
