"""
group_study.py -- [2026-09-29] The user's plan: look back over the saved days, separate the
stock-days into pattern groups, then inside each group compare the winners with the
losers and find the point where they start to differ.

Unit: every stock-day on the saved 9:28/9:30 lists (bars/ + flow/ + 14-day reference).
Group: the shape of the first 45 minutes, decided at 10:15 (known in real time):
  1 OPENING DRIVE   >= +2% over the open, above VWAP, within 1% of the day's high
  2 UP, PULLED BACK >= +1% over the open, above VWAP, more than 1% off the high
  3 V RECOVERY      fell >= 2% under the open, back above VWAP
  4 SPIKE & FADE    first-30-min high >= +2% over the open, now below VWAP
  5 SELLING SLIDE   <= -2% under the open, below VWAP
  6 WEAK DRIFT      below VWAP (the rest)
  7 FLAT / CHOP     above VWAP, the rest
Outcome after 10:15: winner = +2% or more from the 10:15 price to the close,
loser = -2% or worse, flat in between.
Divergence: for each measure, every 5 minutes from 10:15 to 12:30, AUC winners vs losers
(0.5 = no difference); the first minute with AUC >= 0.70 (or <= 0.30) that ALSO holds in
both halves of the days (>= 0.62 / <= 0.38 each). For each such point: how much of the
winners' final move they had already made by then.
"""
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")
sys.path.insert(0, "/var/www/screener/trade1/sim/analysis")
import reference as REF
from entry_v2 import _r2_eff
from separation import auc

ET = ZoneInfo("America/New_York")
T = 45                       # 10:15
ONLY = None                  # optional day filter (set by the caller)
MINUTES = list(range(T, 181, 5))   # 10:15 .. 12:30
GROUPS = ["1 OPENING DRIVE", "2 UP, PULLED BACK", "3 V RECOVERY", "4 SPIKE & FADE",
          "5 SELLING SLIDE", "6 WEAK DRIFT", "7 FLAT / CHOP"]
MEASURES = {
    "ret_since_1015": "price change since 10:15 %",
    "vwap_pct": "price vs VWAP %",
    "ret_open": "price vs open %",
    "off_high": "price vs day high %",
    "range_pos": "position in day range (0-1)",
    "ret15": "last 15 min change %",
    "buy15": "buy share, last 15 min",
    "buy_cum": "buy share since the open",
    "rvol": "RVOL so far (vs 14-day by minute)",
    "vol15x": "last 15 min volume vs normal",
    "r2": "smoothness (R2, 15 bars)",
    "eff": "straightness (15 bars)",
    "vs_prev_high": "price vs yesterday's high %",
}


def stock_days():
    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"]:
            fp = HERE / "flow" / day / f"{s}.json"
            ref = REF.load(day, s)
            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 or not ref:
                continue
            if ONLY and not ONLY(day):
                continue
            # [2026-09-30] history days have no buy/sell flow (no ticks) -- neutral 50/50
            fl = json.loads(fp.read_text()) if fp.exists() else {"buy": [1.0] * 390, "sell": [1.0] * 390}
            yield day, s, bars, fl, ref


def per_minute(bars, fl, ref):
    """Forward-filled per-minute series + the measures at each minute."""
    by = {}
    for b in bars:
        t = datetime.fromtimestamp(b[0], ET)
        by[(t.hour - 9) * 60 + t.minute - 30] = b
    o = bars[0][1]
    close, hi, lo, vol, pv = [], [], [], [], []
    last, h, l, cv, cpv = o, o, o, 0.0, 0.0
    for m in range(390):
        b = by.get(m)
        v = 0.0
        if b:
            last = b[4]
            h, l = max(h, b[2]), min(l, b[3])
            v = b[5]
            cv += v
            cpv += (b[2] + b[3] + b[4]) / 3 * v
        close.append(last)
        hi.append(h)
        lo.append(l)
        vol.append(v)
        pv.append(cpv / cv if cv else last)
    buy, sell = fl["buy"], fl["sell"]
    out = {}
    for m in MINUTES:
        c = close[m]
        bv, sv = sum(buy[m - 14:m + 1]), sum(sell[m - 14:m + 1])
        cb, cs = sum(buy[:m + 1]), sum(sell[:m + 1])
        norm15 = sum(ref["vol_per_min"][m - 14:m + 1])
        r2, eff = _r2_eff(close[m - 14:m + 1])
        out[m] = {
            "ret_since_1015": (c / close[T] - 1) * 100,
            "vwap_pct": (c / pv[m] - 1) * 100,
            "ret_open": (c / o - 1) * 100,
            "off_high": (c / hi[m] - 1) * 100,
            "range_pos": (c - lo[m]) / (hi[m] - lo[m]) if hi[m] > lo[m] else 0.5,
            "ret15": (c / close[m - 15] - 1) * 100,
            "buy15": bv / (bv + sv) if bv + sv else 0.5,
            "buy_cum": cb / (cb + cs) if cb + cs else 0.5,
            "rvol": sum(vol[:m + 1]) / max(ref["cum_vol"][m], 1),
            "vol15x": sum(vol[m - 14:m + 1]) / norm15 if norm15 else 1,
            "r2": r2,
            "eff": eff,
            "vs_prev_high": (c / ref["prev_high"] - 1) * 100 if ref.get("prev_high") else 0,
        }
    first30_hi = max(hi[:30])
    return {"open": o, "close": close, "hi": hi, "lo": lo, "vwap": pv, "m": out,
            "first30_hi_pct": (first30_hi / o - 1) * 100,
            "low_pct": (lo[T] / o - 1) * 100,
            "final": close[385]}


def group_of(x):
    f = x["m"][T]
    if f["ret_open"] >= 2 and f["vwap_pct"] >= 0 and f["off_high"] >= -1:
        return GROUPS[0]
    if f["ret_open"] >= 1 and f["vwap_pct"] >= 0:
        return GROUPS[1]
    if x["low_pct"] <= -2 and f["vwap_pct"] >= 0:
        return GROUPS[2]
    if x["first30_hi_pct"] >= 2 and f["vwap_pct"] < 0:
        return GROUPS[3]
    if f["ret_open"] <= -2 and f["vwap_pct"] < 0:
        return GROUPS[4]
    if f["vwap_pct"] < 0:
        return GROUPS[5]
    return GROUPS[6]


def main():
    rows = []
    for day, s, bars, fl, ref in stock_days():
        x = per_minute(bars, fl, ref)
        after = (x["final"] / x["close"][T] - 1) * 100
        rows.append({"day": day, "sym": s, "group": group_of(x), "after": after,
                     "oc": (x["final"] / x["open"] - 1) * 100,
                     "cls": "W" if after >= 2 else "L" if after <= -2 else "F", "x": x})
    days = sorted({r["day"] for r in rows})
    h1 = set(days[:len(days) // 2])
    print(f"{len(rows)} stock-days, {len(days)} days ({days[0]} .. {days[-1]}); halves split at {days[len(days) // 2]}\n")

    print(f"{'group at 10:15':20} {'n':>4} {'win':>4} {'flat':>4} {'lose':>4} {'avg 10:15->close':>17} {'med':>6} {'big (open->close >=5%)':>23}")
    for g in GROUPS:
        R = [r for r in rows if r["group"] == g]
        if not R:
            continue
        a = [r["after"] for r in R]
        print(f"{g:20} {len(R):4} {sum(r['cls'] == 'W' for r in R):4} {sum(r['cls'] == 'F' for r in R):4} "
              f"{sum(r['cls'] == 'L' for r in R):4} {st.mean(a):+16.2f}% {st.median(a):+5.2f}% {sum(r['oc'] >= 5 for r in R):23}")

    for g in GROUPS:
        R = [r for r in rows if r["group"] == g]
        W = [r for r in R if r["cls"] == "W"]
        L = [r for r in R if r["cls"] == "L"]
        print(f"\n=== {g}: {len(R)} stock-days, {len(W)} winners, {len(L)} losers ===")
        if len(W) < 6 or len(L) < 6:
            print("   too few winners or losers to compare")
            continue
        print(f"   winners: " + ", ".join(f"{r['sym']} {r['day'][5:]} {r['after']:+.0f}%" for r in sorted(W, key=lambda r: -r["after"])[:12]))
        print(f"   losers:  " + ", ".join(f"{r['sym']} {r['day'][5:]} {r['after']:+.0f}%" for r in sorted(L, key=lambda r: r["after"])[:12]))
        wfinal = st.median(r["after"] for r in W)
        print(f"   {'measure':34} {'at 10:15 AUC':>12}  {'first divergence (both halves)':34} {'winners / losers there':>24} {'winners move made':>18}")
        for k, name in MEASURES.items():
            a0 = auc([r["x"]["m"][T][k] for r in W], [r["x"]["m"][T][k] for r in L])
            found = None
            for m in MINUTES:
                w = [r["x"]["m"][m][k] for r in W]
                lo_ = [r["x"]["m"][m][k] for r in L]
                a = auc(w, lo_)
                a1 = auc([r["x"]["m"][m][k] for r in W if r["day"] in h1], [r["x"]["m"][m][k] for r in L if r["day"] in h1])
                a2 = auc([r["x"]["m"][m][k] for r in W if r["day"] not in h1], [r["x"]["m"][m][k] for r in L if r["day"] not in h1])
                if (a >= .70 and a1 >= .62 and a2 >= .62) or (a <= .30 and a1 <= .38 and a2 <= .38):
                    made = st.median(r["x"]["m"][m]["ret_since_1015"] for r in W)
                    found = (m, a, a1, a2, st.median(w), st.median(lo_), made)
                    break
            if found:
                m, a, a1, a2, mw, ml, made = found
                t = f"{9 + (30 + m) // 60}:{(30 + m) % 60:02d}"
                print(f"   {name:34} {a0:12.2f}  {t} AUC {a:.2f} ({a1:.2f}/{a2:.2f}){'':8} {mw:+11.2f} / {ml:+.2f} "
                      f"{made:+7.2f}% of {wfinal:+.1f}%")
            else:
                print(f"   {name:34} {a0:12.2f}  none by 12:30")


if __name__ == "__main__" and len(sys.argv) == 1:
    main()


def rule_check():
    """[2026-09-29] Each divergence as a rule: at its minute, stocks in the group on the
    winners' side of the midpoint vs the rest -- result FROM THAT MINUTE to the close."""
    rows = []
    for day, s, bars, fl, ref in stock_days():
        x = per_minute(bars, fl, ref)
        rows.append({"day": day, "sym": s, "group": group_of(x), "x": x})
    days = sorted({r["day"] for r in rows})
    h1 = set(days[:len(days) // 2])
    tests = [  # group, measure, minute, direction (+1: higher is the winners' side)
        (GROUPS[0], "off_high", 50, +1), (GROUPS[0], "eff", 60, +1), (GROUPS[0], "r2", 55, +1), (GROUPS[0], "vs_prev_high", 55, +1),
        (GROUPS[1], "r2", 45, +1), (GROUPS[1], "eff", 45, +1), (GROUPS[1], "rvol", 45, -1), (GROUPS[1], "off_high", 45, +1), (GROUPS[1], "ret15", 45, +1),
        (GROUPS[2], "rvol", 45, -1), (GROUPS[2], "vol15x", 45, -1), (GROUPS[2], "ret_open", 45, +1), (GROUPS[2], "r2", 55, +1),
        (GROUPS[3], "eff", 60, +1), (GROUPS[3], "ret15", 60, +1), (GROUPS[3], "r2", 60, +1),
        (GROUPS[4], "ret15", 60, +1), (GROUPS[4], "eff", 60, +1), (GROUPS[4], "buy15", 70, +1),
        (GROUPS[5], "ret15", 60, +1), (GROUPS[5], "buy15", 55, +1), (GROUPS[5], "eff", 60, +1),
    ]
    print(f"\n{'group':18} {'rule (at minute)':38} {'n sig':>5} {'sig -> close':>12} {'halves':>14} {'rest -> close':>13} {'sig up':>6}")
    for g, k, m, sgn in tests:
        R = [r for r in rows if r["group"] == g]
        vals = sorted(r["x"]["m"][m][k] for r in R)
        # threshold: midpoint of winners' and losers' medians at that minute (from main())
        W = [r for r in R if (r["x"]["final"] / r["x"]["close"][T] - 1) * 100 >= 2]
        L = [r for r in R if (r["x"]["final"] / r["x"]["close"][T] - 1) * 100 <= -2]
        thr = (st.median(r["x"]["m"][m][k] for r in W) + st.median(r["x"]["m"][m][k] for r in L)) / 2
        fwd = lambda r: (r["x"]["final"] / r["x"]["close"][m] - 1) * 100
        sig = [r for r in R if (r["x"]["m"][m][k] - thr) * sgn > 0]
        rest = [r for r in R if r not in sig]
        s1 = [fwd(r) for r in sig if r["day"] in h1]
        s2 = [fwd(r) for r in sig if r["day"] not in h1]
        t = f"{9 + (30 + m) // 60}:{(30 + m) % 60:02d}"
        rule = f"{k} {'>' if sgn > 0 else '<'} {thr:+.2f} at {t}"
        print(f"{g:18} {rule:38} {len(sig):5} {st.mean(fwd(r) for r in sig):+11.2f}% "
              f"{(st.mean(s1) if s1 else 0):+6.2f}/{(st.mean(s2) if s2 else 0):+.2f}% {st.mean(fwd(r) for r in rest):+12.2f}% "
              f"{sum(fwd(r) > 0 for r in sig) / len(sig) * 100:5.0f}%")


if __name__ == "__main__" and len(sys.argv) > 1 and sys.argv[1] == "rules":
    rule_check()
