"""
shadow_outcomes.py

[2026-09-26] Labels entry_rules.py shadow snapshots with what price did
afterwards, from 1-minute bars (after the close, or over replayed days):

  returns_pct   close of the first bar ending >= t0 + h, for h in HORIZONS_S
  mfe / mae     best / worst price vs p0 over bars starting >= t0, up to 60 min
  plan_result   'target' / 'stop' / 'neither' -- which of the plan's own
                target and stop (from bot_context.plan) price reached first
                before 15:55 (a bar touching both counts as 'stop')

Usage:
  python shadow_outcomes.py logs/shadow/shadow_entry_2026-09-28.jsonl
      bars from Alpaca REST (live days)
  label_file(path, bars_for)  -- from code, with your own bars source
      (bars_for(symbol, date) -> [{"t": aware dt of bar START, "o","h","l","c","v"}])
Writes <file>.outcomes.jsonl next to the input.
"""
import json
import sys
from datetime import datetime, timedelta, timezone
from pathlib import Path
from zoneinfo import ZoneInfo

ET = ZoneInfo("America/New_York")
HORIZONS_S = [60, 300, 900, 1800, 3600]
BAR = timedelta(minutes=1)


def _pct(a, b):
    return (a / b - 1.0) * 100.0 if a is not None and b else None


def label_snapshot(snap: dict, bars: list) -> dict:
    t0 = datetime.fromisoformat(snap["ts"])
    p0 = snap.get("price")
    out = {"type": "outcome", "id": snap["id"], "symbol": snap["symbol"], "event": snap.get("event"),
           "t0": snap["ts"], "p0": p0, "returns_pct": {}, "mfe_pct": None, "mae_pct": None,
           "plan_result": None, "minutes_to_result": None}
    if not p0:
        return out
    eod = t0.astimezone(ET).replace(hour=15, minute=55, second=0, microsecond=0)
    after = [b for b in bars if b["t"] >= t0 and b["t"] < eod]
    for h in HORIZONS_S:
        tgt = t0 + timedelta(seconds=h)
        b = next((b for b in after if b["t"] + BAR >= tgt), None)
        out["returns_pct"][str(h)] = round(_pct(b["c"], p0), 4) if b else None
    win = [b for b in after if b["t"] < t0 + timedelta(seconds=max(HORIZONS_S))]
    if win:
        out["mfe_pct"] = round(_pct(max(b["h"] for b in win), p0), 4)
        out["mae_pct"] = round(_pct(min(b["l"] for b in win), p0), 4)
    plan = ((snap.get("bot") or {}).get("plan") or {})
    stop, target = plan.get("stop"), plan.get("target")
    if stop and target and stop < p0 < target:
        out["plan_result"] = "neither"
        for b in after:
            if b["l"] <= stop:
                out["plan_result"] = "stop"
            elif b["h"] >= target:
                out["plan_result"] = "target"
            if out["plan_result"] != "neither":
                out["minutes_to_result"] = round((b["t"] + BAR - t0).total_seconds() / 60, 1)
                break
    return out


def label_file(path, bars_for) -> Path:
    path = Path(path)
    snaps = [json.loads(l) for l in open(path) if '"type":"snapshot"' in l]
    cache = {}
    out_path = path.with_suffix(".outcomes.jsonl")
    with open(out_path, "w") as f:
        for s in snaps:
            if s.get("error"):
                continue
            d = datetime.fromisoformat(s["ts"]).astimezone(ET).date()
            key = (s["symbol"], d)
            if key not in cache:
                cache[key] = bars_for(s["symbol"], d)
            f.write(json.dumps(label_snapshot(s, cache[key]), separators=(",", ":")) + "\n")
    return out_path


def alpaca_bars_for(symbol, d):
    sys.path.insert(0, str(Path(__file__).resolve().parent))
    from alpaca_client import get_client
    c = get_client()
    start = datetime(d.year, d.month, d.day, 9, 30, tzinfo=ET)
    raw = c.get_minute_bars(symbol, start=start, end=start.replace(hour=16), limit=1000)
    return [{"t": b.timestamp, "o": float(b.open), "h": float(b.high), "l": float(b.low),
             "c": float(b.close), "v": float(b.volume)} for b in raw]


if __name__ == "__main__":
    for p in sys.argv[1:]:
        print("wrote", label_file(p, alpaca_bars_for))
