#!/usr/bin/env python3
from __future__ import annotations

import argparse
import itertools
import importlib.util
import io
import math
import sys
from dataclasses import dataclass, asdict
from pathlib import Path
from typing import Iterable

import matplotlib
import numpy as np
import pandas as pd
import scipy.signal
from openpyxl import Workbook
from openpyxl.styles import Font

import matplotlib.pyplot as plt


"""
Examples
--------
Default: run peak search and lattice-parameter guessing
    python guess_lattice_parameters.py sample.TXT

Peak search only: save peak information to an Excel file
    python guess_lattice_parameters.py sample.TXT --mode search --threshold 1000

Guess only: read a peak-information file saved by --mode search
    python guess_lattice_parameters.py sample-peaks.xlsx --mode guess --crystal-system cubic

Legacy positional mode is also accepted:
    python guess_lattice_parameters.py search sample.TXT
"""


CU_KA1 = 1.54056
CU_KA2 = 1.54439
ANGLE_EPS = 1.0e-12

SUPPORTED_CRYSTAL_SYSTEMS = ["hexagonal", "orthorhombic", "tetragonal", "cubic"]
CRYSTAL_SYSTEM_ALIASES = {
    "auto": "auto",
    "all": "auto",
    "any": "auto",
    "hex": "hexagonal",
    "hexagonal": "hexagonal",
    "orth": "orthorhombic",
    "ortho": "orthorhombic",
    "orthorhombic": "orthorhombic",
    "tet": "tetragonal",
    "tetragonal": "tetragonal",
    "cub": "cubic",
    "cubic": "cubic",
}
KNOWN_MODES = {"search", "guess", "all", "compare", "calc", "index", "seaerch"}
SUPPORTED_METHODS = ["combinatorial", "grid", "both"]


@dataclass
class Peak:
    index: int
    two_theta: float
    intensity: float
    intensity_raw: float
    fwhm_deg: float
    inv_d2: float
    d: float
    q: float
    ka_role: str = ""
    ka_pair_index: int | None = None
    source: str = "detected"


@dataclass
class Candidate:
    crystal_system: str
    ls_code: int
    params: dict[str, float]
    score_matches: int
    score_rms_rel: float
    selected_matches: list[dict]
    all_matches: list[dict]


def bragg_d(two_theta_deg: float, wavelength: float = CU_KA1) -> float:
    theta = math.radians(two_theta_deg / 2.0)
    s = math.sin(theta)
    if s <= 0:
        return float("inf")
    return wavelength / (2.0 * s)


def inv_d2_from_two_theta(two_theta_deg: float, wavelength: float = CU_KA1) -> float:
    d = bragg_d(two_theta_deg, wavelength)
    if not math.isfinite(d) or d <= 0:
        return float("nan")
    return 1.0 / (d * d)


def two_theta_from_d(d: float, wavelength: float = CU_KA1) -> float | None:
    x = wavelength / (2.0 * d)
    if not (0.0 < x < 1.0):
        return None
    return math.degrees(2.0 * math.asin(x))


def two_theta_from_inv_d2(inv_d2: float, wavelength: float = CU_KA1) -> float | None:
    if inv_d2 <= 0:
        return None
    return two_theta_from_d(1.0 / math.sqrt(inv_d2), wavelength)


def sanitize_stem(text: str) -> str:
    return "".join(ch if ch.isalnum() or ch in "-_." else "_" for ch in text).strip("_") or "output"


def read_xy_data(path: Path) -> tuple[np.ndarray, np.ndarray]:
    suffix = path.suffix.lower()
    if suffix in {".xlsx", ".xls"}:
        df = pd.read_excel(path)
        if df.shape[1] < 2:
            raise ValueError("Need at least two columns in spreadsheet.")
        x = pd.to_numeric(df.iloc[:, 0], errors="coerce").to_numpy()
        y = pd.to_numeric(df.iloc[:, 1], errors="coerce").to_numpy()
    else:
        df = pd.read_csv(path, sep=None, engine="python", header=None, comment="#")
        if df.shape[1] < 2:
            raise ValueError("Need at least two columns in text data.")
        x = pd.to_numeric(df.iloc[:, 0], errors="coerce").to_numpy()
        y = pd.to_numeric(df.iloc[:, 1], errors="coerce").to_numpy()
    mask = np.isfinite(x) & np.isfinite(y)
    x = x[mask].astype(float)
    y = y[mask].astype(float)
    order = np.argsort(x)
    x = x[order]
    y = y[order]
    return x, y


def _interp_zero(x1: float, y1: float, x2: float, y2: float) -> float:
    if abs(y2 - y1) <= ANGLE_EPS:
        return 0.5 * (x1 + x2)
    return float(x1 - y1 * (x2 - x1) / (y2 - y1))


def _find_zero_crossing_bounds(
    x: np.ndarray,
    ydiff3: np.ndarray,
    idx: int,
) -> tuple[tuple[int, float] | None, tuple[int, float] | None]:
    left_info = None
    right_info = None

    for i in range(idx, 0, -1):
        y1 = float(ydiff3[i - 1])
        y2 = float(ydiff3[i])
        if y1 == 0.0 or y1 * y2 <= 0.0:
            xz = _interp_zero(float(x[i - 1]), y1, float(x[i]), y2)
            left_info = (i, xz)
            break

    for i in range(idx + 1, len(x)):
        y1 = float(ydiff3[i - 1])
        y2 = float(ydiff3[i])
        if y2 == 0.0 or y1 * y2 <= 0.0:
            xz = _interp_zero(float(x[i - 1]), y1, float(x[i]), y2)
            right_info = (i, xz)
            break

    return left_info, right_info


def estimate_fwhm(
    x: np.ndarray,
    ysmooth: np.ndarray,
    idx: int,
    ydiff3: np.ndarray | None = None,
) -> float | None:
    """
    Estimate FWHM with priority on local regions bounded by third-derivative zero crossings.
    """
    ytop = float(ysmooth[idx])
    if not math.isfinite(ytop) or ytop <= 0.0:
        return None

    left_info = None
    right_info = None
    if ydiff3 is not None:
        left_info, right_info = _find_zero_crossing_bounds(x, ydiff3, idx)

    if left_info is not None and right_info is not None:
        il, xl0 = left_info
        ir, xr0 = right_info
        if il < idx < ir:
            ybg_left = float(ysmooth[max(il - 1, 0)])
            ybg_right = float(ysmooth[min(ir, len(ysmooth) - 1)])
            ybg = min(ybg_left, ybg_right)
            half = ybg + 0.5 * (ytop - ybg)

            xl = None
            for i in range(idx, il - 1, -1):
                if i <= 0:
                    break
                y1 = float(ysmooth[i - 1])
                y2 = float(ysmooth[i])
                if (y1 - half) == 0.0 or (y1 - half) * (y2 - half) <= 0.0:
                    xl = _interp_zero(float(x[i - 1]), y1 - half, float(x[i]), y2 - half)
                    break

            xr = None
            for i in range(idx + 1, ir + 1):
                if i >= len(x):
                    break
                y1 = float(ysmooth[i - 1])
                y2 = float(ysmooth[i])
                if (y2 - half) == 0.0 or (y1 - half) * (y2 - half) <= 0.0:
                    xr = _interp_zero(float(x[i - 1]), y1 - half, float(x[i]), y2 - half)
                    break

            if xl is not None and xr is not None and xr > xl:
                return float(xr - xl)

            width_like = float(xr0 - xl0)
            if width_like > 0.0:
                return width_like

    lo = max(0, idx - 10)
    hi = min(len(x), idx + 11)
    local_bg = float(np.min(ysmooth[lo:hi]))
    half = local_bg + 0.5 * (ytop - local_bg)

    il = idx
    while il > 0 and ysmooth[il] > half:
        il -= 1
    ir = idx
    while ir < len(x) - 1 and ysmooth[ir] > half:
        ir += 1

    if il >= idx or ir <= idx:
        return None

    if abs(float(ysmooth[il + 1]) - float(ysmooth[il])) <= ANGLE_EPS:
        xl = float(x[il])
    else:
        xl = float(np.interp(half, [ysmooth[il], ysmooth[il + 1]], [x[il], x[il + 1]]))

    if abs(float(ysmooth[ir]) - float(ysmooth[ir - 1])) <= ANGLE_EPS:
        xr = float(x[ir])
    else:
        xr = float(np.interp(half, [ysmooth[ir - 1], ysmooth[ir]], [x[ir - 1], x[ir]]))

    return float(xr - xl) if xr > xl else None


def peak_search_deriv3(
    x: np.ndarray,
    y: np.ndarray,
    nsmooth: int = 11,
    norder: int = 4,
    threshold: float = 1000.0,
    ydiff1_threshold: float = 1.0e-2,
    fwhm_min_deg: float = 0.06,
    fwhm_max_deg: float = 1.0,
    is_print: bool = True,
) -> tuple[list[Peak], dict[str, np.ndarray | float]]:
    if len(x) < max(nsmooth + 2, 7):
        raise ValueError("Not enough points.")
    if nsmooth % 2 == 0:
        nsmooth += 1
    if nsmooth <= norder:
        nsmooth = norder + 3 if (norder + 3) % 2 == 1 else norder + 4

    h = float(np.median(np.diff(x)))
    ysmooth = scipy.signal.savgol_filter(y, nsmooth, norder, deriv=0)
    ydiff1 = scipy.signal.savgol_filter(y, nsmooth, norder, deriv=1) / h
    ydiff2 = scipy.signal.savgol_filter(ydiff1, nsmooth, norder, deriv=1) / h
    ydiff3 = scipy.signal.savgol_filter(ydiff2, nsmooth, norder, deriv=1) / h
    ydiff3 = scipy.signal.savgol_filter(ydiff3, nsmooth, norder, deriv=0)

    diff1_ratio = np.abs(ydiff1 / np.maximum(ysmooth, 1.0e-5))
    max_diff1_ratio = float(np.max(diff1_ratio))
    diff1_ratio_th = max_diff1_ratio * ydiff1_threshold

    peaks: list[Peak] = []
    last_accept_idx = -10**9
    visited_idx: set[int] = set()

    if is_print:
        print("")
        print("=== Peak search diagnostics ===")
        print(f"nsmooth          : {nsmooth}")
        print(f"norder           : {norder}")
        print(f"threshold        : {threshold}")
        print(f"FWHM window      : {fwhm_min_deg} - {fwhm_max_deg} deg")
        print(f"|dy/dx|/y thresh : {diff1_ratio_th:.6g} (diagnostic only)")
        print("")

    for i in range(1, len(x)):
        if ydiff3[i - 1] == 0.0 or ydiff3[i - 1] * ydiff3[i] <= 0.0:
            lo = max(0, i - max(3, nsmooth // 2))
            hi = min(len(x), i + max(4, nsmooth // 2 + 1))
            idx = lo + int(np.argmax(ysmooth[lo:hi]))
            if idx in visited_idx:
                continue
            visited_idx.add(idx)
            if idx - last_accept_idx <= 1:
                continue

            ytop = float(ysmooth[idx])
            d1ratio = float(abs(ydiff1[idx]) / max(ysmooth[idx], 1.0e-5))
            d2 = float(ydiff2[idx])

            reason = None
            fwhm = estimate_fwhm(x, ysmooth, idx, ydiff3=ydiff3)
            if ytop < threshold:
                reason = "too weak"
            elif d2 >= 0.0 and (fwhm is None or fwhm < fwhm_min_deg):
                reason = "minimum-like"
            elif fwhm is None:
                reason = "fwhm unavailable"
            elif fwhm < fwhm_min_deg:
                reason = f"FWHM too small ({fwhm:.4f})"
            elif fwhm > fwhm_max_deg:
                reason = f"FWHM too large ({fwhm:.4f})"

            if is_print:
                status = "accepted" if reason is None else f"excluded: {reason}"
                print(
                    f"2theta={x[idx]:8.3f}  I={ytop:10.3f}  FWHM={fwhm if fwhm is not None else float('nan'):7.4f}  "
                    f"|dy/dx|/y={d1ratio:10.5g}  d2={d2:10.5g}  -> {status}"
                )

            if reason is not None:
                continue

            inv_d2 = inv_d2_from_two_theta(float(x[idx]), CU_KA1)
            d = 1.0 / math.sqrt(inv_d2)
            q = 2.0 * math.pi / d
            peaks.append(
                Peak(
                    index=int(idx),
                    two_theta=float(x[idx]),
                    intensity=float(ysmooth[idx]),
                    intensity_raw=float(y[idx]),
                    fwhm_deg=float(fwhm),
                    inv_d2=float(inv_d2),
                    d=float(d),
                    q=float(q),
                )
            )
            last_accept_idx = idx

    # merge very close duplicates
    merged: list[Peak] = []
    for p in peaks:
        if merged and abs(p.two_theta - merged[-1].two_theta) < max(p.fwhm_deg, merged[-1].fwhm_deg) * 0.4:
            if p.intensity > merged[-1].intensity:
                merged[-1] = p
        else:
            merged.append(p)

    info = {
        "ysmooth": ysmooth,
        "ydiff1": ydiff1,
        "ydiff2": ydiff2,
        "ydiff3": ydiff3,
        "max_diff1_ratio": max_diff1_ratio,
        "diff1_ratio_th": diff1_ratio_th,
    }
    return merged, info


def assign_ka2(peaks: list[Peak], tol_deg: float = 0.06, ratio_min: float = 0.18, ratio_max: float = 0.75) -> list[Peak]:
    for p in peaks:
        p.ka_role = ""
        p.ka_pair_index = None
    # stronger/narrower first
    order = sorted(range(len(peaks)), key=lambda i: peaks[i].intensity, reverse=True)
    used_secondary = set()
    for i in order:
        p1 = peaks[i]
        d = bragg_d(p1.two_theta, CU_KA1)
        tt2 = two_theta_from_d(d, CU_KA2)
        if tt2 is None:
            continue
        best_j = None
        best_delta = None
        for j, p2 in enumerate(peaks):
            if j == i or j in used_secondary:
                continue
            if p2.two_theta <= p1.two_theta:
                continue
            delta = abs(p2.two_theta - tt2)
            ratio = p2.intensity / max(p1.intensity, 1e-9)
            if delta <= tol_deg and ratio_min <= ratio <= ratio_max:
                if best_delta is None or delta < best_delta:
                    best_delta = delta
                    best_j = j
        if best_j is not None:
            peaks[i].ka_role = "ka1"
            peaks[i].ka_pair_index = best_j
            peaks[best_j].ka_role = "ka2"
            peaks[best_j].ka_pair_index = i
            used_secondary.add(best_j)
    return peaks


def system_rows(system: str, hmax: int = 12) -> list[tuple[tuple[int, ...], tuple[int, int, int]]]:
    rows: list[tuple[tuple[int, ...], tuple[int, int, int]]] = []
    seen: set[tuple[int, ...]] = set()
    for h in range(hmax + 1):
        for k in range(hmax + 1):
            for l in range(hmax + 1):
                if h == 0 and k == 0 and l == 0:
                    continue
                hkl = (h, k, l)
                if system == "cubic":
                    feat = (h * h + k * k + l * l,)
                    hkl = tuple(sorted((h, k, l), reverse=True))
                elif system == "tetragonal":
                    feat = (h * h + k * k, l * l)
                    hk = tuple(sorted((h, k), reverse=True))
                    hkl = (hk[0], hk[1], l)
                elif system == "hexagonal":
                    feat = (h * h + h * k + k * k, l * l)
                elif system == "orthorhombic":
                    feat = (h * h, k * k, l * l)
                else:
                    raise ValueError(system)
                if all(v == 0 for v in feat):
                    continue
                if feat in seen:
                    continue
                seen.add(feat)
                rows.append((feat, hkl))
    rows.sort(key=lambda t: (sum(t[0]), t[0]))
    return rows


def params_from_beta(system: str, beta: np.ndarray) -> dict[str, float]:
    if system == "cubic":
        a = 1.0 / math.sqrt(beta[0])
        return {"a": a, "b": a, "c": a, "alpha": 90.0, "beta": 90.0, "gamma": 90.0}
    if system == "tetragonal":
        a = 1.0 / math.sqrt(beta[0])
        c = 1.0 / math.sqrt(beta[1])
        return {"a": a, "b": a, "c": c, "alpha": 90.0, "beta": 90.0, "gamma": 90.0}
    if system == "hexagonal":
        a = math.sqrt(4.0 / (3.0 * beta[0]))
        c = 1.0 / math.sqrt(beta[1])
        return {"a": a, "b": a, "c": c, "alpha": 90.0, "beta": 90.0, "gamma": 120.0}
    if system == "orthorhombic":
        a = 1.0 / math.sqrt(beta[0])
        b = 1.0 / math.sqrt(beta[1])
        c = 1.0 / math.sqrt(beta[2])
        return {"a": a, "b": b, "c": c, "alpha": 90.0, "beta": 90.0, "gamma": 90.0}
    raise ValueError(system)


def predicted_inv_d2(system: str, feat: Iterable[int], beta: np.ndarray) -> float:
    if system in {"cubic", "tetragonal", "hexagonal", "orthorhombic"}:
        return float(np.dot(np.asarray(tuple(feat), dtype=float), beta))
    raise ValueError(system)


def score_system(system: str, selected: list[Peak], all_peaks: list[Peak], hmax: int = 12) -> Candidate | None:
    rows = system_rows(system, hmax=hmax)
    p = len(rows[0][0])
    pool = min(len(rows), 12)
    first_n = min(len(selected), max(p + 2, 5))
    if len(selected) < p:
        return None

    ls_map = {"orthorhombic": 4, "tetragonal": 5, "cubic": 6, "hexagonal": 8}
    best = None
    obs = np.array([pk.inv_d2 for pk in selected], dtype=float)
    for obs_idx in __import__("itertools").combinations(range(first_n), p):
        y = np.array([obs[i] for i in obs_idx], dtype=float)
        for cand_idx in __import__("itertools").combinations(range(pool), p):
            X = np.array([rows[i][0] for i in cand_idx], dtype=float)
            try:
                beta = np.linalg.solve(X, y)
            except np.linalg.LinAlgError:
                continue
            if np.any(beta <= 0):
                continue
            pred_lines = np.array([predicted_inv_d2(system, feat, beta) for feat, _hkl in rows], dtype=float)
            selected_matches = []
            used = set()
            rels = []
            for pk in selected:
                rel = np.abs(pred_lines - pk.inv_d2) / np.maximum(pk.inv_d2, 1e-12)
                j = int(np.argmin(rel))
                if rel[j] <= 0.08 and j not in used:
                    used.add(j)
                    selected_matches.append({
                        "two_theta_obs_deg": pk.two_theta,
                        "inv_d2_obs": pk.inv_d2,
                        "h": rows[j][1][0],
                        "k": rows[j][1][1],
                        "l": rows[j][1][2],
                        "inv_d2_calc": float(pred_lines[j]),
                        "rel_error_inv_d2": float(rel[j]),
                    })
                    rels.append(float(rel[j]))
            if not selected_matches:
                continue
            rms = float(np.sqrt(np.mean(np.square(rels))))
            score = (len(selected_matches), -rms)
            if best is None or score > best["score"]:
                # assign all peaks
                all_matches = []
                for pk in all_peaks:
                    rel = np.abs(pred_lines - pk.inv_d2) / np.maximum(pk.inv_d2, 1e-12)
                    j = int(np.argmin(rel))
                    matched = rel[j] <= 0.10
                    tt_calc = two_theta_from_inv_d2(float(pred_lines[j]), CU_KA1) if matched else None
                    rel_t = abs((tt_calc - pk.two_theta) / pk.two_theta) if (matched and tt_calc and pk.two_theta != 0) else None
                    all_matches.append({
                        "two_theta_obs_deg": pk.two_theta,
                        "intensity": pk.intensity,
                        "ka_role": pk.ka_role,
                        "fwhm_deg": pk.fwhm_deg,
                        "inv_d2_obs": pk.inv_d2,
                        "h": rows[j][1][0] if matched else None,
                        "k": rows[j][1][1] if matched else None,
                        "l": rows[j][1][2] if matched else None,
                        "inv_d2_calc": float(pred_lines[j]) if matched else None,
                        "rel_error_inv_d2": float(rel[j]) if matched else None,
                        "two_theta_calc_deg": tt_calc if matched else None,
                        "rel_error_2theta": rel_t if matched else None,
                        "matched": bool(matched),
                    })
                best = {
                    "score": score,
                    "beta": beta,
                    "params": params_from_beta(system, beta),
                    "selected_matches": selected_matches,
                    "all_matches": all_matches,
                }

    if best is None:
        return None
    return Candidate(
        crystal_system=system,
        ls_code=ls_map[system],
        params=best["params"],
        score_matches=best["score"][0],
        score_rms_rel=float(-best["score"][1]),
        selected_matches=best["selected_matches"],
        all_matches=best["all_matches"],
    )


def import_lsq_module(path: Path):
    spec = importlib.util.spec_from_file_location("lsq_latt2_imported", path)
    if spec is None or spec.loader is None:
        raise ImportError(f"Cannot load {path}")
    mod = importlib.util.module_from_spec(spec)
    sys.modules[spec.name] = mod
    spec.loader.exec_module(mod)
    return mod


def refine_with_lsq(lsq_script: Path, candidate: Candidate, matched_rows: list[dict], wavelength: float = CU_KA1):
    mod = import_lsq_module(lsq_script)
    obs_list = []
    serial = 0
    for row in matched_rows:
        if not row.get("matched"):
            continue
        serial += 1
        h, k, l = int(row["h"]), int(row["k"]), int(row["l"])
        design = mod.lattice_design_row(candidate.ls_code, h, k, l)
        obs_list.append(
            mod.Observation(
                serial_index=serial,
                h=h,
                k=k,
                l=l,
                raw_position=float(row["two_theta_obs_deg"]),
                transformed_position=float(row["inv_d2_obs"]),
                weight=1.0,
                design_row=design,
                wavelength=wavelength,
            )
        )
    if not obs_list:
        return None

    fit = mod.weighted_least_squares(obs_list, io.StringIO())
    cell = mod.derive_cell_constants(candidate.ls_code, fit.parameters, fit.parameter_sigma)

    refined_rows = []
    for obs, fitted, resid in zip(fit.kept_observations, fit.fitted_y, fit.residuals):
        tt_calc = two_theta_from_inv_d2(float(fitted), wavelength)
        rel_inv = abs(resid) / max(obs.transformed_position, 1e-12)
        rel_tt = abs((tt_calc - obs.raw_position) / obs.raw_position) if tt_calc is not None and obs.raw_position != 0 else None
        refined_rows.append({
            "serial": obs.serial_index,
            "h": obs.h,
            "k": obs.k,
            "l": obs.l,
            "two_theta_obs_deg": obs.raw_position,
            "two_theta_calc_deg": tt_calc,
            "rel_error_2theta": rel_tt,
            "inv_d2_obs": obs.transformed_position,
            "inv_d2_calc": float(fitted),
            "rel_error_inv_d2": float(rel_inv),
            "residual_inv_d2": float(resid),
        })

    rms_inv = float(np.sqrt(np.mean([r["rel_error_inv_d2"] ** 2 for r in refined_rows]))) if refined_rows else None
    rms_tt = float(np.sqrt(np.mean([r["rel_error_2theta"] ** 2 for r in refined_rows if r["rel_error_2theta"] is not None]))) if refined_rows else None

    summary = {
        "crystal_system": candidate.crystal_system,
        "ls_code": candidate.ls_code,
        "a": float(cell.direct_lengths[0]),
        "b": float(cell.direct_lengths[1]),
        "c": float(cell.direct_lengths[2]),
        "alpha": float(cell.direct_angles_deg[0]),
        "beta": float(cell.direct_angles_deg[1]),
        "gamma": float(cell.direct_angles_deg[2]),
        "sigma_a": float(cell.direct_length_sigma[0]),
        "sigma_b": float(cell.direct_length_sigma[1]),
        "sigma_c": float(cell.direct_length_sigma[2]),
        "sigma_alpha": float(cell.direct_angle_sigma_deg[0]),
        "sigma_beta": float(cell.direct_angle_sigma_deg[1]),
        "sigma_gamma": float(cell.direct_angle_sigma_deg[2]),
        "n_used": len(refined_rows),
        "rms_rel_error_inv_d2": rms_inv,
        "rms_rel_error_2theta": rms_tt,
    }
    return summary, refined_rows



def _canonical_column_name(name: object) -> str:
    s = str(name).strip().lower()
    s = s.replace("θ", "theta")
    s = s.replace("°", "deg")
    for ch in " -/.()[]{}":
        s = s.replace(ch, "_")
    while "__" in s:
        s = s.replace("__", "_")
    return s.strip("_")


def _find_column(df: pd.DataFrame, candidates: set[str]) -> str | None:
    for col in df.columns:
        if _canonical_column_name(col) in candidates:
            return str(col)
    return None


def _read_peak_dataframe(path: Path) -> pd.DataFrame:
    suffix = path.suffix.lower()
    if suffix in {".xlsx", ".xls"}:
        xls = pd.ExcelFile(path)
        preferred = ["peaks", "all_detected_peaks", "selected_for_indexing"]
        sheet = None
        lower_map = {str(s).lower(): s for s in xls.sheet_names}
        for name in preferred:
            if name.lower() in lower_map:
                sheet = lower_map[name.lower()]
                break
        if sheet is None:
            sheet = xls.sheet_names[0]
        return pd.read_excel(path, sheet_name=sheet)
    return pd.read_csv(path, sep=None, engine="python", comment="#")


def read_peaks_file(path: Path, wavelength: float = CU_KA1) -> list[Peak]:
    """Read a peak-information file and return Peak objects.

    The preferred input is the Excel file created by this script in --mode search.
    A simple CSV/text table is also accepted if it contains a diffraction-angle
    column such as two_theta_deg, 2theta, angle, or position.
    """
    df = _read_peak_dataframe(path)
    if df.empty:
        raise ValueError(f"Peak file is empty: {path}")

    two_theta_col = _find_column(df, {
        "two_theta_deg", "two_theta", "twotheta", "2theta", "2theta_deg",
        "theta2", "angle", "angle_deg", "position", "position_deg",
        "two_theta_obs_deg",
    })
    # The canonical form of "inv_d2_A-2" becomes "inv_d2_a_2".
    inv_d2_col = _find_column(df, {
        "inv_d2", "inv_d2_obs", "inv_d2_calc", "inv_d2_a_2", "inv_d2_a2",
        "one_over_d2", "one_over_d_2", "d_star2",
    })
    intensity_col = _find_column(df, {"intensity", "i", "counts", "count", "y"})
    intensity_raw_col = _find_column(df, {"intensity_raw", "raw_intensity", "i_raw", "counts_raw"})
    fwhm_col = _find_column(df, {"fwhm_deg", "fwhm", "fwhm_degree"})
    role_col = _find_column(df, {"ka_role", "role", "kalpha_role"})
    pair_col = _find_column(df, {"ka_pair_peak_id", "pair_peak_id", "ka_pair_index"})

    if two_theta_col is None and inv_d2_col is None:
        # Last fallback: use the first numeric column as 2theta.
        for col in df.columns:
            values = pd.to_numeric(df[col], errors="coerce")
            if values.notna().sum() > 0:
                two_theta_col = str(col)
                break
    if two_theta_col is None and inv_d2_col is None:
        raise ValueError(
            "Cannot find diffraction-angle column in peak file. "
            "Use two_theta_deg, 2theta, angle, or inv_d2_A-2."
        )

    peaks: list[Peak] = []
    for row_index, row in df.iterrows():
        two_theta = float("nan")
        inv_d2 = float("nan")
        if two_theta_col is not None:
            two_theta = float(pd.to_numeric(row.get(two_theta_col), errors="coerce"))
        if inv_d2_col is not None:
            inv_d2 = float(pd.to_numeric(row.get(inv_d2_col), errors="coerce"))

        if not math.isfinite(two_theta):
            if math.isfinite(inv_d2) and inv_d2 > 0:
                tt = two_theta_from_inv_d2(inv_d2, wavelength)
                two_theta = float(tt) if tt is not None else float("nan")
        if not math.isfinite(inv_d2) and math.isfinite(two_theta):
            inv_d2 = inv_d2_from_two_theta(two_theta, wavelength)
        if not (math.isfinite(two_theta) and math.isfinite(inv_d2) and inv_d2 > 0):
            continue

        intensity = float(pd.to_numeric(row.get(intensity_col), errors="coerce")) if intensity_col else 1.0
        if not math.isfinite(intensity):
            intensity = 1.0
        intensity_raw = float(pd.to_numeric(row.get(intensity_raw_col), errors="coerce")) if intensity_raw_col else intensity
        if not math.isfinite(intensity_raw):
            intensity_raw = intensity
        fwhm = float(pd.to_numeric(row.get(fwhm_col), errors="coerce")) if fwhm_col else float("nan")
        if not math.isfinite(fwhm):
            fwhm = 0.0

        d = 1.0 / math.sqrt(inv_d2)
        q = 2.0 * math.pi / d
        role = ""
        if role_col is not None and not pd.isna(row.get(role_col)):
            role = str(row.get(role_col)).strip().lower()
        pair_index = None
        if pair_col is not None and not pd.isna(row.get(pair_col)):
            try:
                pair_id = int(row.get(pair_col))
                if pair_id > 0:
                    pair_index = pair_id - 1
            except Exception:
                pair_index = None

        peaks.append(
            Peak(
                index=int(row_index),
                two_theta=float(two_theta),
                intensity=float(intensity),
                intensity_raw=float(intensity_raw),
                fwhm_deg=float(fwhm),
                inv_d2=float(inv_d2),
                d=float(d),
                q=float(q),
                ka_role=role,
                ka_pair_index=pair_index,
                source=str(path),
            )
        )

    peaks.sort(key=lambda p: p.two_theta)
    if not peaks:
        raise ValueError(f"No usable peaks were found in: {path}")
    if not any(p.ka_role for p in peaks):
        assign_ka2(peaks)
    return peaks


def normalize_mode(mode: str) -> str:
    m = mode.strip().lower()
    if m == "seaerch":
        return "search"
    if m == "index":
        # Backward compatibility: old "index" mode did search + indexing.
        return "all"
    if m not in {"search", "guess", "all", "compare", "calc"}:
        raise ValueError(f"Unsupported mode: {mode}")
    return m


def parse_crystal_systems(text: str) -> list[str]:
    systems: list[str] = []
    for raw in str(text).replace(";", ",").split(","):
        key = raw.strip().lower()
        if not key:
            continue
        if key not in CRYSTAL_SYSTEM_ALIASES:
            allowed = ", ".join(["auto"] + SUPPORTED_CRYSTAL_SYSTEMS)
            raise ValueError(f"Unsupported crystal system: {raw}. Allowed: {allowed}")
        value = CRYSTAL_SYSTEM_ALIASES[key]
        if value == "auto":
            return list(SUPPORTED_CRYSTAL_SYSTEMS)
        if value not in systems:
            systems.append(value)
    return systems or list(SUPPORTED_CRYSTAL_SYSTEMS)


def default_peak_file_for_input(infile: Path) -> Path:
    stem = sanitize_stem(infile.stem)
    return infile.with_name(f"{stem}-peaks.xlsx")


def strip_peak_suffix(stem: str) -> str:
    for suffix in ("-peaks", "_peaks", ".peaks", "-peak", "_peak"):
        if stem.lower().endswith(suffix):
            return stem[: -len(suffix)]
    return stem


def default_guess_file_for_peak_file(peak_file: Path, base_input: Path | None = None) -> Path:
    if base_input is not None:
        stem = sanitize_stem(base_input.stem)
        return base_input.with_name(f"{stem}-guess.xlsx")
    stem = sanitize_stem(strip_peak_suffix(peak_file.stem))
    return peak_file.with_name(f"{stem}-guess.xlsx")


def peaks_to_frame(peaks: list[Peak]) -> pd.DataFrame:
    rows = []
    for i, p in enumerate(peaks, 1):
        rows.append({
            "peak_id": i,
            "two_theta_deg": p.two_theta,
            "intensity": p.intensity,
            "intensity_raw": p.intensity_raw,
            "fwhm_deg": p.fwhm_deg,
            "d_A": p.d,
            "inv_d2_A-2": p.inv_d2,
            "q_A-1": p.q,
            "ka_role": p.ka_role,
            "ka_pair_peak_id": (p.ka_pair_index + 1) if p.ka_pair_index is not None else None,
        })
    return pd.DataFrame(rows)


def print_peak_table(peaks: list[Peak], title: str) -> None:
    print("")
    print(title)
    print(f"{'id':>3s} {'2theta':>9s} {'I':>11s} {'FWHM':>7s} {'role':>5s}")
    for i, p in enumerate(peaks, 1):
        print(f"{i:3d} {p.two_theta:9.3f} {p.intensity:11.2f} {p.fwhm_deg:7.3f} {p.ka_role or '-':>5s}")


def save_workbook(path: Path, sheets: dict[str, pd.DataFrame]) -> None:
    wb = Workbook()
    first = True
    for name, df in sheets.items():
        ws = wb.active if first else wb.create_sheet(title=name[:31])
        ws.title = name[:31]
        first = False
        for c, col in enumerate(df.columns, 1):
            cell = ws.cell(1, c, col)
            cell.font = Font(bold=True)
        for r, (_, row) in enumerate(df.iterrows(), 2):
            for c, val in enumerate(row, 1):
                ws.cell(r, c, None if (pd.isna(val) if not isinstance(val, str) else False) else val)
        ws.freeze_panes = "A2"
    wb.save(path)



def _maybe_install_mplcursors(artists: list, hover_texts: list[str]) -> None:
    """Attach mplcursors hover annotations when mplcursors is available."""
    if not artists:
        return
    try:
        import mplcursors  # type: ignore
    except Exception:
        print("mplcursors is not installed; hover annotations are disabled.")
        print("Install with: pip install mplcursors")
        return

    text_by_artist = {artist: text for artist, text in zip(artists, hover_texts)}
    cursor = mplcursors.cursor(artists, hover=True)

    @cursor.connect("add")
    def _on_add(sel):  # pragma: no cover - interactive callback
        sel.annotation.set_text(text_by_artist.get(sel.artist, ""))
        sel.annotation.get_bbox_patch().set(alpha=0.92)


def _peak_hover_text(p: Peak, hkl_text: str | None = None, calc_two_theta: float | None = None, resid: float | None = None) -> str:
    parts = []
    if hkl_text:
        parts.append(f"hkl: {hkl_text}")
    parts.append(f"2theta(obs): {p.two_theta:.5f} deg")
    if calc_two_theta is not None and math.isfinite(calc_two_theta):
        parts.append(f"2theta(calc): {calc_two_theta:.5f} deg")
    if resid is not None and math.isfinite(resid):
        parts.append(f"resid: {resid:+.5f} deg")
    parts.append(f"intensity: {p.intensity:.6g}")
    parts.append(f"FWHM: {p.fwhm_deg:.5g} deg")
    if p.ka_role:
        parts.append(f"role: {p.ka_role}")
    return "\n".join(parts)


def plot_results(
    x: np.ndarray,
    y: np.ndarray,
    info: dict[str, np.ndarray | float],
    peaks: list[Peak],
    outpath: Path | None,
    show: bool = False,
    save: bool = True,
) -> None:
    if not show and not save:
        return
    ys = info["ysmooth"]
    fig, ax = plt.subplots(figsize=(12, 7))
    ax.plot(x, y, lw=0.8, label="input")
    ax.plot(x, ys, lw=1.0, label="smoothed")
    hover_artists = []
    hover_texts = []
    ymin = float(np.nanmin(y)) if len(y) else 0.0
    ymax = float(np.nanmax([np.nanmax(y), np.nanmax(ys)])) if len(y) else 1.0
    for p in peaks:
        color = "tab:red" if p.ka_role == "ka2" else "tab:green"
        line = ax.axvline(p.two_theta, color=color, lw=0.9, alpha=0.85, picker=5)
        line.set_gid(f"peak_{p.two_theta:.5f}")
        hover_artists.append(line)
        hover_texts.append(_peak_hover_text(p))
        ax.text(p.two_theta, p.intensity, f"{p.two_theta:.2f}", rotation=90, va="bottom", fontsize=7)
    ax.set_ylim(min(ymin, 0.0), ymax * 1.05 if ymax > 0 else ymax + 1.0)
    ax.set_xlabel("2Theta (deg)")
    ax.set_ylabel("Intensity")
    ax.legend()
    fig.tight_layout()
    if save and outpath is not None:
        fig.savefig(outpath, dpi=160)
    if show:
        _maybe_install_mplcursors(hover_artists, hover_texts)
        plt.show()
#        plt.pause(0.1)
#        input("\nPress ENTER to terminate>>\n")
    if save and outpath is not None:
        plt.close(fig)


def beta_from_params(system: str, params: dict[str, float]) -> np.ndarray:
    system = system.lower()
    if system == "cubic":
        return np.array([1.0 / (params["a"] ** 2)], dtype=float)
    if system == "tetragonal":
        return np.array([1.0 / (params["a"] ** 2), 1.0 / (params["c"] ** 2)], dtype=float)
    if system == "hexagonal":
        return np.array([4.0 / (3.0 * params["a"] ** 2), 1.0 / (params["c"] ** 2)], dtype=float)
    if system == "orthorhombic":
        return np.array([1.0 / (params["a"] ** 2), 1.0 / (params["b"] ** 2), 1.0 / (params["c"] ** 2)], dtype=float)
    raise ValueError(system)


def candidate_to_summary_row(rank: int, c: Candidate, method: str = "") -> dict[str, object]:
    row: dict[str, object] = {
        "rank": rank,
        "method": method,
        "crystal_system": c.crystal_system,
        "score_matches": c.score_matches,
        "score_rms_rel": c.score_rms_rel,
    }
    row.update(c.params)
    return row


def candidate_from_summary_row(row: pd.Series) -> Candidate:
    system = str(row.get("crystal_system", row.get("system", ""))).strip().lower()
    if system not in SUPPORTED_CRYSTAL_SYSTEMS:
        raise ValueError(f"Unsupported crystal system in candidate summary: {system}")
    params = {
        key: float(row[key])
        for key in ["a", "b", "c", "alpha", "beta", "gamma"]
        if key in row.index and pd.notna(row[key])
    }
    if "a" not in params:
        raise ValueError("Candidate summary does not contain lattice constant a.")
    if system in {"tetragonal", "hexagonal"} and "c" not in params:
        raise ValueError("Candidate summary does not contain lattice constant c.")
    if system == "orthorhombic" and not all(k in params for k in ["b", "c"]):
        raise ValueError("Candidate summary does not contain lattice constants b and c.")
    params.setdefault("b", params.get("a", float("nan")))
    params.setdefault("c", params.get("a", float("nan")))
    params.setdefault("alpha", 90.0)
    params.setdefault("beta", 90.0)
    params.setdefault("gamma", 120.0 if system == "hexagonal" else 90.0)
    ls_map = {"orthorhombic": 4, "tetragonal": 5, "cubic": 6, "hexagonal": 8}
    return Candidate(
        crystal_system=system,
        ls_code=ls_map[system],
        params=params,
        score_matches=int(row.get("score_matches", 0) or 0),
        score_rms_rel=float(row.get("score_rms_rel", float("nan"))),
        selected_matches=[],
        all_matches=[],
    )

def has_explicit_lattice_params(args: argparse.Namespace) -> bool:
    return any(getattr(args, name, None) is not None for name in ["a", "b", "c"])


def candidate_from_lattice_args(args: argparse.Namespace) -> Candidate:
    systems = parse_crystal_systems(args.crystal_system)
    if len(systems) != 1:
        raise ValueError("Explicit lattice constants require one --crystal-system, not auto or a comma-separated list.")
    system = systems[0]
    a = getattr(args, "a", None)
    b = getattr(args, "b", None)
    c = getattr(args, "c", None)
    if a is None:
        raise ValueError("Explicit lattice comparison requires --a.")

    params: dict[str, float] = {
        "a": float(a),
        "alpha": float(getattr(args, "alpha", 90.0) if getattr(args, "alpha", None) is not None else 90.0),
        "beta": float(getattr(args, "beta", 90.0) if getattr(args, "beta", None) is not None else 90.0),
        "gamma": float(getattr(args, "gamma", 120.0 if system == "hexagonal" else 90.0) if getattr(args, "gamma", None) is not None else (120.0 if system == "hexagonal" else 90.0)),
    }
    if system == "cubic":
        params["b"] = float(a)
        params["c"] = float(a if c is None else c)
        # Keep cubic internally cubic even if --b/--c were accidentally supplied.
        params["b"] = params["a"]
        params["c"] = params["a"]
    elif system in {"tetragonal", "hexagonal"}:
        if c is None:
            raise ValueError(f"{system} comparison requires --c in addition to --a.")
        params["b"] = float(a)
        params["c"] = float(c)
        if system == "hexagonal":
            params["gamma"] = 120.0
    elif system == "orthorhombic":
        if b is None or c is None:
            raise ValueError("orthorhombic comparison requires --a, --b, and --c.")
        params["b"] = float(b)
        params["c"] = float(c)
    else:
        raise ValueError(f"Unsupported crystal system: {system}")

    for key in ["a", "b", "c"]:
        if not math.isfinite(params[key]) or params[key] <= 0:
            raise ValueError(f"Lattice constant {key} must be positive: {params[key]}")

    ls_map = {"orthorhombic": 4, "tetragonal": 5, "cubic": 6, "hexagonal": 8}
    return Candidate(
        crystal_system=system,
        ls_code=ls_map[system],
        params=params,
        score_matches=0,
        score_rms_rel=float("nan"),
        selected_matches=[],
        all_matches=[],
    )


def filter_theoretical_lines_for_plot(
    theoretical_df: pd.DataFrame,
    mode: str = "near",
    near_deg: float = 0.5,
) -> pd.DataFrame:
    if theoretical_df.empty:
        return theoretical_df
    mode = str(mode).strip().lower()
    if mode == "all":
        return theoretical_df
    if mode == "none":
        return theoretical_df.iloc[0:0].copy()
    if mode == "matched":
        if "matched" not in theoretical_df.columns:
            return theoretical_df.iloc[0:0].copy()
        return theoretical_df[theoretical_df["matched"].astype(bool)].copy()
    if mode == "near":
        if "nearest_resid_two_theta_deg" not in theoretical_df.columns:
            if "matched" in theoretical_df.columns:
                return theoretical_df[theoretical_df["matched"].astype(bool)].copy()
            return theoretical_df.iloc[0:0].copy()
        resid = pd.to_numeric(theoretical_df["nearest_resid_two_theta_deg"], errors="coerce").abs()
        return theoretical_df[resid <= float(near_deg)].copy()
    raise ValueError(f"Unsupported --plot-theory: {mode}")



def estimate_lattice_limit_from_peaks(
    peaks: list[Peak],
    wavelength: float = CU_KA1,
    factor: float = 1.2,
) -> tuple[float, float, float]:
    """Return (minimum 2theta, d at that angle, default lattice upper limit).

    The estimate is intentionally simple: the largest visible lattice spacing is
    approximated by the d-spacing of the lowest-angle non-Ka2 peak. A safety
    factor is then applied because low-angle reflections can be weak or missing.
    """
    usable = [
        p for p in peaks
        if p.ka_role != "ka2" and math.isfinite(p.two_theta) and p.two_theta > 0.0
    ]
    if not usable:
        usable = [p for p in peaks if math.isfinite(p.two_theta) and p.two_theta > 0.0]
    if not usable:
        raise ValueError("Cannot estimate lattice upper limit because no finite positive peak angle is available.")
    min_two_theta = min(p.two_theta for p in usable)
    dmax = bragg_d(min_two_theta, wavelength)
    if not math.isfinite(dmax) or dmax <= 0.0:
        raise ValueError(f"Cannot estimate lattice upper limit from minimum 2theta={min_two_theta}.")
    return float(min_two_theta), float(dmax), float(dmax * factor)


def resolve_lattice_max_limit(args: argparse.Namespace, peaks: list[Peak]) -> float | None:
    """Resolve the lattice-constant upper limit used for guessing.

    --lattice-max / --max-lattice gives an explicit limit. If omitted, the
    default is factor * d(minimum observed 2theta). Passing a non-positive
    explicit value disables this filter.
    """
    explicit = getattr(args, "lattice_max", None)
    if explicit is not None:
        explicit = float(explicit)
        if explicit <= 0.0:
            print("Lattice max limit : disabled (--lattice-max <= 0)")
            return None
        print(f"Lattice max limit : {explicit:.6g} Å (user specified)")
        return explicit

    factor = float(getattr(args, "lattice_max_factor", 1.2))
    if factor <= 0.0:
        raise ValueError("--lattice-max-factor must be positive.")
    min_tt, dmax, limit = estimate_lattice_limit_from_peaks(peaks, wavelength=args.wavelength, factor=factor)
    print(f"Minimum peak angle: {min_tt:.6g} deg 2theta")
    print(f"d(min 2theta)    : {dmax:.6g} Å")
    print(f"Lattice max limit : {limit:.6g} Å (= d(min 2theta) x {factor:.6g})")
    return limit


def params_within_lattice_max(params: dict[str, float], system: str, lattice_max: float | None) -> bool:
    if lattice_max is None:
        return True
    keys = ["a"]
    if system in {"tetragonal", "hexagonal"}:
        keys = ["a", "c"]
    elif system == "orthorhombic":
        keys = ["a", "b", "c"]
    for key in keys:
        val = params.get(key, None)
        if val is None:
            continue
        try:
            v = float(val)
        except Exception:
            continue
        if math.isfinite(v) and v > lattice_max * (1.0 + 1.0e-12):
            return False
    return True


def filter_candidates_by_lattice_max(
    candidates: list[Candidate],
    lattice_max: float | None,
    label: str = "candidates",
    is_print: bool = True,
) -> list[Candidate]:
    if lattice_max is None:
        return candidates
    kept = [c for c in candidates if params_within_lattice_max(c.params, c.crystal_system, lattice_max)]
    n_removed = len(candidates) - len(kept)
    if is_print and n_removed > 0:
        print(f"Excluded by lattice max ({lattice_max:.6g} Å): {n_removed} {label}")
    return kept


def theoretical_lines_for_candidate(
    candidate: Candidate,
    hmax: int = 12,
    wavelength: float = CU_KA1,
    xmin: float | None = None,
    xmax: float | None = None,
) -> pd.DataFrame:
    beta = beta_from_params(candidate.crystal_system, candidate.params)
    rows = []
    for feat, (h, k, l) in system_rows(candidate.crystal_system, hmax=hmax):
        inv_d2 = predicted_inv_d2(candidate.crystal_system, feat, beta)
        tt = two_theta_from_inv_d2(inv_d2, wavelength)
        if tt is None or not math.isfinite(tt):
            continue
        if xmin is not None and tt < xmin:
            continue
        if xmax is not None and tt > xmax:
            continue
        d = 1.0 / math.sqrt(inv_d2)
        rows.append({
            "h": h,
            "k": k,
            "l": l,
            "hkl": f"{h}{k}{l}",
            "two_theta_calc_deg": float(tt),
            "d_calc_A": float(d),
            "inv_d2_calc": float(inv_d2),
        })
    return pd.DataFrame(rows).sort_values("two_theta_calc_deg").reset_index(drop=True)


def compare_candidate_with_peaks(
    candidate: Candidate,
    peaks: list[Peak],
    hmax: int = 12,
    wavelength: float = CU_KA1,
    xmin: float | None = None,
    xmax: float | None = None,
    tolerance_deg: float = 0.25,
) -> tuple[pd.DataFrame, pd.DataFrame]:
    if xmin is None and peaks:
        xmin = max(0.0, min(p.two_theta for p in peaks) - 2.0)
    if xmax is None and peaks:
        xmax = max(p.two_theta for p in peaks) + 2.0
    calc_df = theoretical_lines_for_candidate(candidate, hmax=hmax, wavelength=wavelength, xmin=xmin, xmax=xmax)
    if calc_df.empty:
        return calc_df, pd.DataFrame()
    calc_tt = calc_df["two_theta_calc_deg"].to_numpy(dtype=float)

    rows = []
    used_calc: set[int] = set()
    for peak_id, p in enumerate(sorted(peaks, key=lambda pk: pk.two_theta), 1):
        diffs = calc_tt - p.two_theta
        order = np.argsort(np.abs(diffs))
        best_idx = None
        for idx in order:
            if int(idx) not in used_calc:
                best_idx = int(idx)
                break
        if best_idx is None:
            continue
        used_calc.add(best_idx)
        calc = calc_df.iloc[best_idx]
        resid = p.two_theta - float(calc["two_theta_calc_deg"])
        matched = abs(resid) <= tolerance_deg
        rows.append({
            "peak_id": peak_id,
            "two_theta_obs_deg": p.two_theta,
            "intensity": p.intensity,
            "fwhm_deg": p.fwhm_deg,
            "ka_role": p.ka_role,
            "h": int(calc["h"]),
            "k": int(calc["k"]),
            "l": int(calc["l"]),
            "hkl": str(calc["hkl"]),
            "two_theta_calc_deg": float(calc["two_theta_calc_deg"]),
            "resid_two_theta_deg": float(resid),
            "abs_resid_two_theta_deg": float(abs(resid)),
            "matched": bool(matched),
        })

    assign_df = pd.DataFrame(rows)
    calc_df = calc_df.copy()
    if not assign_df.empty:
        obs_by_hkl = assign_df.set_index("hkl")
        calc_df["nearest_peak_id"] = calc_df["hkl"].map(obs_by_hkl["peak_id"].to_dict())
        calc_df["nearest_two_theta_obs_deg"] = calc_df["hkl"].map(obs_by_hkl["two_theta_obs_deg"].to_dict())
        calc_df["nearest_intensity"] = calc_df["hkl"].map(obs_by_hkl["intensity"].to_dict())
        calc_df["nearest_fwhm_deg"] = calc_df["hkl"].map(obs_by_hkl["fwhm_deg"].to_dict())
        calc_df["nearest_resid_two_theta_deg"] = calc_df["hkl"].map(obs_by_hkl["resid_two_theta_deg"].to_dict())
        calc_df["matched"] = calc_df["hkl"].map(obs_by_hkl["matched"].to_dict()).fillna(False)
    else:
        calc_df["matched"] = False
    return calc_df, assign_df


def plot_compare_results(
    x: np.ndarray | None,
    y: np.ndarray | None,
    peaks: list[Peak],
    theoretical_df: pd.DataFrame,
    assignment_df: pd.DataFrame,
    candidate: Candidate,
    xmin: float,
    xmax: float,
    outpath: Path | None,
    show: bool = False,
    save: bool = True,
) -> None:
    if not show and not save:
        return
    fig, ax = plt.subplots(figsize=(12, 7))
    if x is not None and y is not None and len(x) and len(y):
        ax.plot(x, y, lw=0.8, label="input XRD")
        ymax = float(np.nanmax(y)) if np.isfinite(y).any() else 1.0
    else:
        ymax = max([p.intensity for p in peaks] + [1.0])
        markerline, stemlines, baseline = ax.stem(
            [p.two_theta for p in peaks],
            [p.intensity for p in peaks],
            basefmt=" ",
            linefmt="C0-",
            markerfmt="C0o",
            label="observed peaks",
        )
        try:
            plt.setp(stemlines, linewidth=1.0)
            plt.setp(markerline, markersize=4)
        except Exception:
            pass

    hover_artists = []
    hover_texts = []
    assign_by_obs = {}
    if not assignment_df.empty:
        for _, row in assignment_df.iterrows():
            assign_by_obs[round(float(row["two_theta_obs_deg"]), 6)] = row

    for p in sorted(peaks, key=lambda pk: pk.two_theta):
        key = round(float(p.two_theta), 6)
        row = assign_by_obs.get(key)
        hkl_text = None
        calc_tt = None
        resid = None
        if row is not None:
            hkl_text = f"{int(row['h'])}{int(row['k'])}{int(row['l'])}"
            calc_tt = float(row["two_theta_calc_deg"])
            resid = float(row["resid_two_theta_deg"])
        line = ax.axvline(p.two_theta, color="tab:green", lw=0.3, linestyle="--", alpha=0.75, picker=5)
        hover_artists.append(line)
        hover_texts.append(_peak_hover_text(p, hkl_text=hkl_text, calc_two_theta=calc_tt, resid=resid))

    theory_height = ymax * 0.16 if ymax > 0 else 1.0
    for _, row in theoretical_df.iterrows():
        tt = float(row["two_theta_calc_deg"])
        matched = bool(row.get("matched", False))
        line = ax.vlines(tt, ymin=0, ymax=theory_height, colors="tab:red", lw=1.0 if matched else 0.7, alpha=0.85 if matched else 0.45)
        # mplcursors accepts LineCollection from vlines.
        hkl_text = f"{int(row['h'])}{int(row['k'])}{int(row['l'])}"
        hover = [
            f"hkl: {hkl_text}",
            f"2theta(calc): {tt:.5f} deg",
        ]
        if pd.notna(row.get("nearest_two_theta_obs_deg", np.nan)):
            hover.append(f"2theta(obs): {float(row['nearest_two_theta_obs_deg']):.5f} deg")
        if pd.notna(row.get("nearest_resid_two_theta_deg", np.nan)):
            hover.append(f"resid: {float(row['nearest_resid_two_theta_deg']):+.5f} deg")
        if pd.notna(row.get("nearest_intensity", np.nan)):
            hover.append(f"intensity: {float(row['nearest_intensity']):.6g}")
        if pd.notna(row.get("nearest_fwhm_deg", np.nan)):
            hover.append(f"FWHM: {float(row['nearest_fwhm_deg']):.5g} deg")
        hover_artists.append(line)
        hover_texts.append("\n".join(hover))
        if matched:
            ax.text(tt, theory_height, hkl_text, rotation=90, va="bottom", fontsize=7)

    title_params = ", ".join(f"{k}={v:.5g}" for k, v in candidate.params.items() if k in {"a", "b", "c"})
    ax.set_title(f"{candidate.crystal_system}: {title_params}")
    ax.set_xlabel("2Theta (deg)")
    ax.set_ylabel("Intensity")
    ax.set_xlim([xmin, xmax])
    ax.legend(loc="best")
    fig.tight_layout()
    if save and outpath is not None:
        fig.savefig(outpath, dpi=160)
    if show:
        _maybe_install_mplcursors(hover_artists, hover_texts)
        plt.pause(0.1)
        input("\nPress ENTER to terminate>>\n")
    if save and outpath is not None:
        plt.close(fig)


def frange_values(vmin: float, vmax: float, step: float) -> np.ndarray:
    if step <= 0:
        raise ValueError("Grid step must be positive.")
    if vmax < vmin:
        raise ValueError("Grid max must be >= min.")
    n = int(math.floor((vmax - vmin) / step + 0.5)) + 1
    vals = vmin + step * np.arange(max(n, 1))
    return vals[vals <= vmax + 1e-12]


def _param_grid_for_system(system: str, ranges: dict[str, tuple[float, float, float]]) -> Iterable[dict[str, float]]:
    avec = frange_values(*ranges["a"])
    if system == "cubic":
        for a in avec:
            yield {"a": float(a), "b": float(a), "c": float(a), "alpha": 90.0, "beta": 90.0, "gamma": 90.0}
        return
    cvec = frange_values(*ranges["c"])
    if system == "tetragonal":
        for a in avec:
            for c in cvec:
                yield {"a": float(a), "b": float(a), "c": float(c), "alpha": 90.0, "beta": 90.0, "gamma": 90.0}
        return
    if system == "hexagonal":
        for a in avec:
            for c in cvec:
                yield {"a": float(a), "b": float(a), "c": float(c), "alpha": 90.0, "beta": 90.0, "gamma": 120.0}
        return
    if system == "orthorhombic":
        bvec = frange_values(*ranges["b"])
        for a in avec:
            for b in bvec:
                for c in cvec:
                    yield {"a": float(a), "b": float(b), "c": float(c), "alpha": 90.0, "beta": 90.0, "gamma": 90.0}
        return
    raise ValueError(system)


def _ranges_from_args_for_system(args: argparse.Namespace, system: str) -> dict[str, tuple[float, float, float]] | None:
    if args.amin is None or args.amax is None:
        return None
    ranges = {"a": (float(args.amin), float(args.amax), float(args.astep))}
    if system in {"tetragonal", "hexagonal", "orthorhombic"}:
        if args.cmin is None or args.cmax is None:
            return None
        ranges["c"] = (float(args.cmin), float(args.cmax), float(args.cstep))
    if system == "orthorhombic":
        if args.bmin is None or args.bmax is None:
            return None
        ranges["b"] = (float(args.bmin), float(args.bmax), float(args.bstep))
    return ranges


def _ranges_around_candidate(candidate: Candidate, margin: float, args: argparse.Namespace) -> dict[str, tuple[float, float, float]]:
    params = candidate.params
    ranges: dict[str, tuple[float, float, float]] = {}
    for key, step_name in [("a", "astep"), ("b", "bstep"), ("c", "cstep")]:
        if key not in params or not math.isfinite(float(params[key])):
            continue
        value = float(params[key])
        step = float(getattr(args, step_name))
        ranges[key] = (max(0.1, value * (1.0 - margin)), value * (1.0 + margin), step)
    if candidate.crystal_system in {"cubic"}:
        ranges = {"a": ranges["a"]}
    elif candidate.crystal_system in {"tetragonal", "hexagonal"}:
        ranges = {"a": ranges["a"], "c": ranges["c"]}
    elif candidate.crystal_system == "orthorhombic":
        ranges = {"a": ranges["a"], "b": ranges["b"], "c": ranges["c"]}
    return ranges


def _assign_peaks_for_params(
    peaks: list[Peak],
    system: str,
    params: dict[str, float],
    hmax: int,
    wavelength: float,
    tolerance_deg: float,
) -> tuple[list[dict], float, float, int]:
    cand = Candidate(system, {"orthorhombic": 4, "tetragonal": 5, "cubic": 6, "hexagonal": 8}[system], params, 0, 0.0, [], [])
    calc_df = theoretical_lines_for_candidate(cand, hmax=hmax, wavelength=wavelength)
    if calc_df.empty:
        return [], float("inf"), float("inf"), 0
    calc_tt = calc_df["two_theta_calc_deg"].to_numpy(dtype=float)
    used: set[int] = set()
    rows = []
    residuals = []
    rels = []
    for pk in sorted(peaks, key=lambda p: p.two_theta):
        order = np.argsort(np.abs(calc_tt - pk.two_theta))
        chosen = None
        for idx in order:
            if int(idx) not in used:
                chosen = int(idx)
                break
        if chosen is None:
            continue
        used.add(chosen)
        calc = calc_df.iloc[chosen]
        resid = pk.two_theta - float(calc["two_theta_calc_deg"])
        matched = abs(resid) <= tolerance_deg
        if matched:
            residuals.append(float(resid))
            rel = abs(float(calc["inv_d2_calc"]) - pk.inv_d2) / max(pk.inv_d2, 1e-12)
            rels.append(float(rel))
        rows.append({
            "two_theta_obs_deg": pk.two_theta,
            "intensity": pk.intensity,
            "ka_role": pk.ka_role,
            "fwhm_deg": pk.fwhm_deg,
            "inv_d2_obs": pk.inv_d2,
            "h": int(calc["h"]) if matched else None,
            "k": int(calc["k"]) if matched else None,
            "l": int(calc["l"]) if matched else None,
            "inv_d2_calc": float(calc["inv_d2_calc"]) if matched else None,
            "rel_error_inv_d2": rels[-1] if matched else None,
            "two_theta_calc_deg": float(calc["two_theta_calc_deg"]) if matched else None,
            "rel_error_2theta": abs(resid) / max(pk.two_theta, 1e-12) if matched else None,
            "resid_two_theta_deg": float(resid) if matched else None,
            "matched": bool(matched),
        })
    nmatched = len(residuals)
    rms_tt = float(np.sqrt(np.mean(np.square(residuals)))) if residuals else float("inf")
    rms_rel = float(np.sqrt(np.mean(np.square(rels)))) if rels else float("inf")
    return rows, rms_tt, rms_rel, nmatched


def _refit_params_from_assignment(system: str, rows: list[dict]) -> dict[str, float] | None:
    matched = [r for r in rows if r.get("matched")]
    if not matched:
        return None
    y = np.array([float(r["inv_d2_obs"]) for r in matched], dtype=float)
    feats = []
    for r in matched:
        h, k, l = int(r["h"]), int(r["k"]), int(r["l"])
        if system == "cubic":
            feats.append((h*h + k*k + l*l,))
        elif system == "tetragonal":
            feats.append((h*h + k*k, l*l))
        elif system == "hexagonal":
            feats.append((h*h + h*k + k*k, l*l))
        elif system == "orthorhombic":
            feats.append((h*h, k*k, l*l))
        else:
            raise ValueError(system)
    X = np.array(feats, dtype=float)
    if len(y) < X.shape[1]:
        return None
    coef, _, rank, _ = np.linalg.lstsq(X, y, rcond=None)
    if rank < X.shape[1] or np.any(coef <= 0):
        return None
    return params_from_beta(system, coef)


def score_system_grid(
    system: str,
    selected: list[Peak],
    all_peaks: list[Peak],
    ranges: dict[str, tuple[float, float, float]],
    hmax: int = 12,
    wavelength: float = CU_KA1,
    tolerance_deg: float = 0.25,
    refine: int = 2,
    keep: int = 50,
    penalty_unassigned: float = 0.5,
    lattice_max: float | None = None,
) -> list[Candidate]:
    candidates: list[Candidate] = []
    ls_map = {"orthorhombic": 4, "tetragonal": 5, "cubic": 6, "hexagonal": 8}
    for params0 in _param_grid_for_system(system, ranges):
        if not params_within_lattice_max(params0, system, lattice_max):
            continue
        params = dict(params0)
        rows = []
        for _ in range(max(refine, 1)):
            rows, _rms_tt, _rms_rel, nmatched = _assign_peaks_for_params(
                selected, system, params, hmax, wavelength, tolerance_deg
            )
            if nmatched < len(beta_from_params(system, params)):
                rows = []
                break
            refit = _refit_params_from_assignment(system, rows)
            if refit is None:
                rows = []
                break
            params = refit
        if not rows:
            continue
        if not params_within_lattice_max(params, system, lattice_max):
            continue
        selected_rows, rms_tt, rms_rel, nmatched = _assign_peaks_for_params(
            selected, system, params, hmax, wavelength, tolerance_deg
        )
        if nmatched == 0:
            continue
        all_rows, _all_rms_tt, all_rms_rel, all_nmatched = _assign_peaks_for_params(
            all_peaks, system, params, hmax, wavelength, tolerance_deg
        )
        # Store a penalty-adjusted value in a private key for sorting only.
        penalty_score = rms_tt + penalty_unassigned * (len(selected) - nmatched)
        cand = Candidate(
            crystal_system=system,
            ls_code=ls_map[system],
            params=params,
            score_matches=int(all_nmatched),
            score_rms_rel=float(all_rms_rel if math.isfinite(all_rms_rel) else rms_rel),
            selected_matches=[r for r in selected_rows if r.get("matched")],
            all_matches=all_rows,
        )
        cand.params = dict(cand.params)
        cand.params["_grid_penalty_score"] = float(penalty_score)
        candidates.append(cand)
        candidates.sort(key=lambda c: (c.params.get("_grid_penalty_score", float("inf")), -c.score_matches, c.score_rms_rel))
        if len(candidates) > keep:
            candidates = candidates[:keep]
    for c in candidates:
        c.params.pop("_grid_penalty_score", None)
    candidates.sort(key=lambda c: (-c.score_matches, c.score_rms_rel))
    return candidates


def deduplicate_candidates(candidates: list[Candidate], ndigits: int = 4) -> list[Candidate]:
    seen = set()
    out = []
    for c in candidates:
        key = (c.crystal_system,) + tuple(round(float(c.params.get(k, 0.0)), ndigits) for k in ["a", "b", "c"])
        if key in seen:
            continue
        seen.add(key)
        out.append(c)
    return out


def guess_candidates(
    args: argparse.Namespace,
    peaks: list[Peak],
    crystal_systems: list[str],
) -> tuple[list[Candidate], list[Peak]]:
    selected = [p for p in sorted(peaks, key=lambda p: p.intensity, reverse=True) if p.ka_role != "ka2"][:args.npeak]
    method = args.method.lower()
    if method not in SUPPORTED_METHODS:
        raise ValueError(f"Unsupported method: {args.method}")
    candidates: list[Candidate] = []

    if method in {"combinatorial", "both", "grid"}:
        # For --method grid without explicit ranges, these candidates become seeds.
        seed_candidates: list[Candidate] = []
        for system in crystal_systems:
            cand = score_system(system, sorted(selected, key=lambda p: p.two_theta), sorted(peaks, key=lambda p: p.two_theta), hmax=args.hmax)
            if cand is not None:
                seed_candidates.append(cand)
        seed_candidates.sort(key=lambda c: (-c.score_matches, c.score_rms_rel))
        seed_candidates = filter_candidates_by_lattice_max(
            seed_candidates, getattr(args, "resolved_lattice_max", None), label="combinatorial seed candidates"
        )
    else:
        seed_candidates = []

    if method == "combinatorial":
        candidates = seed_candidates
    elif method in {"grid", "both"}:
        if method == "both":
            candidates.extend(seed_candidates)
        for system in crystal_systems:
            ranges = _ranges_from_args_for_system(args, system)
            if ranges is not None:
                grid_cands = score_system_grid(
                    system, selected, peaks, ranges,
                    hmax=args.hmax,
                    wavelength=args.wavelength,
                    tolerance_deg=args.tolerance_deg,
                    refine=args.grid_refine,
                    keep=args.keep,
                    penalty_unassigned=args.penalty_unassigned,
                    lattice_max=getattr(args, "resolved_lattice_max", None),
                )
                candidates.extend(grid_cands)
                continue

            # No explicit range: use top combinatorial candidates of the same system.
            seeds = [c for c in seed_candidates if c.crystal_system == system][:args.grid_seed_top]
            if not seeds and method == "grid":
                print(f"No explicit grid range and no seed candidate for {system}; skipped.")
            for seed in seeds:
                ranges = _ranges_around_candidate(seed, args.grid_margin, args)
                grid_cands = score_system_grid(
                    system, selected, peaks, ranges,
                    hmax=args.hmax,
                    wavelength=args.wavelength,
                    tolerance_deg=args.tolerance_deg,
                    refine=args.grid_refine,
                    keep=args.keep,
                    penalty_unassigned=args.penalty_unassigned,
                    lattice_max=getattr(args, "resolved_lattice_max", None),
                )
                candidates.extend(grid_cands)
    candidates = filter_candidates_by_lattice_max(
        candidates, getattr(args, "resolved_lattice_max", None), label="final candidates", is_print=False
    )
    candidates.sort(key=lambda c: (-c.score_matches, c.score_rms_rel))
    return deduplicate_candidates(candidates), selected


def read_candidate_from_guess_file(path: Path, rank: int) -> Candidate:
    suffix = path.suffix.lower()
    if suffix not in {".xlsx", ".xls"}:
        raise ValueError("Candidate-rank comparison needs an Excel guess file or recomputation from peaks.")
    xls = pd.ExcelFile(path)
    lower_map = {str(s).lower(): s for s in xls.sheet_names}
    if "candidate_summary" not in lower_map:
        raise ValueError(f"candidate_summary sheet was not found in {path}")
    df = pd.read_excel(path, sheet_name=lower_map["candidate_summary"])
    if df.empty:
        raise ValueError(f"candidate_summary is empty in {path}")
    if "rank" in df.columns:
        row_df = df[df["rank"] == rank]
        if row_df.empty:
            raise ValueError(f"Candidate rank {rank} was not found in {path}")
        row = row_df.iloc[0]
    else:
        if rank < 1 or rank > len(df):
            raise ValueError(f"Candidate rank {rank} is outside 1..{len(df)}")
        row = df.iloc[rank - 1]
    return candidate_from_summary_row(row)


def try_read_raw_xy_for_plot(path: Path) -> tuple[np.ndarray | None, np.ndarray | None]:
    try:
        # Avoid treating peak/guess workbooks as raw XRD data.
        if path.suffix.lower() in {".xlsx", ".xls"}:
            xls = pd.ExcelFile(path)
            names = {str(s).lower() for s in xls.sheet_names}
            if names & {"peaks", "all_detected_peaks", "candidate_summary", "theoretical_lines"}:
                return None, None
        return read_xy_data(path)
    except Exception:
        return None, None


def looks_like_raw_xrd_file(path: Path, min_rows: int = 50) -> bool:
    """Heuristic to avoid interpreting a dense raw XRD XY file as a peak list."""
    try:
        suffix = path.suffix.lower()
        if suffix in {".xlsx", ".xls"}:
            xls = pd.ExcelFile(path)
            names = {str(s).lower() for s in xls.sheet_names}
            if names & {"peaks", "all_detected_peaks", "candidate_summary", "theoretical_lines"}:
                return False
            df = pd.read_excel(path)
        else:
            df = pd.read_csv(path, sep=None, engine="python", comment="#")
        if df.empty or df.shape[1] < 2:
            return False
        peak_cols = {
            "two_theta_deg", "two_theta", "twotheta", "2theta", "2theta_deg",
            "theta2", "angle", "angle_deg", "position", "position_deg",
            "two_theta_obs_deg", "inv_d2", "inv_d2_obs", "inv_d2_a_2",
        }
        if any(_canonical_column_name(c) in peak_cols for c in df.columns):
            return False
        numeric_cols = 0
        for col in df.columns[:4]:
            if pd.to_numeric(df[col], errors="coerce").notna().sum() >= min_rows:
                numeric_cols += 1
        return len(df) >= min_rows and numeric_cols >= 2
    except Exception:
        return False


def build_argument_parser() -> argparse.ArgumentParser:
    p = argparse.ArgumentParser(description="Peak search and lattice-parameter guessing for powder XRD")
    p.add_argument(
        "arg1",
        help=(
            "input data file, peak-information file, or guess workbook. "
            "Legacy usage 'search infile' is also accepted."
        ),
    )
    p.add_argument(
        "arg2",
        nargs="?",
        help="input file for legacy positional mode, e.g. 'search sample.TXT'",
    )
    p.add_argument(
        "--mode",
        type=str,
        default="all",
        choices=["search", "guess", "all", "compare", "calc"],
        help=(
            "search: peak search only; guess: guess lattice constants from peak file; "
            "all: search + guess (default); compare: compare a ranked candidate with observed peaks; calc: compare user-specified lattice constants"
        ),
    )
    p.add_argument(
        "--method",
        type=str,
        default="combinatorial",
        choices=SUPPORTED_METHODS,
        help=(
            "lattice-parameter guessing method. combinatorial: fast hkl-combination search; "
            "grid: grid search using explicit ranges or combinatorial seeds; both: merge both results."
        ),
    )
    p.add_argument(
        "--crystal-system", "--system",
        type=str,
        default="auto",
        help=(
            "crystal system to test. Use auto/all, cubic, tetragonal, hexagonal, "
            "orthorhombic, or comma-separated values. Default: auto"
        ),
    )
    p.add_argument("--peak-file", type=Path, default=None, help="peak-information Excel/CSV file path")
    p.add_argument("--guess-file", type=Path, default=None, help="guess workbook used by --mode compare")
    p.add_argument("--output-file", type=Path, default=None, help="output workbook for guess/all/compare results")
    p.add_argument("--plot-file", type=Path, default=None, help="output graph path")
    p.add_argument("--candidate-rank", type=int, default=1, help="candidate rank used by --mode compare")
    p.add_argument("--wavelength", type=float, default=CU_KA1, help=f"X-ray wavelength in Angstrom. Default: Cu Kα1 = {CU_KA1}")

    # Explicit lattice constants for --mode calc, or for --mode compare without a guess workbook.
    p.add_argument("--a", "--lattice-a", dest="a", type=float, default=None, help="explicit lattice constant a in Angstrom")
    p.add_argument("--b", "--lattice-b", dest="b", type=float, default=None, help="explicit lattice constant b in Angstrom")
    p.add_argument("--c", "--lattice-c", dest="c", type=float, default=None, help="explicit lattice constant c in Angstrom")
    p.add_argument("--alpha", type=float, default=None, help="explicit alpha angle in degrees; currently stored in output")
    p.add_argument("--beta", type=float, default=None, help="explicit beta angle in degrees; currently stored in output")
    p.add_argument("--gamma", type=float, default=None, help="explicit gamma angle in degrees; currently stored in output")

    # Raw XRD / peak-search options.
    p.add_argument("--xmin", type=float, default=None)
    p.add_argument("--xmax", type=float, default=None)
    p.add_argument("--threshold", type=float, default=1000.0)
    p.add_argument("--nsmooth", type=int, default=11)
    p.add_argument("--norder", type=int, default=4)
    p.add_argument("--ydiff1-threshold", type=float, default=1.0e-2)
    p.add_argument("--fwhm-min-deg", type=float, default=0.06)
    p.add_argument("--fwhm-max-deg", type=float, default=1.0)

    # Indexing / comparison options.
    p.add_argument("--npeak", type=int, default=10)
    p.add_argument("--hmax", type=int, default=12)
    p.add_argument(
        "--lattice-max", "--max-lattice", dest="lattice_max",
        type=float, default=None,
        help=(
            "upper limit for guessed lattice constants in Angstrom. "
            "Default: d(minimum observed non-Ka2 2theta) x --lattice-max-factor. "
            "Set <=0 to disable the limit."
        ),
    )
    p.add_argument(
        "--lattice-max-factor",
        type=float, default=1.2,
        help="factor multiplied by d(minimum observed 2theta) when --lattice-max is omitted. Default: 1.2",
    )
    p.add_argument("--tolerance-deg", type=float, default=0.25, help="2theta assignment tolerance for grid/compare")
    p.add_argument("--penalty-unassigned", type=float, default=0.5)
    p.add_argument("--keep", type=int, default=200, help="number of grid candidates kept internally")
    p.add_argument("--grid-refine", type=int, default=2, help="assignment/refit iterations for grid search")
    p.add_argument("--grid-margin", type=float, default=0.05, help="relative margin around combinatorial seed for grid search")
    p.add_argument("--grid-seed-top", type=int, default=3, help="number of combinatorial seeds per system used for local grid search")

    # Explicit grid ranges. If omitted, grid search can use combinatorial seeds.
    p.add_argument("--amin", type=float, default=None)
    p.add_argument("--amax", type=float, default=None)
    p.add_argument("--astep", type=float, default=0.02)
    p.add_argument("--bmin", type=float, default=None)
    p.add_argument("--bmax", type=float, default=None)
    p.add_argument("--bstep", type=float, default=0.02)
    p.add_argument("--cmin", type=float, default=None)
    p.add_argument("--cmax", type=float, default=None)
    p.add_argument("--cstep", type=float, default=0.02)

    p.add_argument("--lsq-script", type=Path, default=Path("lsq_latt2.py"))
    p.add_argument("--show", type=int, default=0, choices=[0, 1], help="show Matplotlib graph window")
    p.add_argument("--save", "--save-plot", dest="save_plot", type=int, default=1, choices=[0, 1], help="save graph image")
    p.add_argument(
        "--plot-theory",
        type=str,
        default="near",
        choices=["near", "matched", "all", "none"],
        help=(
            "which calculated theoretical lines to draw in compare/calc plots. "
            "near draws only lines assigned to observed peaks within --plot-near-deg; "
            "all can be very dense for large unit cells. Default: near"
        ),
    )
    p.add_argument("--plot-near-deg", type=float, default=0.5, help="2theta window used by --plot-theory near")
    p.add_argument("--no-show", action="store_true", help="legacy option; equivalent to --show 0")
    return p

def resolve_mode_and_input(args: argparse.Namespace) -> tuple[str, Path]:
    first = str(args.arg1)
    if args.arg2 is not None and first.lower() in KNOWN_MODES:
        mode = normalize_mode(first)
        infile = Path(args.arg2)
    else:
        if args.arg2 is not None:
            raise SystemExit(
                "Unexpected second positional argument. Use either:\n"
                "  python guess_lattice_parameters.py sample.TXT --mode all\n"
                "or legacy style:\n"
                "  python guess_lattice_parameters.py search sample.TXT"
            )
        mode = normalize_mode(args.mode)
        infile = Path(args.arg1)
    return mode, infile


def apply_x_range(x: np.ndarray, y: np.ndarray, xmin: float | None, xmax: float | None) -> tuple[np.ndarray, np.ndarray]:
    if xmin is not None:
        mask = x >= xmin
        x, y = x[mask], y[mask]
    if xmax is not None:
        mask = x <= xmax
        x, y = x[mask], y[mask]
    if len(x) == 0:
        raise ValueError("No data points remain after applying xmin/xmax.")
    return x, y



def run_search(args: argparse.Namespace, infile: Path) -> tuple[list[Peak], Path, Path | None]:
    x, y = read_xy_data(infile)
    x, y = apply_x_range(x, y, args.xmin, args.xmax)

    print(f"Input file        : {infile}")
    print("Mode              : search")
    print(f"2theta range      : {x.min():.3f} - {x.max():.3f}")
    print(f"Threshold         : {args.threshold}")
    print(f"nsmooth           : {args.nsmooth}")
    print(f"norder            : {args.norder}")
    print(f"FWHM range        : {args.fwhm_min_deg} - {args.fwhm_max_deg} deg")

    peaks, info = peak_search_deriv3(
        x, y,
        nsmooth=args.nsmooth,
        norder=args.norder,
        threshold=args.threshold,
        ydiff1_threshold=args.ydiff1_threshold,
        fwhm_min_deg=args.fwhm_min_deg,
        fwhm_max_deg=args.fwhm_max_deg,
        is_print=True,
    )
    assign_ka2(peaks)

    print_peak_table(peaks, "Detected peaks")
    n_ka2 = sum(1 for p in peaks if p.ka_role == "ka2")
    print("")
    print(f"Detected Cu Kα2 peaks : {n_ka2}")
    print(f"Non-Kα2 peaks         : {sum(1 for p in peaks if p.ka_role != 'ka2')}")

    stem = sanitize_stem(infile.stem)
    plot_path = args.plot_file if args.plot_file is not None else infile.with_name(f"{stem}-peaksearch.png")
    peak_file = args.peak_file if args.peak_file is not None else default_peak_file_for_input(infile)

    save_plot = bool(args.save_plot)
    show_plot = bool(args.show) and not bool(args.no_show)
    plot_results(x, y, info, peaks, plot_path, show=show_plot, save=save_plot)
    save_workbook(peak_file, {"peaks": peaks_to_frame(peaks)})
    print("")
    print(f"Saved peak file   : {peak_file}")
    if save_plot:
        print(f"Saved plot        : {plot_path}")
    return peaks, peak_file, plot_path if save_plot else None


def run_guess(args: argparse.Namespace, peak_file: Path, base_input: Path | None = None) -> Path:
    crystal_systems = parse_crystal_systems(args.crystal_system)
    peaks = read_peaks_file(peak_file, wavelength=args.wavelength)

    print("")
    print(f"Peak file         : {peak_file}")
    print("Mode              : guess")
    print(f"Method            : {args.method}")
    print(f"Crystal system(s) : {', '.join(crystal_systems)}")
    print(f"npeak             : {args.npeak}")
    print(f"hmax              : {args.hmax}")

    print_peak_table(peaks, "Peaks read from file")
    n_ka2 = sum(1 for p in peaks if p.ka_role == "ka2")
    print("")
    print(f"Cu Kα2 peaks in file/estimate : {n_ka2}")
    print(f"Non-Kα2 peaks                : {sum(1 for p in peaks if p.ka_role != 'ka2')}")
    print("")
    args.resolved_lattice_max = resolve_lattice_max_limit(args, peaks)

    candidates, selected = guess_candidates(args, peaks, crystal_systems)
    print("")
    print(f"Guess uses strongest non-Kα2 peaks: {len(selected)}/{args.npeak}")
    print_peak_table(selected, "Selected peaks for lattice-parameter guessing")

    print("")
    print("Lattice-parameter candidates")
    if not candidates:
        print("No candidate was found. Try increasing --npeak, --hmax, or using --crystal-system auto.")
    for i, cand in enumerate(candidates[:10], 1):
        print(f"{i:2d}. {cand.crystal_system:12s} matches={cand.score_matches:2d} rms_rel={cand.score_rms_rel:.6e} params={cand.params}")

    refined_summary = None
    refined_rows: list[dict] = []
    if candidates and args.lsq_script.exists():
        refined = refine_with_lsq(args.lsq_script, candidates[0], [r for r in candidates[0].all_matches if r["matched"]], wavelength=args.wavelength)
        if refined is not None:
            refined_summary, refined_rows = refined
            print("")
            print("Refined lattice constants")
            print(f"  crystal_system : {refined_summary['crystal_system']}")
            print(f"  a = {refined_summary['a']:.6f} ± {refined_summary['sigma_a']:.6f} Å")
            print(f"  b = {refined_summary['b']:.6f} ± {refined_summary['sigma_b']:.6f} Å")
            print(f"  c = {refined_summary['c']:.6f} ± {refined_summary['sigma_c']:.6f} Å")
            print(f"  alpha = {refined_summary['alpha']:.6f} ± {refined_summary['sigma_alpha']:.6f} deg")
            print(f"  beta  = {refined_summary['beta']:.6f} ± {refined_summary['sigma_beta']:.6f} deg")
            print(f"  gamma = {refined_summary['gamma']:.6f} ± {refined_summary['sigma_gamma']:.6f} deg")
            print(f"  rms_rel_error_inv_d2 = {refined_summary['rms_rel_error_inv_d2']:.6e}")
            print(f"  rms_rel_error_2theta = {refined_summary['rms_rel_error_2theta']:.6e}")
    elif candidates:
        print("")
        print(f"Refinement skipped because lsq script was not found: {args.lsq_script}")

    sheets: dict[str, pd.DataFrame] = {
        "all_detected_peaks": peaks_to_frame(peaks),
        "selected_for_guess": peaks_to_frame(selected),
    }
    if candidates:
        cand_rows = [candidate_to_summary_row(i, c, method=args.method) for i, c in enumerate(candidates, 1)]
        sheets["candidate_summary"] = pd.DataFrame(cand_rows)
        sheets["all_peaks_assignment"] = pd.DataFrame(candidates[0].all_matches)
        sheets["selected_assignment"] = pd.DataFrame(candidates[0].selected_matches)
    if refined_summary:
        sheets["refined_summary"] = pd.DataFrame([refined_summary])
        sheets["refined_matched"] = pd.DataFrame(refined_rows)

    out_xlsx = args.output_file if args.output_file is not None else default_guess_file_for_peak_file(peak_file, base_input)
    save_workbook(out_xlsx, sheets)
    print("")
    print(f"Saved guess file  : {out_xlsx}")
    return out_xlsx


def _load_peaks_for_compare(args: argparse.Namespace, infile: Path) -> tuple[list[Peak], np.ndarray | None, np.ndarray | None, Path | None]:
    if args.peak_file is not None:
        peaks = read_peaks_file(args.peak_file, wavelength=args.wavelength)
        x, y = try_read_raw_xy_for_plot(infile)
        return peaks, x, y, args.peak_file

    # Dense two-column files are usually raw XRD profiles, not peak lists.
    if looks_like_raw_xrd_file(infile):
        peaks, peak_file, _plot = run_search(args, infile)
        x, y = try_read_raw_xy_for_plot(infile)
        return peaks, x, y, peak_file

    # If the input is a guess workbook, use its all_detected_peaks sheet if present.
    if infile.suffix.lower() in {".xlsx", ".xls"}:
        try:
            xls = pd.ExcelFile(infile)
            names = {str(s).lower() for s in xls.sheet_names}
            if "all_detected_peaks" in names or "peaks" in names:
                peaks = read_peaks_file(infile, wavelength=args.wavelength)
                return peaks, None, None, infile
        except Exception:
            pass

    # Otherwise first try peak-file interpretation; if that fails, treat as raw XRD and run search.
    try:
        peaks = read_peaks_file(infile, wavelength=args.wavelength)
        return peaks, None, None, infile
    except Exception:
        peaks, peak_file, _plot = run_search(args, infile)
        x, y = try_read_raw_xy_for_plot(infile)
        return peaks, x, y, peak_file


def run_compare(args: argparse.Namespace, infile: Path) -> Path:
    explicit_lattice = has_explicit_lattice_params(args)
    mode_label = "calc" if normalize_mode(args.mode) == "calc" or explicit_lattice else "compare"
    print("")
    print(f"Input file        : {infile}")
    print(f"Mode              : {mode_label}")
    if explicit_lattice:
        print("Candidate source  : explicit lattice constants")
    else:
        print(f"Candidate rank    : {args.candidate_rank}")

    peaks, x, y, peak_file_used = _load_peaks_for_compare(args, infile)
    if args.xmin is not None or args.xmax is not None:
        peaks = [p for p in peaks if (args.xmin is None or p.two_theta >= args.xmin) and (args.xmax is None or p.two_theta <= args.xmax)]
    if not peaks:
        raise ValueError("No peaks are available for comparison.")

    guess_file = args.guess_file
    if guess_file is None and infile.suffix.lower() in {".xlsx", ".xls"}:
        try:
            xls = pd.ExcelFile(infile)
            if "candidate_summary" in {str(s).lower() for s in xls.sheet_names}:
                guess_file = infile
        except Exception:
            guess_file = None

    if explicit_lattice:
        candidate = candidate_from_lattice_args(args)
    elif guess_file is not None:
        candidate = read_candidate_from_guess_file(guess_file, args.candidate_rank)
    else:
        crystal_systems = parse_crystal_systems(args.crystal_system)
        if not hasattr(args, "resolved_lattice_max"):
            print("")
            args.resolved_lattice_max = resolve_lattice_max_limit(args, peaks)
        candidates, _selected = guess_candidates(args, peaks, crystal_systems)
        if not candidates:
            raise ValueError("No candidate was found for comparison. Provide --guess-file, explicit --a/--b/--c, or adjust guess options.")
        if args.candidate_rank < 1 or args.candidate_rank > len(candidates):
            raise ValueError(f"Candidate rank {args.candidate_rank} is outside 1..{len(candidates)}")
        candidate = candidates[args.candidate_rank - 1]

    xmin = args.xmin
    xmax = args.xmax
    if xmin is None and x is not None and len(x):
        xmin = float(np.nanmin(x))
    if xmax is None and x is not None and len(x):
        xmax = float(np.nanmax(x))

    theoretical_df, assignment_df = compare_candidate_with_peaks(
        candidate,
        peaks,
        hmax=args.hmax,
        wavelength=args.wavelength,
        xmin=xmin,
        xmax=xmax,
        tolerance_deg=args.tolerance_deg,
    )

    print("")
    print("Candidate used")
    print(f"  crystal_system : {candidate.crystal_system}")
    for k in ["a", "b", "c", "alpha", "beta", "gamma"]:
        if k in candidate.params:
            print(f"  {k:5s} = {candidate.params[k]:.8g}")

    if not assignment_df.empty:
        matched = assignment_df[assignment_df["matched"] == True]
        print("")
        print(f"Observed peaks compared : {len(assignment_df)}")
        print(f"Matched within tolerance: {len(matched)}")
        if len(matched):
            rms = float(np.sqrt(np.mean(np.square(matched["resid_two_theta_deg"].to_numpy(dtype=float)))))
            print(f"RMS residual 2theta    : {rms:.6g} deg")
        cols = ["peak_id", "two_theta_obs_deg", "hkl", "two_theta_calc_deg", "resid_two_theta_deg", "intensity", "fwhm_deg", "matched"]
        with pd.option_context("display.max_rows", 300, "display.width", 180):
            print(assignment_df[cols].to_string(index=False))

    stem = sanitize_stem(strip_peak_suffix((peak_file_used or infile).stem))
    suffix = "calc" if explicit_lattice else f"compare-rank{args.candidate_rank}"
    out_xlsx = args.output_file if args.output_file is not None else (peak_file_used or infile).with_name(f"{stem}-{suffix}.xlsx")
    plot_path = args.plot_file if args.plot_file is not None else out_xlsx.with_suffix(".png")

    sheets = {
        "candidate_used": pd.DataFrame([candidate_to_summary_row(args.candidate_rank if not explicit_lattice else 0, candidate, method="explicit" if explicit_lattice else "")]),
        "observed_peaks": peaks_to_frame(peaks),
        "theoretical_lines": theoretical_df,
        "observed_vs_calc": assignment_df,
    }
    save_workbook(out_xlsx, sheets)

    save_plot = bool(args.save_plot)
    show_plot = bool(args.show) and not bool(args.no_show)
    plot_theoretical_df = filter_theoretical_lines_for_plot(
        theoretical_df,
        mode=args.plot_theory,
        near_deg=args.plot_near_deg if args.plot_near_deg is not None else args.tolerance_deg,
    )
    plot_compare_results(x, y, peaks, plot_theoretical_df, assignment_df, 
                candidate, xmin=xmin, xmax=xmax,
                outpath=plot_path, show=show_plot, save=save_plot)

    print("")
    print(f"Theoretical lines : {len(theoretical_df)} calculated, {len(plot_theoretical_df)} drawn (--plot-theory {args.plot_theory})")
    print(f"Saved compare file: {out_xlsx}")
    if save_plot:
        print(f"Saved plot        : {plot_path}")
    if show_plot:
        print("Hover annotations : enabled when mplcursors is installed")
    return out_xlsx


def main() -> int:
    args = build_argument_parser().parse_args()
    try:
        if args.no_show:
            args.show = 0
        mode, infile = resolve_mode_and_input(args)
        print(f"Requested mode    : {mode}")
        if mode == "search":
            run_search(args, infile)
            return 0
        if mode == "guess":
            peak_file = args.peak_file if args.peak_file is not None else infile
            run_guess(args, peak_file)
            return 0
        if mode in {"compare", "calc"}:
            if mode == "calc" and not has_explicit_lattice_params(args):
                raise ValueError("--mode calc requires explicit lattice constants, for example --crystal-system cubic --a 3.905")
            run_compare(args, infile)
            return 0
        if mode == "all":
            _peaks, peak_file, _plot_path = run_search(args, infile)
            guess_file = run_guess(args, peak_file, base_input=infile)
            # --mode all remains search + guess only. Use --mode compare with --guess-file to compare a selected rank.
            return 0
        raise SystemExit(f"Unsupported mode: {mode}")
    except Exception as exc:
        print(f"ERROR: {exc}", file=sys.stderr)
        raise

    input("\nPress ENTER to terminate>>\n")
    
if __name__ == "__main__":
    raise SystemExit(main())
