#!/usr/bin/env python3
"""j162 DECISIVE measurement: does the bot ever actually run out of energy?

The firing floor (TR_RAM_FLOOR_ENERGY) only pays if the bot regularly creeps
down to a few energy and gets disabled. This answers that from the ALREADY
RECORDED closed-loop corpus, state only:

  A) self energy AT DEATH (the reserve we actually held when the killing blow
     landed) -- the floor's entire claim
  B) how long we stay at energy <= 0 (isDisabled) before the round ends
  C) recovery: how often self energy RISES tick-over-tick, and from what level
     (the only refill in the game is +3*power per landed bullet hit, so a rise
     is a landed hit -- this is "can we climb back out by shooting")
  D) what a floor at {3,5,10,20} would cost: % ticks suppressed, run length, and
     the heat-limited ceiling on how much energy it could possibly save

NO battle, NO server, NO counterfactual replay, NO damage estimate (the offline
harness scored 0/6 on closed-loop questions, docs/offline_harness_trust.md).

Usage:  python3 common_libs/tests/measure_ram_exhaustion [glob-dir]
"""
import glob, json, os, statistics, sys
from array import array
from multiprocessing import Pool

ROOTS = sys.argv[1:] or ["/tmp"]

# Tank Royale gun heat: heat += 1 + power/5 and the gun cools 0.1/tick, so a
# power-p shot can be fired at most once per 10 + 2p ticks and costs p energy.
# The cost bracket is therefore p/(10+2p) energy per tick, from 0.0098 at the
# cheapest legal shot (0.1) to 0.1875 at the most expensive (3.0).
def per_tick(power):
    return power / (10.0 + 2.0 * power)


def num(line, key):
    i = line.find('"' + key + '":')
    if i < 0:
        return None
    i += len(key) + 3
    j = line.find(',', i)
    if j < 0:
        j = line.find('}', i)
    try:
        return float(line[i:j])
    except ValueError:
        return None


def load(path):
    """[(self, enemy)] per tick, with round boundaries from the round map."""
    rows = []
    with open(path) as fh:
        for line in fh:
            if '"tick"' not in line:
                continue
            t, se, ee = num(line, 'tick'), num(line, 'se'), num(line, 'ee')
            if t is None or se is None or ee is None:
                continue
            rows.append((t, se, ee))
    if not rows:
        return []
    rf = path.replace(".jsonl", ".jsonl.rounds.json")
    bounds = []
    if os.path.exists(rf):
        try:
            for r in json.load(open(rf))["rounds"]:
                bounds.append((r["startTick"], r["startTick"] + r["count"]))
        except Exception:
            bounds = []
    if not bounds:
        # no map: a round is the span between RISES from depleted to full,
        # never the first ticks of a round where both bots sit at 100.
        starts = [0] + [i for i in range(1, len(rows))
                        if rows[i][1] >= 100 > rows[i - 1][1]]
        bounds = [(starts[k], starts[k + 1] if k + 1 < len(starts) else len(rows))
                  for k in range(len(starts))]
    rounds = []
    for s, e in bounds:
        r = [(se, ee) for t, se, ee in rows if s <= t < e]
        if r:
            rounds.append(r)
    return rounds


def corpus():
    files = []
    for root in ROOTS:
        for f in glob.glob(os.path.join(root, "**", "*.jsonl"), recursive=True):
            if f.endswith(".events.jsonl"):
                continue
            try:
                with open(f) as fh:
                    first = fh.readline()
            except OSError:
                continue
            if '"closed_loop":true' not in first.replace(" ", ""):
                continue
            files.append(f)
    out = []
    for r in Pool(8).imap(load, sorted(files), chunksize=32):
        out += r
    return sorted(files), out


def pct(sorted_x, q):
    if not sorted_x:
        return 0.0
    i = q * (len(sorted_x) - 1)
    lo, hi = int(i), min(int(i) + 1, len(sorted_x) - 1)
    return sorted_x[lo] + (sorted_x[hi] - sorted_x[lo]) * (i - lo)


def main():
    files, rounds = corpus()
    N = sum(len(r) for r in rounds)
    print(f"recordings={len(files)}  rounds={len(rounds)}  ticks={N}\n")

    # ---- A) how each round ends, and the reserve held at that moment --------
    self_dead = enemy_dead = both_dead = alive_end = 0
    last_alive = []          # self energy on the last tick we were alive
    death_tick = []          # self energy on the tick we crossed 0 (can be < 0)
    zero_runs = []           # ticks spent at self energy <= 0 before round end
    over = []                # reserve that would have absorbed the killing blow
    for r in rounds:
        sd = ed = None
        for i, (a, b) in enumerate(r):
            if sd is None and a <= 0:
                sd = i
            if ed is None and b <= 0:
                ed = i
            if sd is not None and ed is not None:
                break
        if sd is None and ed is None:
            alive_end += 1
            continue
        if sd is not None and ed is not None:
            both_dead += 1
        elif sd is not None:
            self_dead += 1
        else:
            enemy_dead += 1
        if sd is not None:
            last_alive.append(r[sd - 1][0] if sd > 0 else r[0][0])
            death_tick.append(r[sd][0])
            over.append(-r[sd][0])
            j = len(r)
            while j > sd and r[j - 1][0] <= 0:
                j -= 1
            zero_runs.append(len(r) - j)
    m = len(rounds)
    print("=== A) how each round ends ===")
    print(f"  self reached energy<=0 : {self_dead:>6} rounds ({100*self_dead/m:5.1f}%)")
    print(f"  only the enemy did     : {enemy_dead:>6} rounds ({100*enemy_dead/m:5.1f}%)")
    print(f"  both in the same round : {both_dead:>6} rounds ({100*both_dead/m:5.1f}%)")
    print(f"  neither (truncated)    : {alive_end:>6} rounds ({100*alive_end/m:5.1f}%)")

    print("\n=== B) SELF ENERGY AT DEATH (last value above 0 before the kill) ===")
    s = sorted(last_alive)
    if s:
        print(f"  n={len(s)}  min {s[0]:.2f}  p10 {pct(s,.10):.2f}  median {pct(s,.5):.2f}"
              f"  mean {statistics.fmean(s):.2f}  p90 {pct(s,.90):.2f}  max {s[-1]:.2f}")
        for t in (0, 1, 3, 5, 10, 20):
            c = sum(1 for x in s if x <= t)
            print(f"    <= {t:>2} energy: {c:>6} ({100*c/len(s):5.1f}% of self deaths,"
                  f" {100*c/m:5.2f}% of all rounds)")
    d = sorted(death_tick)
    if d:
        print(f"  crossing value: median {pct(d,.5):.2f}  p10 {pct(d,.10):.2f}"
              f"  p90 {pct(d,.90):.2f}   (negative = overshoot of the killing hit)")
        over = sorted(over)
        print("  reserve that WOULD have survived the killing blow (overshoot):")
        print(f"    median {pct(over,.5):.2f}  p75 {pct(over,.75):.2f}"
              f"  p90 {pct(over,.90):.2f}  p99 {pct(over,.99):.2f}  max {over[-1]:.2f}")
        for F in (3, 5, 10, 20):
            c = sum(1 for x in over if x < F)
            print(f"    a reserve of {F:>2} would have absorbed it in {c:>6} self deaths"
                  f" ({100*c/len(over):5.1f}%)")

    print("\n=== C) time spent at energy<=0 (isDisabled) before the round ends ===")
    z = sorted(zero_runs)
    if z:
        print(f"  ticks disabled: median {pct(z,.5):.0f}  p90 {pct(z,.9):.0f}"
              f"  max {z[-1]}   total {sum(z)} of {N} ticks"
              f" ({100*sum(z)/N:.4f}%)")

    # ---- D) recovery: energy RISES tick-over-tick = a landed bullet hit -----
    rises, pre = 0, []
    pre_low = {20: 0, 10: 0, 5: 0, 3: 0}
    tot_ticks = 0
    for r in rounds:
        for i in range(1, len(r)):
            tot_ticks += 1
            if r[i][0] - r[i - 1][0] > 0.01:
                rises += 1
                pre.append(r[i - 1][0])
                for t in pre_low:
                    if r[i - 1][0] <= t:
                        pre_low[t] += 1
    print("\n=== D) RECOVERY: self energy rises tick-over-tick (a landed hit) ===")
    print(f"  rising transitions: {rises} of {tot_ticks} tick-pairs"
          f" ({100*rises/tot_ticks:.3f}%), i.e. ~{rises/len(rounds):.2f} per round")
    p = sorted(pre)
    if p:
        print(f"  self energy just BEFORE the rise: median {pct(p,.5):.2f}"
              f"  p10 {pct(p,.10):.2f}  p90 {pct(p,.90):.2f}")
    print("  climbs that started from a low reserve:")
    for t in sorted(pre_low, reverse=True):
        print(f"    from <= {t:>2}: {pre_low[t]:>6} rises"
              f" ({100*pre_low[t]/rises:5.2f}% of rises)")

    # the decisive conditional: sitting low, do we climb back out or die?
    # "death" counts the ONE tick that crosses 0. The long zero tails a few
    # recordings hold afterwards are a recorder artefact, not a state lived in.
    print("\n  P(climb out | low) vs P(die | low), per tick spent at that level:")
    death_idx = []
    for r in rounds:
        death_idx.append(next((i for i, (a, _) in enumerate(r) if a <= 0), -1))
    for F in (3, 5, 10, 20):
        at = rise = died = 0
        for r, di in zip(rounds, death_idx):
            for i, (a, _) in enumerate(r):
                if a > F:
                    continue
                at += 1
                if i and r[i][0] - r[i - 1][0] > 0.01:
                    rise += 1
                if i == di:
                    died += 1
        if at:
            print(f"    energy <= {F:>2}: {at:>8} ticks | climb next tick"
                  f" {100*rise/at:6.3f}%  | killed on this tick {100*died/at:6.3f}%"
                  f"  -> dying is {died/max(1,rise):.1f}x more likely than recovering")

    # ---- E) what the floor would cost ---------------------------------------
    print("\n=== E) COST of TR_RAM_FLOOR_ENERGY: ticks where a new shot is blocked ===")
    print(f"{'floor':>5} {'%ticks':>7} {'rounds':>7} {'med run':>8} {'p90 run':>8}"
          f" {'max run':>8} {'energy saved, corpus (0.1..3.0 p)':>34}"
          f" {'per med run @1.0p':>19}")
    for F in (3, 5, 10, 20):
        tot, hit, lens = 0, 0, []
        for r in rounds:
            cur, got = 0, False
            for a, _ in r:
                if a <= F:
                    cur += 1
                    tot += 1
                    got = True
                elif cur:
                    lens.append(cur)
                    cur = 0
            if cur:
                lens.append(cur)
            hit += 1 if got else 0
        lens.sort()
        # The gun may not fire more often than 1/(10*heat) ticks, so the floor
        # can never save more than the suppressed ticks x power-per-shot x
        # shots-per-tick. Report the bracket: 0.1 power (cheapest legal shot) to
        # 3.0 power (most expensive legal shot).
        med = pct(lens, .5) if lens else 0
        lo, hi = tot * per_tick(0.1), tot * per_tick(3.0)
        mid = med * per_tick(1.0)
        print(f"{F:>5} {100*tot/N:>6.2f}% {hit:>7} {med:>8.0f} "
              f"{pct(lens,.9) if lens else 0:>8.0f} {lens[-1] if lens else 0:>8}"
              f"   {lo:>7.0f} .. {hi:>7.0f}   {mid:>6.2f}")
    print("  energy saved over the WHOLE corpus, heat-limited: the 0.1..3.0 power")
    print("  bracket, then the p=1.0 column = what one median suppressed RUN is worth")
    print("  (1.0 power is the mode of the measured landed-hit histogram).")
    print("  Median run lengths 34/53/89/139 ticks; one 1.0-power landed hit = 3.0.")
    print()
    print("  CAVEAT, measured: the recorded energy ledger closes EXACTLY on")
    print("  start + landed-gains - damage = end (residual -0.00 over 35065 rounds),")
    print("  i.e. these captures DO NOT charge the firepower cost. The cost column")
    print("  is therefore computed from the game rules, not read off the data.")


if __name__ == "__main__":
    main()
