"""Artifact generation for the standalone wormhole proof experiment."""

from __future__ import annotations

import csv
import json
import math
from dataclasses import asdict
from pathlib import Path
from typing import Iterable, List

import numpy as np
from tabulate import tabulate

from .config import ExperimentConfig, ExperimentIO
from .core.benchmarks.quadratic import TwinQuadraticBenchmark
from .core.sweep import SweepResult
from .core.wormhole import RegimeResult, SimulationResult


def _sanitize(value: Any) -> Any:
    if isinstance(value, float):
        if math.isfinite(value):
            return value
        return None
    if isinstance(value, list):
        return [_sanitize(v) for v in value]
    if isinstance(value, dict):
        return {k: _sanitize(v) for k, v in value.items()}
    return value


def _regime_to_row(seed: int, regime: RegimeResult) -> dict:
    m = regime.metrics
    return {
        "seed": seed,
        "regime": regime.name,
        "throughput": regime.throughput,
        "commits": regime.commits,
        "invariant_ok": regime.invariant_ok,
        "budget_error": regime.budget_error,
        "debt_balance": regime.debt_balance,
        "defect_rate": regime.defect_rate,
        "relaxation_steps": m.relaxation_steps,
        "soft_active_rank_mean": m.soft_active_rank_mean,
        "percolation_span_mean": m.percolation_span_mean,
        "update_ratio_mean": m.update_ratio_mean,
        "arrival_delay_mean": m.arrival_delay_mean,
        "shock_delay_mean": m.shock_delay_mean,
        "hp_match_fraction": m.hp_match_fraction,
        "hp_success": m.hp_success,
        "notes": regime.notes,
    }


def _hp_trace_rows(seed: int, regime: RegimeResult) -> Iterable[dict]:
    if regime.trace is None:
        return []
    return (
        {
            "seed": seed,
            "regime": regime.name,
            "packet_index": entry.packet_index,
            "step": entry.step,
            "match_fraction": entry.match_fraction,
            "l2_error": entry.l2_error,
            "cosine": entry.cosine,
        }
        for entry in regime.trace.hp_trace
    )


def write_metrics_json(
    io: ExperimentIO,
    config: ExperimentConfig,
    simulations: List[SimulationResult],
    sweep: SweepResult,
    benchmark_metrics: dict,
) -> Path:
    payload = {
        "config": {
            "wormhole": asdict(config.wormhole),
            "sweep": asdict(config.sweep),
            "bridge": asdict(config.bridge),
            "seeds": config.seeds,
        },
        "wormhole": [
            {
                "seed": sim.seed,
                "regimes": [
                    {
                        "name": regime.name,
                        "throughput": regime.throughput,
                        "commits": regime.commits,
                        "invariant_ok": regime.invariant_ok,
                        "budget_error": regime.budget_error,
                        "debt_balance": regime.debt_balance,
                        "defect_rate": regime.defect_rate,
                        "notes": regime.notes,
                        "metrics": asdict(regime.metrics),
                    }
                    for regime in sim.regimes
                ],
            }
            for sim in simulations
        ],
        "sweep": {
            "analysis": asdict(sweep.analysis),
            "points": [asdict(point) for point in sweep.points],
        },
        "benchmark": benchmark_metrics,
    }
    path = io.output_dir / io.json_filename
    path.write_text(json.dumps(payload, indent=2), encoding="utf-8")
    return path


def write_markdown_summary(
    io: ExperimentIO,
    simulations: List[SimulationResult],
    sweep: SweepResult,
    benchmark_metrics: dict,
) -> Path:
    wormhole_rows = [_regime_to_row(sim.seed, regime) for sim in simulations for regime in sim.regimes]
    table = tabulate(
        [
            (
                row["seed"],
                row["regime"],
                f"{row['throughput']:.4f}",
                row["commits"],
                "yes" if row["invariant_ok"] else "no",
                f"{row['defect_rate']*100:.1f}%",
                f"{row['hp_match_fraction']:.3f}" if np.isfinite(row["hp_match_fraction"]) else "nan",
                row["notes"],
            )
            for row in wormhole_rows
        ],
        headers=["seed", "regime", "throughput", "commits", "invariants", "defects", "HP match", "notes"],
        tablefmt="github",
    )

    sweep_table = tabulate(
        [
            (
                f"{p.lambda_h:.3f}",
                f"{p.sigma:.4f}",
                f"{p.soft_rank_fraction:.3f}",
                f"{p.throughput:.4f}",
                f"{p.relaxation_steps:.1f}",
                "yes" if p.invariant_ok else "no",
            )
            for p in sweep.points
        ],
        headers=["lambda_h", "sigma", "soft_frac", "throughput", "relax_steps", "invariants"],
        tablefmt="github",
    )

    benchmark_table = tabulate(
        [
            ("alignment_error", f"{benchmark_metrics['alignment_error']:.4e}"),
            ("transport_mean_error", f"{benchmark_metrics['transport_mean_error']:.4e}"),
            ("loss_gain_mean", f"{benchmark_metrics['loss_gain_mean']:.4f}"),
            ("loss_gain_std", f"{benchmark_metrics['loss_gain_std']:.4f}"),
        ],
        headers=["metric", "value"],
        tablefmt="github",
    )

    lines = [
        "# Wormhole Proof Experiment Summary",
        "",
        "## Wormhole Regimes",
        table,
        "",
        "## Critical Sweep",
        sweep_table,
        "",
        "### Sweep Analysis",
        json.dumps(asdict(sweep.analysis), indent=2),
        "",
        "## Canonical Bridge Benchmark",
        benchmark_table,
    ]

    path = io.output_dir / io.markdown_filename
    path.write_text("\n".join(lines), encoding="utf-8")
    return path


def write_sweep_csv(io: ExperimentIO, sweep: SweepResult) -> Path:
    path = io.output_dir / io.sweep_csv
    with path.open("w", newline="", encoding="utf-8") as fp:
        writer = csv.DictWriter(
            fp,
            fieldnames=[
                "lambda_h",
                "sigma",
                "soft_rank_mean",
                "soft_rank_fraction",
                "throughput",
                "relaxation_steps",
                "percolation_mean",
                "percolation_max",
                "update_ratio_mean",
                "candidate_norm_mean",
                "invariant_ok",
            ],
        )
        writer.writeheader()
        for point in sweep.points:
            writer.writerow(asdict(point))
    return path


def write_hp_trace_csv(io: ExperimentIO, simulations: List[SimulationResult]) -> Path:
    path = io.output_dir / io.hp_csv
    rows = [
        row
        for sim in simulations
        for regime in sim.regimes
        for row in _hp_trace_rows(sim.seed, regime)
    ]

    with path.open("w", newline="", encoding="utf-8") as fp:
        if not rows:
            fp.write("")
            return path
        writer = csv.DictWriter(fp, fieldnames=list(rows[0].keys()))
        writer.writeheader()
        writer.writerows(rows)
    return path


def write_wireframe_payload(
    path: Path,
    simulations: List[SimulationResult],
    sweep: SweepResult,
    benchmark_metrics: dict,
) -> Path:
    def _geometry_payload(regime: RegimeResult) -> Dict[str, Any] | None:
        geom = regime.geometry
        if geom is None:
            return None
        limit = 256
        def tail(data: List[List[float]]) -> List[List[float]]:
            return data[-limit:]

        return {
            "state_a": tail(geom.state_a),
            "state_b": tail(geom.state_b),
            "candidate": tail(geom.candidate),
            "donor_basis": geom.donor_basis[-limit:],
            "receiver_basis": geom.receiver_basis[-limit:],
            "canonical_z": tail(geom.canonical_z),
            "hp_decode": tail(geom.hp_decode),
            "hp_target": tail(geom.hp_target),
        }

    payload = {
        "wormhole": [
            {
                "seed": sim.seed,
                "regimes": [
                    {
                        "name": regime.name,
                        "throughput": regime.throughput,
                        "budget_error": regime.budget_error,
                        "debt_balance": regime.debt_balance,
                        "metrics": {
                            "soft_rank_mean": regime.metrics.soft_active_rank_mean,
                            "hp_match": regime.metrics.hp_match_fraction,
                            "shock_delay_mean": regime.metrics.shock_delay_mean,
                        },
                        "hp_trace": [
                            {
                                "packet": entry.packet_index,
                                "step": entry.step,
                                "match": entry.match_fraction,
                            }
                            for entry in (regime.trace.hp_trace if regime.trace else [])
                        ],
                        "geometry": _geometry_payload(regime),
                    }
                    for regime in sim.regimes
                ],
            }
            for sim in simulations
        ],
        "sweep": {
            "lambda_h": [p.lambda_h for p in sweep.points],
            "throughput": [p.throughput for p in sweep.points],
            "sigma": [p.sigma for p in sweep.points],
            "soft_frac": [p.soft_rank_fraction for p in sweep.points],
            "lambda_crit_est": sweep.analysis.lambda_crit_est,
        },
        "bridge": {
            "alignment_error": benchmark_metrics.get("alignment_error"),
            "loss_gain_mean": benchmark_metrics.get("loss_gain_mean"),
            "loss_gain_std": benchmark_metrics.get("loss_gain_std"),
        },
    }
    path.parent.mkdir(parents=True, exist_ok=True)
    path.write_text(json.dumps(_sanitize(payload), indent=2), encoding="utf-8")
    return path
