import random
from typing import Any, List, Mapping

from core import clamp, require
from graph import SimpleGraph, rewire_edges
from schema_types import RewireEvent, ShockEvent


def apply_shocks(states: List[int], g: SimpleGraph, stress_cfg: Mapping[str, Any], t: int, rng: random.Random) -> List[ShockEvent]:
    shocks = stress_cfg.get("shocks", [])
    if not shocks:
        return []

    applied: List[ShockEvent] = []
    for s in shocks:
        if int(require(s, "t", "stress.shocks[]")) != int(t):
            continue
        frac = float(require(s, "frac_flip", "stress.shocks[]"))
        mode = str(s.get("mode", "random")).lower()

        frac = clamp(frac, 0.0, 1.0)
        k = int(round(frac * g.n))
        if k <= 0:
            applied.append({"type": "shock", "t": int(t), "k": 0, "mode": mode})
            continue

        if mode == "random":
            idx = rng.sample(range(g.n), k=min(k, g.n))
        elif mode == "high_degree":
            deg = [(i, len(g.adj[i])) for i in range(g.n)]
            deg.sort(key=lambda x: x[1], reverse=True)
            idx = [i for i, _ in deg[: min(k, g.n)]]
        elif mode == "low_degree":
            deg = [(i, len(g.adj[i])) for i in range(g.n)]
            deg.sort(key=lambda x: x[1])
            idx = [i for i, _ in deg[: min(k, g.n)]]
        else:
            raise ValueError(f"Unsupported shock mode: {mode}")

        for i in idx:
            states[i] = 0 if states[i] else 1

        applied.append({"type": "shock", "t": int(t), "k": int(len(idx)), "mode": mode})

    return applied


def apply_rewires(g: SimpleGraph, stress_cfg: Mapping[str, Any], t: int, rng: random.Random) -> List[RewireEvent]:
    rewires = stress_cfg.get("rewires", [])
    if not rewires:
        return []

    max_tries = int(stress_cfg.get("rewire_max_tries", 25))
    applied: List[RewireEvent] = []
    for r in rewires:
        if int(require(r, "t", "stress.rewires[]")) != int(t):
            continue
        prob = float(require(r, "prob", "stress.rewires[]"))
        prob = clamp(prob, 0.0, 1.0)
        n = rewire_edges(g, prob=prob, rng=rng, max_tries=max_tries)
        applied.append({"type": "rewire", "t": int(t), "prob": float(prob), "rewired": int(n)})

    return applied
