#!/usr/bin/env python3
# SPDX-License-Identifier: MIT
"""Restricted Hartree-Fock for helium in one-centre s-Gaussian bases.

The program evaluates all overlap, kinetic, nuclear-attraction, and
electron-repulsion integrals analytically. It solves the Roothaan equations by
symmetric orthogonalization, iterates the closed-shell density to
self-consistency, and writes deterministic convergence and validation data.

NumPy is the only dependency. Energies are in Hartree atomic units.
"""

from __future__ import annotations

import argparse
import csv
import json
import platform
from dataclasses import dataclass
from pathlib import Path

import numpy as np


HELIUM_HARTREE_FOCK_LIMIT = -2.8616799956122388
HELIUM_EXACT_NONREL = -2.9037243770341196

# This insertion order creates nested spaces. The complete set is the
# even-tempered sequence 0.2 * 3**k, k = 0,...,7.
PUBLISHED_EXPONENT_ORDER = (0.6, 1.8, 0.2, 5.4, 16.2, 48.6, 145.8, 437.4)


@dataclass(frozen=True)
class IntegralSet:
    overlap: np.ndarray
    kinetic: np.ndarray
    nuclear: np.ndarray
    core: np.ndarray
    eri: np.ndarray


@dataclass(frozen=True)
class IterationRecord:
    iteration: int
    total_energy: float
    delta_energy: float
    density_rms: float
    commutator_norm: float


@dataclass(frozen=True)
class SCFResult:
    exponents: tuple[float, ...]
    integrals: IntegralSet
    density: np.ndarray
    fock: np.ndarray
    occupied_coefficients: np.ndarray
    occupied_energy: float
    total_energy: float
    kinetic_energy: float
    nuclear_energy: float
    electron_repulsion_energy: float
    iterations: int
    density_rms: float
    commutator_norm: float
    eigenpair_residual: float
    electron_count: float
    overlap_condition: float
    history: tuple[IterationRecord, ...]


def gaussian_normalization(alpha: float) -> float:
    if alpha <= 0.0:
        raise ValueError("Gaussian exponents must be positive")
    return (2.0 * alpha / np.pi) ** 0.75


def build_integrals(
    exponents: tuple[float, ...],
    nuclear_charge: float = 2.0,
) -> IntegralSet:
    """Return analytic one-centre normalized s-Gaussian integrals."""
    if not exponents:
        raise ValueError("at least one exponent is required")
    if nuclear_charge <= 0.0:
        raise ValueError("nuclear charge must be positive")

    alpha = np.asarray(exponents, dtype=float)
    normalization = np.array(
        [gaussian_normalization(value) for value in alpha],
        dtype=float,
    )
    dimension = len(alpha)
    overlap = np.empty((dimension, dimension), dtype=float)
    kinetic = np.empty_like(overlap)
    nuclear = np.empty_like(overlap)
    eri = np.empty(
        (dimension, dimension, dimension, dimension),
        dtype=float,
    )

    for mu in range(dimension):
        for nu in range(dimension):
            p = alpha[mu] + alpha[nu]
            prefactor = normalization[mu] * normalization[nu]
            overlap[mu, nu] = prefactor * (np.pi / p) ** 1.5
            kinetic[mu, nu] = (
                overlap[mu, nu]
                * 3.0
                * alpha[mu]
                * alpha[nu]
                / p
            )
            nuclear[mu, nu] = (
                -nuclear_charge
                * prefactor
                * 2.0
                * np.pi
                / p
            )

            for lam in range(dimension):
                for sig in range(dimension):
                    q = alpha[lam] + alpha[sig]
                    eri[mu, nu, lam, sig] = (
                        prefactor
                        * normalization[lam]
                        * normalization[sig]
                        * 2.0
                        * np.pi**2.5
                        / (p * q * np.sqrt(p + q))
                    )

    return IntegralSet(
        overlap=overlap,
        kinetic=kinetic,
        nuclear=nuclear,
        core=kinetic + nuclear,
        eri=eri,
    )


def symmetric_orthogonalizer(
    overlap: np.ndarray,
    *,
    eigenvalue_floor: float = 1.0e-11,
) -> tuple[np.ndarray, np.ndarray]:
    eigenvalues, eigenvectors = np.linalg.eigh(overlap)
    if eigenvalues[0] <= eigenvalue_floor:
        raise np.linalg.LinAlgError(
            "overlap matrix is singular or below the configured eigenvalue floor"
        )
    inverse_sqrt = (
        eigenvectors
        @ np.diag(eigenvalues ** -0.5)
        @ eigenvectors.T
    )
    return inverse_sqrt, eigenvalues


def solve_generalized_symmetric(
    matrix: np.ndarray,
    overlap: np.ndarray,
    orthogonalizer: np.ndarray,
) -> tuple[np.ndarray, np.ndarray]:
    transformed = orthogonalizer.T @ matrix @ orthogonalizer
    transformed = 0.5 * (transformed + transformed.T)
    eigenvalues, transformed_vectors = np.linalg.eigh(transformed)
    coefficients = orthogonalizer @ transformed_vectors
    return eigenvalues, coefficients


def density_from_occupied(coefficients: np.ndarray) -> np.ndarray:
    return 2.0 * np.outer(coefficients, coefficients)


def build_fock(
    core: np.ndarray,
    eri: np.ndarray,
    density: np.ndarray,
) -> np.ndarray:
    coulomb = np.einsum("ls,mnls->mn", density, eri, optimize=True)
    exchange = np.einsum("ls,mlns->mn", density, eri, optimize=True)
    fock = core + coulomb - 0.5 * exchange
    return 0.5 * (fock + fock.T)


def electronic_energy(
    density: np.ndarray,
    core: np.ndarray,
    fock: np.ndarray,
) -> float:
    return 0.5 * float(np.einsum("mn,mn", density, core + fock))


def commutator_residual(
    fock: np.ndarray,
    density: np.ndarray,
    overlap: np.ndarray,
) -> float:
    residual = fock @ density @ overlap - overlap @ density @ fock
    return float(np.linalg.norm(residual))


def run_scf(
    exponents: tuple[float, ...],
    *,
    nuclear_charge: float = 2.0,
    density_tolerance: float = 1.0e-12,
    energy_tolerance: float = 1.0e-13,
    max_iterations: int = 256,
) -> SCFResult:
    """Solve closed-shell two-electron RHF with an undamped fixed-point loop."""
    integrals = build_integrals(exponents, nuclear_charge)
    overlap = integrals.overlap
    orthogonalizer, overlap_eigenvalues = symmetric_orthogonalizer(overlap)

    _, core_coefficients = solve_generalized_symmetric(
        integrals.core,
        overlap,
        orthogonalizer,
    )
    occupied = core_coefficients[:, 0]
    density = density_from_occupied(occupied)
    previous_energy: float | None = None
    history: list[IterationRecord] = []

    for iteration in range(1, max_iterations + 1):
        input_fock = build_fock(integrals.core, integrals.eri, density)
        _, coefficients = solve_generalized_symmetric(
            input_fock,
            overlap,
            orthogonalizer,
        )
        new_occupied = coefficients[:, 0]
        new_density = density_from_occupied(new_occupied)
        self_fock = build_fock(
            integrals.core,
            integrals.eri,
            new_density,
        )
        total_energy = electronic_energy(
            new_density,
            integrals.core,
            self_fock,
        )
        density_rms = float(np.sqrt(np.mean((new_density - density) ** 2)))
        delta_energy = (
            float("nan")
            if previous_energy is None
            else total_energy - previous_energy
        )
        commutator_norm = commutator_residual(
            self_fock,
            new_density,
            overlap,
        )
        history.append(
            IterationRecord(
                iteration=iteration,
                total_energy=total_energy,
                delta_energy=delta_energy,
                density_rms=density_rms,
                commutator_norm=commutator_norm,
            )
        )

        energy_converged = (
            previous_energy is not None
            and abs(delta_energy) < energy_tolerance
        )
        density = new_density
        occupied = new_occupied
        previous_energy = total_energy
        if density_rms < density_tolerance and energy_converged:
            break
    else:
        raise RuntimeError("SCF did not converge within max_iterations")

    # Rebuild once from the converged density and solve its Roothaan equation
    # for final orbital diagnostics.
    fock = build_fock(integrals.core, integrals.eri, density)
    orbital_energies, coefficients = solve_generalized_symmetric(
        fock,
        overlap,
        orthogonalizer,
    )
    occupied = coefficients[:, 0]
    final_density = density_from_occupied(occupied)
    final_fock = build_fock(
        integrals.core,
        integrals.eri,
        final_density,
    )
    occupied_energy = float(occupied @ final_fock @ occupied)
    total_energy = electronic_energy(
        final_density,
        integrals.core,
        final_fock,
    )
    kinetic_energy = float(
        np.einsum("mn,mn", final_density, integrals.kinetic)
    )
    nuclear_energy = float(
        np.einsum("mn,mn", final_density, integrals.nuclear)
    )
    electron_repulsion_energy = total_energy - kinetic_energy - nuclear_energy
    generalized_residual = (
        final_fock @ occupied
        - overlap @ occupied * occupied_energy
    )
    density_rms = float(
        np.sqrt(np.mean((final_density - density) ** 2))
    )
    commutator_norm = commutator_residual(
        final_fock,
        final_density,
        overlap,
    )

    return SCFResult(
        exponents=exponents,
        integrals=integrals,
        density=final_density,
        fock=final_fock,
        occupied_coefficients=occupied,
        occupied_energy=occupied_energy,
        total_energy=total_energy,
        kinetic_energy=kinetic_energy,
        nuclear_energy=nuclear_energy,
        electron_repulsion_energy=electron_repulsion_energy,
        iterations=iteration,
        density_rms=density_rms,
        commutator_norm=commutator_norm,
        eigenpair_residual=float(np.linalg.norm(generalized_residual)),
        electron_count=float(np.trace(final_density @ overlap)),
        overlap_condition=float(
            overlap_eigenvalues[-1] / overlap_eigenvalues[0]
        ),
        history=tuple(history),
    )


def integral_symmetry_error(eri: np.ndarray) -> float:
    errors = [
        np.max(np.abs(eri - eri.transpose(1, 0, 2, 3))),
        np.max(np.abs(eri - eri.transpose(0, 1, 3, 2))),
        np.max(np.abs(eri - eri.transpose(2, 3, 0, 1))),
    ]
    return float(max(errors))


def orbital_relations(result: SCFResult) -> dict[str, float]:
    c = result.occupied_coefficients
    h_occ = float(c @ result.integrals.core @ c)
    coulomb = float(
        np.einsum(
            "m,n,l,s,mnls",
            c,
            c,
            c,
            c,
            result.integrals.eri,
            optimize=True,
        )
    )
    return {
        "h_occ_Eh": h_occ,
        "J_occ_Eh": coulomb,
        "two_h_plus_J_Eh": 2.0 * h_occ + coulomb,
        "h_plus_J_Eh": h_occ + coulomb,
        "two_epsilon_minus_J_Eh": 2.0 * result.occupied_energy - coulomb,
    }


def validate_results(results: list[SCFResult]) -> dict[str, float | bool]:
    final = results[-1]
    relations = orbital_relations(final)
    energies = np.array([result.total_energy for result in results])
    monotonic_violation = float(np.max(np.diff(energies)))
    overlap_orthonormality = abs(
        float(
            final.occupied_coefficients
            @ final.integrals.overlap
            @ final.occupied_coefficients
        )
        - 1.0
    )
    density_idempotency = float(
        np.linalg.norm(
            final.density
            @ final.integrals.overlap
            @ final.density
            - 2.0 * final.density
        )
    )
    eri_error = integral_symmetry_error(final.integrals.eri)
    total_relation_error = abs(
        relations["two_h_plus_J_Eh"] - final.total_energy
    )
    orbital_relation_error = abs(
        relations["h_plus_J_Eh"] - final.occupied_energy
    )
    double_counting_error = abs(
        relations["two_epsilon_minus_J_Eh"] - final.total_energy
    )

    checks = {
        "basis_energy_nonincreasing": monotonic_violation < 5.0e-12,
        "overlap_normalization": overlap_orthonormality < 5.0e-13,
        "density_electron_count": abs(final.electron_count - 2.0) < 5.0e-12,
        "density_metric_idempotency": density_idempotency < 2.0e-11,
        "eri_permutation_symmetry": eri_error < 5.0e-13,
        "fock_commutator": final.commutator_norm < 2.0e-10,
        "generalized_eigenpair": final.eigenpair_residual < 2.0e-10,
        "closed_shell_total_relation": total_relation_error < 2.0e-11,
        "closed_shell_orbital_relation": orbital_relation_error < 2.0e-11,
        "coulomb_double_counting_relation": double_counting_error < 2.0e-11,
        "finite_basis_above_hf_limit": (
            final.total_energy >= HELIUM_HARTREE_FOCK_LIMIT
        ),
        "finite_basis_above_exact": final.total_energy >= HELIUM_EXACT_NONREL,
    }
    failed = [name for name, passed in checks.items() if not passed]
    if failed:
        raise AssertionError("validation failed: " + ", ".join(failed))

    return {
        **checks,
        "basis_monotonic_max_delta_Eh": monotonic_violation,
        "overlap_normalization_error": overlap_orthonormality,
        "density_idempotency_norm": density_idempotency,
        "eri_symmetry_max_error": eri_error,
        "total_relation_error_Eh": total_relation_error,
        "orbital_relation_error_Eh": orbital_relation_error,
        "double_counting_relation_error_Eh": double_counting_error,
        "hf_basis_error_Eh": final.total_energy - HELIUM_HARTREE_FOCK_LIMIT,
        "finite_basis_gap_to_exact_Eh": (
            final.total_energy - HELIUM_EXACT_NONREL
        ),
        "hartree_fock_correlation_energy_magnitude_Eh": (
            HELIUM_HARTREE_FOCK_LIMIT - HELIUM_EXACT_NONREL
        ),
    }


def write_basis_convergence(
    path: Path,
    results: list[SCFResult],
) -> None:
    fieldnames = [
        "basis_size",
        "exponents",
        "total_energy_Eh",
        "occupied_energy_Eh",
        "kinetic_Eh",
        "nuclear_Eh",
        "electron_repulsion_Eh",
        "iterations",
        "density_rms",
        "commutator_norm",
        "eigenpair_residual",
        "electron_count",
        "overlap_condition",
        "gap_to_HF_limit_Eh",
        "gap_to_exact_Eh",
    ]
    with path.open("w", newline="", encoding="utf-8") as handle:
        writer = csv.DictWriter(handle, fieldnames=fieldnames)
        writer.writeheader()
        for result in results:
            writer.writerow(
                {
                    "basis_size": len(result.exponents),
                    "exponents": ";".join(
                        f"{value:.12g}" for value in result.exponents
                    ),
                    "total_energy_Eh": f"{result.total_energy:.16g}",
                    "occupied_energy_Eh": f"{result.occupied_energy:.16g}",
                    "kinetic_Eh": f"{result.kinetic_energy:.16g}",
                    "nuclear_Eh": f"{result.nuclear_energy:.16g}",
                    "electron_repulsion_Eh": (
                        f"{result.electron_repulsion_energy:.16g}"
                    ),
                    "iterations": result.iterations,
                    "density_rms": f"{result.density_rms:.16g}",
                    "commutator_norm": f"{result.commutator_norm:.16g}",
                    "eigenpair_residual": (
                        f"{result.eigenpair_residual:.16g}"
                    ),
                    "electron_count": f"{result.electron_count:.16g}",
                    "overlap_condition": (
                        f"{result.overlap_condition:.16g}"
                    ),
                    "gap_to_HF_limit_Eh": (
                        f"{result.total_energy-HELIUM_HARTREE_FOCK_LIMIT:.16g}"
                    ),
                    "gap_to_exact_Eh": (
                        f"{result.total_energy-HELIUM_EXACT_NONREL:.16g}"
                    ),
                }
            )


def write_scf_history(path: Path, result: SCFResult) -> None:
    with path.open("w", newline="", encoding="utf-8") as handle:
        writer = csv.writer(handle)
        writer.writerow(
            [
                "iteration",
                "total_energy_Eh",
                "delta_energy_Eh",
                "density_rms",
                "commutator_norm",
            ]
        )
        for record in result.history:
            writer.writerow(
                [
                    record.iteration,
                    f"{record.total_energy:.16g}",
                    (
                        ""
                        if np.isnan(record.delta_energy)
                        else f"{record.delta_energy:.16g}"
                    ),
                    f"{record.density_rms:.16g}",
                    f"{record.commutator_norm:.16g}",
                ]
            )


def matrix_payload(result: SCFResult) -> dict[str, object]:
    relations = orbital_relations(result)
    return {
        "basis": {
            "type": "normalized uncontracted one-centre s Gaussians",
            "exponents": list(result.exponents),
        },
        "overlap": result.integrals.overlap.tolist(),
        "kinetic": result.integrals.kinetic.tolist(),
        "nuclear_attraction": result.integrals.nuclear.tolist(),
        "core_hamiltonian": result.integrals.core.tolist(),
        "density": result.density.tolist(),
        "fock": result.fock.tolist(),
        "occupied_coefficients": result.occupied_coefficients.tolist(),
        "occupied_energy_Eh": result.occupied_energy,
        "total_energy_Eh": result.total_energy,
        "electron_count": result.electron_count,
        "relations": relations,
    }


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--nuclear-charge", type=float, default=2.0)
    parser.add_argument(
        "--max-basis-size",
        type=int,
        default=len(PUBLISHED_EXPONENT_ORDER),
    )
    parser.add_argument("--density-tolerance", type=float, default=1.0e-12)
    parser.add_argument("--energy-tolerance", type=float, default=1.0e-13)
    parser.add_argument(
        "--output-dir",
        type=Path,
        default=Path("hartree-fock-helium-output"),
    )
    return parser.parse_args()


def main() -> None:
    args = parse_args()
    if args.nuclear_charge <= 0.0:
        raise ValueError("nuclear charge must be positive")
    if not 2 <= args.max_basis_size <= len(PUBLISHED_EXPONENT_ORDER):
        raise ValueError(
            f"max-basis-size must be between 2 and {len(PUBLISHED_EXPONENT_ORDER)}"
        )

    results = [
        run_scf(
            tuple(PUBLISHED_EXPONENT_ORDER[:basis_size]),
            nuclear_charge=args.nuclear_charge,
            density_tolerance=args.density_tolerance,
            energy_tolerance=args.energy_tolerance,
        )
        for basis_size in range(1, args.max_basis_size + 1)
    ]
    validation = validate_results(results)
    final = results[-1]
    minimal = results[1]

    args.output_dir.mkdir(parents=True, exist_ok=True)
    basis_path = args.output_dir / "hartree-fock-helium-basis-convergence.csv"
    history_path = args.output_dir / "hartree-fock-helium-scf-history.csv"
    matrices_path = args.output_dir / "hartree-fock-helium-minimal-matrices.json"
    metadata_path = args.output_dir / "hartree-fock-helium-metadata.json"

    write_basis_convergence(basis_path, results)
    write_scf_history(history_path, final)
    matrices_path.write_text(
        json.dumps(matrix_payload(minimal), indent=2, sort_keys=True) + "\n",
        encoding="utf-8",
    )
    metadata = {
        "model": "closed-shell helium RHF in one-centre s-Gaussian bases",
        "units": "Hartree atomic units",
        "nuclear_charge": args.nuclear_charge,
        "published_exponent_order": list(
            PUBLISHED_EXPONENT_ORDER[: args.max_basis_size]
        ),
        "final_basis_size": len(final.exponents),
        "minimal_demonstrator_basis_size": len(minimal.exponents),
        "density_tolerance": args.density_tolerance,
        "energy_tolerance_Eh": args.energy_tolerance,
        "final_total_energy_Eh": final.total_energy,
        "final_occupied_energy_Eh": final.occupied_energy,
        "final_iterations": final.iterations,
        "reference_values": {
            "helium_hartree_fock_limit": {
                "energy_Eh": HELIUM_HARTREE_FOCK_LIMIT,
                "source_doi": "10.2477/jccjie.2024-0032",
            },
            "helium_exact_nonrelativistic_clamped_nucleus": {
                "energy_Eh": HELIUM_EXACT_NONREL,
                "source_doi": "10.1002/qua.10344",
            },
        },
        "validation": validation,
        "python": platform.python_version(),
        "numpy": np.__version__,
        "platform": platform.platform(),
        "random_seed": None,
        "license": "MIT",
        "outputs": [
            basis_path.name,
            history_path.name,
            matrices_path.name,
            metadata_path.name,
        ],
    }
    metadata_path.write_text(
        json.dumps(metadata, indent=2, sort_keys=True) + "\n",
        encoding="utf-8",
    )

    print(f"minimal basis energy = {minimal.total_energy:.15f} Eh")
    print(f"final basis energy   = {final.total_energy:.15f} Eh")
    print(f"occupied energy      = {final.occupied_energy:.15f} Eh")
    print(f"SCF iterations       = {final.iterations}")
    print(f"electron count       = {final.electron_count:.15f}")
    print(f"commutator norm      = {final.commutator_norm:.6e}")
    print(f"eigenpair residual   = {final.eigenpair_residual:.6e}")
    print(f"gap to HF limit      = {validation['hf_basis_error_Eh']:.12e} Eh")
    print(f"outputs               = {args.output_dir.resolve()}")
    print("validation            = all checks passed")


if __name__ == "__main__":
    main()
