#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
tkcif_reader.py

Thin compatibility layer for CIF reading.

Public API
----------
read_structure(path)
    Return a pymatgen Structure object.

read_structure(path, return_info=True)
    Return (Structure, CIFReadInfo).

read_ase_atoms(path)
    Return an ASE Atoms object using ase.io.read.

read_ase_atoms(path, return_info=True)
    Return (ase.Atoms, CIFReadInfo).

read_tkcrystal(path)
    Return a legacy tkCrystal object using tkCIF.

read_tkcrystal(path, return_info=True)
    Return (tkCrystal, CIFReadInfo).

Backend order for read_structure()
----------------------------------
1. pymatgen direct CIF reading
2. tkcif_base normalization by tkcif_normalize.py, then pymatgen
3. ASE -> pymatgen Structure
4. legacy tkCIF -> tkCrystal -> pymatgen Structure

Notes
-----
- This module intentionally catches ImportError/Exception for optional backends.
- Legacy tkCIF support is a best-effort scaffold. You may need to adjust
  tkcrystal_to_pymatgen_structure() for your exact tkCrystal API.
"""

from __future__ import annotations

import importlib
import tempfile
import warnings
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Iterable


class CIFReadError(RuntimeError):
    """Raised when all CIF reading backends fail."""


@dataclass
class CIFReadInfo:
    path: str
    backend: str = ""
    attempted_backends: list[str] = field(default_factory=list)
    unavailable_backends: list[str] = field(default_factory=list)
    errors: list[str] = field(default_factory=list)
    warnings: list[str] = field(default_factory=list)
    normalized: bool = False
    encoding: str | None = None
    fallback_object: Any | None = None

    def add_error(self, backend: str, exc: BaseException) -> None:
        self.errors.append(f"{backend}: {type(exc).__name__}: {exc}")

    def add_unavailable(self, backend: str, reason: str) -> None:
        self.unavailable_backends.append(f"{backend}: {reason}")

    def short_report(self) -> str:
        lines = [
            f"path: {self.path}",
            f"backend: {self.backend or '(none)'}",
            f"attempted: {', '.join(self.attempted_backends) if self.attempted_backends else '(none)'}",
        ]
        if self.unavailable_backends:
            lines.append("unavailable:")
            lines.extend(f"  - {x}" for x in self.unavailable_backends)
        if self.warnings:
            lines.append("warnings:")
            lines.extend(f"  - {x}" for x in self.warnings)
        if self.errors:
            lines.append("errors:")
            lines.extend(f"  - {x}" for x in self.errors)
        return "\n".join(lines)


def read_structure(
    path: str | Path,
    *,
    primitive: bool = False,
    return_info: bool = False,
    backend_order: Iterable[str] | None = None,
) -> Any:
    """
    Read one CIF file and return a pymatgen Structure object.

    Parameters
    ----------
    path:
        CIF path.
    primitive:
        Passed to pymatgen parser. False preserves the CIF cell.
    return_info:
        If True, return (structure, CIFReadInfo).
        If False, return structure only.
    backend_order:
        Optional backend names.
        Available names:
            "pymatgen"
            "tkcif_base+pymatgen"
            "ase+pymatgen"
            "tkcif_legacy+pymatgen"
    """
    path = Path(path)
    info = CIFReadInfo(path=str(path))

    if backend_order is None:
        backend_order = (
            "pymatgen",
            "tkcif_base+pymatgen",
            "ase+pymatgen",
            "tkcif_legacy+pymatgen",
        )

    for backend_name in backend_order:
        info.attempted_backends.append(backend_name)
        try:
            if backend_name == "pymatgen":
                structure, backend_warnings = _read_structure_pymatgen(path, primitive=primitive)
                info.backend = backend_name
                info.warnings.extend(backend_warnings)
                return (structure, info) if return_info else structure

            if backend_name == "tkcif_base+pymatgen":
                structure, backend_warnings, encoding = _read_structure_tkcif_base_pymatgen(
                    path,
                    primitive=primitive,
                )
                info.backend = backend_name
                info.normalized = True
                info.encoding = encoding
                info.warnings.extend(backend_warnings)
                return (structure, info) if return_info else structure


            if backend_name == "ase+pymatgen":
                atoms, ase_info = read_ase_atoms(path, return_info=True)
                structure = ase_atoms_to_pymatgen_structure(atoms)
                info.backend = backend_name
                info.fallback_object = atoms
                info.warnings.extend(ase_info.warnings)
                info.errors.extend(ase_info.errors)
                return (structure, info) if return_info else structure

            if backend_name == "tkcif_legacy+pymatgen":
                cry, legacy_info = read_tkcrystal(path, return_info=True)
                structure = tkcrystal_to_pymatgen_structure(cry)
                info.backend = backend_name
                info.fallback_object = cry
                info.warnings.extend(legacy_info.warnings)
                info.errors.extend(legacy_info.errors)
                return (structure, info) if return_info else structure

            info.add_unavailable(backend_name, "unknown backend name")

        except ImportError as exc:
            info.add_unavailable(backend_name, str(exc))
            continue
        except Exception as exc:
            info.add_error(backend_name, exc)
            continue

    raise CIFReadError(
        "Failed to read CIF by all available backends.\n" + info.short_report()
    )


def read_ase_atoms(
    path: str | Path,
    *,
    return_info: bool = False,
    index: int | str | None = None,
    format: str | None = None,
) -> Any:
    """
    Read one CIF file using ASE and return an ase.Atoms object.

    This function is independent of pymatgen.  It can be used as a diagnostic
    or as a fallback source for read_structure() through ase+pymatgen.

    Parameters
    ----------
    path:
        CIF path.
    return_info:
        If True, return (ase.Atoms, CIFReadInfo).
    index:
        Passed to ase.io.read.  Default None usually means the first readable
        configuration.  Use ":" if you want to experiment with all structures
        outside this wrapper.
    format:
        Optional ASE format.  If omitted, ASE guesses from the file name.
    """
    path = Path(path)
    info = CIFReadInfo(path=str(path), backend="ase")

    try:
        from ase.io import read as ase_read
    except Exception as exc:
        raise ImportError(f"ase.io.read is not available: {exc}") from exc

    try:
        kwargs = {}
        if index is not None:
            kwargs["index"] = index
        if format is not None:
            kwargs["format"] = format

        with warnings.catch_warnings(record=True) as wlist:
            warnings.simplefilter("always")
            atoms = ase_read(str(path), **kwargs)
            info.warnings.extend(f"{w.category.__name__}: {w.message}" for w in wlist)

        if isinstance(atoms, list):
            if not atoms:
                raise CIFReadError("ASE returned an empty list.")
            info.warnings.append("ASE returned multiple Atoms objects; the first one was used.")
            atoms = atoms[0]

        if atoms is None:
            raise CIFReadError("ASE returned None.")

        return (atoms, info) if return_info else atoms

    except Exception as exc:
        info.add_error("ase", exc)
        raise



def read_tkcrystal(
    path: str | Path,
    *,
    return_info: bool = False,
    module_candidates: Iterable[str] | None = None,
) -> Any:
    """
    Read one CIF file using legacy tkCIF and return a tkCrystal object.

    This function is independent of pymatgen. It is the emergency path when
    pymatgen cannot be used at all.
    """
    path = Path(path)
    info = CIFReadInfo(path=str(path), backend="tkcif_legacy")

    if module_candidates is None:
        module_candidates = (
            "tklib.tkcrystal.tkcif",
            "tkcrystal.tkcif",
            "tkcif",
        )

    tkCIF = None
    last_import_error: Exception | None = None
    for modname in module_candidates:
        try:
            mod = importlib.import_module(modname)
            tkCIF = getattr(mod, "tkCIF")
            break
        except Exception as exc:
            last_import_error = exc
            info.add_unavailable(modname, f"{type(exc).__name__}: {exc}")

    if tkCIF is None:
        raise ImportError(
            "Could not import legacy tkCIF module. "
            f"Last error: {last_import_error}"
        )

    cif = None
    try:
        cif = tkCIF(str(path))

        try:
            cifdata = cif.ReadCIF(find_valid_structure=1, print_level=0)
        except TypeError:
            # Older or slightly different legacy signature.
            cifdata = cif.ReadCIF()

        if cifdata is None:
            raise CIFReadError("legacy tkCIF returned None from ReadCIF().")

        cry = cifdata.GetCrystal()
        if cry is None:
            raise CIFReadError("legacy tkCIF returned None from GetCrystal().")

        return (cry, info) if return_info else cry

    except Exception as exc:
        info.add_error("tkcif_legacy", exc)
        raise

    finally:
        if cif is not None:
            try:
                cif.Close()
            except Exception:
                pass


def normalize_cif(path: str | Path) -> tuple[str, Any, list[str], str]:
    """
    Normalize one CIF by tkcif_normalize.py.

    Returns
    -------
    normalized_text, document, warnings, encoding
    """
    try:
        from .tkcif_normalize import normalize_cif_file
    except Exception as exc:
        raise ImportError(
            "tkcif.tkcif_normalize.py is not importable. "
            "Check that tkcif is a package and tkcif_normalize.py exists."
        ) from exc

    return normalize_cif_file(Path(path), self_check=True)


def _read_structure_pymatgen(path: Path, *, primitive: bool) -> tuple[Any, list[str]]:
    CifParser = _import_cif_parser()

    captured_warnings: list[str] = []
    with warnings.catch_warnings(record=True) as wlist:
        warnings.simplefilter("always")

        parser = CifParser(str(path))
        structures = _parse_structures_compat(parser, primitive=primitive)

        captured_warnings.extend(
            f"{w.category.__name__}: {w.message}" for w in wlist
        )

    parser_warnings = getattr(parser, "warnings", None)
    if parser_warnings:
        captured_warnings.extend(str(x) for x in parser_warnings)

    if not structures:
        raise CIFReadError("pymatgen returned no structures.")

    return structures[0], captured_warnings


def _read_structure_tkcif_base_pymatgen(
    path: Path,
    *,
    primitive: bool,
) -> tuple[Any, list[str], str]:
    normalized_text, _doc, norm_warnings, encoding = normalize_cif(path)
    structure, pymatgen_warnings = _parse_pymatgen_from_text(
        normalized_text,
        primitive=primitive,
    )
    warnings_all = []
    warnings_all.extend(f"tkcif_base: {x}" for x in norm_warnings)
    warnings_all.extend(f"pymatgen: {x}" for x in pymatgen_warnings)
    return structure, warnings_all, encoding


def _parse_pymatgen_from_text(cif_text: str, *, primitive: bool) -> tuple[Any, list[str]]:
    CifParser = _import_cif_parser()

    captured_warnings: list[str] = []
    with warnings.catch_warnings(record=True) as wlist:
        warnings.simplefilter("always")

        if hasattr(CifParser, "from_str"):
            parser = CifParser.from_str(cif_text)
            structures = _parse_structures_compat(parser, primitive=primitive)
        else:
            tmp_path: str | None = None
            try:
                with tempfile.NamedTemporaryFile(
                    mode="w",
                    suffix=".cif",
                    encoding="utf-8",
                    delete=False,
                ) as fp:
                    fp.write(cif_text)
                    tmp_path = fp.name
                parser = CifParser(tmp_path)
                structures = _parse_structures_compat(parser, primitive=primitive)
            finally:
                if tmp_path is not None:
                    try:
                        Path(tmp_path).unlink()
                    except OSError:
                        pass

        captured_warnings.extend(
            f"{w.category.__name__}: {w.message}" for w in wlist
        )

    parser_warnings = getattr(parser, "warnings", None)
    if parser_warnings:
        captured_warnings.extend(str(x) for x in parser_warnings)

    if not structures:
        raise CIFReadError("pymatgen returned no structures from normalized CIF text.")

    return structures[0], captured_warnings


def _import_cif_parser() -> Any:
    try:
        from pymatgen.io.cif import CifParser
    except Exception as exc:
        raise ImportError(f"pymatgen.io.cif.CifParser is not available: {exc}") from exc
    return CifParser


def _parse_structures_compat(parser: Any, *, primitive: bool) -> list[Any]:
    if hasattr(parser, "parse_structures"):
        return list(parser.parse_structures(primitive=primitive))

    if hasattr(parser, "get_structures"):
        return list(parser.get_structures(primitive=primitive))

    raise CIFReadError("Unsupported pymatgen CifParser: no parse_structures() or get_structures().")


def ase_atoms_to_pymatgen_structure(atoms: Any) -> Any:
    """
    Convert ASE Atoms to pymatgen Structure.

    Preferred path:
        pymatgen.io.ase.AseAtomsAdaptor.get_structure(atoms)

    Fallback path:
        manually construct Structure from atoms.cell, atoms.get_chemical_symbols(),
        atoms.get_scaled_positions(), and periodic boundary flags.

    Notes
    -----
    - CIF occupancy/disorder information may already be simplified by ASE.
    - If the ASE Atoms object is not 3D periodic, this function still tries to
      create a Structure, but appends a warning only at the read_ase_atoms()
      stage.  For molecules or clusters, pymatgen Molecule may be more suitable.
    """
    try:
        from pymatgen.io.ase import AseAtomsAdaptor
        return AseAtomsAdaptor.get_structure(atoms)
    except Exception:
        pass

    try:
        from pymatgen.core import Lattice, Structure
    except Exception as exc:
        raise ImportError(f"pymatgen.core.Lattice/Structure is not available: {exc}") from exc

    try:
        cell = atoms.cell
        lattice = Lattice(cell.array)
    except Exception as exc:
        raise CIFReadError("ASE Atoms -> pymatgen conversion failed: invalid atoms.cell.") from exc

    try:
        species = atoms.get_chemical_symbols()
    except Exception as exc:
        raise CIFReadError("ASE Atoms -> pymatgen conversion failed: get_chemical_symbols() failed.") from exc

    try:
        frac_coords = atoms.get_scaled_positions(wrap=False)
    except Exception as exc:
        raise CIFReadError("ASE Atoms -> pymatgen conversion failed: get_scaled_positions() failed.") from exc

    if not species:
        raise CIFReadError("ASE Atoms contains no atoms.")

    return Structure(
        lattice,
        species,
        frac_coords,
        coords_are_cartesian=False,
        to_unit_cell=False,
        validate_proximity=False,
    )



def tkcrystal_to_pymatgen_structure(cry: Any) -> Any:
    """
    Convert legacy tkCrystal to pymatgen Structure.

    This is intentionally written as an adapter function so that you can adjust
    it to the exact tkCrystal API without touching read_structure().

    Expected tkCrystal-like methods:
        cry.LatticeParameters()
        cry.AtomSiteList()

    Expected atom-site-like methods:
        site.AtomNameOnly()
        site.Position()
        site.Occupancy()
    """
    try:
        from pymatgen.core import Lattice, Structure
    except Exception as exc:
        raise ImportError(f"pymatgen.core.Lattice/Structure is not available: {exc}") from exc

    try:
        latt = cry.LatticeParameters()
    except Exception as exc:
        raise CIFReadError(
            "tkCrystal -> pymatgen conversion failed: "
            "cry.LatticeParameters() is not available."
        ) from exc

    if latt is None or len(latt) < 6:
        raise CIFReadError(f"Invalid lattice parameters from tkCrystal: {latt!r}")

    a, b, c, alpha, beta, gamma = [float(x) for x in latt[:6]]
    lattice = Lattice.from_parameters(a, b, c, alpha, beta, gamma)

    try:
        atom_sites = list(cry.AtomSiteList())
    except Exception as exc:
        raise CIFReadError(
            "tkCrystal -> pymatgen conversion failed: "
            "cry.AtomSiteList() is not available."
        ) from exc

    species: list[Any] = []
    frac_coords: list[list[float]] = []

    for i, site in enumerate(atom_sites):
        try:
            name = site.AtomNameOnly()
        except Exception:
            try:
                name = site.AtomName()
            except Exception as exc:
                raise CIFReadError(f"Cannot get atom species for tkCrystal site #{i}.") from exc

        try:
            pos = site.Position()
        except Exception as exc:
            raise CIFReadError(f"Cannot get fractional position for tkCrystal site #{i}.") from exc

        try:
            occ = float(site.Occupancy())
        except Exception:
            occ = 1.0

        if name is None or str(name).strip() == "":
            raise CIFReadError(f"Empty atom species for tkCrystal site #{i}.")

        try:
            xyz = [float(pos[0]), float(pos[1]), float(pos[2])]
        except Exception as exc:
            raise CIFReadError(f"Invalid fractional position for tkCrystal site #{i}: {pos!r}") from exc

        if abs(occ - 1.0) < 1.0e-12:
            species.append(str(name))
        else:
            species.append({str(name): occ})
        frac_coords.append(xyz)

    if not species:
        raise CIFReadError("tkCrystal contains no atom sites.")

    return Structure(
        lattice,
        species,
        frac_coords,
        coords_are_cartesian=False,
        to_unit_cell=False,
        validate_proximity=False,
    )


def structure_summary_text(structure: Any, *, symprec: float = 1.0e-3) -> str:
    """Return a human-readable summary for a pymatgen Structure."""
    lines: list[str] = []
    comp = structure.composition
    latt = structure.lattice
    a, b, c = latt.abc
    alpha, beta, gamma = latt.angles

    lines.append(f"Formula                  : {comp.formula}")
    lines.append(f"Reduced formula          : {comp.reduced_formula}")
    lines.append(f"Number of sites          : {len(structure)}")
    lines.append(f"Volume [A^3]             : {float(latt.volume):.10g}")
    try:
        lines.append(f"Density [g/cm^3]         : {float(structure.density):.10g}")
    except Exception:
        pass

    lines.append("")
    lines.append("Lattice parameters")
    lines.append(f"  a, b, c [A]             : {a:.10g}, {b:.10g}, {c:.10g}")
    lines.append(f"  alpha, beta, gamma [deg]: {alpha:.10g}, {beta:.10g}, {gamma:.10g}")

    try:
        from pymatgen.symmetry.analyzer import SpacegroupAnalyzer

        sga = SpacegroupAnalyzer(structure, symprec=symprec, angle_tolerance=5.0)
        lines.append("")
        lines.append("Space group estimated by pymatgen/spglib")
        lines.append(f"  Symbol                  : {sga.get_space_group_symbol()}")
        lines.append(f"  Number                  : {sga.get_space_group_number()}")
        lines.append(f"  Crystal system          : {sga.get_crystal_system()}")
    except Exception as exc:
        lines.append("")
        lines.append("Space group estimated by pymatgen/spglib")
        lines.append(f"  Failed                  : {type(exc).__name__}: {exc}")

    lines.append("")
    lines.append("Atomic sites")
    lines.append("  #    species                 frac_x        frac_y        frac_z")
    for i, site in enumerate(structure.sites):
        fx, fy, fz = site.frac_coords
        lines.append(f"  {i:4d} {species_string(site):22s} {fx:12.7f} {fy:12.7f} {fz:12.7f}")

    return "\n".join(lines)


def species_string(site: Any) -> str:
    try:
        parts = []
        for sp, occ in site.species.items():
            occ_f = float(occ)
            if abs(occ_f - 1.0) < 1.0e-12:
                parts.append(str(sp))
            else:
                parts.append(f"{sp}:{occ_f:.4g}")
        return ",".join(parts)
    except Exception:
        return str(getattr(site, "species_string", getattr(site, "species", "?")))


def ase_atoms_summary_text(atoms: Any) -> str:
    """
    Best-effort summary for ASE Atoms.
    """
    lines: list[str] = []
    try:
        formula = atoms.get_chemical_formula()
        lines.append(f"Formula                  : {formula}")
    except Exception:
        pass

    try:
        lines.append(f"Number of atoms          : {len(atoms)}")
    except Exception:
        pass

    try:
        cell = atoms.cell
        lengths = cell.lengths()
        angles = cell.angles()
        volume = atoms.get_volume()
        lines.append(f"Volume [A^3]             : {float(volume):.10g}")
        lines.append("")
        lines.append("Cell parameters")
        lines.append(f"  a, b, c [A]             : {lengths[0]:.10g}, {lengths[1]:.10g}, {lengths[2]:.10g}")
        lines.append(f"  alpha, beta, gamma [deg]: {angles[0]:.10g}, {angles[1]:.10g}, {angles[2]:.10g}")
    except Exception as exc:
        lines.append(f"Cell                     : failed: {type(exc).__name__}: {exc}")

    try:
        pbc = atoms.get_pbc()
        lines.append(f"Periodic boundary flags  : {list(bool(x) for x in pbc)}")
    except Exception:
        pass

    lines.append("")
    lines.append("Atomic sites")
    lines.append("  #    species                 frac_x        frac_y        frac_z")
    try:
        symbols = atoms.get_chemical_symbols()
        scaled = atoms.get_scaled_positions(wrap=False)
        for i, (sym, xyz) in enumerate(zip(symbols, scaled)):
            x, y, z = float(xyz[0]), float(xyz[1]), float(xyz[2])
            lines.append(f"  {i:4d} {str(sym):22s} {x:12.7f} {y:12.7f} {z:12.7f}")
    except Exception as exc:
        lines.append(f"Atomic positions         : failed: {type(exc).__name__}: {exc}")

    return "\n".join(lines)



def tkcrystal_summary_text(cry: Any) -> str:
    """Best-effort summary for legacy tkCrystal."""
    lines: list[str] = []

    for label, method in (
        ("Crystal name", "CrystalName"),
        ("Sample name", "SampleName"),
        ("Chemical formula", "ChemicalComposition"),
    ):
        try:
            val = getattr(cry, method)()
            lines.append(f"{label:<24s}: {val}")
        except Exception:
            pass

    try:
        latt = cry.LatticeParameters()
        lines.append("Lattice parameters")
        lines.append(f"  a, b, c, alpha, beta, gamma: {latt}")
    except Exception as exc:
        lines.append(f"Lattice parameters      : failed: {type(exc).__name__}: {exc}")

    try:
        sites = list(cry.AtomSiteList())
        lines.append(f"Number of atom sites     : {len(sites)}")
        lines.append("")
        lines.append("Atomic sites")
        lines.append("  #    species                 frac_x        frac_y        frac_z       occ")
        for i, site in enumerate(sites):
            try:
                name = site.AtomNameOnly()
            except Exception:
                name = "?"
            try:
                pos = site.Position()
                x, y, z = float(pos[0]), float(pos[1]), float(pos[2])
            except Exception:
                x, y, z = float("nan"), float("nan"), float("nan")
            try:
                occ = float(site.Occupancy())
            except Exception:
                occ = 1.0
            lines.append(f"  {i:4d} {str(name):22s} {x:12.7f} {y:12.7f} {z:12.7f} {occ:9.4f}")
    except Exception as exc:
        lines.append(f"AtomSiteList             : failed: {type(exc).__name__}: {exc}")

    return "\n".join(lines)
