#!/usr/bin/env python3
"""Monte Carlo over parameter uncertainty.

Design variables of the distributed architecture are FIXED (they are a design, not an uncertainty):
  n_domains=100, authority_cap=500 MW, envelope=0.3, ramp=0.2/min.
Uncertain variables are sampled from LOW–HIGH (triangular, log-triangular for scale-like keys).
Certification (p / p_reduction_cert) is applied to BOTH sides.

Compared configurations (all fleet-level, all channels):
  A    : single control plane, no device envelope
  A+   : single plane + server-side aggregate cap (API compromise bounded; backend compromise not)
  B    : 100 domains + device envelope + cap, software-only (floor clamp lost on vendor firmware compromise)
  Bfull: B + independent-monitor floor clamp (vendor_hw_enforce=1) + firmware-line reach cap (top share ≤ 12%)

Run: python3 model/monte_carlo.py → results/monte_carlo_summary.json, results/mc_samples.csv,
     results/tornado_{A,Aplus,B,Bfull}.csv, results/mc_winrate_by_p_reduction.csv
"""

from __future__ import annotations

import json
from pathlib import Path

import numpy as np

from cbr_model import base_params, evaluate_a_plus, evaluate_architecture, load_assumptions, write_csv

ROOT = Path(__file__).resolve().parents[1]
RESULTS = ROOT / "results"
RNG = np.random.default_rng(7)

LOG_KEYS = {"n_devices", "p_compromise", "t_recovery_h", "p_reduction_cert", "p_vendor_rel", "p_backend_rel", "server_cap_mw"}
UNCERTAIN = [
    "n_devices", "device_kw", "remote_reach", "n_vendors", "vendor_top_share", "local_enforce",
    "p_compromise", "p_reduction_cert", "p_vendor_rel", "p_backend_rel", "domain_beta", "server_cap_mw", "t_recovery_h",
    "grid_demand_mw", "fcr_frac", "frr_frac", "severity_x0", "severity_k", "sync_factor",
]
DESIGN_B = dict(n_domains=100.0, authority_cap_mw=500.0, envelope_frac=0.3, ramp_limit_frac_per_min=0.2)


def sample(assump: dict, n: int) -> list[dict]:
    cols = {}
    for k in UNCERTAIN:
        lo, mode, hi = assump[k]["low"], assump[k]["base"], assump[k]["high"]
        if k in LOG_KEYS:
            lo, mode, hi = np.log10(lo), np.log10(mode), np.log10(hi)
        if lo == hi:
            cols[k] = np.full(n, mode)
        else:
            lo2, hi2 = min(lo, hi), max(lo, hi)
            cols[k] = RNG.triangular(lo2, min(max(mode, lo2), hi2), hi2, size=n)
        if k in LOG_KEYS:
            cols[k] = 10 ** cols[k]
    base = base_params(assump)
    out = []
    for i in range(n):
        d = dict(base)
        d.update(DESIGN_B)
        for k in UNCERTAIN:
            d[k] = float(cols[k][i])
        out.append(d)
    return out


def configs(d: dict) -> dict[str, dict]:
    pc = d["p_compromise"] / d["p_reduction_cert"]
    a = evaluate_architecture(d, n_domains=1, local_autonomy=False, envelope_frac=1.0, authority_cap_mw=1e9, p_compromise=pc)
    aplus = evaluate_a_plus(d, p_compromise=pc)
    b = evaluate_architecture({**d, "vendor_hw_enforce": 0.0}, n_domains=d["n_domains"], local_autonomy=True, p_compromise=pc)
    df = dict(d)
    df["vendor_hw_enforce"] = 1.0
    df["vendor_top_share"] = min(d["vendor_top_share"], 0.12)
    df["n_vendors"] = max(d["n_vendors"], 9)
    bfull = evaluate_architecture(df, n_domains=d["n_domains"], local_autonomy=True, p_compromise=pc)
    return {"A": a, "Aplus": aplus, "B": b, "Bfull": bfull}


def main() -> None:
    assump = load_assumptions()
    n = 20_000
    draws = sample(assump, n)
    rows = []
    for d in draws:
        c = configs(d)
        row = {k: d[k] for k in UNCERTAIN}
        for name, r in c.items():
            row[f"risk_{name}"] = r["fleet_systemic_risk_total"]
            row[f"plane_{name}"] = r["fleet_systemic_risk"]
            row[f"vendor_{name}"] = r["fleet_systemic_risk_vendor_channel"]
            row[f"emw_{name}"] = r["fleet_expected_affected_mw_total"]
            row[f"x_{name}"] = r["x_domain"]
            row[f"xv_{name}"] = r["x_vendor"]
        rows.append(row)
    write_csv(RESULTS / "mc_samples.csv", rows[:5000])

    def arr(k):
        return np.array([r[k] for r in rows])

    R = {k: arr(f"risk_{k}") for k in ("A", "Aplus", "B", "Bfull")}
    P = {k: arr(f"plane_{k}") for k in ("A", "Aplus", "B", "Bfull")}
    E = {k: arr(f"emw_{k}") for k in ("A", "Aplus", "B", "Bfull")}

    def cmp(x, y):
        return {"share_lower": float((y < x).mean()), "share_lower_by_10x": float((y * 10 < x).mean()), "median_ratio": float(np.median(x / np.maximum(y, 1e-300)))}

    summary = {
        "n_draws": n,
        "design_B": DESIGN_B,
        "note": "All risks are fleet-level annual probabilities of a wide-area event (UFLS or worse) summed over the control-plane channel (n domains, with common-cause share β) and the vendor channel (all firmware lines). Certification p-reduction applied to all configurations.",
        "medians": {k: {"total": float(np.median(R[k])), "p90": float(np.percentile(R[k], 90)), "plane": float(np.median(P[k])), "vendor": float(np.median(arr(f"vendor_{k}"))), "expected_mw": float(np.median(E[k]))} for k in R},
        "plane_channel_only": {"B_vs_A": cmp(P["A"], P["B"]), "Bfull_vs_A": cmp(P["A"], P["Bfull"]), "Aplus_vs_A": cmp(P["A"], P["Aplus"])},
        "total": {
            "B_vs_A": cmp(R["A"], R["B"]), "Bfull_vs_A": cmp(R["A"], R["Bfull"]), "Aplus_vs_A": cmp(R["A"], R["Aplus"]),
            "Bfull_vs_Aplus": cmp(R["Aplus"], R["Bfull"]), "B_vs_Aplus": cmp(R["Aplus"], R["B"]),
        },
        "expected_mw": {"B_vs_A": cmp(E["A"], E["B"]), "Bfull_vs_A": cmp(E["A"], E["Bfull"]), "Aplus_vs_A": cmp(E["A"], E["Aplus"]), "Bfull_vs_Aplus": cmp(E["Aplus"], E["Bfull"])},
        "x": {k: {"domain_median": float(np.median(arr(f"x_{k}"))), "domain_share_above_1": float((arr(f"x_{k}") > 1).mean()), "vendor_median": float(np.median(arr(f"xv_{k}"))), "vendor_share_above_1": float((arr(f"xv_{k}") > 1).mean())} for k in R},
    }
    # when does Bfull lose to A+?
    lose = [r for r in rows if r["risk_Bfull"] >= r["risk_Aplus"]]
    if lose:
        summary["when_Bfull_loses_to_Aplus"] = {
            "count": len(lose),
            "median_p_vendor_rel": float(np.median([r["p_vendor_rel"] for r in lose])),
            "median_p_backend_rel": float(np.median([r["p_backend_rel"] for r in lose])),
            "median_domain_beta": float(np.median([r["domain_beta"] for r in lose])),
            "median_server_cap_mw": float(np.median([r["server_cap_mw"] for r in lose])),
            "median_n_devices": float(np.median([r["n_devices"] for r in lose])),
            "median_severity_k": float(np.median([r["severity_k"] for r in lose])),
            "note": "Bfull loses when the vendor channel is relatively likely (p_vendor_rel high) and/or the server-side cap is small and hard to bypass (p_backend_rel low), or the fleet is small relative to the grid",
        }
    pr = arr("p_reduction_cert")
    wr = []
    for lo, hi in ((1, 1.5), (1.5, 2.5), (2.5, 4), (4, 7), (7, 10.01)):
        m = (pr >= lo) & (pr < hi)
        if m.sum():
            wr.append({"p_reduction_low": lo, "p_reduction_high": hi, "n": int(m.sum()),
                       "share_B_lower_than_A": float((R["B"][m] < R["A"][m]).mean()), "share_Bfull_lower_than_A": float((R["Bfull"][m] < R["A"][m]).mean()),
                       "share_Aplus_lower_than_A": float((R["Aplus"][m] < R["A"][m]).mean()), "share_Bfull_lower_than_Aplus": float((R["Bfull"][m] < R["Aplus"][m]).mean())})
    write_csv(RESULTS / "mc_winrate_by_p_reduction.csv", wr)

    # tornado (one variable at a time, LOW/HIGH), fleet total risk
    base = base_params(assump)
    base.update(DESIGN_B)
    for arch in ("A", "Aplus", "B", "Bfull"):
        trows = []
        base_val = configs(base)[arch]["fleet_systemic_risk_total"]
        for k in UNCERTAIN:
            vals = {}
            for level in ("low", "high"):
                d = dict(base)
                d[k] = assump[k][level]
                vals[level] = configs(d)[arch]["fleet_systemic_risk_total"]
            trows.append({"variable": k, "low": vals["low"], "high": vals["high"], "base": base_val, "swing": abs(vals["high"] - vals["low"])})
        trows.sort(key=lambda r: -r["swing"])
        write_csv(RESULTS / f"tornado_{arch}.csv", trows)

    (RESULTS / "monte_carlo_summary.json").write_text(json.dumps(summary, ensure_ascii=False, indent=2), encoding="utf-8")
    print(json.dumps(summary, ensure_ascii=False, indent=2))


if __name__ == "__main__":
    main()
