from __future__ import annotations

import argparse
import json
from dataclasses import asdict
from pathlib import Path
from typing import List
import math

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


def format_markdown(results):
    lines = ["# Smooth Wormhole Experiment", ""]
    lines.append("| seed | throughput | packets | lambda_max_mean | cfl_margin_min | hp_match |")
    lines.append("|------|------------|---------|-----------------|----------------|---------|")
    for res in results:
        m = res.metrics
        lines.append(
            f"| {res.seed} | {m.throughput:.4f} | {m.packets} | {m.lambda_max_mean:.4f} | "
            f"{m.cfl_margin_min:.4f} | {m.hp_match:.3f} |"
        )
    lines.append("")
    return "\n".join(lines)


def _sanitize(value):
    if isinstance(value, float):
        return value if math.isfinite(value) else 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 _tail(sequence, limit=256):
    if sequence is None:
        return None
    return sequence[-limit:]


def build_wireframe_payload(results):
    payload = {"wormhole": [], "sweep": {"lambda_h": [], "throughput": [], "sigma": [], "soft_frac": [], "lambda_crit_est": None}, "bridge": {}}

    for res in results:
        metrics = res.metrics
        geometry = geometry_to_payload(res.geometry, 256) if res.geometry else None

        wormhole_entry = {
            "seed": res.seed,
            "regimes": [
                {
                    "name": "wormhole",
                    "throughput": metrics.throughput,
                    "budget_error": 0.0,
                    "debt_balance": 0.0,
                    "metrics": {
                        "hp_match": metrics.hp_match,
                        "hp_packets": metrics.hp_packets,
                        "cfl_margin_min": metrics.cfl_margin_min,
                        "lambda_max_mean": metrics.lambda_max_mean,
                        "packet_energy_mean": metrics.packet_energy_mean,
                    },
                    "hp_trace": [
                        {"packet": entry.packet, "match": entry.match}
                        for entry in res.hp_trace
                    ],
                    "geometry": geometry,
                }
            ],
        }

        payload["wormhole"].append(wormhole_entry)

    return _sanitize(payload)


def main():
    parser = argparse.ArgumentParser(description="Run smooth wormhole experiment")
    parser.add_argument("--output", type=Path, required=True)
    parser.add_argument("--seed", type=int, default=7)
    parser.add_argument("--extra-seeds", type=int, nargs="*", default=[])
    parser.add_argument("--hp-message-magnitude", type=float, default=None)
    parser.add_argument("--wormhole-budget-fraction", type=float, default=None)
    parser.add_argument("--debt-repay-rate", type=float, default=None)
    parser.add_argument("--lambda-c", type=float, default=None)
    args = parser.parse_args()

    seeds = [args.seed, *args.extra_seeds]
    cfg_overrides = {}
    if args.hp_message_magnitude is not None:
        cfg_overrides["hp_message_magnitude"] = args.hp_message_magnitude
    if args.wormhole_budget_fraction is not None:
        cfg_overrides["wormhole_budget_fraction"] = args.wormhole_budget_fraction
    if args.debt_repay_rate is not None:
        cfg_overrides["debt_repay_rate"] = args.debt_repay_rate
    if args.lambda_c is not None:
        cfg_overrides["lambda_c"] = args.lambda_c

    cfg = SmoothConfig(**cfg_overrides)

    results = [run_smooth_experiment(cfg, seed) for seed in seeds]

    output_dir = args.output
    output_dir.mkdir(parents=True, exist_ok=True)

    json_path = output_dir / "smooth_metrics.json"
    json_path.write_text(
        json.dumps(
            {
                "config": asdict(cfg),
                "results": [
                    {
                        "seed": res.seed,
                        "metrics": asdict(res.metrics),
                    }
                    for res in results
                ],
            },
            indent=2,
        ),
        encoding="utf-8",
    )

    md_path = output_dir / "smooth_summary.md"
    md_path.write_text(format_markdown(results), encoding="utf-8")

    wireframe_payload = build_wireframe_payload(results)
    (output_dir / "wireframe.json").write_text(json.dumps(wireframe_payload, indent=2), encoding="utf-8")

    print(f"Metrics written to {json_path}")
    print(f"Summary written to {md_path}")
    print(f"Wireframe feed written to {output_dir / 'wireframe.json'}")


if __name__ == "__main__":
    main()
