#!/usr/bin/env python3
# SPDX-License-Identifier: MIT
"""Minimal LCAO molecular orbitals for the hydrogen molecular ion.

The program evaluates the analytic two-centre hydrogenic-1s overlap,
Coulomb, and resonance integrals, assembles the nonorthogonal Hamiltonian,
solves the generalized eigenproblem by symmetric orthogonalization, and
writes deterministic potential-curve, matrix, density, and validation data.

NumPy is the only dependency. Distances are in bohr and energies in hartree.
"""

from __future__ import annotations

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

import numpy as np


BENCHMARK_SEPARATION_BOHR = 2.0
ACCURATE_GROUND_RE_BOHR = 1.99719332
ACCURATE_GROUND_MINIMUM_EH = -0.6026346191
LCAO_GROUND_RE_BOHR = 2.492830425253425
LCAO_GROUND_MINIMUM_EH = -0.5648309923708077


@dataclass(frozen=True)
class MOResult:
    separation: float
    overlap_scalar: float
    coulomb_integral: float
    resonance_integral: float
    overlap: np.ndarray
    hamiltonian: np.ndarray
    electronic_g: float
    electronic_u: float
    total_g: float
    total_u: float
    coefficients_g: np.ndarray
    coefficients_u: np.ndarray
    residual_g: float
    residual_u: float
    metric_norm_g: float
    metric_norm_u: float
    overlap_condition: float
    analytic_electronic_g: float
    analytic_electronic_u: float


def analytic_integrals(
    separation: float,
) -> tuple[float, float, float]:
    """Return S, J, and K for normalized hydrogenic 1s orbitals."""
    if separation <= 0.0:
        raise ValueError("internuclear separation must be positive")
    exponential = math.exp(-separation)
    overlap = exponential * (
        1.0 + separation + separation * separation / 3.0
    )
    coulomb = (
        -1.0 / separation
        + math.exp(-2.0 * separation)
        * (1.0 + 1.0 / separation)
    )
    resonance = -exponential * (1.0 + separation)
    return overlap, coulomb, resonance


def build_matrices(
    separation: float,
) -> tuple[np.ndarray, np.ndarray, float, float, float]:
    overlap_scalar, coulomb, resonance = analytic_integrals(separation)
    diagonal = -0.5 + coulomb
    off_diagonal = -0.5 * overlap_scalar + resonance
    overlap = np.array(
        [
            [1.0, overlap_scalar],
            [overlap_scalar, 1.0],
        ],
        dtype=float,
    )
    hamiltonian = np.array(
        [
            [diagonal, off_diagonal],
            [off_diagonal, diagonal],
        ],
        dtype=float,
    )
    return overlap, hamiltonian, overlap_scalar, coulomb, resonance


def symmetric_orthogonalizer(
    overlap: np.ndarray,
    *,
    eigenvalue_floor: float = 1.0e-12,
) -> 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 canonicalize_parity_vector(
    coefficients: np.ndarray,
    parity: str,
) -> np.ndarray:
    result = np.asarray(coefficients, dtype=float).copy()
    if parity == "g":
        if float(np.sum(result)) < 0.0:
            result *= -1.0
    elif parity == "u":
        if result[0] < 0.0:
            result *= -1.0
    else:
        raise ValueError("parity must be 'g' or 'u'")
    return result


def solve_molecular_orbitals(separation: float) -> MOResult:
    (
        overlap,
        hamiltonian,
        overlap_scalar,
        coulomb,
        resonance,
    ) = build_matrices(separation)
    orthogonalizer, overlap_eigenvalues = symmetric_orthogonalizer(overlap)
    transformed = orthogonalizer.T @ hamiltonian @ orthogonalizer
    transformed = 0.5 * (transformed + transformed.T)
    eigenvalues, transformed_vectors = np.linalg.eigh(transformed)
    coefficients = orthogonalizer @ transformed_vectors

    states: dict[str, tuple[float, np.ndarray]] = {}
    for energy, vector in zip(eigenvalues, coefficients.T, strict=True):
        parity = "g" if vector[0] * vector[1] > 0.0 else "u"
        states[parity] = (
            float(energy),
            canonicalize_parity_vector(vector, parity),
        )
    if set(states) != {"g", "u"}:
        raise RuntimeError("failed to identify one gerade and one ungerade state")

    electronic_g, coefficients_g = states["g"]
    electronic_u, coefficients_u = states["u"]
    analytic_g = (
        hamiltonian[0, 0] + hamiltonian[0, 1]
    ) / (1.0 + overlap_scalar)
    analytic_u = (
        hamiltonian[0, 0] - hamiltonian[0, 1]
    ) / (1.0 - overlap_scalar)
    residual_g = float(
        np.linalg.norm(
            hamiltonian @ coefficients_g
            - overlap @ coefficients_g * electronic_g
        )
    )
    residual_u = float(
        np.linalg.norm(
            hamiltonian @ coefficients_u
            - overlap @ coefficients_u * electronic_u
        )
    )
    metric_norm_g = float(coefficients_g @ overlap @ coefficients_g)
    metric_norm_u = float(coefficients_u @ overlap @ coefficients_u)

    return MOResult(
        separation=separation,
        overlap_scalar=overlap_scalar,
        coulomb_integral=coulomb,
        resonance_integral=resonance,
        overlap=overlap,
        hamiltonian=hamiltonian,
        electronic_g=electronic_g,
        electronic_u=electronic_u,
        total_g=electronic_g + 1.0 / separation,
        total_u=electronic_u + 1.0 / separation,
        coefficients_g=coefficients_g,
        coefficients_u=coefficients_u,
        residual_g=residual_g,
        residual_u=residual_u,
        metric_norm_g=metric_norm_g,
        metric_norm_u=metric_norm_u,
        overlap_condition=float(
            overlap_eigenvalues[-1] / overlap_eigenvalues[0]
        ),
        analytic_electronic_g=float(analytic_g),
        analytic_electronic_u=float(analytic_u),
    )


def golden_section_minimum(
    function,
    lower: float,
    upper: float,
    *,
    tolerance: float = 1.0e-14,
    max_iterations: int = 256,
) -> tuple[float, float, int]:
    if not lower < upper:
        raise ValueError("minimum bracket must satisfy lower < upper")
    ratio = (math.sqrt(5.0) - 1.0) / 2.0
    left = upper - ratio * (upper - lower)
    right = lower + ratio * (upper - lower)
    value_left = function(left)
    value_right = function(right)

    for iteration in range(1, max_iterations + 1):
        if upper - lower < tolerance:
            break
        if value_left < value_right:
            upper = right
            right = left
            value_right = value_left
            left = upper - ratio * (upper - lower)
            value_left = function(left)
        else:
            lower = left
            left = right
            value_left = value_right
            right = lower + ratio * (upper - lower)
            value_right = function(right)
    else:
        raise RuntimeError("golden-section minimization did not converge")

    location = 0.5 * (lower + upper)
    return location, function(location), iteration


def hydrogenic_1s_axis(z: float, centre: float) -> float:
    return math.exp(-abs(z - centre)) / math.sqrt(math.pi)


def axis_density_rows(
    result: MOResult,
    *,
    extent: float,
    points: int,
) -> list[dict[str, float]]:
    if extent <= result.separation / 2.0:
        raise ValueError("axis extent must lie beyond both nuclei")
    if points < 3 or points % 2 == 0:
        raise ValueError("axis point count must be an odd integer at least 3")
    centre_a = -result.separation / 2.0
    centre_b = result.separation / 2.0
    rows: list[dict[str, float]] = []

    for z_value in np.linspace(-extent, extent, points):
        z = float(z_value)
        phi_a = hydrogenic_1s_axis(z, centre_a)
        phi_b = hydrogenic_1s_axis(z, centre_b)
        psi_g = float(
            result.coefficients_g[0] * phi_a
            + result.coefficients_g[1] * phi_b
        )
        psi_u = float(
            result.coefficients_u[0] * phi_a
            + result.coefficients_u[1] * phi_b
        )
        rows.append(
            {
                "z_bohr": z,
                "phi_A": phi_a,
                "phi_B": phi_b,
                "psi_g": psi_g,
                "psi_u": psi_u,
                "density_g_per_bohr3": psi_g * psi_g,
                "density_u_per_bohr3": psi_u * psi_u,
            }
        )
    return rows


def validate_results(
    curve_results: list[MOResult],
    benchmark: MOResult,
    minimum_location: float,
    minimum_energy: float,
    density_rows: list[dict[str, float]],
) -> dict[str, bool | float]:
    formula_error = max(
        max(
            abs(result.electronic_g - result.analytic_electronic_g),
            abs(result.electronic_u - result.analytic_electronic_u),
        )
        for result in curve_results
    )
    residual_max = max(
        max(result.residual_g, result.residual_u)
        for result in curve_results
    )
    metric_error = max(
        max(
            abs(result.metric_norm_g - 1.0),
            abs(result.metric_norm_u - 1.0),
        )
        for result in curve_results
    )
    parity_error = max(
        max(
            abs(result.coefficients_g[0] - result.coefficients_g[1]),
            abs(result.coefficients_u[0] + result.coefficients_u[1]),
        )
        for result in curve_results
    )
    matrix_symmetry_error = max(
        max(
            float(np.max(np.abs(result.overlap - result.overlap.T))),
            float(
                np.max(
                    np.abs(result.hamiltonian - result.hamiltonian.T)
                )
            ),
        )
        for result in curve_results
    )
    minimum_reference_error_r = abs(
        minimum_location - LCAO_GROUND_RE_BOHR
    )
    minimum_reference_error_e = abs(
        minimum_energy - LCAO_GROUND_MINIMUM_EH
    )
    midpoint = density_rows[len(density_rows) // 2]
    largest_condition = max(
        result.overlap_condition for result in curve_results
    )
    last = curve_results[-1]
    dissociation_error = max(
        abs(last.total_g + 0.5),
        abs(last.total_u + 0.5),
    )
    lcao_well_depth = -0.5 - minimum_energy
    accurate_well_depth = -0.5 - ACCURATE_GROUND_MINIMUM_EH

    checks = {
        "analytic_and_generalized_energies_agree": formula_error < 2.0e-13,
        "generalized_eigenpair_residuals": residual_max < 2.0e-13,
        "metric_normalization": metric_error < 2.0e-13,
        "symmetry_adapted_coefficients": parity_error < 5.0e-11,
        "matrix_symmetry": matrix_symmetry_error < 2.0e-15,
        "overlap_positive_definite": all(
            0.0 < result.overlap_scalar < 1.0
            for result in curve_results
        ),
        "gerade_below_ungerade": all(
            result.total_g < result.total_u
            for result in curve_results
        ),
        "benchmark_overlap": (
            abs(benchmark.overlap_scalar - 0.5864528940253216)
            < 2.0e-15
        ),
        "benchmark_gerade_energy": (
            abs(benchmark.total_g + 0.5537714953184827)
            < 2.0e-15
        ),
        "benchmark_ungerade_energy": (
            abs(benchmark.total_u + 0.16085396559668752)
            < 2.0e-15
        ),
        "minimum_location_reference": minimum_reference_error_r < 2.0e-10,
        "minimum_energy_reference": minimum_reference_error_e < 2.0e-13,
        "ungerade_midpoint_node": abs(midpoint["psi_u"]) < 2.0e-15,
        "gerade_midpoint_density_positive": (
            midpoint["density_g_per_bohr3"] > 0.0
        ),
        "dissociation_limit_at_grid_end": dissociation_error < 4.0e-6,
        "variational_minimum_above_accurate_reference": (
            minimum_energy > ACCURATE_GROUND_MINIMUM_EH
        ),
    }
    failed = [name for name, passed in checks.items() if not passed]
    if failed:
        raise AssertionError("validation failed: " + ", ".join(failed))

    return {
        **{
            name: bool(passed)
            for name, passed in checks.items()
        },
        "analytic_energy_max_error_Eh": float(formula_error),
        "generalized_residual_max": float(residual_max),
        "metric_normalization_max_error": float(metric_error),
        "parity_coefficient_max_error": float(parity_error),
        "matrix_symmetry_max_error": float(matrix_symmetry_error),
        "maximum_overlap_condition": float(largest_condition),
        "minimum_location_reference_error_bohr": (
            float(minimum_reference_error_r)
        ),
        "minimum_energy_reference_error_Eh": (
            float(minimum_reference_error_e)
        ),
        "dissociation_error_at_grid_end_Eh": float(dissociation_error),
        "lcao_well_depth_Eh": float(lcao_well_depth),
        "accurate_reference_well_depth_Eh": float(accurate_well_depth),
        "well_depth_fraction_recovered": (
            float(lcao_well_depth / accurate_well_depth)
        ),
        "equilibrium_distance_relative_error": (
            float(
                (minimum_location - ACCURATE_GROUND_RE_BOHR)
                / ACCURATE_GROUND_RE_BOHR
            )
        ),
    }


def write_curve_data(
    path: Path,
    results: list[MOResult],
) -> None:
    fieldnames = [
        "R_bohr",
        "overlap",
        "coulomb_J_Eh",
        "resonance_K_Eh",
        "electronic_g_Eh",
        "electronic_u_Eh",
        "total_g_Eh",
        "total_u_Eh",
        "splitting_Eh",
        "coefficient_g_A",
        "coefficient_g_B",
        "coefficient_u_A",
        "coefficient_u_B",
        "metric_norm_g",
        "metric_norm_u",
        "residual_g",
        "residual_u",
        "overlap_condition",
    ]
    with path.open("w", newline="", encoding="utf-8") as handle:
        writer = csv.DictWriter(handle, fieldnames=fieldnames)
        writer.writeheader()
        for result in results:
            writer.writerow(
                {
                    "R_bohr": f"{result.separation:.16g}",
                    "overlap": f"{result.overlap_scalar:.16g}",
                    "coulomb_J_Eh": f"{result.coulomb_integral:.16g}",
                    "resonance_K_Eh": (
                        f"{result.resonance_integral:.16g}"
                    ),
                    "electronic_g_Eh": f"{result.electronic_g:.16g}",
                    "electronic_u_Eh": f"{result.electronic_u:.16g}",
                    "total_g_Eh": f"{result.total_g:.16g}",
                    "total_u_Eh": f"{result.total_u:.16g}",
                    "splitting_Eh": (
                        f"{result.total_u-result.total_g:.16g}"
                    ),
                    "coefficient_g_A": (
                        f"{result.coefficients_g[0]:.16g}"
                    ),
                    "coefficient_g_B": (
                        f"{result.coefficients_g[1]:.16g}"
                    ),
                    "coefficient_u_A": (
                        f"{result.coefficients_u[0]:.16g}"
                    ),
                    "coefficient_u_B": (
                        f"{result.coefficients_u[1]:.16g}"
                    ),
                    "metric_norm_g": f"{result.metric_norm_g:.16g}",
                    "metric_norm_u": f"{result.metric_norm_u:.16g}",
                    "residual_g": f"{result.residual_g:.16g}",
                    "residual_u": f"{result.residual_u:.16g}",
                    "overlap_condition": (
                        f"{result.overlap_condition:.16g}"
                    ),
                }
            )


def write_density_data(
    path: Path,
    rows: list[dict[str, float]],
) -> None:
    fieldnames = list(rows[0])
    with path.open("w", newline="", encoding="utf-8") as handle:
        writer = csv.DictWriter(handle, fieldnames=fieldnames)
        writer.writeheader()
        for row in rows:
            writer.writerow(
                {
                    key: f"{value:.16g}"
                    for key, value in row.items()
                }
            )


def matrix_payload(result: MOResult) -> dict[str, object]:
    return {
        "system": "H2+ with clamped protons",
        "separation_bohr": result.separation,
        "basis": {
            "functions": [
                "normalized hydrogenic 1s on nucleus A",
                "normalized hydrogenic 1s on nucleus B",
            ],
            "orbital_exponent_per_bohr": 1.0,
        },
        "integrals": {
            "overlap": result.overlap_scalar,
            "coulomb_J_Eh": result.coulomb_integral,
            "resonance_K_Eh": result.resonance_integral,
        },
        "overlap_matrix": result.overlap.tolist(),
        "electronic_hamiltonian_matrix_Eh": (
            result.hamiltonian.tolist()
        ),
        "states": {
            "gerade": {
                "coefficients": result.coefficients_g.tolist(),
                "electronic_energy_Eh": result.electronic_g,
                "total_energy_Eh": result.total_g,
                "metric_norm": result.metric_norm_g,
                "generalized_residual": result.residual_g,
            },
            "ungerade": {
                "coefficients": result.coefficients_u.tolist(),
                "electronic_energy_Eh": result.electronic_u,
                "total_energy_Eh": result.total_u,
                "metric_norm": result.metric_norm_u,
                "generalized_residual": result.residual_u,
            },
        },
        "nuclear_repulsion_Eh": 1.0 / result.separation,
        "overlap_condition": result.overlap_condition,
        "analytic_cross_check": {
            "electronic_g_Eh": result.analytic_electronic_g,
            "electronic_u_Eh": result.analytic_electronic_u,
        },
    }


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--r-min", type=float, default=0.5)
    parser.add_argument("--r-max", type=float, default=15.0)
    parser.add_argument("--r-points", type=int, default=291)
    parser.add_argument("--density-extent", type=float, default=5.0)
    parser.add_argument("--density-points", type=int, default=401)
    parser.add_argument(
        "--output-dir",
        type=Path,
        default=Path("molecular-orbital-h2plus-output"),
    )
    return parser.parse_args()


def main() -> None:
    args = parse_args()
    if args.r_min <= 0.0 or args.r_max <= args.r_min:
        raise ValueError("require 0 < r-min < r-max")
    if args.r_points < 3:
        raise ValueError("r-points must be at least 3")

    separations = np.linspace(args.r_min, args.r_max, args.r_points)
    curve_results = [
        solve_molecular_orbitals(float(separation))
        for separation in separations
    ]
    benchmark = solve_molecular_orbitals(BENCHMARK_SEPARATION_BOHR)
    minimum_location, minimum_energy, minimum_iterations = (
        golden_section_minimum(
            lambda separation: (
                solve_molecular_orbitals(
                    separation
                ).analytic_electronic_g
                + 1.0 / separation
            ),
            1.0,
            5.0,
        )
    )
    density_rows = axis_density_rows(
        benchmark,
        extent=args.density_extent,
        points=args.density_points,
    )
    validation = validate_results(
        curve_results,
        benchmark,
        minimum_location,
        minimum_energy,
        density_rows,
    )

    args.output_dir.mkdir(parents=True, exist_ok=True)
    curve_path = (
        args.output_dir / "molecular-orbital-h2plus-curves.csv"
    )
    density_path = (
        args.output_dir / "molecular-orbital-h2plus-axis-density.csv"
    )
    matrices_path = (
        args.output_dir / "molecular-orbital-h2plus-matrices.json"
    )
    metadata_path = (
        args.output_dir / "molecular-orbital-h2plus-metadata.json"
    )

    write_curve_data(curve_path, curve_results)
    write_density_data(density_path, density_rows)
    matrices_path.write_text(
        json.dumps(matrix_payload(benchmark), indent=2, sort_keys=True)
        + "\n",
        encoding="utf-8",
    )
    metadata = {
        "model": (
            "H2+ minimal LCAO with fixed hydrogenic 1s orbitals"
        ),
        "units": {
            "distance": "bohr",
            "energy": "hartree",
            "axis_density": "inverse bohr cubed",
        },
        "basis": {
            "size": 2,
            "type": "normalized two-centre hydrogenic 1s Slater orbitals",
            "orbital_exponent_per_bohr": 1.0,
        },
        "grid": {
            "r_min_bohr": args.r_min,
            "r_max_bohr": args.r_max,
            "r_points": args.r_points,
            "density_separation_bohr": BENCHMARK_SEPARATION_BOHR,
            "density_extent_bohr": args.density_extent,
            "density_points": args.density_points,
        },
        "lcao_minimum": {
            "R_bohr": minimum_location,
            "total_energy_Eh": minimum_energy,
            "well_depth_Eh": -0.5 - minimum_energy,
            "optimizer": "golden-section search on 1 <= R <= 5",
            "iterations": minimum_iterations,
        },
        "reference_values": {
            "accurate_ground_minimum": {
                "R_bohr": ACCURATE_GROUND_RE_BOHR,
                "total_energy_Eh": ACCURATE_GROUND_MINIMUM_EH,
                "source_doi": "10.1002/slct.202102509",
            },
        },
        "validation": validation,
        "python": platform.python_version(),
        "numpy": np.__version__,
        "platform": platform.platform(),
        "random_seed": None,
        "license": "MIT",
        "outputs": [
            curve_path.name,
            density_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(
        "R=2 gerade total energy   = "
        f"{benchmark.total_g:.15f} Eh"
    )
    print(
        "R=2 ungerade total energy = "
        f"{benchmark.total_u:.15f} Eh"
    )
    print(
        "minimal-LCAO minimum      = "
        f"{minimum_energy:.15f} Eh at R={minimum_location:.12f} bohr"
    )
    print(
        "well-depth fraction       = "
        f"{validation['well_depth_fraction_recovered']:.6f}"
    )
    print(
        "max generalized residual  = "
        f"{validation['generalized_residual_max']:.6e}"
    )
    print(
        "max analytic energy error = "
        f"{validation['analytic_energy_max_error_Eh']:.6e} Eh"
    )
    print(f"outputs                    = {args.output_dir.resolve()}")
    print("validation                 = all checks passed")


if __name__ == "__main__":
    main()
