"""
analyze_trades.py -- [2026-09-30] The user: "evaluate each trade today, analyze how you
could have made money on each one of them and add it to the play book".

For every trade the bot made on a day:
  - official 1-min bars 9:30-16:00, VWAP, volume vs the 14-day normal for that minute
  - the bot's trade: entry/exit, best price reached after the buy (MFE) and worst (MAE),
    what holding to the close would have made
  - the RIGHT trade (hindsight), buying and selling on 1-min closes:
      (a) while the bot was watching the stock (its scanner windows)
      (b) at any time of the day
  - measured signals at the right entry and right exit (what the bot would need to see)
  - pattern labels for the entry and the exit
Each trade becomes one CASE, appended to data/playbook/cases.jsonl (the playbook /
case library the user wants the bot to learn from), plus a chart per trade.

    python3 analyze_trades.py 2026-09-30                    the trade bot's trades that day
    python3 analyze_trades.py all --source sip --no-charts   [2026-09-30] every day of the
        sip_bot's trades (/var/www/screener/premarket, READ ONLY); it logs no watch windows,
        so its "right trade" window is the hour around its buy (right_watched) + the day
"""
import json
import sys
from datetime import datetime
from pathlib import Path
from zoneinfo import ZoneInfo

import matplotlib
matplotlib.use("Agg")
import matplotlib.dates as md
import matplotlib.pyplot as plt

TRADE = Path("/var/www/screener/trade")
sys.path.insert(0, str(TRADE))
from alpaca_client import get_client
import reference

ET = ZoneInfo("America/New_York")
OUT = TRADE / "reports" / "playbook"
CASES = TRADE / "data" / "playbook" / "cases.jsonl"


def et(ts):
    return datetime.fromisoformat(ts).astimezone(ET) if isinstance(ts, str) else ts.astimezone(ET)


def hm(d):
    return d.strftime("%H:%M")


def watch_windows(day, sym):
    scans = []
    for line in open(TRADE / "logs" / f"trade_bot_{day}.log"):
        if "[SCAN]" in line and "candidates selected" in line:
            scans.append((line[11:16], f"'{sym}'" in line))
    out, start = [], None
    for t, on in scans:
        if on and start is None:
            start = t
        if not on and start is not None:
            out.append((start, t))
            start = None
    if start is not None:
        out.append((start, "15:15"))
    return out


def load_bars(client, sym, day, ref):
    d = datetime.fromisoformat(day).replace(tzinfo=ET)
    raw = client.get_minute_bars(sym, start=d.replace(hour=9, minute=30), end=d.replace(hour=16), limit=1000)
    B, pv, v = [], 0.0, 0.0
    for b in raw:
        t = b.timestamp.astimezone(ET)
        m = (t.hour - 9) * 60 + t.minute - 30
        if not 0 <= m < 390:
            continue
        x = {"t": t, "m": m, "o": float(b.open), "h": float(b.high), "l": float(b.low), "c": float(b.close),
             "v": float(b.volume)}
        pv += (x["h"] + x["l"] + x["c"]) / 3 * x["v"]
        v += x["v"]
        x["vwap"] = pv / v
        x["volx"] = x["v"] / max(ref["vol_per_min"][m], 1) if ref else None
        B.append(x)
    return B


def best_trade(B, lo_m, hi_m, exit_by=385):
    """Best buy-on-close / sell-on-a-later-close pair, buy minute in [lo_m, hi_m]."""
    best, run_min = None, None
    for j, b in enumerate(B):
        if b["m"] > exit_by:
            break
        if run_min is not None:
            g = b["c"] / B[run_min]["c"] - 1
            if best is None or g > best[2]:
                best = (run_min, j, g)
        if lo_m <= b["m"] <= hi_m and (run_min is None or b["c"] < B[run_min]["c"]):
            run_min = j
    return best


def entry_signals(B, i, open_px):
    b = B[i]
    prior = B[:i + 1]
    low_j = min(range(len(prior)), key=lambda k: prior[k]["l"])
    hi_j = max(range(len(prior)), key=lambda k: prior[k]["h"])
    crossed = any(B[k]["c"] <= B[k]["vwap"] for k in range(max(0, i - 5), i)) and b["c"] > b["vwap"]
    last15 = B[max(0, i - 14):i + 1]
    vx = [x["volx"] for x in last15 if x["volx"] is not None]
    base_lo = min(x["l"] for x in B[max(0, i - 30):i + 1])
    s = {
        "time": hm(b["t"]), "price": b["c"],
        "vs_open_pct": round((b["c"] / open_px - 1) * 100, 2),
        "vs_vwap_pct": round((b["c"] / b["vwap"] - 1) * 100, 2),
        "day_low": prior[low_j]["l"], "day_low_time": hm(prior[low_j]["t"]),
        "min_since_low": i - low_j,
        "low_held_open": prior[low_j]["l"] >= open_px * 0.995,
        "vwap_reclaim_last5": crossed,
        "high_so_far": prior[hi_j]["h"], "below_high_pct": round((b["c"] / prior[hi_j]["h"] - 1) * 100, 2),
        "vol15x": round(sum(vx) / len(vx), 2) if vx else None,
        "max_volx_last30": round(max((x["volx"] or 0) for x in B[max(0, i - 30):i + 1]), 1),
        "no_new_low_min": i - max(k for k in range(len(prior)) if prior[k]["l"] <= base_lo + 1e-9),
    }
    tags = []
    if s["vwap_reclaim_last5"]:
        tags.append("VWAP reclaim")
    if s["min_since_low"] >= 20 and s["vs_vwap_pct"] < 1.5:
        tags.append(f"base after low ({s['min_since_low']} min)")
    if s["low_held_open"] and i <= 15:
        tags.append("opening dip held the open")
    if -0.5 <= s["vs_vwap_pct"] <= 0.5 and not s["vwap_reclaim_last5"]:
        tags.append("pullback to VWAP held")
    if s["below_high_pct"] >= -0.2 and s["vs_open_pct"] > 1:
        tags.append("breakout to a new high")
    s["tags"] = tags or ["other"]
    return s


def exit_signals(B, j, entry_i):
    b = B[j]
    rng = b["h"] - b["l"]
    near = B[max(0, j - 5):j + 6]
    run_hi = max(x["h"] for x in B[entry_i:j + 1])
    s = {
        "time": hm(b["t"]), "price": b["c"],
        "vs_vwap_pct": round((b["c"] / b["vwap"] - 1) * 100, 2),
        "bar_close_pos": round((b["c"] - b["l"]) / rng, 2) if rng else None,
        "bar_volx": round(b["volx"], 1) if b["volx"] is not None else None,
        "max_volx_pm5": round(max((x["volx"] or 0) for x in near), 1),
        "at_run_high": b["h"] >= run_hi - 1e-9,
        "next30_low_pct": round((min(x["l"] for x in B[j + 1:j + 31]) / b["c"] - 1) * 100, 2) if j + 1 < len(B) else None,
    }
    tags = []
    if s["max_volx_pm5"] >= 10:
        tags.append(f"volume climax ({s['max_volx_pm5']}x)")
    if s["bar_close_pos"] is not None and s["bar_close_pos"] <= 0.35 and (s["bar_volx"] or 0) >= 3:
        tags.append("rejection bar on heavy volume")
    highs = [x["h"] for x in B[entry_i:j + 1]]
    if len(highs) > 20:
        prev_top = max(highs[:-10])
        if abs(max(highs[-10:]) / prev_top - 1) < 0.003:
            tags.append("double top")
    if s["vs_vwap_pct"] >= 3:
        tags.append(f"stretched {s['vs_vwap_pct']}% over VWAP")
    if b["m"] >= 380:
        tags.append("end of day")
    s["tags"] = tags or ["top (no clear signal)"]
    return s


def chart(sym, day, B, t, case, path):
    fig, (ax, axv) = plt.subplots(2, 1, figsize=(12, 6.6), gridspec_kw={"height_ratios": [3, 1]}, sharex=True)
    ts = [x["t"].replace(tzinfo=None) for x in B]
    ax.plot(ts, [x["c"] for x in B], color="#1f3a5f", lw=1, label="price")
    ax.plot(ts, [x["vwap"] for x in B], color="#c0392b", lw=1.1, label="VWAP")
    ax.axhline(B[0]["o"], color="grey", ls=":", lw=.8)
    for a, z in case["watched"]:
        ax.axvspan(datetime.fromisoformat(f"{day}T{a}"), datetime.fromisoformat(f"{day}T{z}"), color="#eaf2fb", zorder=0)

    def mark(tm, p, col, mk, txt, dy):
        tt = datetime.fromisoformat(f"{day}T{tm}")
        ax.scatter([tt], [p], color=col, marker=mk, s=130, zorder=5, edgecolor="black", linewidth=.6)
        ax.annotate(txt, (tt, p), xytext=(8, dy), textcoords="offset points", fontsize=7.5, color=col,
                    bbox=dict(boxstyle="round", fc="white", ec=col, lw=.6))
    bt = case["bot"]
    mark(bt["entry_time"], bt["entry_price"], "#27ae60", "o", f"BOT BUY {bt['entry_price']:.2f}", 18)
    mark(bt["exit_time"], bt["exit_price"], "#e74c3c", "o", f"BOT SELL {bt['exit_price']:.2f} ({bt['pl']:+.0f}$)", -22)
    rw = case["right_watched"]
    if rw:
        mark(rw["entry"]["time"], rw["entry"]["price"], "#16a085", "*", "RIGHT BUY " + ", ".join(rw["entry"]["tags"]), -34)
        mark(rw["exit"]["time"], rw["exit"]["price"], "#8e44ad", "*", "RIGHT SELL " + ", ".join(rw["exit"]["tags"]), 22)
    ax.set_title(f"{sym} {day} — bot {bt['pl']:+.2f}$ | right trade while watched "
                 f"{(rw['gain_pct'] if rw else 0):+.1f}% | blue = on the bot's watchlist", fontsize=9.5)
    ax.grid(alpha=.3)
    ax.legend(loc="best", fontsize=7)
    axv.bar(ts, [x["v"] for x in B], width=1 / 1440,
            color=["#8e44ad" if (x["volx"] or 0) >= 5 else "#1f3a5f" for x in B])
    axv.grid(alpha=.3)
    axv.xaxis.set_major_formatter(md.DateFormatter("%H:%M"))
    plt.tight_layout()
    plt.savefig(path, dpi=110)
    plt.close()


SIP = Path("/var/www/screener/premarket/data/trades")


def load_trades(day, source):
    if source == "sip":
        out = []
        for l in open(SIP / f"{day}_trades.jsonl"):
            try:
                t = json.loads(l)
            except Exception:
                continue
            if t.get("status") != "closed" or t.get("exit_price") is None or not t.get("exit_time"):
                continue
            out.append({"symbol": t["symbol"], "entry_price": float(t["entry_price"]),
                        "exit_price": float(t["exit_price"]), "qty": float(t["shares"]),
                        "pl_dollars": round((float(t["exit_price"]) - float(t["entry_price"])) * float(t["shares"]), 2),
                        "entry_time": t["entry_time"], "exit_time": t["exit_time"],
                        "exit_reason": t.get("exit_reason") or "", "reasons": [t.get("entry_reason", "")[:200]],
                        "stop_price": t.get("initial_stop")})
        return out
    return [json.loads(l) for l in open(TRADE / "data" / "trades" / f"{day}_trades.jsonl")]


def main(day, source="trade", charts=True):
    client = get_client()
    trades = load_trades(day, source)
    if not trades:
        return []
    need = sorted({t["symbol"] for t in trades if not reference.load(day, t["symbol"])})
    if need:
        try:
            reference.build(datetime.fromisoformat(day).date(), need, verbose=False)
            reference._CACHE.clear() if hasattr(reference, "_CACHE") else None
        except Exception as e:
            print("   reference build failed", e)
    cases = []
    for t in sorted(trades, key=lambda x: x["entry_time"]):
        sym = t["symbol"]
        ref = reference.load(day, sym)
        try:
            B = load_bars(client, sym, day, ref)
        except Exception as e:
            print(f"{sym}: bars failed {e}")
            continue
        if len(B) < 30:
            continue
        ein, eout = et(t["entry_time"]), et(t["exit_time"])
        ei = next(k for k, x in enumerate(B) if x["t"] >= ein.replace(second=0, microsecond=0))
        xo = next((k for k, x in enumerate(B) if x["t"] >= eout.replace(second=0, microsecond=0)), len(B) - 1)
        after = B[ei:]
        mfe_j = max(range(ei, len(B)), key=lambda k: B[k]["h"])
        mae = min(x["l"] for x in B[ei:xo + 1])
        if source == "sip":
            a = max(ein.hour * 60 + ein.minute - 60, 9 * 60 + 30)
            z = min(ein.hour * 60 + ein.minute + 60, 15 * 60 + 15)
            win = [(f"{a // 60:02d}:{a % 60:02d}", f"{z // 60:02d}:{z % 60:02d}")]
        else:
            win = watch_windows(day, sym)
        mins = lambda s: (int(s[:2]) - 9) * 60 + int(s[3:]) - 30
        rw = None
        best_w = None
        for a, z in win:
            bt = best_trade(B, mins(a), min(mins(z), 345))
            if bt and (best_w is None or bt[2] > best_w[2]):
                best_w = bt
        if best_w:
            rw = {"entry": entry_signals(B, best_w[0], B[0]["o"]), "exit": exit_signals(B, best_w[1], best_w[0]),
                  "gain_pct": round(best_w[2] * 100, 2)}
        bd = best_trade(B, 0, 345)
        rd = {"entry": entry_signals(B, bd[0], B[0]["o"]), "exit": exit_signals(B, bd[1], bd[0]),
              "gain_pct": round(bd[2] * 100, 2)} if bd else None
        case = {
            "source": "sip_bot" if source == "sip" else "trade_bot",
            "date": day, "symbol": sym, "watched": win,
            "day": {"open": B[0]["o"], "high": max(x["h"] for x in B), "low": min(x["l"] for x in B),
                    "close": B[-1]["c"], "open_to_close_pct": round((B[-1]["c"] / B[0]["o"] - 1) * 100, 2)},
            "bot": {"entry_time": hm(ein), "entry_price": t["entry_price"], "exit_time": hm(eout),
                    "exit_price": t["exit_price"], "pl": t["pl_dollars"], "qty": t["qty"],
                    "exit_reason": t["exit_reason"][:90], "entry_reasons": t.get("reasons"),
                    "entry_signals": entry_signals(B, ei, B[0]["o"]),
                    "mfe_price": B[mfe_j]["h"], "mfe_time": hm(B[mfe_j]["t"]),
                    "mfe_pct": round((B[mfe_j]["h"] / t["entry_price"] - 1) * 100, 2),
                    "mae_pct": round((mae / t["entry_price"] - 1) * 100, 2),
                    "hold_to_close_pl": round((B[-1]["c"] - t["entry_price"]) * t["qty"], 2)},
            "right_watched": rw, "right_day": rd,
        }
        cases.append(case)
        if charts:
            chart(sym, day, B, t, case, OUT / f"{day}_{sym}_{hm(ein).replace(':', '')}.png")
        print(f"{sym}: bot {t['pl_dollars']:+.2f}  MFE {case['bot']['mfe_pct']:+.2f}% at {case['bot']['mfe_time']}  "
              f"hold {case['bot']['hold_to_close_pl']:+.2f}  | right (watched) "
              + (f"{rw['entry']['time']} {rw['entry']['price']:.2f} -> {rw['exit']['time']} {rw['exit']['price']:.2f} "
                 f"{rw['gain_pct']:+.2f}% [{', '.join(rw['entry']['tags'])} | {', '.join(rw['exit']['tags'])}]" if rw else "-")
              + (f"  | right (day) {rd['entry']['time']}->{rd['exit']['time']} {rd['gain_pct']:+.2f}%" if rd else ""), flush=True)
    CASES.parent.mkdir(parents=True, exist_ok=True)
    src = "sip_bot" if source == "sip" else "trade_bot"
    old = [json.loads(l) for l in open(CASES)] if CASES.exists() else []
    keep = [c for c in old if not (c["date"] == day and c.get("source", "trade_bot") == src)]
    for c in keep:   # keep hand-written lessons when a day is re-run
        pass
    lessons = {(c["date"], c["symbol"], c["bot"]["entry_time"]): c for c in old
               if c["date"] == day and c.get("source", "trade_bot") == src and c.get("suggestion")}
    for c in cases:
        o = lessons.get((c["date"], c["symbol"], c["bot"]["entry_time"]))
        if o:
            for k in ("what_happened", "suggestion", "rule_candidates"):
                c[k] = o[k]
    with open(CASES, "w") as f:
        for c in keep + cases:
            f.write(json.dumps(c, default=str) + "\n")
    if source != "sip":
        json.dump(cases, open(OUT / f"{day}_cases.json", "w"), indent=1, default=str)
    print(f"{day} {src}: {len(cases)} cases -> {CASES}", flush=True)
    return cases


if __name__ == "__main__":
    src = "sip" if "--source" in sys.argv and sys.argv[sys.argv.index("--source") + 1] == "sip" else "trade"
    charts = "--no-charts" not in sys.argv
    if sys.argv[1] == "all":
        folder = SIP if src == "sip" else TRADE / "data" / "trades"
        days = sorted(p.name[:10] for p in folder.glob("*_trades.jsonl"))
    else:
        days = [sys.argv[1]]
    for d in days:
        main(d, src, charts)
