#!/usr/bin/env python3
"""Correlated (common-cause) failure of DER fleets.

Part 1 — analytic beta-factor view (IEC 61508-6 / NUREG/CR-5485 style):
  each device fails independently with probability (1-β)·p per year, and in addition the whole
  group sharing a component (cloud, firmware line, certificate, API) fails together with
  probability β·p.  Independent-failure arithmetic (p^N) is then meaningless for the tail.

Part 2 — Monte Carlo over shared components:
  fleet = vendors (shares) × {vendor cloud, firmware line} + aggregator domains + individual
  devices.  Each year every shared component is compromised with its own probability; affected
  capacity is the union of everything reachable through compromised components, after the
  envelope/enforcement that survives the compromise.  Compared fleets:
    monoculture      : 1 vendor, 1 cloud, 1 firmware line, no autonomy
    diverse          : 8 vendors (Zipf), 100 domains, no autonomy
    diverse+autonomy : as above, software envelope (lost on firmware compromise)
    diverse+hw       : as above, hardware-enforced envelope survives firmware compromise

Run: python3 model/correlated_failure.py → results/ccf_*.csv, results/ccf_summary.json
"""

from __future__ import annotations

import json
import math
from pathlib import Path

import numpy as np

from cbr_model import base_params, load_assumptions, total_capacity_mw, vendor_shares, write_csv

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


def analytic_table(n_devices: int, p_dev: float) -> list[dict]:
    rows = []
    for beta in (0.0, 0.001, 0.01, 0.05, 0.1):
        # probability that at least 10% of the fleet is compromised in the same year
        k = int(0.1 * n_devices)
        if beta == 0.0:
            # normal approximation to Binomial tail; astronomically small for p_dev << 0.1
            mu, sd = n_devices * p_dev, math.sqrt(n_devices * p_dev * (1 - p_dev))
            z = (k - mu) / sd if sd > 0 else float("inf")
            p_tail = 0.5 * math.erfc(z / math.sqrt(2)) if z < 38 else 0.0
        else:
            p_tail = beta * p_dev  # dominated by the common-cause term
        rows.append({
            "beta": beta,
            "p_device_per_year": p_dev,
            "expected_compromised_fraction": p_dev,  # same mean regardless of beta
            "p_at_least_10pct_simultaneously": p_tail,
            "interpretation": "independent arithmetic" if beta == 0 else "common cause dominates the tail",
        })
    return rows


def simulate_fleet(
    years: int,
    c_total: float,
    shares: list[float],
    n_domains: int,
    p_vendor_cloud: float,
    p_firmware: float,
    p_domain: float,
    p_device: float,
    reach: float,
    envelope: float,
    rho_sw: float,
    rho_hw: float,
    n_devices: int,
) -> np.ndarray:
    """Return affected MW per simulated year (max simultaneous hostile-controllable capacity)."""
    v = len(shares)
    shares_arr = np.array(shares)
    vendor_cap = c_total * shares_arr * reach
    affected = np.zeros(years)
    # vendor clouds: envelope enforced by device firmware survives (sw+hw)
    cloud_hits = RNG.random((years, v)) < p_vendor_cloud
    enforce_cloud = 1.0 - (rho_sw + rho_hw) * (1.0 - envelope)
    affected += (cloud_hits * vendor_cap * enforce_cloud).sum(axis=1)
    # firmware lines: software envelope is gone, only hardware envelope survives
    fw_hits = RNG.random((years, v)) < p_firmware
    enforce_fw = 1.0 - rho_hw * (1.0 - envelope)
    fw_aff = fw_hits * vendor_cap * enforce_fw
    # union with cloud hits of the same vendor: take the larger of the two per vendor
    cloud_aff = cloud_hits * vendor_cap * enforce_cloud
    affected = np.maximum(cloud_aff, fw_aff).sum(axis=1)
    # aggregator domains (independent credentials), devices split equally
    dom_cap = c_total * reach / n_domains
    dom_hits = RNG.binomial(n_domains, p_domain, size=years)
    affected += dom_hits * dom_cap * enforce_cloud
    # individual devices (independent)
    dev_hits = RNG.binomial(n_devices, p_device, size=years)
    affected += dev_hits * (c_total / n_devices) * reach
    return np.minimum(affected, c_total * reach)


def summarize(name: str, aff: np.ndarray, fcr: float, c_total: float) -> dict:
    return {
        "fleet": name,
        "mean_affected_mw": float(aff.mean()),
        "p50_mw": float(np.percentile(aff, 50)),
        "p90_mw": float(np.percentile(aff, 90)),
        "p99_mw": float(np.percentile(aff, 99)),
        "p999_mw": float(np.percentile(aff, 99.9)),
        "max_mw": float(aff.max()),
        "p_exceed_fcr_per_year": float((aff > fcr).mean()),
        "p_exceed_1_5_fcr_per_year": float((aff > 1.5 * fcr).mean()),
        "p_exceed_10pct_fleet_per_year": float((aff > 0.1 * c_total).mean()),
    }


def main() -> None:
    assump = load_assumptions()
    p = base_params(assump)
    n_dev = int(p["n_devices"])
    c_total = total_capacity_mw(p["n_devices"], p["device_kw"])
    fcr = p["grid_demand_mw"] * p["fcr_frac"]

    write_csv(RESULTS / "ccf_analytic.csv", analytic_table(n_dev, 1e-4))

    years = 200_000
    p_cloud, p_fw, p_dom, p_dev = p["p_compromise"], p["p_compromise"] / 3, p["p_compromise"], 1e-5
    fleets = {
        "monoculture (1 vendor, 1 cloud, 1 firmware, no autonomy)": dict(shares=[1.0], n_domains=1, envelope=1.0, rho_sw=0.0, rho_hw=0.0),
        "diverse (8 vendors, 100 domains, no autonomy)": dict(shares=vendor_shares(p["n_vendors"], p["vendor_top_share"]), n_domains=100, envelope=1.0, rho_sw=0.0, rho_hw=0.0),
        "diverse + software envelope (lost on firmware compromise)": dict(shares=vendor_shares(p["n_vendors"], p["vendor_top_share"]), n_domains=100, envelope=p["envelope_frac"], rho_sw=p["local_enforce"], rho_hw=0.0),
        "diverse + hardware-enforced envelope": dict(shares=vendor_shares(p["n_vendors"], p["vendor_top_share"]), n_domains=100, envelope=p["envelope_frac"], rho_sw=0.0, rho_hw=p["local_enforce"]),
        "monoculture + hardware-enforced envelope": dict(shares=[1.0], n_domains=1, envelope=p["envelope_frac"], rho_sw=0.0, rho_hw=p["local_enforce"]),
        "very diverse (20 vendors, top 15%) + hardware envelope": dict(shares=vendor_shares(20, 0.15), n_domains=1000, envelope=p["envelope_frac"], rho_sw=0.0, rho_hw=p["local_enforce"]),
    }
    rows, hist = [], {}
    for name, kw in fleets.items():
        aff = simulate_fleet(years, c_total, kw["shares"], kw["n_domains"], p_cloud, p_fw, p_dom, p_dev, p["remote_reach"], kw["envelope"], kw["rho_sw"], kw["rho_hw"], n_dev)
        rows.append(summarize(name, aff, fcr, c_total))
        counts, edges = np.histogram(aff, bins=np.logspace(0, math.log10(c_total) + 0.1, 40))
        hist[name] = {"edges_mw": edges.tolist(), "counts": counts.tolist()}
    write_csv(RESULTS / "ccf_fleets.csv", rows)
    (RESULTS / "ccf_summary.json").write_text(json.dumps({
        "years_simulated": years,
        "p_vendor_cloud": p_cloud, "p_firmware": p_fw, "p_domain": p_dom, "p_device": p_dev,
        "fcr_mw": fcr, "c_total_mw": c_total, "fleets": rows, "histograms": hist,
    }, ensure_ascii=False, indent=1), encoding="utf-8")

    # vendor-count sweep: P(exceed FCR) vs number of vendors, hardware envelope on/off
    sweep = []
    for nv in (1, 2, 3, 5, 8, 12, 20):
        for hw in (False, True):
            sh = vendor_shares(nv, max(1.0 / nv, 0.3 if nv <= 3 else 1.5 / nv))
            aff = simulate_fleet(50_000, c_total, sh, 100, p_cloud, p_fw, p_dom, p_dev, p["remote_reach"], p["envelope_frac"] if hw else 1.0, 0.0, p["local_enforce"] if hw else 0.0, n_dev)
            sweep.append({"n_vendors": nv, "top_share": sh[0], "hardware_envelope": hw, "p_exceed_fcr_per_year": float((aff > fcr).mean()), "p99_mw": float(np.percentile(aff, 99))})
    write_csv(RESULTS / "ccf_vendor_sweep.csv", sweep)

    for r in rows:
        print(f"{r['fleet']:<62} p99={r['p99_mw']:>9,.0f} MW  P(>FCR)={r['p_exceed_fcr_per_year']:.4f}/yr")


if __name__ == "__main__":
    main()
