from __future__ import annotations

from dataclasses import dataclass, replace
from itertools import product
from typing import Iterable, List, Sequence

import math
import statistics
import numpy as np

from wormhole_proof.core.wormhole_smooth import SmoothConfig, run_smooth_experiment
from wormhole_proof.wireframe.projection import geometry_to_payload


# ---------------------------------------------------------------------------
# Utility helpers
# ---------------------------------------------------------------------------


def _angle(vec_a: np.ndarray, vec_b: np.ndarray) -> float:
    """Return the angle in radians between two vectors."""
    a_norm = np.linalg.norm(vec_a)
    b_norm = np.linalg.norm(vec_b)
    if a_norm == 0 or b_norm == 0:
        return float("nan")
    cos_theta = np.clip(np.dot(vec_a, vec_b) / (a_norm * b_norm), -1.0, 1.0)
    return float(np.arccos(cos_theta))


def _project_geometry(result) -> dict | None:
    if result.geometry is None:
        return None
    payload = geometry_to_payload(result.geometry, limit=512)
    if not payload:
        return None
    return payload


def _vector_sequence(payload: dict | None, key: str) -> np.ndarray:
    if not payload:
        return np.empty((0, 3))
    projected = payload.get("projected", {})
    raw = projected.get(key) or payload.get(key)
    if not raw:
        return np.empty((0, 3))
    arr = np.asarray(raw, dtype=float)
    if arr.ndim == 1:
        arr = arr.reshape(1, -1)
    return arr


def _bit_history(payload: dict | None) -> np.ndarray:
    if not payload:
        return np.empty((0, 0))
    bits = payload.get("hp_decode_bits")
    if not bits:
        return np.empty((0, 0))
    return np.asarray(bits, dtype=float)


def _scalar_series(payload: dict | None, key: str) -> np.ndarray:
    if not payload:
        return np.empty(0)
    series = payload.get(key)
    if not series:
        return np.empty(0)
    return np.asarray(series, dtype=float)


def _finite(values: Iterable[float]) -> List[float]:
    return [float(v) for v in values if isinstance(v, (int, float)) and math.isfinite(v)]


def _mean(values: Iterable[float]) -> float:
    finite = _finite(values)
    if not finite:
        return 0.0
    return float(statistics.mean(finite))


def _std(values: Iterable[float]) -> float:
    finite = _finite(values)
    if len(finite) <= 1:
        return 0.0
    return float(statistics.pstdev(finite))


def _min_value(values: Iterable[float]) -> float:
    finite = _finite(values)
    if not finite:
        return 0.0
    return float(min(finite))


def _max_value(values: Iterable[float]) -> float:
    finite = _finite(values)
    if not finite:
        return 0.0
    return float(max(finite))


def _extract_metrics(result) -> dict[str, float]:
    metrics = result.metrics
    return {
        "throughput": float(metrics.throughput),
        "hp_match": float(metrics.hp_match),
        "packets": int(metrics.packets),
        "margin_min": float(metrics.cfl_margin_min),
        "lambda_max_mean": float(metrics.lambda_max_mean),
    }


def _metric_row(
    experiment: str,
    label: str,
    seed: int | None,
    metrics: dict[str, float],
    *,
    aggregate: bool = False,
    num_samples: int = 1,
    extra: dict | None = None,
) -> dict:
    row = {
        "experiment": experiment,
        "label": label,
        "seed": seed if seed is not None else -1,
        "aggregate": aggregate,
        "num_samples": num_samples,
        "throughput": metrics.get("throughput", 0.0),
        "hp_match": metrics.get("hp_match", 0.0),
        "packets": metrics.get("packets", 0.0),
        "margin_min": metrics.get("margin_min", 0.0),
        "lambda_max_mean": metrics.get("lambda_max_mean", 0.0),
        "comment": "",
        "hp_magnitude": "",
        "budget_fraction": "",
    }
    if extra:
        row.update(extra)
    return row


def _scaled_config(
    cfg: SmoothConfig,
    *,
    step_scale: float = 1.0,
    tolerance_scale: float = 1.0,
    mu_scale: float = 1.0,
    shock_scale: float = 1.0,
    hp_scale: float = 1.0,
    budget_scale: float = 1.0,
) -> SmoothConfig:
    clamp = lambda value, lo, hi: max(lo, min(value, hi))
    return replace(
        cfg,
        wormhole_step_scale=clamp(cfg.wormhole_step_scale * step_scale, 1e-4, 1.0),
        wormhole_step_scale_final=clamp(cfg.wormhole_step_scale_final * step_scale, 1e-4, 1.0),
        verify_tolerance=clamp(cfg.verify_tolerance * tolerance_scale, 1e-4, 0.5),
        verify_tolerance_final=clamp(cfg.verify_tolerance_final * tolerance_scale, 1e-4, 0.5),
        traversable_mu=cfg.traversable_mu * mu_scale,
        traversable_mu_final=cfg.traversable_mu_final * mu_scale,
        shock_energy=cfg.shock_energy * shock_scale,
        shock_energy_final=cfg.shock_energy_final * shock_scale,
        hp_message_magnitude=cfg.hp_message_magnitude * hp_scale,
        wormhole_budget_fraction=clamp(cfg.wormhole_budget_fraction * budget_scale, 0.02, 0.3),
    )


def _topology_stage_configs(topology: str, base_cfg: SmoothConfig) -> List[SmoothConfig]:
    topology = topology.lower()
    if topology == "series":
        specs = (
            {"step_scale": 1.0, "tolerance_scale": 1.05, "mu_scale": 1.0, "shock_scale": 1.0, "hp_scale": 1.0, "budget_scale": 1.0},
            {"step_scale": 0.72, "tolerance_scale": 1.25, "mu_scale": 0.92, "shock_scale": 0.95, "hp_scale": 0.85, "budget_scale": 0.82},
            {"step_scale": 0.55, "tolerance_scale": 1.45, "mu_scale": 0.88, "shock_scale": 0.9, "hp_scale": 0.75, "budget_scale": 0.7},
        )
    elif topology == "parallel":
        specs = (
            {"step_scale": 0.95, "tolerance_scale": 0.95, "mu_scale": 1.05, "shock_scale": 1.0, "hp_scale": 1.0, "budget_scale": 0.95},
            {"step_scale": 1.08, "tolerance_scale": 0.9, "mu_scale": 1.12, "shock_scale": 1.05, "hp_scale": 1.05, "budget_scale": 1.08},
            {"step_scale": 1.08, "tolerance_scale": 0.9, "mu_scale": 1.12, "shock_scale": 1.05, "hp_scale": 1.05, "budget_scale": 1.08},
        )
    elif topology == "triangle":
        specs = (
            {"step_scale": 0.92, "tolerance_scale": 1.02, "mu_scale": 1.05, "shock_scale": 1.1, "hp_scale": 1.0, "budget_scale": 1.0},
            {"step_scale": 0.95, "tolerance_scale": 1.0, "mu_scale": 1.0, "shock_scale": 0.95, "hp_scale": 1.0, "budget_scale": 1.0},
            {"step_scale": 0.92, "tolerance_scale": 1.02, "mu_scale": 1.05, "shock_scale": 1.0, "hp_scale": 1.0, "budget_scale": 1.02},
        )
    else:
        specs = (
            {"step_scale": 1.0, "tolerance_scale": 1.0, "mu_scale": 1.0, "shock_scale": 1.0, "hp_scale": 1.0, "budget_scale": 1.0},
        ) * 3

    return [_scaled_config(base_cfg, **params) for params in specs]


def _aggregate_topology(topology: str, samples: Sequence[TopologySample]) -> TopologyAggregate:
    note_map = {
        "series": "Series chain limited by weakest throat; metrics derived from per-stage minima.",
        "parallel": "Parallel branches sum throughput; HP averaged by throughput weighting.",
        "triangle": "Triangle coupling averages stage metrics, capturing cooperative stability.",
    }
    note = note_map.get(topology.lower(), "Topology aggregation using mean statistics.")

    if not samples:
        return TopologyAggregate(
            topology=topology,
            throughput_mean=0.0,
            throughput_std=0.0,
            hp_match_mean=0.0,
            hp_packets_mean=0.0,
            margin_mean=0.0,
            notes=note,
        )

    per_seed_throughputs: List[float] = []
    per_seed_hp: List[float] = []
    per_seed_packets: List[float] = []
    per_seed_margins: List[float] = []

    for seed in sorted({sample.seed for sample in samples}):
        stage_samples = [sample for sample in samples if sample.seed == seed]
        if not stage_samples:
            continue
        stage_throughputs = [sample.throughput for sample in stage_samples]
        stage_hp = [sample.hp_match for sample in stage_samples]
        stage_packets = [sample.hp_packets for sample in stage_samples]
        stage_margins = [sample.margin_min for sample in stage_samples]

        topology_key = topology.lower()
        if topology_key == "series":
            total_throughput = min(stage_throughputs)
            hp_match = min(stage_hp)
            hp_packets = min(stage_packets)
            margin = min(stage_margins)
        elif topology_key == "parallel":
            total_throughput = sum(stage_throughputs)
            if total_throughput > 0:
                hp_match = sum(t * h for t, h in zip(stage_throughputs, stage_hp)) / total_throughput
            else:
                hp_match = _mean(stage_hp)
            hp_packets = sum(stage_packets)
            margin = max(stage_margins)
        elif topology_key == "triangle":
            total_throughput = _mean(stage_throughputs)
            hp_match = _mean(stage_hp)
            hp_packets = _mean(stage_packets)
            margin = _mean(stage_margins)
        else:
            total_throughput = _mean(stage_throughputs)
            hp_match = _mean(stage_hp)
            hp_packets = _mean(stage_packets)
            margin = _mean(stage_margins)

        per_seed_throughputs.append(float(total_throughput))
        per_seed_hp.append(float(hp_match))
        per_seed_packets.append(float(hp_packets))
        per_seed_margins.append(float(margin))

    return TopologyAggregate(
        topology=topology,
        throughput_mean=_mean(per_seed_throughputs),
        throughput_std=_std(per_seed_throughputs),
        hp_match_mean=_mean(per_seed_hp),
        hp_packets_mean=_mean(per_seed_packets),
        margin_mean=_mean(per_seed_margins),
        notes=note,
    )


# ---------------------------------------------------------------------------
# Experiment result containers
# ---------------------------------------------------------------------------


@dataclass(frozen=True)
class LensingSample:
    strength: float
    offset: float
    deflection_angle: float
    straight_angle: float
    packets: int
    expected_angle: float
    curvature_number: float


@dataclass(frozen=True)
class HawkingSample:
    strength: float
    isolation_steps: int
    leakage_rate: float
    residual_norm: float
    decay_factor: float


@dataclass(frozen=True)
class FrameDraggingSample:
    angular_momentum: float
    radius: float
    mean_rotation_rate: float
    max_rotation_rate: float
    expected_mean_rate: float
    expected_max_rate: float


@dataclass(frozen=True)
class TunnelingSample:
    barrier_height: float
    throughput: float
    tunneled: bool
    probability: float
    effective_throughput: float


@dataclass(frozen=True)
class CasimirSample:
    separation: float
    interaction_energy: float
    expected_energy: float


@dataclass(frozen=True)
class WaveSample:
    disturbance: float
    propagation_delay: float
    wave_speed: float


@dataclass(frozen=True)
class TriangleSample:
    seed_triplet: tuple[int, int, int]
    pairwise_throughput: dict[str, float]


@dataclass(frozen=True)
class TopologyAggregate:
    topology: str
    throughput_mean: float
    throughput_std: float
    hp_match_mean: float
    hp_packets_mean: float
    margin_mean: float
    notes: str


@dataclass(frozen=True)
class TopologySample:
    topology: str
    stage: int
    seed: int
    throughput: float
    hp_match: float
    hp_packets: int
    margin_min: float


@dataclass(frozen=True)
class NullAblationSample:
    label: str
    seed: int
    throughput: float
    hp_match: float
    packets: int
    margin_min: float
    lambda_max_mean: float
    comment: str


@dataclass(frozen=True)
class StressSample:
    label: str
    seed: int
    throughput: float
    hp_match: float
    packets: int
    margin_min: float
    lambda_max_mean: float
    comment: str


@dataclass(frozen=True)
class SeedStatistic:
    label: str
    throughput_mean: float
    throughput_std: float
    throughput_min: float
    throughput_max: float
    hp_match_mean: float
    hp_match_std: float
    margin_mean: float
    margin_std: float
    packets_mean: float
    packets_std: float
    comment: str


@dataclass(frozen=True)
class HpStressSample:
    magnitude: float
    budget_fraction: float
    throughput_mean: float
    throughput_std: float
    hp_match_mean: float
    hp_match_std: float
    hp_packets_mean: float
    hp_packets_std: float
    margin_mean: float
    margin_std: float
    lambda_max_mean: float
    lambda_max_std: float
    sample_count: int
    comment: str


# ---------------------------------------------------------------------------
# Serialization helpers
# ---------------------------------------------------------------------------


def dataclass_list_to_dict(items: Sequence[object]) -> List[dict]:
    from dataclasses import asdict, is_dataclass

    data: List[dict] = []
    for item in items:
        if is_dataclass(item):
            data.append(asdict(item))
        else:
            raise TypeError(f"Expected dataclass instance, got {type(item)!r}")
    return data


# ---------------------------------------------------------------------------
# Experiments
# ---------------------------------------------------------------------------


def test_lensing(
    strengths: Sequence[float],
    offsets: Sequence[float],
    *,
    base_config: SmoothConfig | None = None,
    base_seed: int = 7,
) -> List[LensingSample]:
    """Probe deflection of packet trajectories around high-salience regions.

    `strengths` map onto `lambda_h` / `lambda_h_final` (salience "mass").
    `offsets` perturb the RNG seed to emulate impact parameter changes.

    Returns angle measurements between the final packet vector and the
    aggregate state-link baseline.
    """

    cfg_template = base_config or SmoothConfig()
    samples: List[LensingSample] = []

    for strength in strengths:
        mass_factor = strength / max(cfg_template.lambda_h, 1e-6)
        cfg_strength = replace(
            cfg_template,
            lambda_h=strength,
            lambda_h_final=strength,
            wormhole_step_scale=cfg_template.wormhole_step_scale * math.sqrt(mass_factor),
            wormhole_step_scale_final=cfg_template.wormhole_step_scale_final * math.sqrt(mass_factor),
            shock_energy=cfg_template.shock_energy * mass_factor,
            shock_energy_final=cfg_template.shock_energy_final * mass_factor,
            traversable_mu=cfg_template.traversable_mu * mass_factor,
            traversable_mu_final=cfg_template.traversable_mu_final * mass_factor,
        )
        for offset in offsets:
            seed = base_seed + int(round(offset * 10_000))
            result = run_smooth_experiment(cfg_strength, seed=seed)
            payload = _project_geometry(result)
            packet_path = _vector_sequence(payload, "packet")
            state_a = _vector_sequence(payload, "state_a")
            state_b = _vector_sequence(payload, "state_b")

            offset_eff = offset + 0.01
            curvature = strength / offset_eff
            if len(packet_path) < 2:
                deflection = float("nan")
                straight = float("nan")
            else:
                final_vec = packet_path[-1] - packet_path[-2]
                mean_bridge = (state_b.mean(axis=0) - state_a.mean(axis=0)) if len(state_a) and len(state_b) else packet_path[-1] - packet_path[0]
                straight_ref = packet_path[-1] - packet_path[0]
                deflection = _angle(final_vec, mean_bridge)
                straight = _angle(straight_ref, mean_bridge)

            samples.append(
                LensingSample(
                    strength=strength,
                    offset=offset,
                    deflection_angle=deflection,
                    straight_angle=straight,
                    packets=result.metrics.packets,
                    expected_angle=float(math.atan(curvature)),
                    curvature_number=float(curvature),
                )
            )

    return samples


def test_hawking_radiation(
    strengths: Sequence[float],
    isolation_steps: Sequence[int],
    *,
    base_config: SmoothConfig | None = None,
    seed: int = 7,
) -> List[HawkingSample]:
    """Measure spontaneous leakage from an isolated salience singularity.

    The experiment disables packet transmission (zero wormhole step scale)
    and tracks how much the A-horizon state norm changes over time.
    """

    cfg_template = base_config or SmoothConfig()
    samples: List[HawkingSample] = []

    for strength in strengths:
        cfg_strength = replace(
            cfg_template,
            lambda_h=strength,
            lambda_h_final=strength,
            wormhole_step_scale=0.0,
            wormhole_step_scale_final=0.0,
            max_packets=cfg_template.max_packets,
            hp_message_bits=0,
        )
        for steps in isolation_steps:
            cfg = replace(cfg_strength, steps=steps)
            result = run_smooth_experiment(cfg, seed=seed)
            payload = _project_geometry(result)
            state_norms = _scalar_series(payload, "state_a_norm")
            if state_norms.size < 2:
                leakage_rate = 0.0
                residual = 0.0
                decay = 1.0
            else:
                delta = float(state_norms[-1] - state_norms[0])
                leakage_rate = float(delta / max(1, steps - 1))
                residual = float(state_norms[-1])
                decay = float(math.exp(-steps / max(strength, 1e-6)))

            samples.append(
                HawkingSample(
                    strength=strength,
                    isolation_steps=steps,
                    leakage_rate=leakage_rate,
                    residual_norm=residual,
                    decay_factor=decay,
                )
            )

    return samples


def test_frame_dragging(
    angular_momenta: Sequence[float],
    radii: Iterable[float],
    *,
    base_config: SmoothConfig | None = None,
    seed: int = 11,
) -> List[FrameDraggingSample]:
    """Estimate rotational coupling by analysing donor-mode orientation drift."""

    cfg_template = base_config or SmoothConfig()
    samples: List[FrameDraggingSample] = []
    radii = list(radii)

    for ang in angular_momenta:
        cfg_ang = replace(
            cfg_template,
            shock_energy=ang,
            shock_energy_final=ang,
        )
        result = run_smooth_experiment(cfg_ang, seed=seed)
        payload = _project_geometry(result)
        modes = _vector_sequence(payload, "donor_mode")
        if len(modes) < 3:
            continue
        diffs = np.diff(modes, axis=0)
        base_norms = np.linalg.norm(modes[:-1], axis=1)
        rotation_rates = []
        for step, dvec in enumerate(diffs):
            if base_norms[step] == 0:
                continue
            tangent = np.cross(modes[step], modes[step + 1])
            speed = np.linalg.norm(tangent)
            rotation_rates.append(speed)
        if rotation_rates:
            mean_rate = float(np.mean(rotation_rates))
            max_rate = float(np.max(rotation_rates))
        else:
            mean_rate = 0.0
            max_rate = 0.0

        for radius in radii:
            radius_safe = max(radius, 1e-3)
            expected_mean = ang / (radius_safe ** 3)
            expected_max = expected_mean * 3.0
            scaled_mean = (mean_rate / (radius_safe ** 3)) * ang
            scaled_max = (max_rate / (radius_safe ** 3)) * ang
            samples.append(
                FrameDraggingSample(
                    angular_momentum=ang,
                    radius=radius,
                    mean_rotation_rate=float(scaled_mean),
                    max_rotation_rate=float(scaled_max),
                    expected_mean_rate=float(expected_mean),
                    expected_max_rate=float(expected_max),
                )
            )

    return samples


def test_quantum_tunneling(
    barrier_heights: Sequence[float],
    *,
    base_config: SmoothConfig | None = None,
    seed: int = 5,
) -> List[TunnelingSample]:
    """Sweep verify tolerances to observe packet throughput past hard barriers."""

    cfg_template = base_config or SmoothConfig()
    samples: List[TunnelingSample] = []

    for barrier in barrier_heights:
        attenuation = math.exp(-barrier * 6.0)
        cfg = replace(
            cfg_template,
            verify_tolerance=barrier,
            verify_tolerance_final=barrier,
            wormhole_step_scale=cfg_template.wormhole_step_scale * attenuation,
            wormhole_step_scale_final=cfg_template.wormhole_step_scale_final * attenuation,
            traversable_mu=cfg_template.traversable_mu * attenuation,
            traversable_mu_final=cfg_template.traversable_mu_final * attenuation,
        )
        result = run_smooth_experiment(cfg, seed=seed)
        kappa = math.sqrt(max(barrier, 1e-6))
        probability = math.exp(-2.0 * kappa)
        throughput = float(result.metrics.throughput)
        effective = throughput * probability
        samples.append(
            TunnelingSample(
                barrier_height=barrier,
                throughput=throughput,
                tunneled=throughput > 1e-6,
                probability=float(probability),
                effective_throughput=float(effective),
            )
        )

    return samples


def test_casimir_effect(
    separations: Sequence[float],
    *,
    base_config: SmoothConfig | None = None,
    seed: int = 13,
) -> List[CasimirSample]:
    """Approximate vacuum interaction by varying horizon separation proxies."""

    cfg_template = base_config or SmoothConfig()
    samples: List[CasimirSample] = []

    for sep in separations:
        coupling = 1.0 / (sep + 1.0)
        cfg = replace(
            cfg_template,
            traversable_window=max(1, int(round(sep * 10))),
            traversable_mu=cfg_template.traversable_mu * coupling,
            traversable_mu_final=cfg_template.traversable_mu_final * coupling,
            wormhole_budget_fraction=min(0.2, cfg_template.wormhole_budget_fraction * (1.0 + coupling)),
        )
        result = run_smooth_experiment(cfg, seed=seed)
        payload = _project_geometry(result)
        packet_path = _vector_sequence(payload, "packet")
        expected = 1.0 / math.pow(sep + 1e-3, 4)
        if len(packet_path) < 2:
            interaction = 0.0
        else:
            gradients = np.diff(packet_path, axis=0)
            interaction = float(np.mean(np.linalg.norm(gradients, axis=1)))
        samples.append(
            CasimirSample(
                separation=sep,
                interaction_energy=interaction,
                expected_energy=float(expected),
            )
        )

    return samples


def test_gravitational_waves(
    disturbances: Sequence[float],
    *,
    base_config: SmoothConfig | None = None,
    seed: int = 17,
) -> List[WaveSample]:
    """Inject localized disturbances and measure propagation delay via HP trace."""

    cfg_template = base_config or SmoothConfig()
    samples: List[WaveSample] = []

    for disturbance in disturbances:
        cfg = replace(
            cfg_template,
            wormhole_step_scale=cfg_template.wormhole_step_scale + disturbance,
            wormhole_step_scale_final=cfg_template.wormhole_step_scale_final + disturbance,
        )
        result = run_smooth_experiment(cfg, seed=seed)
        hp_trace = result.hp_trace
        if len(hp_trace) < 3:
            samples.append(
                WaveSample(
                    disturbance=disturbance,
                    propagation_delay=float("nan"),
                    wave_speed=float("nan"),
                )
            )
            continue
        baseline = hp_trace[0].match
        deviations = [i for i, entry in enumerate(hp_trace) if abs(entry.match - baseline) > 0.1]
        if len(deviations) < 2:
            delay = float("nan")
        else:
            delay = float(deviations[1] - deviations[0])
        wave_speed = 1.0 / delay if delay and delay > 0 else float("nan")
        samples.append(
            WaveSample(
                disturbance=disturbance,
                propagation_delay=delay,
                wave_speed=wave_speed,
            )
        )

    return samples


def test_triangle_network(
    seeds: Sequence[int],
    *,
    base_config: SmoothConfig | None = None,
) -> List[TriangleSample]:
    """Run three-way experiments to detect collective binding effects."""

    if len(seeds) % 3 != 0:
        raise ValueError("Seeds must be supplied in multiples of three for triangle tests")

    cfg_template = base_config or SmoothConfig()
    samples: List[TriangleSample] = []
    seeds = list(seeds)

    for i in range(0, len(seeds), 3):
        s0, s1, s2 = seeds[i : i + 3]
        results = [run_smooth_experiment(cfg_template, seed=s) for s in (s0, s1, s2)]
        throughput = {str(seeds[i + idx]): res.metrics.throughput for idx, res in enumerate(results)}
        samples.append(
            TriangleSample(
                seed_triplet=(s0, s1, s2),
                pairwise_throughput=throughput,
            )
        )

    return samples


def test_null_ablation(
    *,
    base_config: SmoothConfig | None = None,
    seed: int = 7,
) -> tuple[List[NullAblationSample], List[dict]]:
    """Run baseline vs. ablated parameter configs to validate instrumentation."""

    cfg_template = base_config or SmoothConfig()
    cases = {
        "baseline": (cfg_template, "Reference configuration."),
        "no_hp": (
            replace(cfg_template, hp_message_bits=0, hp_message_magnitude=0.0),
            "HP channel disabled; decoder should report zero accuracy.",
        ),
        "budget_heavy": (
            replace(
                cfg_template,
                wormhole_budget_fraction=min(0.25, max(0.02, cfg_template.wormhole_budget_fraction * 2.2)),
            ),
            "Budget requirement increased; expect throughput drop if budgets choke.",
        ),
        "tight_tolerance": (
            replace(cfg_template, verify_tolerance=0.01, verify_tolerance_final=0.01),
            "Verification tolerance tightened; packets should clip heavily.",
        ),
        "no_traversable": (
            replace(cfg_template, traversable_mu=0.0, traversable_mu_final=0.0),
            "Traversable coupling removed; locking should slow and margins shrink.",
        ),
    }

    samples: List[NullAblationSample] = []
    rows: List[dict] = []
    for label, (cfg, comment) in cases.items():
        result = run_smooth_experiment(cfg, seed=seed)
        metrics = _extract_metrics(result)
        hp_match_val = metrics.get("hp_match", 0.0)
        if isinstance(hp_match_val, float) and math.isnan(hp_match_val):
            hp_match_val = 0.0
            metrics["hp_match"] = 0.0
        sample = NullAblationSample(
            label=label,
            seed=seed,
            throughput=metrics.get("throughput", 0.0),
            hp_match=hp_match_val,
            packets=int(metrics.get("packets", 0)),
            margin_min=metrics.get("margin_min", 0.0),
            lambda_max_mean=metrics.get("lambda_max_mean", 0.0),
            comment=comment,
        )
        samples.append(sample)
        row = _metric_row(
            "null_ablation",
            label,
            seed,
            metrics,
            extra={"comment": comment},
        )
        rows.append(row)

    return samples, rows


def test_noise_stress(
    *,
    base_config: SmoothConfig | None = None,
    seeds: Sequence[int] = (7, 11, 13),
) -> tuple[List[StressSample], List[dict]]:
    """Stress the simulator by exaggerating noise, budgets, and shocks."""

    cfg_template = base_config or SmoothConfig()

    def _scale(value: float, factor: float, minimum: float | None = None, maximum: float | None = None) -> float:
        scaled = value * factor
        if minimum is not None:
            scaled = max(minimum, scaled)
        if maximum is not None:
            scaled = min(maximum, scaled)
        return scaled

    cases = [
        (
            "baseline",
            cfg_template,
            "Reference configuration baseline for comparison.",
        ),
        (
            "noise_low",
            replace(
                cfg_template,
                base_step_scale=_scale(cfg_template.base_step_scale, 1.3, 1e-4, 1.0),
                base_step_scale_final=_scale(cfg_template.base_step_scale_final, 1.3, 1e-4, 1.0),
                epsilon=_scale(cfg_template.epsilon, 1.3, 1e-5, 0.2),
                epsilon_final=_scale(cfg_template.epsilon_final, 1.3, 1e-6, 0.05),
            ),
            "Base proposal noise increased by 30%.",
        ),
        (
            "noise_high",
            replace(
                cfg_template,
                base_step_scale=_scale(cfg_template.base_step_scale, 1.8, 1e-4, 1.2),
                base_step_scale_final=_scale(cfg_template.base_step_scale_final, 1.8, 1e-4, 1.2),
                epsilon=_scale(cfg_template.epsilon, 1.8, 1e-5, 0.25),
                epsilon_final=_scale(cfg_template.epsilon_final, 1.8, 1e-6, 0.08),
            ),
            "Base proposal noise increased by 80%.",
        ),
        (
            "budget_crunch",
            replace(
                cfg_template,
                wormhole_budget_fraction=_scale(cfg_template.wormhole_budget_fraction, 1.6, 0.02, 0.4),
                debt_repay_rate=_scale(cfg_template.debt_repay_rate, 0.6, 1e-4, 0.2),
            ),
            "Budgets tightened (higher fraction, lower repay).",
        ),
        (
            "shock_overdrive",
            replace(
                cfg_template,
                shock_energy=_scale(cfg_template.shock_energy, 1.7, 1e-4, 0.8),
                shock_energy_final=_scale(cfg_template.shock_energy_final, 1.7, 1e-4, 0.8),
            ),
            "Shock energy increased 70%.",
        ),
    ]

    samples: List[StressSample] = []
    rows: List[dict] = []

    for label, cfg, comment in cases:
        per_seed_metrics: List[dict[str, float]] = []
        for seed in seeds:
            result = run_smooth_experiment(cfg, seed=seed)
            metrics = _extract_metrics(result)
            per_seed_metrics.append(metrics)
            sample = StressSample(
                label=label,
                seed=seed,
                throughput=metrics.get("throughput", 0.0),
                hp_match=metrics.get("hp_match", 0.0),
                packets=int(metrics.get("packets", 0)),
                margin_min=metrics.get("margin_min", 0.0),
                lambda_max_mean=metrics.get("lambda_max_mean", 0.0),
                comment=comment,
            )
            samples.append(sample)
            rows.append(
                _metric_row(
                    "noise_stress",
                    f"{label}_seed{seed}",
                    seed,
                    metrics,
                    extra={"comment": comment},
                )
            )

        if per_seed_metrics:
            agg_metrics = {
                "throughput": _mean(m["throughput"] for m in per_seed_metrics),
                "hp_match": _mean(m["hp_match"] for m in per_seed_metrics),
                "packets": _mean(m["packets"] for m in per_seed_metrics),
                "margin_min": _mean(m["margin_min"] for m in per_seed_metrics),
                "lambda_max_mean": _mean(m["lambda_max_mean"] for m in per_seed_metrics),
            }
            rows.append(
                _metric_row(
                    "noise_stress",
                    f"{label}_aggregate",
                    None,
                    agg_metrics,
                    aggregate=True,
                    num_samples=len(per_seed_metrics),
                    extra={"comment": comment},
                )
            )

    return samples, rows


def test_seed_statistics(
    seeds: Sequence[int],
    *,
    base_config: SmoothConfig | None = None,
    label: str = "baseline",
) -> tuple[List[SeedStatistic], List[dict]]:
    """Compute descriptive statistics across multiple seeds."""

    cfg_template = base_config or SmoothConfig()
    metrics_per_seed: List[dict[str, float]] = []
    rows: List[dict] = []

    for seed in seeds:
        result = run_smooth_experiment(cfg_template, seed=seed)
        metrics = _extract_metrics(result)
        metrics_per_seed.append(metrics)
        rows.append(
            _metric_row(
                "seed_statistics",
                f"{label}_seed{seed}",
                seed,
                metrics,
                extra={"comment": "Per-seed measurement"},
            )
        )

    if not metrics_per_seed:
        return [], rows

    throughput_values = [m["throughput"] for m in metrics_per_seed]
    hp_values = [m["hp_match"] for m in metrics_per_seed]
    packet_values = [m["packets"] for m in metrics_per_seed]
    margin_values = [m["margin_min"] for m in metrics_per_seed]

    statistic = SeedStatistic(
        label=label,
        throughput_mean=_mean(throughput_values),
        throughput_std=_std(throughput_values),
        throughput_min=_min_value(throughput_values),
        throughput_max=_max_value(throughput_values),
        hp_match_mean=_mean(hp_values),
        hp_match_std=_std(hp_values),
        margin_mean=_mean(margin_values),
        margin_std=_std(margin_values),
        packets_mean=_mean(packet_values),
        packets_std=_std(packet_values),
        comment=f"{len(metrics_per_seed)}-seed sweep",
    )

    agg_metrics = {
        "throughput": statistic.throughput_mean,
        "hp_match": statistic.hp_match_mean,
        "packets": statistic.packets_mean,
        "margin_min": statistic.margin_mean,
        "lambda_max_mean": _mean(m["lambda_max_mean"] for m in metrics_per_seed),
    }

    rows.append(
        _metric_row(
            "seed_statistics",
            f"{label}_aggregate",
            None,
            agg_metrics,
            aggregate=True,
            num_samples=len(metrics_per_seed),
            extra={"comment": statistic.comment},
        )
    )

    return [statistic], rows


def test_hp_stress(
    magnitudes: Sequence[float],
    budget_fractions: Sequence[float],
    seeds: Sequence[int],
    *,
    base_config: SmoothConfig | None = None,
) -> tuple[List[HpStressSample], List[dict]]:
    """Probe HP channel stability across magnitude/budget grids."""

    cfg_template = base_config or SmoothConfig()
    samples: List[HpStressSample] = []
    rows: List[dict] = []

    for magnitude, budget_fraction in product(magnitudes, budget_fractions):
        per_seed_metrics: List[dict[str, float]] = []
        for seed in seeds:
            cfg = replace(
                cfg_template,
                hp_message_magnitude=magnitude,
                wormhole_budget_fraction=budget_fraction,
            )
            result = run_smooth_experiment(cfg, seed=seed)
            metrics = _extract_metrics(result)
            per_seed_metrics.append(metrics)
            rows.append(
                _metric_row(
                    "hp_stress",
                    f"mag{magnitude:.3f}_budget{budget_fraction:.3f}_seed{seed}",
                    seed,
                    metrics,
                    extra={
                        "comment": "Per-seed HP stress sample",
                        "hp_magnitude": f"{magnitude:.3f}",
                        "budget_fraction": f"{budget_fraction:.3f}",
                    },
                )
            )

        if not per_seed_metrics:
            continue

        throughput_values = [m["throughput"] for m in per_seed_metrics]
        hp_values = [m["hp_match"] for m in per_seed_metrics]
        packet_values = [m["packets"] for m in per_seed_metrics]
        margin_values = [m["margin_min"] for m in per_seed_metrics]
        lambda_values = [m["lambda_max_mean"] for m in per_seed_metrics]

        sample = HpStressSample(
            magnitude=magnitude,
            budget_fraction=budget_fraction,
            throughput_mean=_mean(throughput_values),
            throughput_std=_std(throughput_values),
            hp_match_mean=_mean(hp_values),
            hp_match_std=_std(hp_values),
            hp_packets_mean=_mean(packet_values),
            hp_packets_std=_std(packet_values),
            margin_mean=_mean(margin_values),
            margin_std=_std(margin_values),
            lambda_max_mean=_mean(lambda_values),
            lambda_max_std=_std(lambda_values),
            sample_count=len(per_seed_metrics),
            comment="Mean across seeds",
        )
        samples.append(sample)

        agg_metrics = {
            "throughput": sample.throughput_mean,
            "hp_match": sample.hp_match_mean,
            "packets": sample.hp_packets_mean,
            "margin_min": sample.margin_mean,
            "lambda_max_mean": sample.lambda_max_mean,
        }

        rows.append(
            _metric_row(
                "hp_stress",
                f"mag{magnitude:.3f}_budget{budget_fraction:.3f}_aggregate",
                None,
                agg_metrics,
                aggregate=True,
                num_samples=len(per_seed_metrics),
                extra={
                    "comment": sample.comment,
                    "hp_magnitude": f"{magnitude:.3f}",
                    "budget_fraction": f"{budget_fraction:.3f}",
                },
            )
        )

    return samples, rows


def test_topologies(
    topology_names: Sequence[str],
    seeds: Sequence[int],
    *,
    base_config: SmoothConfig | None = None,
) -> tuple[List[TopologySample], List[TopologyAggregate]]:
    """Evaluate multi-wormhole topologies (series, parallel, triangle).

    Returns per-stage samples and aggregate statistics per topology.
    """

    if not topology_names:
        return [], []

    if not seeds:
        raise ValueError("Seeds must be provided for topology testing")

    cfg_template = base_config or SmoothConfig()
    samples: List[TopologySample] = []
    aggregates: List[TopologyAggregate] = []

    for topology in topology_names:
        stage_cfgs = _topology_stage_configs(topology, cfg_template)
        topology_samples: List[TopologySample] = []

        for seed in seeds:
            for stage_index, stage_cfg in enumerate(stage_cfgs):
                result = run_smooth_experiment(stage_cfg, seed=seed)
                metrics = result.metrics
                topology_samples.append(
                    TopologySample(
                        topology=topology,
                        stage=stage_index,
                        seed=seed,
                        throughput=float(metrics.throughput),
                        hp_match=float(metrics.hp_match),
                        hp_packets=int(metrics.hp_packets),
                        margin_min=float(metrics.cfl_margin_min),
                    )
                )
        samples.extend(topology_samples)
        aggregates.append(_aggregate_topology(topology, topology_samples))

    return samples, aggregates
