#!/usr/bin/env python3
"""Attack simulation: the same hostile command applied to a centralized and a distributed
autonomous DER architecture, evaluated with a single-area frequency model.

Model (per-unit on system demand S):
  2H * d(Δf_pu)/dt = -ΔP_attack(t) + ΔP_fcr(t) + ΔP_ufls(t) - D * Δf_pu
  FCR: first-order lag (time constant T_fcr) towards droop demand  -K * Δf  saturating at ±FCR_max
  UFLS: stage trips when f < threshold (latched); shed = fraction of demand
  Attack: ΔP_attack(t) ramps from 0 to the architecture's effective blast radius at the ramp
          limit allowed by the command envelope (step if no envelope)

Outputs per scenario: affected MW, max RoCoF, frequency nadir, UFLS shed MW, time to recover
within ±0.1 Hz of nominal, systemic severity class.

Run: python3 model/grid_sim.py  → results/attack_scenarios.csv, results/attack_traces.csv
"""

from __future__ import annotations

import csv
from dataclasses import dataclass
from pathlib import Path

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

ROOT = Path(__file__).resolve().parents[1]
RESULTS = ROOT / "results"
F0 = 50.0


@dataclass
class Grid:
    demand_mw: float
    h_s: float
    fcr_mw: float
    fcr_tau_s: float
    damping: float
    ufls_stages: list[tuple[float, float]]  # (threshold Hz, shed fraction of demand)
    frr_mw: float = 0.0
    frr_tau_s: float = 90.0
    droop_gain: float = 20.0  # pu MW per pu Hz (5% droop)
    gen_trip_hz: float = 47.5  # below this, assume cascading generator/inverter disconnection


def simulate(grid: Grid, attack_mw: float, ramp_mw_per_s: float | None, t_end: float = 300.0, dt: float = 0.01, step_mw: float = 0.0) -> dict:
    """attack_mw = total hostile change. `step_mw` of it is applied instantly at t=1 s (devices whose
    limits are not enforced or whose ramp limit died with the firmware); the rest ramps at ramp_mw_per_s."""
    s = grid.demand_mw
    f_pu = 0.0  # Δf in pu of 50 Hz
    fcr_pu = 0.0
    frr_pu = 0.0
    shed_pu = 0.0
    tripped = [False] * len(grid.ufls_stages)
    max_rocof = 0.0
    nadir = 0.0
    recovered_at = None
    collapse = False
    trace = []
    n = int(t_end / dt)
    attack_pu_target = attack_mw / s
    fcr_max_pu = grid.fcr_mw / s
    frr_max_pu = grid.frr_mw / s
    for i in range(n + 1):
        t = i * dt
        if ramp_mw_per_s is None:
            attack_pu = attack_pu_target if t >= 1.0 else 0.0
        else:
            attack_pu = min(attack_pu_target, max(0.0, (step_mw + (t - 1.0) * ramp_mw_per_s) / s)) if t >= 1.0 else 0.0
        # FCR demand with droop and saturation
        fcr_demand = max(-fcr_max_pu, min(fcr_max_pu, -grid.droop_gain * f_pu))
        fcr_pu += (fcr_demand - fcr_pu) * dt / grid.fcr_tau_s
        # FRR / LFC: integral-type restoration towards the residual imbalance, saturating
        frr_target = max(-frr_max_pu, min(frr_max_pu, attack_pu - shed_pu))
        frr_pu += (frr_target - frr_pu) * dt / grid.frr_tau_s
        # UFLS stages
        f_hz = F0 * (1 + f_pu)
        for k, (thr, frac) in enumerate(grid.ufls_stages):
            if not tripped[k] and f_hz < thr:
                tripped[k] = True
                shed_pu += frac
        if f_hz < grid.gen_trip_hz:
            collapse = True
        d_f = (-attack_pu + fcr_pu + frr_pu + shed_pu - grid.damping * f_pu) / (2 * grid.h_s)
        rocof_hz_s = d_f * F0
        max_rocof = max(max_rocof, abs(rocof_hz_s))
        f_pu += d_f * dt
        nadir = min(nadir, F0 * f_pu)
        if t > 1.0 and recovered_at is None and abs(F0 * f_pu) < 0.1 and attack_pu >= attack_pu_target * 0.999:
            recovered_at = t
        if i % int(0.5 / dt) == 0:
            trace.append((t, F0 * (1 + f_pu), attack_pu * s, (fcr_pu + frr_pu) * s, shed_pu * s))
        if collapse:
            break
    return {
        "affected_mw": attack_mw,
        "max_rocof_hz_s": max_rocof,
        "nadir_hz": F0 + nadir,
        "ufls_shed_mw": shed_pu * s,
        "ufls_triggered": any(tripped),
        "collapse": collapse,
        "recovered_within_0_1hz_s": recovered_at if recovered_at is not None else float("nan"),
        "final_dev_hz": F0 * f_pu,
        "trace": trace,
    }


def classify(r: dict) -> str:
    if r["collapse"]:
        return "cascade / collapse"
    if r["ufls_triggered"]:
        return "UFLS (customers shed)"
    if r["nadir_hz"] < 49.8 or r["final_dev_hz"] < -0.2:
        return "degraded (reserves exhausted or nadir < 49.8 Hz)"
    return "contained"


def main() -> None:
    assump = load_assumptions()
    p = base_params(assump)
    grid = Grid(
        demand_mw=p["grid_demand_mw"],
        h_s=p["inertia_h_s"],
        fcr_mw=p["grid_demand_mw"] * p["fcr_frac"],
        fcr_tau_s=p["fcr_time_const_s"],
        damping=p["load_damping"],
        frr_mw=p["grid_demand_mw"] * p["frr_frac"],
        frr_tau_s=p["frr_time_const_s"],
        ufls_stages=[(p["ufls_first_hz"], p["ufls_first_frac"]), (p["ufls_first_hz"] - 0.3, 0.10), (p["ufls_first_hz"] - 0.6, 0.10)],
    )
    c_total = total_capacity_mw(p["n_devices"], p["device_kw"])
    shares = vendor_shares(p["n_vendors"], p["vendor_top_share"])
    reach = p["remote_reach"]
    e, rho, cap = p["envelope_frac"], p["local_enforce"], p["authority_cap_mw"]
    ramp_frac = p["ramp_limit_frac_per_min"]

    def ramp_for(attack_mw: float, base_capacity_mw: float, autonomy: bool):
        if not autonomy:
            return None  # step
        return base_capacity_mw * ramp_frac / 60.0  # MW/s allowed by envelope on the moved fleet

    hostile = 0.5  # "reduce output 50%" command

    def eff_mw(raw_total: float, k_domains: int, autonomy: bool, use_cap: bool, enforce: float, cmd: float) -> float:
        """Effective moved capacity. Cap is per compromised domain; `enforce` is the local
        enforcement reliability that survives the compromise (0 when the attacker owns the firmware
        and the envelope is implemented only in that firmware)."""
        per_domain = raw_total / k_domains
        cap_eff = cap if (autonomy and use_cap) else 1e12
        return k_domains * cbr_effective_mw(per_domain, reach, cap_eff, e, enforce, autonomy) * cmd

    def step_part_mw(raw_total: float, k_domains: int, autonomy: bool, use_cap: bool, enforce: float, cmd: float, ramp_survives: bool) -> float:
        """Part of the moved capacity that changes as a step: devices whose local limits are not
        enforced (1-enforce) move fully and instantly; if the ramp limit itself did not survive the
        compromise (firmware), everything is a step."""
        if not autonomy:
            return eff_mw(raw_total, k_domains, autonomy, use_cap, enforce, cmd)
        if not ramp_survives:
            return eff_mw(raw_total, k_domains, autonomy, use_cap, enforce, cmd)
        per_domain = raw_total / k_domains
        cap_eff = cap if use_cap else 1e12
        capped = min(per_domain, cap_eff)
        return k_domains * reach * capped * (1.0 - enforce) * cmd

    scenarios = [
        # name, description, raw capacity, k compromised domains, autonomy, use_cap, enforce, command fraction, ramp_survives
        ("A1 cloud compromise, national plane, 50% cut", "Scenario A: single control plane, no envelope; attacker cuts 50% of all reachable DER at once", c_total, 1, False, False, 0.0, hostile, False),
        ("A2 cloud compromise, national plane, full OFF", "Scenario A: same plane, attacker sends OFF to everything reachable", c_total, 1, False, False, 0.0, 1.0, False),
        ("A3 vendor cloud compromise (largest vendor), full OFF", "Scenario A: one vendor cloud (30% share), no envelope", c_total * shares[0], 1, False, False, 0.0, 1.0, False),
        ("A4 single plane with server-side cap 1,500 MW, API credential compromise", "Scenario A+: DERMS enforces an aggregate authority cap per window (independent safety monitor); attacker holds API credentials only", 1500.0 / reach, 1, False, False, 0.0, 1.0, False),
        ("H1 one regional domain of 10, full OFF", "Hierarchical: one of 10 regional DERMS compromised, no envelope", c_total / 10, 1, False, False, 0.0, 1.0, False),
        ("F1 one domain of 100, full OFF", "Federated: one of 100 distribution-level domains compromised, no envelope", c_total / 100, 1, False, False, 0.0, 1.0, False),
        ("B1 one domain of 100, autonomy ON", "Scenario B: one domain compromised; devices enforce envelope e=0.3 (ρ=0.95), ramp, cap 500 MW", c_total / 100, 1, True, True, rho, 1.0, True),
        ("B2 national plane compromised, autonomy ON (no cap)", "Scenario B': attacker owns the national plane; devices enforce envelope and ramp but there is no authority cap (5% non-enforced part moves as a step)", c_total, 1, True, False, rho, 1.0, True),
        ("B3 national plane compromised, cap per scope but one credential opens all scopes", "Scenario B'': per-domain cap exists but a single compromise yields all 100 scopes' credentials → cap does not bind", c_total, 100, True, True, rho, 1.0, True),
        ("B4 10 of 100 domains compromised simultaneously", "Scenario B stress: correlated compromise of 10 domains (shared IdP), autonomy ON, cap per domain", c_total / 10, 10, True, True, rho, 1.0, True),
        ("B5a largest vendor firmware compromised, software-only autonomy", "Scenario B stress: vendor firmware/OTA compromise; envelope lived in that firmware so it is gone (ρ→0)", c_total * shares[0], 1, True, False, 0.0, 1.0, False),
        ("B5b largest vendor firmware compromised, independent-monitor floor clamp survives (ramp lost)", "Scenario B stress: same, but the output floor is clamped by an independent monitor MCU / remote-OFF class disabled in hardware; the ramp limit lives in the main firmware and is lost → step of 30%", c_total * shares[0], 1, True, False, rho, 1.0, False),
        ("B5b' hypothetical: floor clamp and ramp both survive firmware compromise", "Reference only: shows what a hardware-enforced ramp would buy if it existed", c_total * shares[0], 1, True, False, rho, 1.0, True),
        ("B5c firmware-line reach capped at 12% of fleet, floor clamp survives", "Scenario B: firmware-line (keys/OTA) reach limited to 12% of fleet; floor clamp survives, ramp lost → step", c_total * 0.12, 1, True, False, rho, 1.0, False),
        ("B5d firmware-line reach capped at 4% of fleet, software-only", "Scenario B: software-only limits need firmware-line reach ≤ ~4% of fleet to stay within FCR", c_total * 0.04, 1, True, False, 0.0, 1.0, False),
    ]
    rows, traces = [], []
    for name, desc, raw, k, autonomy, use_cap, enforce, cmd, ramp_survives in scenarios:
        eff = eff_mw(raw, k, autonomy, use_cap, enforce, cmd)
        step = step_part_mw(raw, k, autonomy, use_cap, enforce, cmd, ramp_survives)
        ramp_applies = autonomy and ramp_survives and step < eff
        s_fast = (step + 0.0) / eff if eff > 0 else 1.0  # fraction acting as a step
        ramped = eff - step
        # section-6 definition: fast part against FCR, whole against FCR+FRR; ramped part contributes
        # within the 30 s window at ramp rate
        in_window = step + (min(ramped, ramp_frac * raw / 60.0 * 30.0) if ramp_applies else ramped)
        x6 = max(in_window / grid.fcr_mw, eff / (grid.fcr_mw + grid.frr_mw))
        r = simulate(grid, eff, ramp_for(eff, eff, ramp_applies), t_end=1500.0 if ramp_applies else 300.0, step_mw=step)
        rows.append({
            "scenario": name,
            "description": desc,
            "raw_capacity_mw": raw,
            "affected_mw": eff,
            "affected_fraction_of_fleet": eff / c_total,
            "step_part_mw": step,
            "x_section6": x6,
            "x_affected_over_fcr": eff / grid.fcr_mw,
            "x_affected_over_fcr_plus_frr": eff / (grid.fcr_mw + grid.frr_mw),
            "ramp_limited": ramp_applies,
            "max_rocof_hz_s": r["max_rocof_hz_s"],
            "nadir_hz": r["nadir_hz"],
            "ufls_shed_mw": r["ufls_shed_mw"],
            "recovered_within_0_1hz_s": r["recovered_within_0_1hz_s"],
            "class": classify(r),
        })
        for t, f, a, fcr, shed in r["trace"]:
            traces.append({"scenario": name, "t_s": t, "f_hz": f, "attack_mw": a, "fcr_mw": fcr, "shed_mw": shed})
    write_csv(RESULTS / "attack_scenarios.csv", rows)
    write_csv(RESULTS / "attack_traces.csv", traces)

    # sweep: step loss vs nadir for the base grid, to show where UFLS starts (grid sensitivity)
    sweep = []
    for mw in (250, 500, 1000, 1350, 1500, 2000, 2500, 3000, 4000, 5000, 7500, 10000):
        r = simulate(grid, mw, None)
        sweep.append({"step_loss_mw": mw, "x_over_fcr": mw / grid.fcr_mw, "nadir_hz": r["nadir_hz"], "max_rocof_hz_s": r["max_rocof_hz_s"], "ufls_shed_mw": r["ufls_shed_mw"], "class": classify(r)})
    write_csv(RESULTS / "grid_step_sweep.csv", sweep)

    # other grids for context (parameters documented in research/grid_sensitivity_calibration.md)
    others = {
        "Hokkaido 2018-like (3,090 MW, H=4, FCR 5%)": Grid(3090, 4.0, 155, 10, 1.5, [(48.5, 0.42)], frr_mw=155),
        "Continental Europe-like (400 GW, H=4, FCR 3,000 MW)": Grid(400000, 4.0, 3000, 10, 1.5, [(49.0, 0.05), (48.7, 0.10), (48.4, 0.10)], frr_mw=20000),
    }
    ctx = []
    for gname, g in others.items():
        for mw in (1160, 2660, 3000, 4500, 10000):
            r = simulate(g, mw, None)
            ctx.append({"grid": gname, "step_loss_mw": mw, "x_over_fcr": mw / g.fcr_mw, "nadir_hz": r["nadir_hz"], "max_rocof_hz_s": r["max_rocof_hz_s"], "class": classify(r)})
    write_csv(RESULTS / "grid_context_cases.csv", ctx)

    for r in rows:
        print(f"{r['scenario'][:60]:<60} {r['affected_mw']:>9,.0f} MW (step {r['step_part_mw']:>7,.0f})  x6={r['x_section6']:>5.2f}  RoCoF={r['max_rocof_hz_s']:.2f}  nadir={r['nadir_hz']:.2f}  {r['class']}")


if __name__ == "__main__":
    main()
