"""
vwap_reclaim_study.py -- [2026-09-29] The user's reversal idea from AGEN 9/29: an unusual
stock spikes early (don't buy the spike), slides below VWAP, keep watching; buy as soon
as price crosses back above VWAP.

Every stock-day on the saved lists (bars/bars_<date>.json.gz, the 9:28/9:30 top 30),
regular session 1-min bars, looking back only:
  spike    : high of the first `spike_min` minutes >= open * (1 + spike_pct)
  pullback : after the spike high, a 1-min close >= dip_pct below VWAP
  trigger  : the first 1-min close back above VWAP after the pullback (confirm=2: two
             closes in a row), before 15:00; buy at the NEXT bar's open
  unusual  : optional, RVOL so far (vs the 14-day by-minute normal) >= min_rvol at the trigger
Outcomes per event: to 15:55; best / worst within 60 min; 1% stop vs 2% target which first;
stop under the pullback low (else hold to 15:55); vs every stock at the same minute (excess).
"""
import gzip
import json
import statistics as st
import sys
from datetime import datetime
from pathlib import Path
from zoneinfo import ZoneInfo

HERE = Path(__file__).resolve().parent
sys.path.insert(0, "/var/www/screener/trade")
import reference as REF
from entry_v2 import _r2_eff

ET = ZoneInfo("America/New_York")


def load_days():
    out = []
    for p in sorted(HERE.glob("bars/bars_2026-*.json.gz")):
        d = json.load(gzip.open(p, "rt"))
        day = d["date"]
        for s in d["top30_current"]:
            bars = [b for b in d["symbols"].get(s, {}).get("bars", [])
                    if "09:30" <= datetime.fromtimestamp(b[0], ET).strftime("%H:%M") < "16:00"]
            if len(bars) < 200:
                continue
            out.append((day, s, bars))
    return out


def minute(ts):
    t = datetime.fromtimestamp(ts, ET)
    return (t.hour - 9) * 60 + t.minute - 30


def events(days, spike_pct, dip_pct, confirm, spike_min=30, min_rvol=None, min_trigger_m=0, min_below=0):
    ev = []
    for day, s, bars in days:
        o = bars[0][1]
        ref = REF.load(day, s) if min_rvol else None
        if min_rvol and not ref:
            continue
        pv = v = 0.0
        hi, hi_i, armed, dipped, run, low_after, below = 0, None, False, False, 0, None, 0
        for i, b in enumerate(bars):
            m = minute(b[0])
            pv += (b[2] + b[3] + b[4]) / 3 * b[5]
            v += b[5]
            vwap = pv / v if v else b[4]
            if m < spike_min and b[2] > hi:
                hi, hi_i = b[2], i
            if not armed and m >= 1 and hi >= o * (1 + spike_pct / 100):
                armed = True
            if not armed:
                continue
            if not dipped:
                if i > hi_i and b[4] <= vwap * (1 - dip_pct / 100):
                    dipped, low_after, low_i = True, b[3], i
                continue
            if b[3] < low_after:
                low_after, low_i = b[3], i
            below = below + 1 if b[4] <= vwap else below
            run = run + 1 if b[4] > vwap else 0
            if run >= confirm and (m < min_trigger_m or below < min_below):
                run = 0 if m < min_trigger_m else run
                if m < min_trigger_m or below < min_below:
                    continue
            if run >= confirm:
                if m >= 330 or i + 1 >= len(bars):
                    break
                if min_rvol:
                    cum = sum(x[5] for x in bars[:i + 1])
                    if cum / max(ref["cum_vol"][min(m, 389)], 1) < min_rvol:
                        break
                e = bars[i + 1][1]
                after = [x for x in bars[i + 1:] if minute(x[0]) <= 385]
                close = after[-1][4]
                w60 = [x for x in after if minute(x[0]) <= m + 61]
                first = None
                for x in after:
                    if x[3] <= e * 0.99:
                        first = "stop"
                        break
                    if x[2] >= e * 1.02:
                        first = "target"
                        break
                stop_pl = None
                for x in after:
                    if x[3] <= low_after - 0.01:
                        stop_pl = (low_after - 0.01) / e - 1
                        break
                ref2 = ref or REF.load(day, s) or {}
                levels = [("open", o)]
                if ref2.get("prev_high"):
                    levels += [("prev high", ref2["prev_high"]), ("prev close", ref2["prev_close"])]
                if ref2.get("sr14"):
                    levels += [("14d high", ref2["sr14"]["high"])]
                    levels += [("14d band", lv["hi"]) for lv in ref2["sr14"]["levels"]]
                    levels += [("14d band", lv["lo"]) for lv in ref2["sr14"]["levels"]]
                # the pullback low "holds" a level: the low sits within 0-1% above it
                # (or at most 0.3% under it -- a wick)
                held = [n for n, lv in levels if lv * 0.997 <= low_after <= lv * 1.01]
                r2, eff = _r2_eff([x[4] for x in bars[max(0, i - 14):i + 1]])
                ev.append({"day": day, "sym": s, "m": m + 1, "entry": e, "open": o, "spike_hi": hi,
                           "held": held, "base_min": i - low_i, "r2": r2, "eff": eff,
                           "low": low_after, "to_close": (close / e - 1) * 100,
                           "best60": (max(x[2] for x in w60) / e - 1) * 100,
                           "worst60": (min(x[3] for x in w60) / e - 1) * 100,
                           "first": first or "neither",
                           "lowstop": (stop_pl if stop_pl is not None else close / e - 1) * 100,
                           "stop_dist": (1 - (low_after - 0.01) / e) * 100})
                break
    return ev


def baseline(days):
    """Average to-close return of every stock at each minute (for the excess)."""
    acc = {}
    for day, s, bars in days:
        close = [x for x in bars if minute(x[0]) <= 385][-1][4]
        for b in bars:
            acc.setdefault(minute(b[0]), []).append((close / b[1] - 1) * 100)
    return {m: st.mean(v) for m, v in acc.items()}


def summary(ev, base, label):
    if not ev:
        print(f"{label:48} n=0")
        return
    h1 = sorted({e["day"] for e in ev})
    half = set(h1[:len(h1) // 2])
    tc = [e["to_close"] for e in ev]
    ex = [e["to_close"] - base.get(e["m"], 0) for e in ev]
    ex1 = [x for x, e in zip(ex, ev) if e["day"] in half]
    ex2 = [x for x, e in zip(ex, ev) if e["day"] not in half]
    print(f"{label:48} n={len(ev):3}  to close {st.mean(tc):+5.2f}% (med {st.median(tc):+5.2f}, up {sum(x > 0 for x in tc) / len(tc) * 100:3.0f}%)"
          f"  excess {st.mean(ex):+5.2f}% [halves {st.mean(ex1) if ex1 else 0:+.2f}/{st.mean(ex2) if ex2 else 0:+.2f}]"
          f"  best60 {st.mean(e['best60'] for e in ev):+4.1f}  worst60 {st.mean(e['worst60'] for e in ev):+4.1f}"
          f"  +2% before -1%: {sum(e['first'] == 'target' for e in ev)}/{sum(e['first'] == 'stop' for e in ev)}"
          f"  low-stop {st.mean(e['lowstop'] for e in ev):+5.2f}% (stop {st.median(e['stop_dist'] for e in ev):.1f}% away)")


def filters_test(days, base):
    for sp, dip in ((2, 0.5), (2, 1.0), (3, 1.0)):
        pool = events(days, sp, dip, 1, min_trigger_m=15, min_below=10)
        print(f"\n=== pool: spike>={sp}% dip>={dip}% after 9:45, >=10 min below VWAP ===")
        summary(pool, base, "all reclaims (no extra filter)")
        F = {
            "1 low holds a level (any)": lambda e: bool(e["held"]),
            "1 low holds prev high / 14d high": lambda e: any(h in ("prev high", "14d high") for h in e["held"]),
            "1 low holds the open": lambda e: "open" in e["held"],
            "2 base: no new low for >=20 min": lambda e: e["base_min"] >= 20,
            "2 base: no new low for >=30 min": lambda e: e["base_min"] >= 30,
            "3 steady reclaim (R2>=0.7, eff>=0.5)": lambda e: e["r2"] >= 0.7 and e["eff"] >= 0.5,
            "1+2 level + base>=20": lambda e: bool(e["held"]) and e["base_min"] >= 20,
            "1+3 level + steady": lambda e: bool(e["held"]) and e["r2"] >= 0.7 and e["eff"] >= 0.5,
            "2+3 base>=20 + steady": lambda e: e["base_min"] >= 20 and e["r2"] >= 0.7 and e["eff"] >= 0.5,
            "1+2+3 all three": lambda e: bool(e["held"]) and e["base_min"] >= 20 and e["r2"] >= 0.7 and e["eff"] >= 0.5,
            "NOT 1 (low holds no level)": lambda e: not e["held"],
        }
        for name, fn in F.items():
            summary([e for e in pool if fn(e)], base, name)
    pool = events(days, 2, 0.5, 1, min_trigger_m=15, min_below=10)
    best = [e for e in pool if e["base_min"] >= 20 and e["r2"] >= 0.7 and e["eff"] >= 0.5]
    print("\nbase>=20 + steady (spike>=2% dip>=0.5%) -- every event:")
    for e in best:
        print(f"  {e['day']} {e['sym']:5} {9 + (30 + e['m']) // 60}:{(30 + e['m']) % 60:02d} buy {e['entry']:.2f} (open {e['open']:.2f}, low {e['low']:.2f}, base {e['base_min']}m, R2 {e['r2']:.2f} eff {e['eff']:.2f}, holds {','.join(e['held']) or '-'})"
              f"  to close {e['to_close']:+5.1f}%  best60 {e['best60']:+4.1f}  worst60 {e['worst60']:+4.1f}  {e['first']}")


if __name__ == "__main__" and len(sys.argv) > 1 and sys.argv[1] == "filters":
    days = load_days()
    filters_test(days, baseline(days))
elif __name__ == "__main__":
    days = load_days()
    base = baseline(days)
    print(f"{len(days)} stock-days, {len({d for d, _, _ in days})} days\n")
    for sp in (2, 3, 5):
        for dip in (0.5, 1.0, 1.5):
            for cf in (1, 2):
                summary(events(days, sp, dip, cf), base, f"spike>={sp}% dip>={dip}% confirm {cf}")
        print()
    print("with unusual volume (RVOL so far >= 1.5 at the trigger):")
    for sp in (2, 3):
        for dip in (0.5, 1.0):
            summary(events(days, sp, dip, 1, min_rvol=1.5), base, f"spike>={sp}% dip>={dip}% confirm 1 rvol1.5")
    print("\nlike AGEN: buy after 9:45, at least N minutes spent below VWAP first:")
    for sp in (2, 3):
        for dip in (0.5, 1.0):
            for mb in (10, 20, 30):
                summary(events(days, sp, dip, 1, min_trigger_m=15, min_below=mb), base, f"spike>={sp}% dip>={dip}% >=9:45 below>={mb}m")
    for mb in (10, 30):
        summary(events(days, 2, 1.0, 1, min_rvol=1.5, min_trigger_m=15, min_below=mb), base, f"spike>=2% dip>=1% >=9:45 below>={mb}m rvol1.5")
    ev = events(days, 2, 1.0, 1, min_trigger_m=15, min_below=30)
    print("\nspike>=2% dip>=1% after 9:45, >=30 min below VWAP -- every event:")
    for e in ev:
        print(f"  {e['day']} {e['sym']:5} {9 + (30 + e['m']) // 60}:{(30 + e['m']) % 60:02d} buy {e['entry']:.2f} (open {e['open']:.2f}, spike {e['spike_hi']:.2f}, low {e['low']:.2f})"
              f"  to close {e['to_close']:+5.1f}%  best60 {e['best60']:+4.1f}  worst60 {e['worst60']:+4.1f}  {e['first']}")
