"""
separation.py -- [2026-09-27] Which values, at the moment of a pick,
separate picks that went on to rise (to 15:55) from picks that fell?
Pick-time indicators (entry_rules snapshot) + what was known at 9:28.
AUC = chance a random winning pick has a HIGHER value than a random losing
pick (0.50 = no separation; 0.60+ / 0.40- worth a look). Checked on both
halves of the days so a one-period fluke doesn't pass.
Usage: python3 sim/analysis/separation.py data/simulations/<run>.json
"""
import json
import statistics as st
import sys

sys.path.insert(0, __file__.rsplit("/", 1)[0])
import pick_report as P

WF = {(r["date"], r["symbol"]): r for r in json.load(open(
    "/var/www/screener/trade/reports/open_window/winner_features.json"))}
MORNING = ["price", "score", "rvol", "pm_volume", "gap_pct", "pm_strength_pct", "dist_to_resistance_pct",
           "daily_atr_pct", "open_vs_20d_high_pct", "open_vs_prev_high_pct", "pos_in_5d_range",
           "prev_day_chg_pct", "chg_5d_pct", "days_on_list_before"]


def features(t, day):
    i = t["plan"]["indicators"]
    f = P.flat(t, day)
    parts = i.get("vas_parts") or {}
    f.update({"vas_" + k: v for k, v in parts.items()})
    f.update(volume_pressure=i.get("volume_pressure"), projected_rvol=i.get("projected_rvol"),
             accel_momentum=i.get("accel_momentum"), trades_30s=i.get("trades_30s"),
             ema_stack=None if i.get("ema_stack") is None else float(i["ema_stack"]),
             trend_accelerating=float(i.get("trend_3bar") == "ACCELERATING"))
    w = WF.get((day, t["symbol"]), {})
    for k in MORNING:
        f["928_" + k] = w.get(k)
    for k in ("big_winner_before", "on_list_prev_day"):
        f["928_" + k] = None if w.get(k) is None else float(w[k])
    f.pop("trend_3bar", None)
    return f


def auc(a, b):
    if not a or not b:
        return None
    allv = sorted([(x, 1) for x in a] + [(x, 0) for x in b])
    rank, i, ranks = 1, 0, []
    while i < len(allv):
        j = i
        while j < len(allv) and allv[j][0] == allv[i][0]:
            j += 1
        r = (rank + rank + (j - i) - 1) / 2
        ranks += [(r, allv[k][1]) for k in range(i, j)]
        rank += j - i
        i = j
    sa = sum(r for r, lab in ranks if lab)
    return (sa - len(a) * (len(a) + 1) / 2) / (len(a) * len(b))


def main():
    res = json.load(open(sys.argv[1]))
    days = sorted(res["days"])
    half = set(days[:len(days) // 2])
    rows = []
    for d in days:
        for t in res["days"][d]["trades"]:
            rows.append((d, t["pl_dollars"] > 0, t["pl_pct"], features(t, d)))
    keys = sorted({k for *_, f in rows for k in f})
    out = []
    for k in keys:
        def split(sel):
            a = [f[k] for d, ok, r, f in rows if ok and sel(d) and isinstance(f.get(k), (int, float))]
            b = [f[k] for d, ok, r, f in rows if not ok and sel(d) and isinstance(f.get(k), (int, float))]
            return a, b
        a, b = split(lambda d: True)
        a1, b1 = split(lambda d: d in half)
        a2, b2 = split(lambda d: d not in half)
        A, A1, A2 = auc(a, b), auc(a1, b1), auc(a2, b2)
        if A is None:
            continue
        out.append((abs(A - 0.5), k, A, A1, A2, st.median(a), st.median(b), len(a) + len(b)))
    out.sort(reverse=True)
    print(f"{len(rows)} picks ({sum(ok for _, ok, *_ in rows)} rose, {sum(not ok for _, ok, *_ in rows)} fell)\n")
    print(f"{'value':30}{'AUC':>6}{'1st half':>9}{'2nd half':>9}{'median rose':>13}{'median fell':>13}{'n':>5}")
    for _, k, A, A1, A2, ma, mb, n in out:
        flag = "  <- consistent" if A1 and A2 and (A1 - 0.5) * (A2 - 0.5) > 0 and min(abs(A1 - .5), abs(A2 - .5)) >= 0.05 else ""
        print(f"{k:30}{A:6.2f}{(A1 or 0):9.2f}{(A2 or 0):9.2f}{ma:13.3f}{mb:13.3f}{n:5}{flag}")
    json.dump([(d, ok, r, f) for d, ok, r, f in rows], open(sys.argv[1].replace(".json", "_features.json"), "w"))


if __name__ == "__main__":
    main()
