"""Continuity-strain critical sweep for the standalone proof."""

from __future__ import annotations

from dataclasses import dataclass
from typing import Dict, Iterable, List

import numpy as np

from ..config import SweepConfig, WormholeConfig
from .wormhole import SimulationResult, run_experiment


@dataclass(frozen=True)
class SweepPoint:
    lambda_h: float
    throughput: float
    relaxation_steps: float
    soft_rank_mean: float
    soft_rank_fraction: float
    sigma: float
    percolation_mean: float
    percolation_max: float
    update_ratio_mean: float
    candidate_norm_mean: float
    invariant_ok: bool


@dataclass(frozen=True)
class SweepAnalysis:
    lambda_crit_est: float | None
    throughput_alpha: float | None
    throughput_logA: float | None
    relaxation_z: float | None
    relaxation_logB: float | None


@dataclass(frozen=True)
class SweepResult:
    points: List[SweepPoint]
    analysis: SweepAnalysis


def compute_sigma(active_rank: float, soft_rank: float) -> float:
    if soft_rank <= 0:
        return 1.0
    frac = np.clip(active_rank / soft_rank, 0.0, 1.0)
    return 1.0 - frac


def fit_power_law(x: np.ndarray, y: np.ndarray) -> Dict[str, float] | None:
    mask = (x > 1e-6) & (y > 1e-9) & np.isfinite(x) & np.isfinite(y)
    if mask.sum() < 3:
        return None
    lx = np.log(x[mask])
    ly = np.log(y[mask])
    try:
        slope, intercept = np.polyfit(lx, ly, 1)
    except np.linalg.LinAlgError:
        return None
    if not np.isfinite(slope) or not np.isfinite(intercept):
        return None
    return {"slope": float(slope), "intercept": float(intercept)}


def analyse(points: List[SweepPoint], critical_fraction: float) -> SweepAnalysis:
    if not points:
        return SweepAnalysis(None, None, None, None, None)

    soft_fracs = np.array([p.soft_rank_fraction for p in points])
    lambdas = np.array([p.lambda_h for p in points])

    below = np.where(soft_fracs <= critical_fraction)[0]
    if below.size == 0:
        lambda_crit = float(lambdas.max())
    else:
        idx = below[0]
        if idx == 0:
            lambda_crit = float(lambdas[0])
        else:
            lambda_crit = float(
                np.interp(
                    critical_fraction,
                    soft_fracs[idx - 1 : idx + 1],
                    lambdas[idx - 1 : idx + 1],
                )
            )

    one_minus_sigma = np.clip(1.0 - np.array([p.sigma for p in points]), 1e-8, 1.0)
    throughput = np.array([p.throughput for p in points])
    relaxation = np.array([p.relaxation_steps for p in points])

    mask_sub = lambdas <= lambda_crit
    throughput_fit = fit_power_law(one_minus_sigma[mask_sub], throughput[mask_sub])
    relaxation_fit = None
    if np.all(relaxation[mask_sub] > 0):
        inv_relax = 1.0 / (relaxation[mask_sub] + 1e-12)
        relaxation_fit = fit_power_law(one_minus_sigma[mask_sub], inv_relax)

    return SweepAnalysis(
        lambda_crit_est=lambda_crit,
        throughput_alpha=throughput_fit["slope"] if throughput_fit else None,
        throughput_logA=throughput_fit["intercept"] if throughput_fit else None,
        relaxation_z=-(relaxation_fit["slope"]) if relaxation_fit else None,
        relaxation_logB=relaxation_fit["intercept"] if relaxation_fit else None,
    )


def run_sweep(
    sweep_cfg: SweepConfig,
    base_cfg: WormholeConfig,
    seed: int,
) -> SweepResult:
    lambda_values = np.linspace(sweep_cfg.lambda_start, sweep_cfg.lambda_end, sweep_cfg.lambda_points)
    points: List[SweepPoint] = []

    for lam in lambda_values:
        cfg = WormholeConfig(
            **{**base_cfg.__dict__, "lambda_h": float(lam)}
        )
        sim = run_experiment(cfg, seed)
        wormhole = next(r for r in sim.regimes if r.name == "wormhole")
        soft_rank_mean = wormhole.metrics.soft_active_rank_mean
        soft_rank_fraction = soft_rank_mean / max(cfg.soft_rank, 1e-6)
        sigma = compute_sigma(soft_rank_mean, cfg.soft_rank)

        points.append(
            SweepPoint(
                lambda_h=float(lam),
                throughput=wormhole.throughput,
                relaxation_steps=wormhole.metrics.relaxation_steps,
                soft_rank_mean=soft_rank_mean,
                soft_rank_fraction=float(soft_rank_fraction),
                sigma=float(sigma),
                percolation_mean=wormhole.metrics.percolation_span_mean,
                percolation_max=wormhole.metrics.percolation_span_max,
                update_ratio_mean=wormhole.metrics.update_ratio_mean,
                candidate_norm_mean=wormhole.metrics.candidate_norm_mean,
                invariant_ok=wormhole.invariant_ok,
            )
        )

    analysis = analyse(points, sweep_cfg.critical_fraction)
    return SweepResult(points=points, analysis=analysis)
