"""Smooth SAL wormhole simulator with SAL-CFL instrumentation."""

from __future__ import annotations

from collections import deque
from dataclasses import dataclass, field
from typing import Deque, Iterable, List, Tuple

import numpy as np


@dataclass
class SmoothConfig:
    """Configuration for the smooth wormhole simulator."""

    # Initial (toy) regime parameters
    dim: int = 64
    epsilon: float = 0.02
    base_step_scale: float = 0.06
    lambda_h: float = 180.0
    wormhole_step_scale: float = 0.25
    verify_tolerance: float = 0.08
    traversable_mu: float = 0.03
    shock_energy: float = 0.04

    # Target regime after warmup
    epsilon_final: float = 1e-3
    base_step_scale_final: float = 0.04
    lambda_h_final: float = 260.0
    wormhole_step_scale_final: float = 0.18
    verify_tolerance_final: float = 0.045
    traversable_mu_final: float = 0.05
    shock_energy_final: float = 0.06

    warmup_steps: int = 200
    steps: int = 500
    horizon_steps: int = 40
    soft_rank: int = 8
    wormhole_budget_fraction: float = 0.12
    lambda_c: float = 12.0
    debt_repay_rate: float = 0.02
    max_packets: int = 72
    traversable_window: int = 4

    hp_message_bits: int = 3
    hp_message_magnitude: float = 0.18
    hp_scramble_wait: int = 5
    hp_decode_packets: int = 6

    budget_total: float = 1.0


@dataclass
class SmoothMetrics:
    throughput: float
    packets: int
    lambda_max_mean: float
    cfl_margin_min: float
    packet_energy_mean: float
    hp_match: float
    hp_packets: int
    relaxation_steps: float


@dataclass
class SmoothResult:
    seed: int
    config: SmoothConfig
    metrics: SmoothMetrics
    geometry: "SmoothGeometry | None"
    hp_trace: List["SmoothTraceEntry"]


@dataclass
class SmoothGeometry:
    state_a: List[List[float]] = field(default_factory=list)
    state_b: List[List[float]] = field(default_factory=list)
    packet: List[List[float]] = field(default_factory=list)
    donor_mode: List[List[float]] = field(default_factory=list)
    receiver_mode: List[List[float]] = field(default_factory=list)
    hp_decode: List[List[float]] = field(default_factory=list)
    hp_target: List[float] = field(default_factory=list)
    hp_decode_bits: List[List[float]] = field(default_factory=list)
    state_a_norm: List[float] = field(default_factory=list)
    state_b_norm: List[float] = field(default_factory=list)


@dataclass
class SmoothTraceEntry:
    packet: int
    match: float


GEOMETRY_HISTORY_LIMIT = 256


def _append_history(buffer: List[List[float]], vector: np.ndarray, limit: int = GEOMETRY_HISTORY_LIMIT) -> None:
    buffer.append(vector.astype(float).tolist())
    if len(buffer) > limit:
        del buffer[0]


def _append_scalar_history(buffer: List[float], value: float, limit: int = GEOMETRY_HISTORY_LIMIT) -> None:
    buffer.append(float(value))
    if len(buffer) > limit:
        del buffer[0]


# ---------------------------------------------------------------------------
# Helpers


def _stack(history: Iterable[np.ndarray], dim: int) -> np.ndarray:
    rows = list(history)
    if not rows:
        return np.zeros((1, dim), dtype=float)
    return np.vstack(rows)


def _soft_modes(
    donor_hist: Iterable[np.ndarray],
    receiver_hist: Iterable[np.ndarray],
    dim: int,
    rank: int,
    rng: np.random.Generator,
) -> Tuple[np.ndarray, np.ndarray]:
    donor_stack = _stack(donor_hist, dim)
    receiver_stack = _stack(receiver_hist, dim)

    cov_d = donor_stack.T @ donor_stack + 1e-9 * np.eye(dim)
    cov_r = receiver_stack.T @ receiver_stack + 1e-9 * np.eye(dim)

    eigvals_d, eigvecs_d = np.linalg.eigh(cov_d)
    eigvals_r, eigvecs_r = np.linalg.eigh(cov_r)
    order_d = np.argsort(eigvals_d)[::-1]
    order_r = np.argsort(eigvals_r)[::-1]
    donor_basis = eigvecs_d[:, order_d[:rank]]
    receiver_basis = eigvecs_r[:, order_r[:rank]]

    if donor_basis.size == 0:
        donor_basis = rng.standard_normal((dim, rank))
        receiver_basis = rng.standard_normal((dim, rank))
        donor_basis, _ = np.linalg.qr(donor_basis)
        receiver_basis, _ = np.linalg.qr(receiver_basis)

    return donor_basis, receiver_basis


def _spectral_radius(matrix: np.ndarray, iters: int = 3) -> float:
    if matrix.size == 0:
        return 0.0
    rng = np.random.default_rng()
    v = rng.normal(size=matrix.shape[1])
    v /= np.linalg.norm(v) + 1e-12
    for _ in range(iters):
        v = matrix @ (matrix.T @ v)
        norm = np.linalg.norm(v) + 1e-12
        v /= norm
    return float(np.linalg.norm(matrix @ v))


def _blend(start: float, end: float, weight: float) -> float:
    weight = np.clip(weight, 0.0, 1.0)
    return start + (end - start) * weight


def _update_state(
    state: np.ndarray,
    proposal: np.ndarray,
    lambda_h: float,
    epsilon: float,
) -> Tuple[np.ndarray, float]:
    eta = 1.0 / (1.0 + lambda_h)
    update = eta * proposal
    norm = np.linalg.norm(update)
    if norm < epsilon:
        update = np.zeros_like(update)
        norm = 0.0
    return state + update, float(norm)


# ---------------------------------------------------------------------------
# Main simulator


def run_smooth_experiment(cfg: SmoothConfig, seed: int) -> SmoothResult:
    rng = np.random.default_rng(seed)

    state_a = np.zeros(cfg.dim, dtype=float)
    state_b = np.zeros(cfg.dim, dtype=float)
    budgets = np.full(2, cfg.budget_total / 2.0, dtype=float)
    debt_balance = 0.0

    history_a: Deque[np.ndarray] = deque(maxlen=cfg.horizon_steps)
    history_b: Deque[np.ndarray] = deque(maxlen=cfg.horizon_steps)
    horizon = np.zeros(2, dtype=int)

    throughput = 0.0
    packets = 0
    first_lock: float | None = None

    hp_matches: List[float] = []
    geometry = SmoothGeometry()
    hp_trace: List[SmoothTraceEntry] = []
    hp_bit_history: List[List[float]] = []

    if cfg.hp_message_bits > 0:
        hp_decode = np.zeros(cfg.hp_message_bits)
        hp_target = rng.choice([-1.0, 1.0], size=cfg.hp_message_bits)
        geometry.hp_target = hp_target[:].tolist()
        hp_threshold = np.zeros(cfg.hp_message_bits)
        hp_gain = np.ones(cfg.hp_message_bits)
    else:
        hp_decode = np.array([])
        hp_target = np.array([])
        hp_threshold = np.array([])
        hp_gain = np.array([])
    hp_scramble = 0
    hp_pending = False

    lambda_max_values: List[float] = []
    cfl_margins: List[float] = []
    packet_energies: List[float] = []

    traversable_window = cfg.traversable_window

    traversable_credit = 0.0

    for step in range(cfg.steps):
        weight = (step / cfg.warmup_steps) if cfg.warmup_steps > 0 else 1.0
        lam_h = _blend(cfg.lambda_h, cfg.lambda_h_final, weight)
        epsilon = _blend(cfg.epsilon, cfg.epsilon_final, weight)
        base_scale = _blend(cfg.base_step_scale, cfg.base_step_scale_final, weight)
        wormhole_scale = _blend(cfg.wormhole_step_scale, cfg.wormhole_step_scale_final, weight)
        verify_tol = _blend(cfg.verify_tolerance, cfg.verify_tolerance_final, weight)
        traversable_mu = _blend(cfg.traversable_mu, cfg.traversable_mu_final, weight)
        shock_energy = _blend(cfg.shock_energy, cfg.shock_energy_final, weight)

        proposal_a = rng.normal(scale=base_scale, size=cfg.dim)
        proposal_b = rng.normal(scale=base_scale, size=cfg.dim)

        state_a, delta_a = _update_state(state_a, proposal_a, lam_h, epsilon)
        state_b, delta_b = _update_state(state_b, proposal_b, lam_h, epsilon)

        _append_scalar_history(geometry.state_a_norm, np.linalg.norm(state_a))
        _append_scalar_history(geometry.state_b_norm, np.linalg.norm(state_b))

        history_a.append(proposal_a)
        history_b.append(proposal_b)

        horizon[0] = horizon[0] + 1 if delta_a < epsilon else 0
        horizon[1] = horizon[1] + 1 if delta_b < epsilon else 0

        if debt_balance > 0:
            repayment = min(cfg.debt_repay_rate, debt_balance)
            debt_balance -= repayment
            budgets += repayment / 2.0

        if horizon.min() < cfg.horizon_steps:
            continue

        if first_lock is None:
            first_lock = float(step)
            if cfg.hp_message_bits > 0:
                hp_vec = np.zeros(cfg.dim)
                idx = rng.choice(cfg.dim, size=cfg.hp_message_bits, replace=False)
                hp_vec[idx] = hp_target
                state_a += hp_vec * cfg.hp_message_magnitude
                history_a.append(hp_vec * cfg.hp_message_magnitude)
                hp_scramble = cfg.hp_scramble_wait

        if hp_scramble > 0:
            hp_scramble -= 1
            continue

        donor_basis, receiver_basis = _soft_modes(history_a, history_b, cfg.dim, cfg.soft_rank, rng)
        spectral_radius = _spectral_radius(receiver_basis.T @ donor_basis)
        lambda_max_values.append(spectral_radius)

        if traversable_mu > 0:
            cross = float(np.dot(state_a, state_b))
            if cross > 0:
                credit = traversable_mu * cross
                traversable_credit = min(
                    traversable_credit + credit, cfg.budget_total
                )
                budgets += credit / 2.0
                debt_balance -= credit

        direction = donor_basis[:, 0]
        if np.linalg.norm(direction) > 0:
            direction /= np.linalg.norm(direction)
        packet = (receiver_basis @ (donor_basis.T @ direction)) * wormhole_scale
        if shock_energy > 0:
            packet *= shock_energy

        packet_norm = float(np.linalg.norm(packet))
        if packet_norm == 0:
            continue

        if packet_norm > verify_tol:
            packet *= verify_tol / packet_norm
            packet_norm = verify_tol

        cfl_margin = 1.0 - wormhole_scale * spectral_radius * (1.0 + cfg.lambda_c * epsilon)
        cfl_margins.append(cfl_margin)

        energy = cfg.lambda_c * packet_norm ** 2
        packet_energies.append(energy)

        if cfg.max_packets > 0 and packets >= cfg.max_packets:
            break

        required_budget = cfg.wormhole_budget_fraction * cfg.budget_total
        if budgets.sum() < required_budget:
            continue
        budgets -= required_budget / 2.0

        shadow = state_b + packet
        delta_shadow = np.linalg.norm(shadow - state_b)
        if delta_shadow > verify_tol * 1.05:
            budgets += required_budget / 2.0
            continue

        state_b = shadow
        throughput += packet_norm
        packets += 1

        debt_balance += energy
        budgets += required_budget / 2.0
        budgets -= energy / 2.0

        if traversable_credit > 0:
            repayment = min(traversable_credit, cfg.debt_repay_rate)
            traversable_credit -= repayment
            budgets += repayment / 2.0

        _append_history(geometry.state_a, state_a)
        _append_history(geometry.state_b, state_b)
        _append_history(geometry.packet, packet)
        _append_history(geometry.donor_mode, donor_basis[:, 0])
        _append_history(geometry.receiver_mode, receiver_basis[:, 0])
        if cfg.hp_message_bits > 0:
            _append_history(geometry.hp_decode, hp_decode[: cfg.hp_message_bits])

        if cfg.hp_message_bits > 0:
            coeffs = receiver_basis.T @ packet
            if cfg.hp_message_bits > 0:
                mode_strength = np.linalg.norm(
                    receiver_basis[:, : cfg.hp_message_bits], axis=0
                )
                norm_ref = mode_strength.max() + 1e-6
                desired_gain = norm_ref / (mode_strength + 1e-6)
                hp_gain = 0.9 * hp_gain + 0.1 * desired_gain

            if cfl_margin < 0.8:
                hp_pending = True
            elif hp_pending and cfl_margin >= 0.85:
                hp_pending = False

            if not hp_pending:
                scaled = coeffs[: cfg.hp_message_bits] * hp_gain
                hp_decode += scaled
                hp_threshold = 0.95 * hp_threshold + 0.05 * np.abs(hp_decode)
                confidence = np.tanh(np.abs(hp_decode) / (hp_threshold + 1e-6))
                decision = np.where(confidence > 0.55, np.sign(hp_decode), 0.0)
                accuracy_vector = (decision == hp_target).astype(float)
                valid = decision != 0.0
                if np.any(valid):
                    hp_match = float(np.mean(accuracy_vector[valid]))
                else:
                    hp_match = 0.0
                hp_matches.append(hp_match)
                hp_trace.append(SmoothTraceEntry(packet=packets, match=hp_match))
                hp_bit_history.append(accuracy_vector.tolist())

    hp_match_final = float(hp_matches[-1]) if hp_matches else float("nan")
    relaxation = first_lock if first_lock is not None else float("nan")

    metrics = SmoothMetrics(
        throughput=throughput,
        packets=packets,
        lambda_max_mean=float(np.mean(lambda_max_values)) if lambda_max_values else 0.0,
        cfl_margin_min=float(np.min(cfl_margins)) if cfl_margins else 0.0,
        packet_energy_mean=float(np.mean(packet_energies)) if packet_energies else 0.0,
        hp_match=hp_match_final,
        hp_packets=len(hp_matches),
        relaxation_steps=relaxation,
    )

    geometry_payload = None
    has_geometry = any(
        (
            geometry.state_a,
            geometry.packet,
            geometry.state_a_norm,
            geometry.state_b_norm,
            geometry.hp_decode,
        )
    )
    if has_geometry:
        geometry_payload = geometry
        geometry_payload.hp_decode = geometry_payload.hp_decode[-GEOMETRY_HISTORY_LIMIT:]
        if hp_bit_history:
            geometry_payload.hp_decode_bits = hp_bit_history[-GEOMETRY_HISTORY_LIMIT:]

    return SmoothResult(
        seed=seed,
        config=cfg,
        metrics=metrics,
        geometry=geometry_payload,
        hp_trace=hp_trace,
    )
