import argparse
import json
from pathlib import Path
from typing import Any, Mapping, Optional, Sequence


def _is_int(x: Any) -> bool:
    return isinstance(x, int) and not isinstance(x, bool)


def _is_number(x: Any) -> bool:
    return isinstance(x, (int, float)) and not isinstance(x, bool)


def _require_obj(x: Any, ctx: str) -> Mapping[str, Any]:
    if not isinstance(x, dict):
        raise ValueError(f"{ctx} must be an object")
    return x


def _require_list(x: Any, ctx: str) -> Sequence[Any]:
    if not isinstance(x, list):
        raise ValueError(f"{ctx} must be a list")
    return x


def _require_key(d: Mapping[str, Any], key: str, ctx: str) -> Any:
    if key not in d:
        raise ValueError(f"Missing required key {ctx}.{key}")
    return d[key]


def _require_str(d: Mapping[str, Any], key: str, ctx: str) -> str:
    v = _require_key(d, key, ctx)
    if not isinstance(v, str):
        raise ValueError(f"{ctx}.{key} must be a string")
    return v


def _require_int(d: Mapping[str, Any], key: str, ctx: str) -> int:
    v = _require_key(d, key, ctx)
    if not _is_int(v):
        raise ValueError(f"{ctx}.{key} must be an int")
    return int(v)


def _require_number(d: Mapping[str, Any], key: str, ctx: str) -> float:
    v = _require_key(d, key, ctx)
    if not _is_number(v):
        raise ValueError(f"{ctx}.{key} must be a number")
    return float(v)


def _require_bool(d: Mapping[str, Any], key: str, ctx: str) -> bool:
    v = _require_key(d, key, ctx)
    if not isinstance(v, bool):
        raise ValueError(f"{ctx}.{key} must be a bool")
    return bool(v)


def _require_opt_bool(d: Mapping[str, Any], key: str, ctx: str) -> Optional[bool]:
    v = _require_key(d, key, ctx)
    if v is None:
        return None
    if not isinstance(v, bool):
        raise ValueError(f"{ctx}.{key} must be a bool or null")
    return bool(v)


def _require_opt_int(d: Mapping[str, Any], key: str, ctx: str) -> Optional[int]:
    v = _require_key(d, key, ctx)
    if v is None:
        return None
    if not _is_int(v):
        raise ValueError(f"{ctx}.{key} must be an int or null")
    return int(v)


def _require_opt_number(d: Mapping[str, Any], key: str, ctx: str) -> Optional[float]:
    v = _require_key(d, key, ctx)
    if v is None:
        return None
    if not _is_number(v):
        raise ValueError(f"{ctx}.{key} must be a number or null")
    return float(v)


def _validate_provenance(prov: Mapping[str, Any], ctx: str) -> None:
    _require_int(prov, "seed", ctx)
    streams = _require_key(prov, "rng_streams", ctx)
    if not isinstance(streams, dict):
        raise ValueError(f"{ctx}.rng_streams must be an object")
    for k, v in streams.items():
        if not isinstance(k, str):
            raise ValueError(f"{ctx}.rng_streams keys must be strings")
        if not _is_int(v):
            raise ValueError(f"{ctx}.rng_streams['{k}'] must be an int")

    _require_str(prov, "rng_stream_method", ctx)
    _require_str(prov, "config_path", ctx)
    _require_str(prov, "config_text_sha256", ctx)
    _require_str(prov, "config_canonical_sha256", ctx)

    py = _require_obj(_require_key(prov, "python", ctx), f"{ctx}.python")
    _require_str(py, "version", f"{ctx}.python")
    _require_str(py, "platform", f"{ctx}.python")


def _validate_event(ev: Mapping[str, Any], ctx: str) -> None:
    typ = _require_str(ev, "type", ctx)
    if typ == "shock":
        _require_int(ev, "t", ctx)
        _require_int(ev, "k", ctx)
        _require_str(ev, "mode", ctx)
        return

    if typ == "rewire":
        _require_int(ev, "t", ctx)
        _require_number(ev, "prob", ctx)
        _require_int(ev, "rewired", ctx)
        return

    if typ == "topology_actuator":
        _require_int(ev, "t", ctx)
        _require_str(ev, "mode", ctx)
        _require_int(ev, "rewired", ctx)
        _require_int(ev, "budget", ctx)
        _require_number(ev, "u", ctx)
        _require_number(ev, "drive", ctx)
        _require_number(ev, "mu_disp_hat", ctx)
        _require_number(ev, "mu_disp_max", ctx)
        _require_int(ev, "low_comm", ctx)
        _require_int(ev, "high_comm", ctx)
        _require_number(ev, "mu_low", ctx)
        _require_number(ev, "mu_high", ctx)
        _require_str(ev, "remove_strategy", ctx)
        return

    raise ValueError(f"{ctx}.type has unsupported value: {typ}")


def _validate_topology_actuator_ts(ts: Mapping[str, Any], ctx: str) -> None:
    active = _require_bool(ts, "active", ctx)
    if active:
        _require_int(ts, "rewired", ctx)
        _require_int(ts, "budget", ctx)
        _require_number(ts, "drive", ctx)
        _require_number(ts, "mu_disp_hat", ctx)
        _require_number(ts, "mu_disp_max", ctx)
        _require_str(ts, "mode", ctx)
        return

    if "cooldown_remaining" in ts and ts["cooldown_remaining"] is not None:
        if not _is_int(ts["cooldown_remaining"]):
            raise ValueError(f"{ctx}.cooldown_remaining must be an int")


def _validate_series_row(row: Mapping[str, Any], ctx: str) -> int:
    t = _require_int(row, "t", ctx)

    events = _require_list(_require_key(row, "events", ctx), f"{ctx}.events")
    for i, ev in enumerate(events):
        _validate_event(_require_obj(ev, f"{ctx}.events[{i}]"), f"{ctx}.events[{i}]")

    graph = _require_obj(_require_key(row, "graph", ctx), f"{ctx}.graph")
    _require_int(graph, "n", f"{ctx}.graph")
    _require_int(graph, "m", f"{ctx}.graph")

    topo_ts_raw = _require_key(row, "topology_actuator", ctx)
    if topo_ts_raw is not None:
        topo_ts = _require_obj(topo_ts_raw, f"{ctx}.topology_actuator")
        _validate_topology_actuator_ts(topo_ts, f"{ctx}.topology_actuator")

    truth = _require_obj(_require_key(row, "truth", ctx), f"{ctx}.truth")
    _require_number(truth, "mu", f"{ctx}.truth")
    _require_number(truth, "eta", f"{ctx}.truth")
    _require_number(truth, "h", f"{ctx}.truth")
    _require_number(truth, "v_mu", f"{ctx}.truth")
    _require_number(truth, "v_eta", f"{ctx}.truth")
    _require_number(truth, "v", f"{ctx}.truth")
    _require_number(truth, "mu_disp", f"{ctx}.truth")
    _require_bool(truth, "in_e", f"{ctx}.truth")

    obs = _require_obj(_require_key(row, "obs", ctx), f"{ctx}.obs")
    _require_number(obs, "cc", f"{ctx}.obs")
    _require_number(obs, "cd", f"{ctx}.obs")
    _require_number(obs, "dd", f"{ctx}.obs")

    det = _require_obj(_require_key(row, "detector", ctx), f"{ctx}.detector")
    for k in (
        "b",
        "risk",
        "b_mu",
        "b_eta",
        "b_v",
        "b_disp",
        "mu_hat",
        "eta_hat",
        "h_hat",
        "mu_disp_hat",
        "v_mu_hat",
        "v_eta_hat",
        "v_hat",
    ):
        _require_number(det, k, f"{ctx}.detector")

    ctrl = _require_obj(_require_key(row, "controller", ctx), f"{ctx}.controller")
    for k in ("u", "u_target", "b_eff", "risk_eff"):
        _require_number(ctrl, k, f"{ctx}.controller")

    dyn = _require_obj(_require_key(row, "dynamics", ctx), f"{ctx}.dynamics")
    for k in ("beta", "p_flip", "theta"):
        _require_number(dyn, k, f"{ctx}.dynamics")

    diag = _require_obj(_require_key(row, "diagnostics", ctx), f"{ctx}.diagnostics")
    _require_str(diag, "binding_hat", f"{ctx}.diagnostics")
    _require_str(diag, "binding_truth", f"{ctx}.diagnostics")
    for k in ("violated_hat", "violated_truth"):
        xs = _require_list(_require_key(diag, k, f"{ctx}.diagnostics"), f"{ctx}.diagnostics.{k}")
        for i, x in enumerate(xs):
            if not isinstance(x, str):
                raise ValueError(f"{ctx}.diagnostics.{k}[{i}] must be a string")

    for k in ("margins_hat", "margins_truth"):
        ms = _require_obj(_require_key(diag, k, f"{ctx}.diagnostics"), f"{ctx}.diagnostics.{k}")
        for kk, vv in ms.items():
            if not isinstance(kk, str) or not _is_number(vv):
                raise ValueError(f"{ctx}.diagnostics.{k} must be mapping[str, number]")

    _require_number(diag, "dmu_hat", f"{ctx}.diagnostics")
    _require_number(diag, "db_mu", f"{ctx}.diagnostics")
    _require_opt_bool(diag, "stuck_condition", f"{ctx}.diagnostics")
    _require_opt_bool(diag, "stuck", f"{ctx}.diagnostics")
    _require_opt_int(diag, "stuck_streak", f"{ctx}.diagnostics")

    meso = _require_obj(_require_key(row, "meso", ctx), f"{ctx}.meso")
    _require_opt_int(meso, "k", f"{ctx}.meso")
    _require_opt_number(meso, "min_total_edges", f"{ctx}.meso")
    mus = _require_key(meso, "truth_mu_by_comm", f"{ctx}.meso")
    if mus is not None:
        mus_list = _require_list(mus, f"{ctx}.meso.truth_mu_by_comm")
        for i, x in enumerate(mus_list):
            if not _is_number(x):
                raise ValueError(f"{ctx}.meso.truth_mu_by_comm[{i}] must be a number")

    return int(t)


def validate_out_dir(out_dir: Path) -> None:
    run_path = out_dir / "run.json"
    summary_path = out_dir / "summary.json"

    if not run_path.exists():
        raise FileNotFoundError(str(run_path))
    if not summary_path.exists():
        raise FileNotFoundError(str(summary_path))

    run = json.loads(run_path.read_text())
    summ = json.loads(summary_path.read_text())

    run_obj = _require_obj(run, "run.json")
    summ_obj = _require_obj(summ, "summary.json")

    _validate_provenance(_require_obj(_require_key(run_obj, "provenance", "run.json"), "run.json.provenance"), "run.json.provenance")
    _validate_provenance(
        _require_obj(_require_key(summ_obj, "provenance", "summary.json"), "summary.json.provenance"),
        "summary.json.provenance",
    )

    if "config" not in run_obj or "summary" not in run_obj or "series" not in run_obj:
        raise ValueError("run.json must have keys: provenance, config, summary, series")
    if "config" not in summ_obj or "summary" not in summ_obj:
        raise ValueError("summary.json must have keys: provenance, config, summary")

    summary = _require_obj(_require_key(run_obj, "summary", "run.json"), "run.json.summary")
    control_steps = int(_require_int(summary, "control_steps", "run.json.summary"))

    series = _require_list(_require_key(run_obj, "series", "run.json"), "run.json.series")
    if len(series) != control_steps:
        raise ValueError(f"run.json.series length ({len(series)}) != summary.control_steps ({control_steps})")

    total_events = 0
    prev_t: Optional[int] = None
    for i, row in enumerate(series):
        rr = _require_obj(row, f"run.json.series[{i}]")
        t = _validate_series_row(rr, f"run.json.series[{i}]")
        if prev_t is None and t != 0:
            raise ValueError("run.json.series[0].t must be 0")
        if prev_t is not None and t != prev_t + 1:
            raise ValueError(f"run.json.series[{i}].t must be previous+1")
        prev_t = t

        evs = rr.get("events", [])
        if isinstance(evs, list):
            total_events += len(evs)

    summary_events = _require_obj(_require_key(summary, "events", "run.json.summary"), "run.json.summary.events")
    events_count = int(_require_int(summary_events, "count", "run.json.summary.events"))
    if total_events != events_count:
        raise ValueError(f"summary.events.count ({events_count}) != sum(len(row.events)) ({total_events})")


def main() -> None:
    p = argparse.ArgumentParser()
    p.add_argument("--out_dir", required=True)
    args = p.parse_args()

    out_dir = Path(args.out_dir).expanduser().resolve()
    validate_out_dir(out_dir)
    print("OK")


if __name__ == "__main__":
    main()
