"""Configuration dataclasses for the wormhole proof experiment."""

from __future__ import annotations

from dataclasses import dataclass, field
from pathlib import Path
from typing import Iterable, List


@dataclass(frozen=True)
class WormholeConfig:
    """Parameters governing the dual-horizon wormhole simulation."""

    steps: int = 400
    dim: int = 64
    horizon_steps: int = 40
    epsilon: float = 0.02
    base_step_scale: float = 0.06
    lambda_h: float = 180.0
    budget_total: float = 1.0
    wormhole_budget_fraction: float = 0.12
    wormhole_step_scale: float = 0.25
    soft_rank: int = 8
    verify_tolerance: float = 0.08
    debt_repay_rate: float = 0.02
    lambda_c: float = 12.0
    max_packets: int = 48
    escrow_factor: float = 0.18
    cfl_margin: float = 0.85
    traversable_mu: float = 0.03
    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
    hp_decode_tolerance: float = 0.1
    shock_energy: float = 0.04
    shock_gap: int = 3
    shock_probe_gap: int = 0
    warmup_steps: int = 200
    lambda_h_final: float | None = 260.0
    epsilon_final: float | None = 1e-3
    base_step_scale_final: float | None = 0.04
    wormhole_step_scale_final: float | None = 0.18
    verify_tolerance_final: float | None = 0.045
    traversable_mu_final: float | None = 0.05
    shock_energy_final: float | None = 0.06


@dataclass(frozen=True)
class SweepConfig:
    """Configuration for continuity-strain critical sweep."""

    lambda_start: float = 120.0
    lambda_end: float = 220.0
    lambda_points: int = 12
    critical_fraction: float = 0.05


@dataclass(frozen=True)
class BridgeConfig:
    """Parameters for the canonical correlation bridge benchmark."""

    dim_donor: int = 8
    dim_receiver: int = 10
    shared_rank: int = 4
    history_samples: int = 256
    noise: float = 0.05
    cca_rank: int = 4
    cca_regularization: float = 1e-4
    trials: int = 64
    step_scale: float = 0.25


@dataclass(frozen=True)
class ExperimentIO:
    """Descriptor for experiment output layout."""

    output_dir: Path
    json_filename: str = "metrics.json"
    markdown_filename: str = "summary.md"
    sweep_csv: str = "sweep.csv"
    hp_csv: str = "hp_trace.csv"

    def ensure_dirs(self) -> None:
        self.output_dir.mkdir(parents=True, exist_ok=True)


@dataclass(frozen=True)
class ExperimentConfig:
    """Top-level configuration bundling all subcomponents."""

    wormhole: WormholeConfig = field(default_factory=WormholeConfig)
    sweep: SweepConfig = field(default_factory=SweepConfig)
    bridge: BridgeConfig = field(default_factory=BridgeConfig)
    seeds: List[int] = field(default_factory=lambda: [7])
    record_hp_trace: bool = True

    def expanded_seeds(self, extra: Iterable[int] | None = None) -> List[int]:
        values = list(self.seeds)
        if extra is not None:
            values.extend(extra)
        return sorted(set(values))
