"""
Copyright (c) Meta Platforms, Inc. and affiliates.

This source code is licensed under the MIT license found in the
LICENSE file in the root directory of this source tree.
"""

from __future__ import annotations

import json
import tempfile
from dataclasses import dataclass
from functools import partial
from pathlib import Path

import ase.io
import numpy as np
import numpy.testing as npt
import pandas as pd
import pytest
from ase import units
from ase.build import bulk
from ase.calculators.emt import EMT
from ase.io import Trajectory
from ase.md.velocitydistribution import MaxwellBoltzmannDistribution
from ase.md.verlet import VelocityVerlet

from fairchem.core.components.calculate import (
    BerendsenNPT,
    BussiThermostat,
    LangevinThermostat,
    MDRunner,
    NoseHooverNVT,
    ParquetTrajectoryWriter,
    TrajectoryFrame,
    VelocityVerletThermostat,
)


@dataclass
class MockMetadata:
    results_dir: str
    checkpoint_dir: str = ""
    preemption_checkpoint_dir: str = ""
    config_path: str = ""
    array_job_num: int = 0


@dataclass
class MockScheduler:
    num_array_jobs: int = 1


@dataclass
class MockJobConfig:
    metadata: MockMetadata
    scheduler: MockScheduler


def _create_mock_job_config(
    results_dir: str,
    checkpoint_dir: str = "",
    preemption_checkpoint_dir: str = "",
) -> MockJobConfig:
    return MockJobConfig(
        metadata=MockMetadata(
            results_dir=results_dir,
            checkpoint_dir=checkpoint_dir,
            preemption_checkpoint_dir=preemption_checkpoint_dir or checkpoint_dir,
        ),
        scheduler=MockScheduler(num_array_jobs=1),
    )


@pytest.fixture()
def cu_atoms():
    atoms = bulk("Cu", cubic=True) * (2, 2, 2)
    atoms.info["sid"] = 123
    np.random.seed(42)
    MaxwellBoltzmannDistribution(atoms, temperature_K=300)
    return atoms


@pytest.fixture()
def results_dir():
    with tempfile.TemporaryDirectory() as tmpdir:
        yield Path(tmpdir)


class TestMDRunner:
    def test_md_correctness_vs_ase(self, cu_atoms, results_dir):
        """
        Verify MDRunner produces identical trajectories to plain ASE.
        """
        atoms_mdrunner = cu_atoms.copy()
        atoms_ase = cu_atoms.copy()

        np.random.seed(42)
        MaxwellBoltzmannDistribution(atoms_mdrunner, temperature_K=300)
        np.random.seed(42)
        MaxwellBoltzmannDistribution(atoms_ase, temperature_K=300)

        steps, interval = 20, 5

        mdrunner_dir = results_dir / "mdrunner"
        mdrunner_dir.mkdir()
        runner = MDRunner(
            calculator=EMT(),
            atoms=atoms_mdrunner,
            thermostat=VelocityVerletThermostat(),
            timestep_fs=1.0,
            steps=steps,
            trajectory_interval=interval,
            log_interval=10,
            trajectory_writer=partial(ParquetTrajectoryWriter, flush_interval=100),
        )
        runner._job_config = _create_mock_job_config(str(mdrunner_dir))
        results = runner.calculate(job_num=0, num_jobs=1)

        # Reference ASE run
        ase_traj_file = results_dir / "ase_traj.traj"
        atoms_ase.calc = EMT()
        dyn_ase = VelocityVerlet(atoms_ase, timestep=1.0 * units.fs)
        traj_ase = Trajectory(str(ase_traj_file), "w", atoms_ase)
        dyn_ase.attach(traj_ase.write, interval=interval)
        dyn_ase.run(steps)
        traj_ase.close()

        traj_df = pd.read_parquet(results["trajectory_file"])
        ase_frames = Trajectory(str(ase_traj_file), "r")
        assert len(traj_df) == len(ase_frames)
        assert (mdrunner_dir / "init_atoms.extxyz").exists()

        for i, ase_atoms in enumerate(ase_frames):
            row = traj_df.iloc[i]
            npt.assert_allclose(
                np.vstack(row["positions"]), ase_atoms.get_positions(), atol=1e-10
            )
            npt.assert_allclose(
                np.vstack(row["velocities"]), ase_atoms.get_velocities(), atol=1e-10
            )
            npt.assert_allclose(
                row["energy"], ase_atoms.get_potential_energy(), atol=1e-10
            )

    @pytest.mark.parametrize(
        "thermostat",
        [
            VelocityVerletThermostat(),
            NoseHooverNVT(temperature_K=300.0, tdamp_fs=25.0),
            BussiThermostat(temperature_K=300.0, taut_fs=25.0),
            LangevinThermostat(temperature_K=300.0, friction_per_fs=0.01),
            BerendsenNPT(
                temperature_K=300.0,
                pressure_bar=1.0,
                taut_fs=500.0,
                taup_fs=1000.0,
                compressibility_bar=1.0 / 140e9,
            ),
        ],
        ids=["VelocityVerlet", "NoseHoover", "Bussi", "Langevin", "BerendsenNPT"],
    )
    @pytest.mark.parametrize(
        "interrupt_at_step",
        [36, 40],
        ids=["non-aligned", "aligned"],
    )
    def test_checkpoint_resume(
        self, cu_atoms, results_dir, thermostat, interrupt_at_step
    ):
        """
        Checkpoint at both a non-aligned and an interval-aligned step, resume,
        and verify the concatenated trajectory stays on a uniform grid with no
        duplicated frames across the restart boundary (regardless of where
        preemption lands relative to trajectory_interval).
        """
        results_dir1 = results_dir / "results1"
        results_dir2 = results_dir / "results2"
        checkpoint_dir = results_dir / "checkpoint"
        results_dir1.mkdir()
        results_dir2.mkdir()

        trajectory_interval = 10
        total_steps = 100

        runner1 = MDRunner(
            calculator=EMT(),
            atoms=cu_atoms.copy(),
            thermostat=thermostat,
            timestep_fs=1.0,
            steps=total_steps,
            trajectory_interval=trajectory_interval,
            log_interval=10,
            trajectory_writer=partial(ParquetTrajectoryWriter, flush_interval=1000),
        )
        runner1._job_config = _create_mock_job_config(str(results_dir1))

        class SimulatedInterrupt(Exception):
            pass

        def interrupt_callback():
            if runner1._dyn.get_number_of_steps() >= interrupt_at_step:
                raise SimulatedInterrupt

        # Run 1: manually drive dynamics so we can attach an interrupt
        try:
            runner1._atoms.calc = runner1.calculator
            runner1._dyn = thermostat.build(runner1._atoms, timestep_fs=1.0)
            parquet_file1 = results_dir1 / "trajectory.parquet"
            runner1._trajectory_writer = ParquetTrajectoryWriter(
                parquet_file1, flush_interval=1000
            )

            def collect_frame():
                step = runner1._dyn.get_number_of_steps()
                if step % trajectory_interval == 0:
                    frame = TrajectoryFrame.from_atoms(
                        runner1._atoms, step=step, time=runner1._dyn.get_time()
                    )
                    runner1._trajectory_writer.append(frame)

            runner1._dyn.attach(collect_frame, interval=1)
            runner1._dyn.attach(interrupt_callback, interval=1)
            runner1._dyn.run(total_steps)
        except SimulatedInterrupt:
            final_positions = runner1._atoms.get_positions().copy()
            final_velocities = runner1._atoms.get_velocities().copy()
            runner1.save_state(str(checkpoint_dir), is_preemption=True)

        df1 = pd.read_parquet(parquet_file1)
        steps1 = list(df1["step"])
        # Run 1 writes every interval-aligned step up to (and including) the
        # interrupt step if that step is itself aligned.
        expected_steps1 = list(range(0, interrupt_at_step + 1, trajectory_interval))
        assert steps1 == expected_steps1

        # Verify checkpoint files
        assert (checkpoint_dir / "checkpoint.extxyz").exists()
        assert (checkpoint_dir / "thermostat_state.json").exists()
        checkpoint_atoms = ase.io.read(str(checkpoint_dir / "checkpoint.extxyz"))
        assert checkpoint_atoms.info["md_step"] == interrupt_at_step

        # Run 2: resume from checkpoint
        runner2 = MDRunner(
            calculator=EMT(),
            atoms=cu_atoms.copy(),
            thermostat=thermostat,
            timestep_fs=1.0,
            steps=total_steps,
            trajectory_interval=trajectory_interval,
            log_interval=10,
            trajectory_writer=partial(ParquetTrajectoryWriter, flush_interval=1000),
        )
        runner2._job_config = _create_mock_job_config(str(results_dir2))
        runner2.load_state(str(checkpoint_dir))

        assert runner2._start_step == interrupt_at_step
        npt.assert_allclose(runner2._atoms.get_positions(), final_positions, atol=1e-8)
        npt.assert_allclose(
            runner2._atoms.get_velocities(), final_velocities, atol=1e-8
        )

        results2 = runner2.calculate(job_num=0, num_jobs=1)
        runner2.write_results(results2, str(results_dir2), job_num=0, num_jobs=1)
        df2 = pd.read_parquet(results2["trajectory_file"])
        steps2 = list(df2["step"])

        # No frame is written at the resume step, so the concatenated
        # trajectory stays on a uniform grid with no duplicates, whether or not
        # the checkpoint step was interval-aligned.
        assert len(steps2) == len(set(steps2))
        all_steps = sorted(steps1 + steps2)
        assert len(all_steps) == len(
            set(all_steps)
        ), f"Duplicate frame(s) across restart boundary: {all_steps}"
        expected = [0, 10, 20, 30, 40, 50, 60, 70, 80, 90, 100]
        assert all_steps == expected, f"Expected {expected}, got {all_steps}"

        # Verify write_results produced valid metadata JSON
        metadata_file = results_dir2 / "metadata.json"
        assert metadata_file.exists()
        with open(metadata_file) as f:
            metadata = json.load(f)
        assert metadata["structure_id"] == cu_atoms.info["sid"]
        assert metadata["total_steps"] == total_steps

    def test_load_state_beyond_total_steps(self, cu_atoms, results_dir):
        """
        When checkpoint step >= total steps, load_state should warn
        and the simulation should end immediately without error.
        """
        checkpoint_dir = results_dir / "checkpoint"
        checkpoint_dir.mkdir()
        md_results_dir = results_dir / "results"
        md_results_dir.mkdir()

        # Create a fake checkpoint with current_step beyond what we'll configure
        atoms = cu_atoms.copy()
        atoms.info["md_step"] = 50
        ase.io.write(str(checkpoint_dir / "checkpoint.extxyz"), atoms, format="extxyz")

        with open(checkpoint_dir / "md_state.json", "w") as f:
            json.dump({"current_step": 50, "total_steps": 100}, f)

        runner = MDRunner(
            calculator=EMT(),
            atoms=cu_atoms.copy(),
            thermostat=VelocityVerletThermostat(),
            timestep_fs=1.0,
            steps=30,  # less than checkpoint step of 50
            trajectory_interval=10,
            log_interval=10,
            trajectory_writer=partial(ParquetTrajectoryWriter, flush_interval=1000),
        )
        runner._job_config = _create_mock_job_config(str(md_results_dir))

        # load_state should not raise, just warn
        runner.load_state(str(checkpoint_dir))
        assert runner._start_step == 50
        assert runner._already_calculated

    def test_stopfair_graceful_stop(self, cu_atoms, results_dir):
        """
        Verify STOPFAIR triggers graceful stop, saves state, and deletes
        the sentinel file.
        """
        md_results_dir = results_dir / "results"
        checkpoint_dir = results_dir / "checkpoints"
        md_results_dir.mkdir()
        checkpoint_dir.mkdir()

        runner = MDRunner(
            calculator=EMT(),
            atoms=cu_atoms.copy(),
            thermostat=NoseHooverNVT(temperature_K=300.0, tdamp_fs=25.0),
            timestep_fs=1.0,
            steps=100,
            trajectory_interval=10,
            heartbeat_interval=20,
            log_interval=10,
            trajectory_writer=partial(ParquetTrajectoryWriter, flush_interval=1000),
        )
        runner._job_config = _create_mock_job_config(
            str(md_results_dir), checkpoint_dir=str(checkpoint_dir)
        )

        stopfair_path = checkpoint_dir.parent / "STOPFAIR"
        stopfair_path.write_text("")

        results = runner.calculate(job_num=0, num_jobs=1)

        assert results["stopped_by_stopfair"] is True
        assert (checkpoint_dir / "checkpoint.extxyz").exists()
        assert (checkpoint_dir / "md_state.json").exists()
        assert not stopfair_path.exists()

        with open(checkpoint_dir / "md_state.json") as f:
            md_state = json.load(f)
        assert md_state["current_step"] == 20

        traj_df = pd.read_parquet(results["trajectory_file"])
        assert list(traj_df["step"]) == [0, 10, 20]

    def test_npt_cell_changes(self, cu_atoms, results_dir):
        """
        Verify that NPT simulation changes the cell volume.
        """
        md_results_dir = results_dir / "results"
        md_results_dir.mkdir()

        atoms = cu_atoms.copy()
        initial_volume = atoms.get_volume()

        # Use a large pressure to drive a noticeable volume change
        thermostat = BerendsenNPT(
            temperature_K=300.0,
            pressure_bar=1e5,
            taut_fs=100.0,
            taup_fs=100.0,
            compressibility_bar=1.0 / 140e9,
        )

        runner = MDRunner(
            calculator=EMT(),
            atoms=atoms,
            thermostat=thermostat,
            timestep_fs=1.0,
            steps=200,
            trajectory_interval=50,
            log_interval=50,
            trajectory_writer=partial(ParquetTrajectoryWriter, flush_interval=1000),
        )
        runner._job_config = _create_mock_job_config(str(md_results_dir))
        runner.calculate(job_num=0, num_jobs=1)

        final_volume = atoms.get_volume()
        assert initial_volume != pytest.approx(final_volume, rel=1e-6)


class TestTrajectoryFrame:
    def test_round_trip_oc20_slab_adsorbate(self, results_dir):
        """
        Run MD on an OC20-style slab+adsorbate and verify the parquet
        trajectory preserves tags, FixAtoms constraints, and charge/spin.
        """
        from ase.constraints import FixAtoms

        from fairchem.core.datasets.common_structures import get_slab_adsorbate

        atoms = get_slab_adsorbate()

        assert np.any(atoms.get_tags() != 0)
        assert len(atoms.constraints) > 0
        assert "charge" in atoms.info
        assert "spin" in atoms.info

        original_tags = atoms.get_tags().copy()
        original_fixed_indices = sorted(atoms.constraints[0].index)

        atoms.calc = EMT()

        md_results_dir = results_dir / "results"
        md_results_dir.mkdir()

        runner = MDRunner(
            calculator=atoms.calc,
            atoms=atoms,
            thermostat=VelocityVerletThermostat(),
            timestep_fs=1.0,
            steps=10,
            trajectory_interval=5,
            log_interval=5,
            trajectory_writer=partial(ParquetTrajectoryWriter, flush_interval=100),
        )
        runner._job_config = _create_mock_job_config(str(md_results_dir))
        results = runner.calculate(job_num=0, num_jobs=1)

        traj_df = pd.read_parquet(results["trajectory_file"])
        assert len(traj_df) > 0

        for _, row in traj_df.iterrows():
            d = row.to_dict()

            # Reconstruct atoms directly from parquet row
            def _to_array(val, dtype=float):
                if isinstance(val, np.ndarray) and val.dtype == object:
                    return np.stack(val).astype(dtype)
                return np.asarray(val, dtype=dtype)

            from ase import Atoms as _Atoms

            reconstructed = _Atoms(
                numbers=_to_array(d["atomic_numbers"], dtype=int),
                positions=_to_array(d["positions"]),
                cell=_to_array(d["cell"]),
                pbc=_to_array(d["pbc"], dtype=bool),
            )
            if d.get("tags") is not None:
                reconstructed.set_tags(_to_array(d["tags"], dtype=int))
            if d.get("velocities") is not None:
                reconstructed.set_velocities(_to_array(d["velocities"]))
            if d.get("fixed") is not None:
                reconstructed.constraints = [
                    FixAtoms(indices=np.where(_to_array(d["fixed"], dtype=bool))[0])
                ]
            if d.get("charge") is not None:
                reconstructed.info["charge"] = d["charge"]
            if d.get("spin") is not None:
                reconstructed.info["spin"] = d["spin"]

            npt.assert_array_equal(reconstructed.get_tags(), original_tags)

            assert len(reconstructed.constraints) == 1
            assert isinstance(reconstructed.constraints[0], FixAtoms)
            npt.assert_array_equal(
                sorted(reconstructed.constraints[0].index), original_fixed_indices
            )

            assert reconstructed.info["charge"] == atoms.info["charge"]
            assert reconstructed.info["spin"] == atoms.info["spin"]
