import math
import json
import platform
import random
import sys
from pathlib import Path
from typing import Any, Dict, List, Optional, Tuple

from controller import SmoothThermostatController
from core import (
    canonical_json,
    clamp,
    parse_mu_disp_mode,
    require,
    require_unit_interval_closed_open,
    require_unit_interval_open_closed,
    resolve_path,
    sha256_hex,
)
from detector import DirichletVolatilityDetector
from dynamics import init_states, run_window
from graph import SimpleGraph, detect_communities
from mapping import dynamics_params_from_u
from metrics import brier_score, log_loss, recovery_times, summarize_u
from observables import (
    community_edge_counts,
    edge_counts,
    edge_metrics_from_counts,
    mu_dispersion,
    mu_by_comm_and_dispersion_from_counts,
    observed_community_counts,
    observed_counts,
)
from stress import apply_rewires, apply_shocks
from topology_actuator import apply_topology_actuator
from schema_types import (
    Event,
    Provenance,
    PythonProvenance,
    RunOutput,
    SeriesRow,
    TopologyActuatorEvent,
    TopologyActuatorTimeseries,
)


def _build_graph(graph_cfg: Dict[str, Any], rng: random.Random) -> SimpleGraph:
    typ = str(require(graph_cfg, "type", "graph")).lower()
    n = int(require(graph_cfg, "n", "graph"))
    if n <= 1:
        raise ValueError("graph.n must be > 1")

    if typ == "erdos_renyi":
        p = float(require(graph_cfg, "p", "graph"))
        return SimpleGraph.erdos_renyi(n=n, p=p, rng=rng)

    if typ == "stochastic_block_model":
        sizes = require(graph_cfg, "sizes", "graph")
        sizes_i = [int(x) for x in sizes]
        k = int(len(sizes_i))
        shuffle_nodes = bool(graph_cfg.get("shuffle_nodes", False))

        if "p_matrix" in graph_cfg:
            p_matrix = require(graph_cfg, "p_matrix", "graph")
        else:
            if ("p_in" in graph_cfg) or ("p_out" in graph_cfg):
                p_in = float(require(graph_cfg, "p_in", "graph"))
                p_out = float(require(graph_cfg, "p_out", "graph"))
            else:
                d = float(require(graph_cfg, "sbm_avg_degree", "graph"))
                mu = float(require(graph_cfg, "sbm_mu", "graph"))
                if mu < 0.0 or mu > 1.0:
                    raise ValueError("graph.sbm_mu must be in [0, 1]")
                if d < 0.0:
                    raise ValueError("graph.sbm_avg_degree must be >= 0")
                if int(sum(sizes_i)) != int(n):
                    raise ValueError("stochastic_block_model requires sum(graph.sizes) == graph.n")

                w = 0.0
                sum_sq = 0.0
                for s in sizes_i:
                    ss = float(s)
                    w += ss * (ss - 1.0) / 2.0
                    sum_sq += ss * ss
                b = (float(n) * float(n) - float(sum_sq)) / 2.0
                m_target = float(d) * float(n) / 2.0

                p_in = 0.0
                p_out = 0.0
                if w > 0.0:
                    p_in = (1.0 - float(mu)) * float(m_target) / float(w)
                if b > 0.0:
                    p_out = float(mu) * float(m_target) / float(b)

            if p_in < 0.0 or p_in > 1.0:
                raise ValueError("graph.p_in must be in [0, 1]")
            if p_out < 0.0 or p_out > 1.0:
                raise ValueError("graph.p_out must be in [0, 1]")

            if k <= 0:
                raise ValueError("stochastic_block_model requires at least one block")
            p_matrix = [[float(p_out) for _ in range(k)] for _ in range(k)]
            for i in range(k):
                p_matrix[i][i] = float(p_in)

        return SimpleGraph.stochastic_block_model(
            n=n,
            sizes=sizes,
            p_matrix=p_matrix,
            rng=rng,
            shuffle_nodes=shuffle_nodes,
        )

    if typ == "ring_lattice":
        k = int(require(graph_cfg, "k", "graph"))
        return SimpleGraph.ring_lattice(n=n, k=k)

    if typ == "barabasi_albert":
        m = int(require(graph_cfg, "m", "graph"))
        return SimpleGraph.barabasi_albert(n=n, m=m, rng=rng)

    raise ValueError(f"Unsupported graph.type: {typ}")


def _truth_labels(mu: float, eta: float, v: float, ethics_cfg: Dict[str, Any], *, mu_disp: float = 0.0) -> Tuple[int, int, int, int, int]:
    mu_min = float(require(ethics_cfg, "mu_min", "ethics"))
    eta_max = float(require(ethics_cfg, "eta_max", "ethics"))
    v_max = float(require(ethics_cfg, "v_max", "ethics"))

    y_disp = 1
    mu_disp_max = ethics_cfg.get("mu_disp_max")
    if mu_disp_max is not None:
        mu_disp_max_f = float(mu_disp_max)
        y_disp = 1 if float(mu_disp) <= float(mu_disp_max_f) else 0

    y_mu = 1 if float(mu) >= mu_min else 0
    y_eta = 1 if float(eta) <= eta_max else 0
    y_v = 1 if float(v) <= v_max else 0
    y_all = 1 if (y_mu == 1 and y_eta == 1 and y_v == 1 and y_disp == 1) else 0
    return y_all, y_mu, y_eta, y_v, y_disp


def run_closed_loop(config: Dict[str, Any], config_path: Path) -> RunOutput:
    seed = int(config.get("seed", 0))

    def _derive_seed(stream: str) -> int:
        h = sha256_hex(f"{int(seed)}:{str(stream)}")
        return int(h[:16], 16)

    rng_seeds: Dict[str, int] = {
        "graph": _derive_seed("graph"),
        "init": _derive_seed("init"),
        "stress": _derive_seed("stress"),
        "observation": _derive_seed("observation"),
        "community_detection": _derive_seed("community_detection"),
        "dynamics": _derive_seed("dynamics"),
        "detector": _derive_seed("detector"),
        "detector_risk": _derive_seed("detector_risk"),
        "topology_actuator": _derive_seed("topology_actuator"),
        "topology_actuator_risk": _derive_seed("topology_actuator_risk"),
    }

    rng_graph = random.Random(int(rng_seeds["graph"]))
    rng_init = random.Random(int(rng_seeds["init"]))
    rng_stress = random.Random(int(rng_seeds["stress"]))
    rng_obs = random.Random(int(rng_seeds["observation"]))
    rng_comm = random.Random(int(rng_seeds["community_detection"]))
    rng_dynamics = random.Random(int(rng_seeds["dynamics"]))
    rng_detector = random.Random(int(rng_seeds["detector"]))
    rng_detector_risk = random.Random(int(rng_seeds["detector_risk"]))
    rng_topo = random.Random(int(rng_seeds["topology_actuator"]))
    rng_topo_risk = random.Random(int(rng_seeds["topology_actuator_risk"]))

    out_dir = resolve_path(config_path, str(require(config, "out_dir", "config")))
    out_dir.mkdir(parents=True, exist_ok=True)

    run_cfg = require(config, "run", "config")
    control_steps = int(require(run_cfg, "control_steps", "run"))
    window = int(require(run_cfg, "window", "run"))

    graph_cfg = require(config, "graph", "config")
    init_cfg = require(config, "init", "config")
    dyn_cfg = require(config, "dynamics", "config")
    mapping_cfg = require(config, "mapping", "config")
    obs_cfg = require(config, "observation", "config")
    det_cfg = require(config, "detector", "config")
    ctrl_cfg = require(config, "controller", "config")
    ethics_cfg = require(config, "ethics", "config")

    mu_min = float(require(ethics_cfg, "mu_min", "ethics"))
    eta_max = float(require(ethics_cfg, "eta_max", "ethics"))
    v_max = float(require(ethics_cfg, "v_max", "ethics"))

    mu_disp_max = ethics_cfg.get("mu_disp_max")

    meso_cfg = config.get("meso", {})
    comm_detect_cfg: Dict[str, Any] = {}
    comm_every_steps = 1
    comm_min_total_edges = 1.0
    if mu_disp_max is not None:
        if not isinstance(meso_cfg, dict):
            raise ValueError("meso must be an object when ethics.mu_disp_max is set")
        cd = meso_cfg.get("community_detection", {})
        if not isinstance(cd, dict):
            raise ValueError("meso.community_detection must be an object")
        comm_detect_cfg = cd
        comm_every_steps = int(comm_detect_cfg.get("every_steps", 1))
        comm_min_total_edges = float(meso_cfg.get("min_total_edges", 1.0))

        if comm_every_steps <= 0:
            raise ValueError("meso.community_detection.every_steps must be > 0")
        if comm_min_total_edges < 0.0:
            raise ValueError("meso.min_total_edges must be >= 0")

    diagnostics_cfg = config.get("diagnostics", {})
    stuck_cfg = diagnostics_cfg.get("stuck") if isinstance(diagnostics_cfg, dict) else None
    stuck_u_above: Optional[float] = None
    stuck_mode: Optional[str] = None
    stuck_mu_hat_below: Optional[float] = None
    stuck_dmu_hat_below: Optional[float] = None
    stuck_b_mu_below: Optional[float] = None
    stuck_db_mu_below: Optional[float] = None
    stuck_consecutive: Optional[int] = None
    if isinstance(stuck_cfg, dict):
        stuck_u_above = require_unit_interval_closed_open(
            require(stuck_cfg, "u_above", "diagnostics.stuck"),
            "diagnostics.stuck.u_above",
        )
        stuck_consecutive = int(require(stuck_cfg, "consecutive", "diagnostics.stuck"))

        if "b_mu_below" in stuck_cfg or "db_mu_below" in stuck_cfg:
            stuck_mode = "b_mu"
            stuck_b_mu_below = require_unit_interval_open_closed(
                require(stuck_cfg, "b_mu_below", "diagnostics.stuck"),
                "diagnostics.stuck.b_mu_below",
            )
            stuck_db_mu_below = float(require(stuck_cfg, "db_mu_below", "diagnostics.stuck"))

            if stuck_db_mu_below < 0.0:
                raise ValueError("diagnostics.stuck.db_mu_below must be >= 0")
        else:
            stuck_mode = "mu_hat"
            stuck_mu_hat_below = require_unit_interval_open_closed(
                require(stuck_cfg, "mu_hat_below", "diagnostics.stuck"),
                "diagnostics.stuck.mu_hat_below",
            )
            stuck_dmu_hat_below = float(require(stuck_cfg, "dmu_hat_below", "diagnostics.stuck"))

            if stuck_dmu_hat_below < 0.0:
                raise ValueError("diagnostics.stuck.dmu_hat_below must be >= 0")
        if stuck_consecutive <= 0:
            raise ValueError("diagnostics.stuck.consecutive must be > 0")

    stress_cfg = config.get("stress", {})

    topology_actuator_cfg = config.get("topology_actuator", {})
    if topology_actuator_cfg is None:
        topology_actuator_cfg = {}
    if not isinstance(topology_actuator_cfg, dict):
        raise ValueError("topology_actuator must be an object")

    top_act_enabled = bool(topology_actuator_cfg.get("enabled", True)) if topology_actuator_cfg else False
    top_act_every_steps: Optional[int] = None
    top_act_cooldown_steps = 0
    if topology_actuator_cfg and top_act_enabled:
        top_act_every_steps = int(require(topology_actuator_cfg, "every_steps", "topology_actuator"))
        if top_act_every_steps <= 0:
            raise ValueError("topology_actuator.every_steps must be > 0")
        top_act_cooldown_steps = int(topology_actuator_cfg.get("cooldown_steps", 0))
        if top_act_cooldown_steps < 0:
            raise ValueError("topology_actuator.cooldown_steps must be >= 0")

    g = _build_graph(graph_cfg, rng=rng_graph)
    p_coop = float(require(init_cfg, "p_coop", "init"))
    states = init_states(g.n, p_coop=p_coop, rng=rng_init)

    detector = DirichletVolatilityDetector(det_cfg, rng=rng_detector, risk_rng=rng_detector_risk)
    controller = SmoothThermostatController(ctrl_cfg)

    truth_cfg = config.get("truth", {})
    truth_alpha_v = float(truth_cfg.get("alpha_v", det_cfg.get("alpha_v", 0.2)))

    mu_prev: Optional[float] = None
    eta_prev: Optional[float] = None
    v_mu_true = 0.0
    v_eta_true = 0.0

    series: List[SeriesRow] = []
    u_series: List[float] = []

    preds_b: List[float] = []
    labels_e: List[int] = []

    preds_mu: List[float] = []
    labels_mu: List[int] = []
    preds_eta: List[float] = []
    labels_eta: List[int] = []
    preds_v: List[float] = []
    labels_v: List[int] = []

    preds_disp: List[float] = []
    labels_disp: List[int] = []

    in_e_flags: List[bool] = []

    constraint_keys = ["mu", "eta", "v"]
    if mu_disp_max is not None:
        constraint_keys.append("mu_disp")

    binding_hat_series: List[str] = []
    binding_truth_series: List[str] = []
    violated_hat_counts = {k: 0 for k in constraint_keys}
    violated_truth_counts = {k: 0 for k in constraint_keys}

    mu_hat_prev: Optional[float] = None
    b_mu_prev: Optional[float] = None
    stuck_streak = 0
    max_stuck_streak = 0
    stuck_flags: List[bool] = []

    k_series: List[int] = []
    truth_mu_disp_series: List[float] = []

    comm_by_node: List[int] = [int(i) for i in range(g.n)]
    comms: List[List[int]] = [[int(i)] for i in range(g.n)]

    events: List[Event] = []

    top_act_rewired_total = 0
    top_act_budget_total = 0
    top_act_event_count = 0
    top_act_steps_active = 0
    top_act_steps_suppressed_cooldown = 0
    top_act_cooldown_until = 0

    for t in range(control_steps):
        step_events: List[Event] = []
        if stress_cfg:
            step_events.extend(apply_shocks(states, g, stress_cfg=stress_cfg, t=t, rng=rng_stress))
            step_events.extend(apply_rewires(g, stress_cfg=stress_cfg, t=t, rng=rng_stress))
        if step_events:
            events.extend(step_events)

        if mu_disp_max is not None:
            if (t % int(comm_every_steps)) == 0:
                comm_by_node, comms = detect_communities(g, cfg=comm_detect_cfg, rng=rng_comm)
            truth_comm_counts = community_edge_counts(states, g.edges, community_by_node=comm_by_node, k=len(comms))
            mu_disp_mode = parse_mu_disp_mode(ethics_cfg)

            if mu_disp_mode == "risk":
                truth_mus: List[float] = []
                truth_risks: List[float] = []
                for c in truth_comm_counts:
                    if float(c.total) < float(comm_min_total_edges):
                        continue
                    mu_c = float(edge_metrics_from_counts(c).mu)
                    truth_mus.append(float(mu_c))
                    truth_risks.append(1.0 if float(mu_c) < float(mu_min) else 0.0)
                truth_mu_by_comm = list(truth_mus)
                truth_mu_disp = float(mu_dispersion(truth_risks))
            else:
                truth_mu_by_comm, truth_mu_disp = mu_by_comm_and_dispersion_from_counts(
                    truth_comm_counts, min_total_edges=float(comm_min_total_edges)
                )

            k_series.append(int(len(comms)))
            truth_mu_disp_series.append(float(truth_mu_disp))
        else:
            truth_mu_by_comm = []
            truth_mu_disp = 0.0

        truth_counts = edge_counts(states, g.edges)
        truth_m = edge_metrics_from_counts(truth_counts)

        if mu_prev is None:
            dmu = 0.0
            deta = 0.0
        else:
            dmu = abs(float(truth_m.mu) - float(mu_prev))
            deta = abs(float(truth_m.eta) - float(eta_prev))

        v_mu_true = (1.0 - truth_alpha_v) * v_mu_true + truth_alpha_v * dmu
        v_eta_true = (1.0 - truth_alpha_v) * v_eta_true + truth_alpha_v * deta
        v_true = float(v_mu_true if v_mu_true >= v_eta_true else v_eta_true)

        y_all, y_mu, y_eta, y_v, y_disp = _truth_labels(
            truth_m.mu, truth_m.eta, v_true, ethics_cfg=ethics_cfg, mu_disp=float(truth_mu_disp)
        )
        in_e = True if y_all == 1 else False

        obs_counts = observed_counts(states, g.edges, obs_cfg=obs_cfg, rng=rng_obs)
        comm_obs = None
        if mu_disp_max is not None:
            comm_obs = observed_community_counts(
                states,
                g.edges,
                community_by_node=comm_by_node,
                k=len(comms),
                obs_cfg=obs_cfg,
                rng=rng_obs,
            )
        det_out = detector.step(
            obs_counts,
            ethics_cfg=ethics_cfg,
            comm_obs=comm_obs,
            comm_min_total_edges=float(comm_min_total_edges),
        )
        ctrl_out = controller.step(det_out.b)

        topo_act_events: List[TopologyActuatorEvent] = []
        topo_act_suppressed_cooldown = False
        topo_act_cooldown_remaining = 0

        if topology_actuator_cfg and top_act_enabled:
            if top_act_cooldown_steps > 0 and int(t) < int(top_act_cooldown_until):
                topo_act_suppressed_cooldown = True
                topo_act_cooldown_remaining = int(top_act_cooldown_until) - int(t)
                if top_act_every_steps is not None and (int(t) % int(top_act_every_steps)) == 0:
                    top_act_steps_suppressed_cooldown += 1
            else:
                topo_act_events = apply_topology_actuator(
                    g,
                    t=int(t),
                    u=float(ctrl_out.u),
                    ethics_cfg=ethics_cfg,
                    meso_cfg=meso_cfg if isinstance(meso_cfg, dict) else {},
                    comms=comms,
                    comm_obs=comm_obs,
                    meso_alpha0=list(detector.meso_alpha0) if hasattr(detector, "meso_alpha0") else None,
                    cfg=topology_actuator_cfg,
                    rng=rng_topo,
                    risk_rng=rng_topo_risk,
                )
                if topo_act_events and top_act_cooldown_steps > 0:
                    top_act_cooldown_until = int(t) + int(top_act_cooldown_steps) + 1
        if topo_act_events:
            step_events.extend(topo_act_events)
            events.extend(topo_act_events)
            top_act_steps_active += 1
            top_act_event_count += int(len(topo_act_events))
            for ev in topo_act_events:
                top_act_rewired_total += int(ev["rewired"])
                top_act_budget_total += int(ev["budget"])

        topo_ts: Optional[TopologyActuatorTimeseries] = None
        if topology_actuator_cfg:
            if topo_act_events:
                ev0 = topo_act_events[0]
                topo_ts = {
                    "active": True,
                    "rewired": int(sum(int(e["rewired"]) for e in topo_act_events)),
                    "budget": int(sum(int(e["budget"]) for e in topo_act_events)),
                    "drive": float(ev0["drive"]),
                    "mu_disp_hat": float(ev0["mu_disp_hat"]),
                    "mu_disp_max": float(ev0["mu_disp_max"]),
                    "mode": str(ev0["mode"]),
                }
            else:
                if topo_act_suppressed_cooldown:
                    topo_ts = {"active": False, "cooldown_remaining": int(topo_act_cooldown_remaining)}
                else:
                    topo_ts = {"active": False}

        margins_hat: Dict[str, float] = {
            "mu": float(det_out.mu_hat) - float(mu_min),
            "eta": float(eta_max) - float(det_out.eta_hat),
            "v": float(v_max) - float(det_out.v_hat),
        }
        if mu_disp_max is not None:
            margins_hat["mu_disp"] = float(mu_disp_max) - float(det_out.mu_disp_hat)
        binding_hat = str(min(margins_hat, key=margins_hat.get))
        violated_hat = [str(k) for (k, v) in margins_hat.items() if float(v) < 0.0]

        margins_truth: Dict[str, float] = {
            "mu": float(truth_m.mu) - float(mu_min),
            "eta": float(eta_max) - float(truth_m.eta),
            "v": float(v_max) - float(v_true),
        }
        if mu_disp_max is not None:
            margins_truth["mu_disp"] = float(mu_disp_max) - float(truth_mu_disp)
        binding_truth = str(min(margins_truth, key=margins_truth.get))
        violated_truth = [str(k) for (k, v) in margins_truth.items() if float(v) < 0.0]

        binding_hat_series.append(binding_hat)
        binding_truth_series.append(binding_truth)
        for k in violated_hat:
            violated_hat_counts[k] = int(violated_hat_counts.get(k, 0)) + 1
        for k in violated_truth:
            violated_truth_counts[k] = int(violated_truth_counts.get(k, 0)) + 1

        dmu_hat = 0.0 if mu_hat_prev is None else abs(float(det_out.mu_hat) - float(mu_hat_prev))
        db_mu = 0.0 if b_mu_prev is None else abs(float(det_out.b_mu) - float(b_mu_prev))
        stuck_condition = None
        stuck = None
        if stuck_consecutive is not None:
            u_c = clamp(float(ctrl_out.u), 0.0, 1.0)
            if str(stuck_mode) == "b_mu":
                b_mu_c = clamp(float(det_out.b_mu), 0.0, 1.0)
                stuck_condition = bool(
                    u_c >= float(stuck_u_above)
                    and b_mu_c <= float(stuck_b_mu_below)
                    and float(db_mu) <= float(stuck_db_mu_below)
                )
            else:
                mu_hat_c = clamp(float(det_out.mu_hat), 0.0, 1.0)
                stuck_condition = bool(
                    u_c >= float(stuck_u_above)
                    and mu_hat_c <= float(stuck_mu_hat_below)
                    and float(dmu_hat) <= float(stuck_dmu_hat_below)
                )
            if stuck_condition:
                stuck_streak += 1
            else:
                stuck_streak = 0
            if stuck_streak > max_stuck_streak:
                max_stuck_streak = stuck_streak
            stuck = bool(stuck_streak >= int(stuck_consecutive))
            stuck_flags.append(bool(stuck))
        mu_hat_prev = float(det_out.mu_hat)
        b_mu_prev = float(det_out.b_mu)

        dyn_params = dynamics_params_from_u(
            ctrl_out.u,
            mapping_cfg=mapping_cfg,
            dynamics_cfg=dyn_cfg,
            b_mu=det_out.b_mu,
            b_v=det_out.b_v,
            mu_hat=det_out.mu_hat,
        )
        run_window(states, g, params=dyn_params, rng=rng_dynamics, sweeps=window)

        series.append(
            {
                "t": int(t),
                "events": step_events,
                "graph": {"n": int(g.n), "m": int(len(g.edges))},
                "topology_actuator": topo_ts,
                "truth": {
                    "mu": float(truth_m.mu),
                    "eta": float(truth_m.eta),
                    "h": float(truth_m.h),
                    "v_mu": float(v_mu_true),
                    "v_eta": float(v_eta_true),
                    "v": float(v_true),
                    "mu_disp": float(truth_mu_disp),
                    "in_e": bool(in_e),
                },
                "obs": {"cc": float(obs_counts.cc), "cd": float(obs_counts.cd), "dd": float(obs_counts.dd)},
                "detector": {
                    "b": float(det_out.b),
                    "risk": float(det_out.risk),
                    "b_mu": float(det_out.b_mu),
                    "b_eta": float(det_out.b_eta),
                    "b_v": float(det_out.b_v),
                    "b_disp": float(det_out.b_disp),
                    "mu_hat": float(det_out.mu_hat),
                    "eta_hat": float(det_out.eta_hat),
                    "h_hat": float(det_out.h_hat),
                    "mu_disp_hat": float(det_out.mu_disp_hat),
                    "v_mu_hat": float(det_out.v_mu_hat),
                    "v_eta_hat": float(det_out.v_eta_hat),
                    "v_hat": float(det_out.v_hat),
                },
                "controller": {
                    "u": float(ctrl_out.u),
                    "u_target": float(ctrl_out.u_target),
                    "b_eff": float(ctrl_out.b_eff),
                    "risk_eff": float(ctrl_out.risk_eff),
                },
                "dynamics": {"beta": float(dyn_params.beta), "p_flip": float(dyn_params.p_flip), "theta": float(dyn_params.theta)},
                "diagnostics": {
                    "binding_hat": binding_hat,
                    "binding_truth": binding_truth,
                    "violated_hat": violated_hat,
                    "violated_truth": violated_truth,
                    "margins_hat": margins_hat,
                    "margins_truth": margins_truth,
                    "dmu_hat": float(dmu_hat),
                    "db_mu": float(db_mu),
                    "stuck_condition": stuck_condition,
                    "stuck": stuck,
                    "stuck_streak": int(stuck_streak) if stuck_consecutive is not None else None,
                },
                "meso": {
                    "k": int(len(comms)) if mu_disp_max is not None else None,
                    "min_total_edges": float(comm_min_total_edges) if mu_disp_max is not None else None,
                    "truth_mu_by_comm": truth_mu_by_comm if mu_disp_max is not None else None,
                },
            }
        )

        u_series.append(float(ctrl_out.u))
        preds_b.append(float(det_out.b))
        labels_e.append(int(y_all))

        preds_mu.append(float(det_out.b_mu))
        labels_mu.append(int(y_mu))
        preds_eta.append(float(det_out.b_eta))
        labels_eta.append(int(y_eta))
        preds_v.append(float(det_out.b_v))
        labels_v.append(int(y_v))

        if mu_disp_max is not None:
            preds_disp.append(float(det_out.b_disp))
            labels_disp.append(int(y_disp))

        in_e_flags.append(bool(in_e))

        mu_prev = float(truth_m.mu)
        eta_prev = float(truth_m.eta)

    shock_times = sorted({int(s.get("t")) for s in stress_cfg.get("shocks", []) if "t" in s})
    rewire_times = sorted({int(r.get("t")) for r in stress_cfg.get("rewires", []) if "t" in r})

    binding_hat_counts = {k: 0 for k in constraint_keys}
    binding_truth_counts = {k: 0 for k in constraint_keys}
    for b in binding_hat_series:
        binding_hat_counts[b] = int(binding_hat_counts.get(b, 0)) + 1
    for b in binding_truth_series:
        binding_truth_counts[b] = int(binding_truth_counts.get(b, 0)) + 1

    denom = float(control_steps if control_steps > 0 else 1)
    diagnostics_summary: Dict[str, Any] = {
        "binding_hat_counts": binding_hat_counts,
        "binding_truth_counts": binding_truth_counts,
        "binding_hat_fraction": {k: float(v) / denom for (k, v) in binding_hat_counts.items()},
        "binding_truth_fraction": {k: float(v) / denom for (k, v) in binding_truth_counts.items()},
        "violated_hat_counts": violated_hat_counts,
        "violated_truth_counts": violated_truth_counts,
        "violated_hat_fraction": {k: float(v) / denom for (k, v) in violated_hat_counts.items()},
        "violated_truth_fraction": {k: float(v) / denom for (k, v) in violated_truth_counts.items()},
    }

    meso_summary = None
    if mu_disp_max is not None:
        k_min = int(min(k_series)) if k_series else None
        k_max = int(max(k_series)) if k_series else None
        k_mean = float(sum(k_series) / float(len(k_series))) if k_series else None
        k_std = None
        if k_series and k_mean is not None:
            k_var = float(sum((float(x) - float(k_mean)) ** 2 for x in k_series)) / float(len(k_series))
            k_std = float(math.sqrt(k_var))

        mu_disp_max_truth = float(max(truth_mu_disp_series)) if truth_mu_disp_series else None
        mu_disp_nonzero_fraction = (
            float(sum(1 for x in truth_mu_disp_series if float(x) > 0.0)) / denom if truth_mu_disp_series else None
        )
        mu_disp_violation_fraction = (
            float(sum(1 for x in truth_mu_disp_series if float(x) > float(mu_disp_max))) / denom
            if truth_mu_disp_series
            else None
        )

        meso_summary = {
            "k": {"min": k_min, "max": k_max, "mean": k_mean, "std": k_std},
            "mu_disp": {
                "max_truth": mu_disp_max_truth,
                "nonzero_fraction": mu_disp_nonzero_fraction,
                "violation_fraction": mu_disp_violation_fraction,
                "mu_disp_max": float(mu_disp_max),
            },
        }

    topology_actuator_summary = None
    if topology_actuator_cfg:
        enabled = bool(topology_actuator_cfg.get("enabled", True))
        rewired_mean_per_event = (
            float(top_act_rewired_total) / float(top_act_event_count) if top_act_event_count > 0 else 0.0
        )
        topology_actuator_summary = {
            "enabled": bool(enabled),
            "cooldown_steps": int(top_act_cooldown_steps),
            "steps_suppressed_cooldown": int(top_act_steps_suppressed_cooldown),
            "fraction_steps_suppressed_cooldown": float(top_act_steps_suppressed_cooldown) / denom,
            "steps_active": int(top_act_steps_active),
            "fraction_steps_active": float(top_act_steps_active) / denom,
            "events": {
                "count": int(top_act_event_count),
                "rewired_total": int(top_act_rewired_total),
                "budget_total": int(top_act_budget_total),
                "rewired_mean_per_event": float(rewired_mean_per_event),
            },
        }

    if stuck_consecutive is not None:
        stuck_first_t = None
        for i, x in enumerate(stuck_flags):
            if x:
                stuck_first_t = int(i)
                break
        diagnostics_summary["stuck"] = {
            "fraction": float(sum(1 for x in stuck_flags if x)) / denom,
            "first_t": stuck_first_t,
            "max_streak": int(max_stuck_streak),
            "mode": str(stuck_mode),
            "u_above": float(stuck_u_above),
            "mu_hat_below": float(stuck_mu_hat_below) if stuck_mu_hat_below is not None else None,
            "dmu_hat_below": float(stuck_dmu_hat_below) if stuck_dmu_hat_below is not None else None,
            "b_mu_below": float(stuck_b_mu_below) if stuck_b_mu_below is not None else None,
            "db_mu_below": float(stuck_db_mu_below) if stuck_db_mu_below is not None else None,
            "consecutive": int(stuck_consecutive),
        }
        diagnostics_summary["u_saturation"] = {
            "above": float(stuck_u_above),
            "fraction": float(sum(1 for u in u_series if float(u) >= float(stuck_u_above))) / denom,
        }

    summary = {
        "control_steps": int(control_steps),
        "window": int(window),
        "graph": {
            "type": str(graph_cfg.get("type")),
            "n": int(g.n),
            "m": int(len(g.edges)),
            "avg_degree": float((2.0 * len(g.edges)) / g.n),
        },
        "time_in_e": float(sum(1 for x in in_e_flags if x) / float(control_steps if control_steps > 0 else 1)),
        "detector": {
            "brier": float(brier_score(preds_b, labels_e)),
            "log_loss": float(log_loss(preds_b, labels_e)),
            "brier_mu": float(brier_score(preds_mu, labels_mu)),
            "brier_eta": float(brier_score(preds_eta, labels_eta)),
            "brier_v": float(brier_score(preds_v, labels_v)),
            "brier_disp": float(brier_score(preds_disp, labels_disp)) if mu_disp_max is not None else None,
        },
        "intervention": summarize_u(u_series),
        "recovery": {
            "shocks": recovery_times(in_e_flags, shock_times),
            "rewires": recovery_times(in_e_flags, rewire_times),
        },
        "events": {"count": int(len(events))},
        "diagnostics": diagnostics_summary,
        "meso": meso_summary,
        "topology_actuator": topology_actuator_summary,
    }

    py_prov: PythonProvenance = {"version": sys.version.split()[0], "platform": platform.platform()}
    provenance: Provenance = {
        "seed": int(seed),
        "rng_streams": {str(k): int(v) for (k, v) in rng_seeds.items()},
        "rng_stream_method": "sha256(seed:stream)[:16]",
        "config_path": str(config_path),
        "config_text_sha256": sha256_hex(config_path.read_text()),
        "config_canonical_sha256": sha256_hex(canonical_json(config)),
        "python": py_prov,
    }

    out: RunOutput = {"provenance": provenance, "config": config, "summary": summary, "series": series}

    (out_dir / "run.json").write_text(json.dumps(out, indent=2, sort_keys=True))
    (out_dir / "summary.json").write_text(
        json.dumps({"provenance": provenance, "config": config, "summary": summary}, indent=2, sort_keys=True)
    )
    return out
