"""
Numerical convergence testing for the FCE simulator.

Every numerical simulation must demonstrate that results converge
as resolution increases. Without convergence testing, reviewers
cannot distinguish physical effects from numerical artifacts.

Tests at resolutions: 100, 200, 400, 800, 1600, 3200 steps.
Computes error E(h) = |observable(h) - observable(h_finest)|.
Fits log(E) vs log(h) to extract convergence order p: E(h) ~ h^p.

Standard numerical analysis requires p > 0 for validity and
p matching the expected order of the integration scheme.

References:
    [1] LeVeque, "Finite Difference Methods for ODEs and PDEs" (2007)
    [2] Hairer et al., "Solving Ordinary Differential Equations I" (1993)
"""

import logging
import numpy as np
from concurrent.futures import ProcessPoolExecutor, as_completed
from dataclasses import dataclass
from typing import List, Optional, Dict, Any, Callable

from .engine import FractalCorrectionEngine, EngineConfig, EvolutionResult
from .system_health import get_default_monitor

logger = logging.getLogger(__name__)


@dataclass
class ConvergencePoint:
    """Result at a single resolution."""
    n_steps: int
    dt: float
    observable_value: float
    runtime_seconds: float


@dataclass
class ConvergenceResult:
    """Complete convergence study results."""
    points: List[ConvergencePoint]
    errors: np.ndarray            # E(h) at each resolution (except finest)
    step_sizes: np.ndarray        # dt at each resolution (except finest)
    convergence_order: float      # Fitted p from E(h) ~ h^p
    r_squared: float              # R^2 of the log-log fit
    richardson_estimate: float    # Richardson extrapolation of true value
    observable_name: str

    @property
    def n_resolutions(self) -> int:
        return len(self.points)

    @property
    def order_flag(self) -> Optional[str]:
        """Flag problematic convergence behavior.

        Returns:
            None if convergence is acceptable, or a warning string.
        """
        if self.convergence_order < 0:
            return "NEGATIVE_ORDER"
        if self.r_squared < 0.5:
            return "POOR_FIT"
        return None

    def summary_table(self) -> str:
        """Format as publication-ready convergence table."""
        lines = []
        lines.append(f"Convergence Study: {self.observable_name}")
        lines.append("=" * 70)
        lines.append(
            f"{'n_steps':>10} {'dt':>12} "
            f"{'Observable':>15} {'Error E(h)':>15} {'Ratio':>10}"
        )
        lines.append("-" * 70)

        finest_val = self.points[-1].observable_value
        prev_error = None

        for i, pt in enumerate(self.points):
            error = abs(pt.observable_value - finest_val)
            error_str = f"{error:.6e}" if i < len(self.points) - 1 else "---"

            if prev_error is not None and error > 1e-15:
                ratio = prev_error / error
                ratio_str = f"{ratio:.2f}"
            else:
                ratio_str = "---"

            lines.append(
                f"{pt.n_steps:>10} {pt.dt:>12.2e} "
                f"{pt.observable_value:>15.8f} "
                f"{error_str:>15} {ratio_str:>10}"
            )

            if error > 1e-15:
                prev_error = error

        lines.append("-" * 70)
        lines.append(f"Convergence order p: {self.convergence_order:.3f}")
        lines.append(f"R^2 of log-log fit: {self.r_squared:.6f}")
        lines.append(f"Richardson extrapolation: {self.richardson_estimate:.10f}")
        lines.append(f"Finest grid value:       {finest_val:.10f}")

        flag = self.order_flag
        if flag:
            lines.append(f"WARNING: {flag}")
            if flag == "NEGATIVE_ORDER":
                lines.append(
                    "  Error increases with resolution -- possible instability"
                )
            elif flag == "POOR_FIT":
                lines.append(
                    "  Log-log fit is poor -- convergence behavior is irregular"
                )

        return "\n".join(lines)


@dataclass
class MultiObservableConvergenceResult:
    """Convergence results for multiple observables simultaneously."""
    results: Dict[str, ConvergenceResult]

    @property
    def all_converged(self) -> bool:
        """True if no observable has a convergence flag."""
        return all(r.order_flag is None for r in self.results.values())

    def summary_table(self) -> str:
        """Format all observables in a single comparison table."""
        lines = []
        lines.append("Multi-Observable Convergence Study")
        lines.append("=" * 75)
        lines.append(
            f"{'Observable':<25} {'Order p':>10} {'R^2':>10} "
            f"{'Richardson':>15} {'Flag':>12}"
        )
        lines.append("-" * 75)

        for name, result in self.results.items():
            flag = result.order_flag or "OK"
            lines.append(
                f"{name:<25} "
                f"{result.convergence_order:>10.3f} "
                f"{result.r_squared:>10.6f} "
                f"{result.richardson_estimate:>15.8f} "
                f"{flag:>12}"
            )

        lines.append("-" * 75)
        status = "ALL CONVERGED" if self.all_converged else "CHECK WARNINGS"
        lines.append(f"Status: {status}")
        return "\n".join(lines)


def _convergence_worker(args: dict) -> dict:
    """
    Top-level worker for parallel convergence resolution runs.

    Must be top-level for pickling by ProcessPoolExecutor.
    """
    import time as time_mod

    n_steps = args["n_steps"]
    total_time = args["total_time"]
    H = args["H"]
    rho_init = args["rho_init"]
    system_params = args["system_params"]
    config_dict = args["config_dict"]

    dt = total_time / n_steps

    config = EngineConfig(**config_dict)
    config.enable_thermal_monitoring = False

    engine = FractalCorrectionEngine(
        hamiltonian=H,
        system_params=system_params,
        config=config,
    )

    t_start = time_mod.perf_counter()
    result = engine.evolve(rho_init, dt, n_steps)
    t_end = time_mod.perf_counter()

    # Default observable: final fidelity
    obs_val = float(result.fidelities[-1])

    return {
        "n_steps": n_steps,
        "dt": dt,
        "observable_value": obs_val,
        "runtime_seconds": t_end - t_start,
    }


class ConvergenceStudy:
    """
    Tests numerical convergence by running at multiple resolutions.

    For standard RK45 integration (used by LindbladEvolver), the
    expected convergence order is p >= 4 for smooth problems.

    Supports parallel execution of independent resolution runs via
    ProcessPoolExecutor with dynamic worker reduction.
    """

    def __init__(
        self,
        hamiltonian: np.ndarray,
        rho_init: np.ndarray,
        total_time: float,
        system_params: Optional[Dict[str, float]] = None,
        config: Optional[EngineConfig] = None,
        observable: Optional[Callable[[EvolutionResult], float]] = None,
        observable_name: str = "final_fidelity",
    ):
        """
        Args:
            hamiltonian: System Hamiltonian
            rho_init: Initial density matrix
            total_time: Total evolution time T (held constant)
            system_params: Decoherence parameters
            config: Engine configuration
            observable: Function extracting a scalar from EvolutionResult.
                        Default: final fidelity.
            observable_name: Label for the observable
        """
        self.H = hamiltonian
        self.rho_init = rho_init
        self.total_time = total_time
        self.system_params = system_params
        self.config = config or EngineConfig()
        self.observable_name = observable_name

        if observable is not None:
            self.observable = observable
        else:
            self.observable = lambda r: float(r.fidelities[-1])

    def run(
        self,
        resolutions: Optional[List[int]] = None,
        parallel: bool = True,
        max_workers: Optional[int] = None,
    ) -> ConvergenceResult:
        """
        Run convergence study across multiple resolutions.

        Args:
            resolutions: List of n_steps values. Default: [100,200,400,800,1600,3200]
            parallel: If True, run resolutions in parallel
            max_workers: Max parallel workers (auto-detected if None)

        Returns:
            ConvergenceResult with errors, fitted order, R^2
        """
        if resolutions is None:
            resolutions = [100, 200, 400, 800, 1600, 3200]

        # Determine worker count
        monitor = get_default_monitor()
        use_parallel = parallel and self.observable is None or (
            parallel and self.observable == (lambda r: float(r.fidelities[-1]))
        )
        # Custom observables cannot be pickled, fall back to sequential
        if parallel and self.observable_name != "final_fidelity":
            use_parallel = False
            logger.info(
                "Convergence study: custom observable, falling back to sequential"
            )

        if use_parallel:
            monitor.wait_until_cool()
            if max_workers is None:
                max_workers = monitor.get_recommended_workers()
            logger.info(
                "Convergence study: %d resolutions, %d workers",
                len(resolutions), max_workers,
            )

        config_dict = {
            "feedback_gain": self.config.feedback_gain,
            "correction_threshold": self.config.correction_threshold,
            "qec_enabled": self.config.qec_enabled,
            "prediction_steps": self.config.prediction_steps,
            "prediction_ds": self.config.prediction_ds,
            "health_window": self.config.health_window,
            "enable_thermal_monitoring": False,
            "health_check_interval": self.config.health_check_interval,
        }

        points = []

        if use_parallel and max_workers and max_workers > 1:
            jobs = [
                {
                    "n_steps": n,
                    "total_time": self.total_time,
                    "H": self.H,
                    "rho_init": self.rho_init,
                    "system_params": self.system_params,
                    "config_dict": config_dict,
                }
                for n in resolutions
            ]

            results_map = {}
            with ProcessPoolExecutor(max_workers=max_workers) as executor:
                futures = {
                    executor.submit(_convergence_worker, job): job["n_steps"]
                    for job in jobs
                }
                for future in as_completed(futures):
                    n = futures[future]
                    try:
                        res = future.result()
                        results_map[res["n_steps"]] = res
                    except Exception as exc:
                        logger.error(
                            "Convergence run failed (n_steps=%d): %s", n, exc
                        )
                monitor.cooldown_between_steps()

            # Reconstruct in resolution order
            for n in resolutions:
                if n in results_map:
                    res = results_map[n]
                    points.append(ConvergencePoint(
                        n_steps=res["n_steps"],
                        dt=res["dt"],
                        observable_value=res["observable_value"],
                        runtime_seconds=res["runtime_seconds"],
                    ))
        else:
            import time as time_mod
            for n_steps in resolutions:
                dt = self.total_time / n_steps

                engine = FractalCorrectionEngine(
                    hamiltonian=self.H,
                    system_params=self.system_params,
                    config=self.config,
                )

                t_start = time_mod.perf_counter()
                result = engine.evolve(self.rho_init, dt, n_steps)
                t_end = time_mod.perf_counter()

                obs_val = self.observable(result)

                points.append(ConvergencePoint(
                    n_steps=n_steps,
                    dt=dt,
                    observable_value=obs_val,
                    runtime_seconds=t_end - t_start,
                ))
                monitor.cooldown_between_steps()

        # Compute errors relative to finest resolution
        finest_val = points[-1].observable_value
        errors = []
        step_sizes = []

        for pt in points[:-1]:
            errors.append(abs(pt.observable_value - finest_val))
            step_sizes.append(pt.dt)

        errors = np.array(errors)
        step_sizes = np.array(step_sizes)

        # Fit convergence order: log(E) = p * log(h) + c
        # Only use points where error > 0
        mask = errors > 1e-15
        if np.sum(mask) >= 2:
            log_h = np.log(step_sizes[mask])
            log_e = np.log(errors[mask])
            coeffs = np.polyfit(log_h, log_e, 1)
            convergence_order = coeffs[0]

            # R^2
            predicted = np.polyval(coeffs, log_h)
            ss_res = np.sum((log_e - predicted) ** 2)
            ss_tot = np.sum((log_e - np.mean(log_e)) ** 2)
            r_squared = 1.0 - ss_res / ss_tot if ss_tot > 0 else 0.0
        else:
            convergence_order = 0.0
            r_squared = 0.0

        # Richardson extrapolation using two finest grids
        if len(points) >= 2 and convergence_order > 0:
            f_h = points[-2].observable_value
            f_h2 = points[-1].observable_value
            r = points[-2].dt / points[-1].dt  # Refinement ratio
            richardson_estimate = (
                r**convergence_order * f_h2 - f_h
            ) / (r**convergence_order - 1.0)
        else:
            richardson_estimate = finest_val

        return ConvergenceResult(
            points=points,
            errors=errors,
            step_sizes=step_sizes,
            convergence_order=convergence_order,
            r_squared=r_squared,
            richardson_estimate=richardson_estimate,
            observable_name=self.observable_name,
        )

    def run_multi(
        self,
        observables: Optional[Dict[str, Callable[[EvolutionResult], float]]] = None,
        resolutions: Optional[List[int]] = None,
    ) -> MultiObservableConvergenceResult:
        """
        Run convergence study for multiple observables simultaneously.

        Runs the engine once per resolution and extracts all observables
        from the stored result, avoiding redundant computation.

        Args:
            observables: Dict of {name: extraction_function}. Default:
                         final_fidelity, mean_curvature, mean_torsion,
                         max_trace_error.
            resolutions: List of n_steps values. Default: [100, 200, 400, 800].

        Returns:
            MultiObservableConvergenceResult with per-observable convergence.
        """
        import time as time_mod

        if resolutions is None:
            resolutions = [100, 200, 400, 800]

        if observables is None:
            observables = {
                'final_fidelity': lambda r: float(r.fidelities[-1]),
                'mean_fidelity': lambda r: float(np.mean(r.fidelities)),
                'hausdorff_dim': lambda r: float(r.hausdorff_dimension),
                'max_trace_error': lambda r: (
                    float(np.max(r.trace_errors))
                    if r.trace_errors is not None else 0.0
                ),
            }

        monitor = get_default_monitor()

        # Run engine once per resolution, extract all observables
        all_obs_values: Dict[str, List[float]] = {
            name: [] for name in observables
        }
        points_data = []

        for n_steps in resolutions:
            dt = self.total_time / n_steps

            engine = FractalCorrectionEngine(
                hamiltonian=self.H,
                system_params=self.system_params,
                config=self.config,
            )

            t_start = time_mod.perf_counter()
            result = engine.evolve(self.rho_init, dt, n_steps)
            t_end = time_mod.perf_counter()

            runtime = t_end - t_start

            for name, extract_fn in observables.items():
                try:
                    val = extract_fn(result)
                except Exception:
                    val = 0.0
                all_obs_values[name].append(val)

            points_data.append({
                'n_steps': n_steps,
                'dt': dt,
                'runtime': runtime,
            })
            monitor.cooldown_between_steps()

        # Build per-observable ConvergenceResult
        results: Dict[str, ConvergenceResult] = {}

        for name in observables:
            values = all_obs_values[name]
            finest_val = values[-1]

            obs_points = [
                ConvergencePoint(
                    n_steps=pd['n_steps'],
                    dt=pd['dt'],
                    observable_value=v,
                    runtime_seconds=pd['runtime'],
                )
                for pd, v in zip(points_data, values)
            ]

            errors = np.array([abs(v - finest_val) for v in values[:-1]])
            step_sizes = np.array([pd['dt'] for pd in points_data[:-1]])

            # Fit convergence order
            mask = errors > 1e-15
            if np.sum(mask) >= 2:
                log_h = np.log(step_sizes[mask])
                log_e = np.log(errors[mask])
                coeffs = np.polyfit(log_h, log_e, 1)
                order = coeffs[0]
                predicted = np.polyval(coeffs, log_h)
                ss_res = np.sum((log_e - predicted) ** 2)
                ss_tot = np.sum((log_e - np.mean(log_e)) ** 2)
                r2 = 1.0 - ss_res / ss_tot if ss_tot > 0 else 0.0
            else:
                order = 0.0
                r2 = 0.0

            # Richardson extrapolation
            if len(values) >= 2 and order > 0:
                f_h = values[-2]
                f_h2 = values[-1]
                r = points_data[-2]['dt'] / points_data[-1]['dt']
                rich = (r**order * f_h2 - f_h) / (r**order - 1.0)
            else:
                rich = finest_val

            results[name] = ConvergenceResult(
                points=obs_points,
                errors=errors,
                step_sizes=step_sizes,
                convergence_order=order,
                r_squared=r2,
                richardson_estimate=rich,
                observable_name=name,
            )

        return MultiObservableConvergenceResult(results=results)
