from __future__ import annotations

import csv
import json
import statistics
import sys
from pathlib import Path

import numpy as np

ROOT = Path(__file__).resolve().parents[1]
if str(ROOT) not in sys.path:
    sys.path.insert(0, str(ROOT))

from wormhole_proof.analysis.relativistic import (
    HpStressSample,
    NullAblationSample,
    SeedStatistic,
    StressSample,
    dataclass_list_to_dict,
    test_casimir_effect,
    test_frame_dragging,
    test_gravitational_waves,
    test_hawking_radiation,
    test_hp_stress,
    test_lensing,
    test_noise_stress,
    test_null_ablation,
    test_quantum_tunneling,
    test_seed_statistics,
    test_topologies,
    test_triangle_network,
)
from wormhole_proof.core.wormhole_smooth import SmoothConfig


def _write_metrics_csv(path: Path, rows: list[dict]) -> None:
    if not rows:
        return
    fieldnames = sorted(set().union(*(row.keys() for row in rows)))
    with path.open("w", newline="", encoding="utf-8") as handle:
        writer = csv.DictWriter(handle, fieldnames=fieldnames)
        writer.writeheader()
        for row in rows:
            writer.writerow(row)


def _topology_rows(aggregates) -> list[dict]:
    rows: list[dict] = []
    for agg in aggregates:
        rows.append(
            {
                "experiment": "topology_aggregate",
                "label": agg.topology,
                "seed": -1,
                "aggregate": True,
                "num_samples": 3,
                "throughput": agg.throughput_mean,
                "throughput_std": agg.throughput_std,
                "hp_match": agg.hp_match_mean,
                "hp_packets": agg.hp_packets_mean,
                "margin_min": agg.margin_mean,
                "comment": agg.notes,
            }
        )
    return rows


def _hp_analysis(samples: list[HpStressSample]) -> dict:
    if not samples:
        return {}
    magnitudes = np.array([sample.magnitude for sample in samples], dtype=float)
    throughput = np.array([sample.throughput_mean for sample in samples], dtype=float)
    hp_match = np.array([sample.hp_match_mean for sample in samples], dtype=float)
    covariance = float(np.cov(magnitudes, throughput)[0, 1]) if magnitudes.size > 1 else 0.0
    correlation = (
        float(statistics.correlation(magnitudes, throughput))
        if magnitudes.size > 1 and not np.allclose(throughput, throughput[0])
        else 0.0
    )
    hp_corr = (
        float(statistics.correlation(magnitudes, hp_match))
        if magnitudes.size > 1 and not np.allclose(hp_match, hp_match[0])
        else 0.0
    )
    slope = float(np.polyfit(magnitudes, throughput, deg=1)[0]) if magnitudes.size > 1 else 0.0
    return {
        "hp_throughput_covariance": covariance,
        "hp_throughput_correlation": correlation,
        "hp_match_correlation": hp_corr,
        "hp_throughput_slope": slope,
    }


def main() -> None:
    output_dir = ROOT / "artifacts" / "relativistic_suite"
    output_dir.mkdir(parents=True, exist_ok=True)

    base_cfg = SmoothConfig(
        hp_message_magnitude=0.26,
        wormhole_budget_fraction=0.10,
        debt_repay_rate=0.03,
    )

    seeds = [7, 11, 13, 17, 19]

    lensing = test_lensing(strengths=[160.0, 200.0, 240.0], offsets=[0.0, 0.01, 0.02], base_config=base_cfg)
    hawking = test_hawking_radiation(strengths=[160.0, 220.0], isolation_steps=[200, 400], base_config=base_cfg)
    frame_drag = test_frame_dragging(angular_momenta=[0.02, 0.05, 0.08], radii=[1.0, 2.0], base_config=base_cfg)
    tunneling = test_quantum_tunneling(barrier_heights=[0.05, 0.08, 0.10, 0.12], base_config=base_cfg)
    casimir = test_casimir_effect(separations=[0.5, 1.0, 1.5], base_config=base_cfg)
    waves = test_gravitational_waves(disturbances=[0.0, 0.02, 0.05], base_config=base_cfg)
    triangle = test_triangle_network(seeds=seeds[:3], base_config=base_cfg)
    topology_samples, topology_aggregates = test_topologies(
        ["series", "parallel", "triangle"],
        seeds=seeds[:3],
        base_config=base_cfg,
    )
    null_samples, null_rows = test_null_ablation(base_config=base_cfg, seed=seeds[0])
    stress_samples, stress_rows = test_noise_stress(base_config=base_cfg, seeds=seeds[:3])
    seed_stats, seed_rows = test_seed_statistics(seeds=seeds, base_config=base_cfg)
    hp_samples, hp_rows = test_hp_stress(
        magnitudes=[0.18, 0.22, 0.26, 0.30],
        budget_fractions=[0.08, 0.10, 0.12],
        seeds=seeds[:3],
        base_config=base_cfg,
    )

    metrics_rows: list[dict] = []
    metrics_rows.extend(null_rows)
    metrics_rows.extend(stress_rows)
    metrics_rows.extend(seed_rows)
    metrics_rows.extend(hp_rows)
    metrics_rows.extend(_topology_rows(topology_aggregates))

    hp_analysis = _hp_analysis(hp_samples)

    payload = {
        "config": base_cfg.__dict__,
        "experiments": {
            "lensing": dataclass_list_to_dict(lensing),
            "hawking": dataclass_list_to_dict(hawking),
            "frame_dragging": dataclass_list_to_dict(frame_drag),
            "quantum_tunneling": dataclass_list_to_dict(tunneling),
            "casimir": dataclass_list_to_dict(casimir),
            "waves": dataclass_list_to_dict(waves),
            "triangle": dataclass_list_to_dict(triangle),
            "topology_samples": dataclass_list_to_dict(topology_samples),
            "topology_aggregates": dataclass_list_to_dict(topology_aggregates),
            "null_ablation": dataclass_list_to_dict(null_samples),
            "noise_stress": dataclass_list_to_dict(stress_samples),
            "seed_statistics": dataclass_list_to_dict(seed_stats),
            "hp_stress": dataclass_list_to_dict(hp_samples),
        },
        "analysis": {
            "hp_channel": hp_analysis,
        },
    }

    output_path = output_dir / "relativistic_suite.json"
    output_path.write_text(json.dumps(payload, indent=2), encoding="utf-8")

    metrics_path = output_dir / "relativistic_metrics.csv"
    _write_metrics_csv(metrics_path, metrics_rows)

    print(f"Experiment bundle written to {output_path}")
    print(f"Metrics table written to {metrics_path}")


if __name__ == "__main__":
    main()
