"""
shadow_analysis.py

[2026-09-26] Labels the shadow snapshots from the bar and tick simulations
with outcomes (shadow_outcomes.label_file) and checks whether the proposed
entry_rules features separate good entries from bad ones:

  A. ENTRY snapshots joined to the simulated trade's real P/L
  B. every decision-change snapshot with a plan stop/target: target-before-
     stop rate and 30/60-minute return, by feature band
"""
import gzip
import json
import sys
from collections import defaultdict
from datetime import datetime, timezone
from pathlib import Path

HERE = Path(__file__).resolve().parent
sys.path.insert(0, str(HERE.parents[1]))
import shadow_outcomes as so

UTC = timezone.utc
SCRATCH = Path("/tmp/claude-0/-var-www-screener-trade/5844eef7-c648-46de-a03c-fe0b1434b640/scratchpad")


def bar_source():
    cache = {}

    def bars_for(symbol, d):
        key = d.isoformat()
        if key not in cache:
            p = HERE / "bars" / f"bars_{key}.json.gz"
            cache[key] = json.load(gzip.open(p, "rt"))["symbols"] if p.exists() else {}
        bl = cache[key].get(symbol, {}).get("bars", [])
        return [{"t": datetime.fromtimestamp(b[0], UTC), "o": b[1], "h": b[2], "l": b[3], "c": b[4], "v": b[5]}
                for b in bl]
    return bars_for


def load(dirpath, bars_for):
    snaps, outs = {}, {}
    for f in sorted(Path(dirpath).glob("shadow_entry_*.jsonl")):
        if f.name.endswith(".outcomes.jsonl"):
            continue
        of = f.with_suffix(".outcomes.jsonl")
        if not of.exists():
            so.label_file(f, bars_for)
        for l in open(f):
            r = json.loads(l)
            if r.get("type") == "snapshot" and "vas" in r:
                snaps[r["id"]] = r
        for l in open(of):
            o = json.loads(l)
            outs[o["id"]] = o
    return snaps, outs


def band(v, edges, fmt="{:g}"):
    if v is None:
        return "n/a"
    for a, b in zip(edges, edges[1:]):
        if a <= v < b:
            return f"{fmt.format(a)}..{fmt.format(b)}"
    return f">={fmt.format(edges[-1])}"


FEATURES = {
    "VAS score": lambda s: band(s["vas"]["score"], [0, 4, 6, 8, 10.01]),
    "shadow verdict": lambda s: s["shadow"]["decision"],
    "beats SPY & IWM": lambda s: (None if not [x for x in s["context"]["rel_strength_pct"].values() if x is not None]
                                  else str(min(x for x in s["context"]["rel_strength_pct"].values()
                                               if x is not None) >= 0)),
    "3-bar volume trend": lambda s: s["volume"]["trend_3bar"],
    "bar RVOL": lambda s: band(s["volume"]["bar_rvol"], [0, 0.75, 1, 1.5, 3]),
    "price vs VWAP %": lambda s: band(s["context"]["price_vs_vwap_pct"], [-99, 0, 1, 2, 4]),
    "imbalance (ticks)": lambda s: band(s["flow"]["imbalance"], [-1.01, -0.1, 0.1, 0.3, 0.6]),
    "spread % (ticks)": lambda s: band(s["flow"]["spread_pct"], [0, 0.2, 0.5, 1]),
}


def report(title, snaps, outs, trades=None):
    print(f"\n==== {title}: {len(snaps)} snapshots")
    # A. entries vs real P/L
    if trades is not None:
        ent = [s for s in snaps.values() if s["event"] == "ENTRY"]
        rows = []
        for s in ent:
            t0 = datetime.fromisoformat(s["ts"])
            tr = next((t for t in trades if t["symbol"] == s["symbol"]
                       and abs((datetime.fromisoformat(t["entry_time"]) - t0).total_seconds()) <= 150), None)
            if tr:
                rows.append((s, tr["pl"]))
        print(f"-- A. {len(rows)} entries joined to trade P/L (total ${sum(p for _, p in rows):+.2f})")
        for name, f in FEATURES.items():
            g = defaultdict(list)
            for s, p in rows:
                g[f(s)].append(p)
            if len(g) > 1:
                print(f"   {name:20} " + " | ".join(f"{k}: n={len(v)} win {sum(x > 0 for x in v)/len(v)*100:.0f}% "
                                                    f"${sum(v):+.0f}" for k, v in sorted(g.items(), key=str)))
    # B. all decision-change snapshots with a plan target/stop
    rows = [(s, outs[i]) for i, s in snaps.items() if i in outs and outs[i].get("plan_result")]
    print(f"-- B. {len(rows)} snapshots with a plan stop/target")
    for name, f in FEATURES.items():
        g = defaultdict(list)
        for s, o in rows:
            g[f(s)].append(o)
        if len(g) > 1:
            parts = []
            for k, v in sorted(g.items(), key=str):
                hit = sum(o["plan_result"] == "target" for o in v); stp = sum(o["plan_result"] == "stop" for o in v)
                r60 = [o["returns_pct"].get("3600") for o in v if o["returns_pct"].get("3600") is not None]
                parts.append(f"{k}: n={len(v)} target {hit/len(v)*100:.0f}% stop {stp/len(v)*100:.0f}% "
                             f"60m {sum(r60)/len(r60):+.2f}%" if r60 else f"{k}: n={len(v)}")
            print(f"   {name:20} " + " | ".join(parts))


def main():
    bf = bar_source()
    for v in ("J", "A", "H"):
        d = HERE / "results" / f"shadow_{v}_current"
        if not d.exists():
            continue
        snaps, outs = load(d, bf)
        res = json.load(open(HERE / "results" / f"{v}_current.json"))
        trades = [t for tr in res["days"].values() for t in tr]
        report(f"20-day bars, setup {v}", snaps, outs, trades)
    d = SCRATCH / "shadow_tick"
    if d.exists():
        snaps, outs = load(d, so.alpaca_bars_for)
        trades = []
        for f in sorted(Path("/var/www/screener/trade/data/simulations").glob("2026-09-*_sim_maxpos5_*_rr20_*_ext_msk25*.json")):
            for t in json.load(open(f))["trades"]:
                t["pl"] = t["pl_dollars"]; trades.append(t)
        report("Tick replay 9/18-9/25, setup J (full order-flow features)", snaps, outs, trades)


if __name__ == "__main__":
    main()
