"""
winner_patterns.py -- [2026-09-27] What do the big winners (closed 5%+
above the 9:30 open) have in common, vs the other stocks on the same
9:28 lists? Only uses information available at 9:28 or at the decision
time stated (10:00 / 10:30) -- nothing from later in the day.

Writes winner_features.json (one row per stock-day) and prints bucket
tables: big-winner rate per band vs the base rate.
"""
import csv
import gzip
import json
import sys
from collections import defaultdict
from datetime import datetime, timedelta
from pathlib import Path
from zoneinfo import ZoneInfo

HERE = Path(__file__).resolve().parent
sys.path.insert(0, str(HERE.parents[1]))
from alpaca_client import get_client
import volatility

ET = ZoneInfo("America/New_York")
CACHE = Path("/var/www/screener/trade/data/simulations/cache")


def hm(ts):
    return datetime.fromtimestamp(ts, ET).strftime("%H:%M")


def session_feats(bars, baseline):
    """bars: [[ts,o,h,l,c,v]] 9:30-16:00. Features at 10:00 and 10:30."""
    out = {}
    o = bars[0][1]
    for label, cut in (("1000", "10:00"), ("1030", "10:30")):
        seg = [b for b in bars if hm(b[0]) < cut]
        if not seg:
            continue
        pv = sum((b[2] + b[3] + b[4]) / 3 * b[5] for b in seg)
        v = sum(b[5] for b in seg)
        vwap = pv / v if v else None
        c = seg[-1][4]
        hi = max(b[2] for b in seg)
        lo = min(b[3] for b in seg)
        out[f"chg_{label}"] = (c / o - 1) * 100
        out[f"vs_vwap_{label}"] = (c / vwap - 1) * 100 if vwap else None
        out[f"off_high_{label}"] = (c / hi - 1) * 100
        out[f"range_{label}"] = (hi / lo - 1) * 100
        out[f"vol_{label}_x_base"] = v / baseline if baseline else None
        # where the open-window low printed relative to the high (dip then rip vs spike then fade)
        hi_i = max(range(len(seg)), key=lambda i: seg[i][2])
        lo_i = min(range(len(seg)), key=lambda i: seg[i][3])
        out[f"low_before_high_{label}"] = lo_i < hi_i
    # higher lows through 10:30: low of 10:00-10:30 above low of 9:30-10:00
    a = [b for b in bars if hm(b[0]) < "10:00"]
    b2 = [b for b in bars if "10:00" <= hm(b[0]) < "10:30"]
    if a and b2:
        out["higher_low_1030"] = min(x[3] for x in b2) > min(x[3] for x in a)
        out["higher_high_1030"] = max(x[2] for x in b2) > max(x[2] for x in a)
    return out


def main():
    client = get_client()
    winners = {(r["date"], r["symbol"]): r for r in json.load(open(HERE / "daily_winners.json"))}
    scan = {(r["date"], r["symbol"]): r for r in csv.DictReader(open(HERE / "scanner_features_20260827_20260924.csv"))}
    days = sorted({d for d, _ in winners})
    rows = []
    seen_big = defaultdict(int)       # symbol -> big-winner days so far
    on_list = defaultdict(int)        # symbol -> days on the list so far
    prev_list = set()
    for day in days:
        syms = [s for d, s in winners if d == day]
        dd = datetime.fromisoformat(day).replace(tzinfo=ET)
        daily = client.get_daily_bars_bulk(syms, dd - timedelta(days=45), dd - timedelta(seconds=1))
        bench = {}
        bp = HERE / "bars" / f"bench_{day}.json.gz"
        bfile = HERE / "bars" / f"bars_{day}.json.gz"
        bdata = json.load(gzip.open(bfile, "rt"))["symbols"] if bfile.exists() else {}
        for bs in ("SPY", "IWM"):
            raw = client.get_minute_bars(bs, start=dd.replace(hour=9, minute=30), end=dd.replace(hour=10, minute=0), limit=100)
            if raw:
                bench[bs] = (float(raw[-1].close) / float(raw[0].open) - 1) * 100
        for s in syms:
            w = winners[(day, s)]
            if s in bdata:
                bars, baseline, lv = bdata[s]["bars"], bdata[s]["baseline"], bdata[s]["levels"]
            else:
                raw = client.get_minute_bars(s, start=dd.replace(hour=9, minute=30), end=dd.replace(hour=16), limit=1000)
                bars = [[int(b.timestamp.timestamp()), float(b.open), float(b.high), float(b.low), float(b.close),
                         float(b.volume)] for b in raw]
                hist = [b for b in daily.get(s, []) if b["t"] < dd]
                baseline = hist[-1]["v"] if hist else None
                lv = {"daily_atr": volatility.daily_atr(hist), **volatility.five_day_reference(hist),
                      "prev_day_high": hist[-1]["h"] if hist else None,
                      "range_20d_high": max(b["h"] for b in hist[-20:]) if hist else None}
            if not bars:
                continue
            hist = [b for b in daily.get(s, []) if b["t"] < dd]
            sc = scan.get((day, s), {})
            f = lambda k: float(sc[k]) if sc.get(k) not in (None, "") else None
            o = w["open"]
            r = {
                "date": day, "symbol": s, "big": w["open_close_pct"] >= 5, "winner": w["open_close_pct"] > 0,
                "open_close_pct": w["open_close_pct"], "price": o,
                "score": f("score"), "rvol": f("rvol"), "pm_volume": f("pm_volume"),
                "gap_pct": (o / w["prev_close"] - 1) * 100,
                "pm_strength_pct": f("pm_strength_pct"), "dist_to_resistance_pct": f("dist_to_resistance_pct"),
                "daily_atr_pct": (lv.get("daily_atr") / o * 100) if lv.get("daily_atr") else f("daily_atr_pct"),
                "open_vs_20d_high_pct": (o / lv["range_20d_high"] - 1) * 100 if lv.get("range_20d_high") else None,
                "open_vs_prev_high_pct": (o / lv["prev_day_high"] - 1) * 100 if lv.get("prev_day_high") else None,
                "pos_in_5d_range": ((o - lv["ref_5d_low"]) / (lv["ref_5d_high"] - lv["ref_5d_low"]) * 100)
                if lv.get("ref_5d_high") and lv.get("ref_5d_low") and lv["ref_5d_high"] > lv["ref_5d_low"] else None,
                "prev_day_chg_pct": (hist[-1]["c"] / hist[-2]["c"] - 1) * 100 if len(hist) >= 2 else None,
                "chg_5d_pct": (hist[-1]["c"] / hist[-6]["c"] - 1) * 100 if len(hist) >= 6 else None,
                "on_list_prev_day": s in prev_list,
                "days_on_list_before": on_list[s],
                "big_winner_before": seen_big[s] > 0,
                "spy_0930_1000": bench.get("SPY"), "iwm_0930_1000": bench.get("IWM"),
                **session_feats(bars, baseline),
            }
            rows.append(r)
        for r in rows:
            if r["date"] == day:
                on_list[r["symbol"]] += 1
                seen_big[r["symbol"]] += r["big"]
        prev_list = set(syms)
        print(day, "done", flush=True)
    json.dump(rows, open(HERE / "winner_features.json", "w"), default=str)


if __name__ == "__main__":
    main()
