"""
check_entry_rules.py -- [2026-09-27] Verifies an entry_rules simulation run:
  1. every pick's saved indicator values really clear every threshold
  2. key indicators recomputed independently from Alpaca's official 1-min
     bars (bar RVOL, price vs VWAP, 5-bar change, relative strength vs
     SPY/IWM) agree with what entry_rules saw at the pick
  3. when in the day each stock first passed
Usage: python3 sim/analysis/check_entry_rules.py data/simulations/entry_rules_21d.json [date ...]
"""
import gzip
import json
import statistics as st
import sys
from datetime import datetime, timedelta, timezone
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parents[2]))
from config_loader import get_config

BARS = Path("/var/www/screener/trade/reports/open_window/bars")
CACHE = Path(get_config()["simulator"]["cache_dir"])
R = get_config()["entry_rules"]["rules"]


def rule_check(ind):
    rs = [x for x in (ind.get("rel_strength_pct") or {}).values() if x is not None]
    checks = {
        "vas>=%s" % R["min_vas"]: ind["vas"] is not None and ind["vas"] >= R["min_vas"],
        "imbalance>=%s" % R["min_imbalance"]: ind["imbalance"] is not None and ind["imbalance"] >= R["min_imbalance"],
        "bar_rvol>=%s" % R["min_bar_rvol"]: ind["bar_rvol"] is not None and ind["bar_rvol"] >= R["min_bar_rvol"],
        "spread<=%s" % R["max_spread_pct"]: ind["spread_pct"] is not None and ind["spread_pct"] <= R["max_spread_pct"],
        "beats SPY and IWM": len(rs) == 2 and min(rs) >= R["min_rel_strength_pct"],
        "not decelerating": ind["trend_3bar"] not in ("DECELERATING", "INSUFFICIENT"),
    }
    if R.get("volume_rising_bars") and ind.get("volume_rising") and ind["volume_rising"].get("bars"):
        checks["volume rising %s bars" % R["volume_rising_bars"]] = ind["volume_rising"].get("ok") is True
    E = get_config()["entry_rules"]
    if E.get("max_vwap_ext_pct") is not None:
        checks["vwap <= +%s%%" % E["max_vwap_ext_pct"]] = ind["price_vs_vwap_pct"] <= E["max_vwap_ext_pct"]
    if E.get("max_up_from_open_pct") is not None and ind.get("up_from_open_pct") is not None:
        checks["open <= +%s%%" % E["max_up_from_open_pct"]] = ind["up_from_open_pct"] <= E["max_up_from_open_pct"]
    return [k for k, ok in checks.items() if not ok]


def official(day, sym):
    p = BARS / f"bars_{day}.json.gz"
    if not hasattr(official, "c"):
        official.c = {}
    if day not in official.c:
        official.c[day] = json.load(gzip.open(p, "rt"))["symbols"] if p.exists() else {}
    bl = official.c[day].get(sym, {}).get("bars")
    return [(datetime.fromtimestamp(b[0], timezone.utc), *b[1:]) for b in bl] if bl else None


def bench(day, sym):
    p = CACHE / day / f"bench_{sym}.json"
    return [(datetime.fromisoformat(r[0]), *r[1:]) for r in json.loads(p.read_text())] if p.exists() else None


def recompute(day, sym, t, price):
    bl = official(day, sym)
    if not bl:
        return None
    done = [b for b in bl if b[0] + timedelta(minutes=1) <= t]
    if len(done) < 7:
        return None
    v = [b[5] for b in done]
    base = v[-21:-1]
    pv = sum((b[2] + b[3] + b[4]) / 3 * b[5] for b in done)
    vwap = pv / sum(v)
    own = (done[-1][4] / done[-6][4] - 1) * 100
    rs = {}
    for bsym in ("SPY", "IWM"):
        bb = [b for b in (bench(day, bsym) or []) if b[0] + timedelta(minutes=1) <= t]
        rs[bsym] = own - (bb[-1][4] / bb[-6][4] - 1) * 100 if len(bb) >= 6 else None
    return {"bar_rvol": v[-1] / (sum(base) / len(base)), "price_vs_vwap_pct": (price / vwap - 1) * 100,
            "own_change_5bar_pct": own, "rs_SPY": rs["SPY"], "rs_IWM": rs["IWM"]}


def main():
    res = json.load(open(sys.argv[1]))
    days = sys.argv[2:] or sorted(res["days"])
    bad, diffs, n = [], {k: [] for k in ("bar_rvol", "price_vs_vwap_pct", "own_change_5bar_pct", "rs_SPY", "rs_IWM")}, 0
    first_pass = []
    for day in days:
        for t in res["days"].get(day, {}).get("trades", []):
            n += 1
            ind = t["plan"]["indicators"]
            fails = rule_check(ind)
            if fails:
                bad.append((day, t["symbol"], fails))
            et = datetime.fromisoformat(t["entry_time"])
            first_pass.append(((et - et.replace(hour=13, minute=30, second=0, microsecond=0)).total_seconds() / 60))
            rc = recompute(day, t["symbol"], et, t["entry_price"])
            if rc:
                mine = {"bar_rvol": ind["bar_rvol"], "price_vs_vwap_pct": ind["price_vs_vwap_pct"],
                        "own_change_5bar_pct": ind["own_change_5bar_pct"],
                        "rs_SPY": (ind["rel_strength_pct"] or {}).get("SPY"), "rs_IWM": (ind["rel_strength_pct"] or {}).get("IWM")}
                for k in diffs:
                    if mine[k] is not None and rc[k] is not None:
                        diffs[k].append((mine[k], rc[k], day, t["symbol"]))
    print(f"{n} picks checked over {len(days)} day(s)\n")
    print(f"1. Threshold check: {n - len(bad)} of {n} picks clear all six rules")
    for b in bad[:20]:
        print("   FAILS:", b)
    print("\n2. entry_rules value vs independent recompute from official 1-min bars")
    for k, v in diffs.items():
        if not v:
            continue
        d = [abs(a - b) for a, b, *_ in v]
        rel = [abs(a - b) / max(abs(b), 1e-9) for a, b, *_ in v]
        agree = sum((x <= 0.25 if k == "bar_rvol" else x <= 0.5) for x in (rel if k == "bar_rvol" else d))
        print(f"   {k:20} n={len(v):4}  median gap {st.median(d):.3f}  close agreement {agree}/{len(v)}")
    worst = sorted(diffs["bar_rvol"], key=lambda x: -abs(x[0] - x[1]) / max(x[1], 1e-9))[:5]
    print("   biggest bar_rvol gaps (entry_rules vs official):", [(d, s, round(a, 2), round(b, 2)) for a, b, d, s in worst])
    print("\n3. When stocks first passed (minutes after 9:30):")
    fp = sorted(first_pass)
    for lim in (10, 15, 30, 60, 120, 390):
        print(f"   by {lim:3} min: {sum(x <= lim for x in fp)}/{len(fp)}")


if __name__ == "__main__":
    main()
