"""
first_reversal.py -- [2026-09-27] At the FIRST REVERSAL after a pick went
positive, do the eventual winners and losers look different?

First reversal = the first finished 1-min bar (after the buy, after price
has been above the buy price) that CLOSES at least 1 x 1-min ATR(14) below
the highest price since the buy. Measured at the end of that bar.
Winner / loser = price at 15:55 above / below the buy price.
Uses Alpaca's official 1-min bars (bars1m_*, bench_*) and the tick cache
(buy pressure, Lee-Ready: trade at/above ask = buy, at/below bid = sell).
Usage: python3 sim/analysis/first_reversal.py data/simulations/entry_rules_21d.json
"""
import json
import os
import statistics as st
import sys
from datetime import datetime, timedelta

sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from separation import auc

CACHE = "/var/www/screener/trade/data/simulations/cache"


def load_bars(day, name):
    return [{"t": datetime.fromisoformat(r[0]), "o": r[1], "h": r[2], "l": r[3], "c": r[4], "v": r[5]}
            for r in json.load(open(f"{CACHE}/{day}/{name}.json"))]


def atr(bars, n=14):
    a = None
    for p, b in zip(bars, bars[1:]):
        tr = max(b["h"] - b["l"], abs(b["h"] - p["c"]), abs(b["l"] - p["c"]))
        a = tr if a is None else a + (tr - a) / n
    return a


def _line_at(f, off):
    f.seek(off)
    if off:
        f.readline()
    line = f.readline()
    return line


def ticks_between(path, t0, t1):
    """Trades+quotes in [t0 - 5 s, t1] via binary search on the time-sorted file."""
    size = os.path.getsize(path)
    with open(path, "rb") as f:
        lo, hi = 0, size
        target = (t0 - timedelta(seconds=5)).isoformat().encode()
        while hi - lo > 4096:
            mid = (lo + hi) // 2
            line = _line_at(f, mid)
            if not line or line[2:2 + len(target)] >= target:
                hi = mid
            else:
                lo = mid
        f.seek(lo)
        if lo:
            f.readline()
        out = []
        for line in f:
            ts, kind, a, b = json.loads(line)
            t = datetime.fromisoformat(ts)
            if t > t1:
                break
            out.append((t, kind, a, b))
        return out


def imbalance(ev, t1, secs):
    q, buy, sell = None, 0.0, 0.0
    for t, kind, a, b in ev:
        if kind == "quote":
            q = (a, b)
        elif t >= t1 - timedelta(seconds=secs) and q:
            bid, ask = q
            if a >= ask:
                buy += b
            elif a <= bid:
                sell += b
            elif a > (bid + ask) / 2:
                buy += b
            elif a < (bid + ask) / 2:
                sell += b
    return (buy - sell) / (buy + sell) if buy + sell else None


def main():
    res = json.load(open(sys.argv[1]))
    rows, never = [], 0
    for day, x in sorted(res["days"].items()):
        bench = {s: load_bars(day, f"bench_{s}") for s in ("SPY", "IWM")}
        for tr in x["trades"]:
            sym, e = tr["symbol"], tr["entry_price"]
            et = datetime.fromisoformat(tr["entry_time"])
            bars = load_bars(day, f"bars1m_{sym}")
            eod = et.replace(hour=19, minute=55, second=0, microsecond=0)
            first_full = et.replace(second=0, microsecond=0) + timedelta(minutes=1)
            hi, went_up, rev = e, False, None
            for i, b in enumerate(bars):
                if b["t"] < first_full or b["t"] + timedelta(minutes=1) > eod:
                    continue
                hi = max(hi, b["h"])
                went_up = went_up or b["h"] > e
                a = atr(bars[max(0, i - 30):i + 1])
                if went_up and a and b["c"] <= hi - a:
                    rev = i
                    break
            if rev is None:
                never += 1
                continue
            b = bars[rev]
            T = b["t"] + timedelta(minutes=1)
            done = bars[:rev + 1]
            v = [z["v"] for z in done]
            vwap = sum((z["h"] + z["l"] + z["c"]) / 3 * z["v"] for z in done) / sum(v)
            since = [z for z in done if z["t"] >= first_full]
            upv = sum(z["v"] for z in since if z["c"] >= z["o"])
            dnv = sum(z["v"] for z in since if z["c"] < z["o"])
            a = atr(done[-31:])
            own5 = (done[-1]["c"] / done[-6]["c"] - 1) * 100 if len(done) >= 6 else None
            rs = []
            for s, bl in bench.items():
                bb = [z for z in bl if z["t"] + timedelta(minutes=1) <= T]
                if len(bb) >= 6 and own5 is not None:
                    rs.append(own5 - (bb[-1]["c"] / bb[-6]["c"] - 1) * 100)
            iwm = [z for z in bench["IWM"] if z["t"] + timedelta(minutes=1) <= T]
            ev = ticks_between(f"{CACHE}/{day}/{sym}.jsonl", T - timedelta(seconds=60), T)
            f = {
                "run_up_pct": (hi / e - 1) * 100,
                "minutes_since_buy": (T - et).total_seconds() / 60,
                "pullback_from_high_pct": (b["c"] / hi - 1) * 100,
                "pullback_in_atr": (hi - b["c"]) / a if a else None,
                "price_vs_buy_pct": (b["c"] / e - 1) * 100,
                "price_vs_stop_R": (b["c"] - tr["stop_price"]) / (e - tr["stop_price"]) if e > tr["stop_price"] else None,
                "vs_vwap_pct": (b["c"] / vwap - 1) * 100,
                "vs_open_pct": (b["c"] / bars[0]["o"] - 1) * 100,
                "vs_day_high_pct": (b["c"] / max(z["h"] for z in done) - 1) * 100,
                "reversal_bar_vol_x_avg": v[-1] / (sum(v[-21:-1]) / len(v[-21:-1])) if len(v) > 5 else None,
                "down_vol_vs_up_vol_since_buy": dnv / upv if upv else None,
                "vol_last3_vs_prev3": sum(v[-3:]) / sum(v[-6:-3]) if len(v) >= 6 and sum(v[-6:-3]) else None,
                "atr_1m_pct": a / b["c"] * 100 if a else None,
                "rel_strength_min_5bar": min(rs) if rs else None,
                "iwm_from_open_pct": (iwm[-1]["c"] / iwm[0]["o"] - 1) * 100 if iwm else None,
                "buy_pressure_30s": imbalance(ev, T, 30),
                "buy_pressure_60s": imbalance(ev, T, 60),
            }
            rows.append((day, sym, tr["pl_dollars"] > 0, f))
        print(day, "done", file=sys.stderr, flush=True)
    W = [f for *_, ok, f in rows if ok]
    L = [f for *_, ok, f in rows if not ok]
    print(f"{len(rows)} picks had a first reversal ({len(W)} ended up, {len(L)} ended down); {never} never did\n")
    days = sorted({d for d, *_ in rows})
    h1 = set(days[:len(days) // 2])
    print(f"{'at the first reversal':32}{'ended up':>10}{'ended down':>11}{'AUC':>6}{'1st half':>9}{'2nd half':>9}")
    out = []
    for k in rows[0][3]:
        a = [f[k] for f in W if f[k] is not None]
        b = [f[k] for f in L if f[k] is not None]
        a1 = [f[k] for d, s, ok, f in rows if ok and d in h1 and f[k] is not None]
        b1 = [f[k] for d, s, ok, f in rows if not ok and d in h1 and f[k] is not None]
        a2 = [f[k] for d, s, ok, f in rows if ok and d not in h1 and f[k] is not None]
        b2 = [f[k] for d, s, ok, f in rows if not ok and d not in h1 and f[k] is not None]
        A = auc(a, b)
        out.append((abs(A - .5), k, st.median(a), st.median(b), A, auc(a1, b1), auc(a2, b2)))
    for _, k, ma, mb, A, A1, A2 in sorted(out, reverse=True):
        flag = "  <- consistent" if (A1 - .5) * (A2 - .5) > 0 and min(abs(A1 - .5), abs(A2 - .5)) >= .05 else ""
        print(f"{k:32}{ma:10.3f}{mb:11.3f}{A:6.2f}{A1:9.2f}{A2:9.2f}{flag}")
    json.dump(rows, open(sys.argv[1].replace(".json", "_first_reversal.json"), "w"), default=str)


if __name__ == "__main__":
    main()
