"""
speed_check.py -- [2026-09-28] Speed / shape measures (after the AGEN check) on any
stock-day, with three signals fixed BEFORE looking at new stocks:

  TURN   (entry)   15-bar acceleration >= 0.5 ATR, straightness >= 0.6, 15-bar R^2 >= 0.7,
                   15-min volume >= 2x its 14-day normal, price above VWAP but <= 2% above
  SMOOTH (hold)    30-bar R^2 >= 0.8 (steady climb)
  CLIMAX (avoid)   5-bar acceleration >= 1.5 ATR, 15-min volume >= 5x normal,
                   >= 4% above VWAP, move from the open >= 3x its normal

For each signal: what price did afterwards (best in the next 60 min, and to the close).
    python3 speed_check.py 2026-09-28 TGB IMMX NWL SG
"""
import json
import sys
from datetime import datetime

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

sys.path.insert(0, "/var/www/screener/trade")
sys.path.insert(0, "/var/www/screener/trade/reports/agen")
import day_report
import reference as REF

OUT = "/var/www/screener/trade/reports/agen"


def measures(day, sym):
    ref = REF.load(day, sym)
    if ref is None:
        REF.build(datetime.fromisoformat(day).date(), [sym], verbose=False)
        REF._CACHE.pop((day, sym), None)
        ref = REF.load(day, sym)
    B = [r for r in day_report.load_minutes(sym, day) if r["session"] == "reg"][:390]
    c = [float(r["close"]) for r in B]
    h = [float(r["high"]) for r in B]
    l = [float(r["low"]) for r in B]
    v = [float(r["volume"]) for r in B]
    vw = [float(r["vwap_regular"]) for r in B]
    t = [r["time_et"] for r in B]
    idx = [(int(x[:2]) - 9) * 60 + int(x[3:]) - 30 for x in t]
    o = float(B[0]["open"])
    atr, a = [None] * len(B), None
    for i in range(1, len(B)):
        tr = max(h[i] - l[i], abs(h[i] - c[i - 1]), abs(l[i] - c[i - 1]))
        a = tr if a is None else a + (tr - a) / 14
        atr[i] = a

    def win(i, n):
        if i < n or not atr[i]:
            return None
        cs = c[i - n:i + 1]
        steps = [cs[k + 1] - cs[k] for k in range(n)]
        net, path = cs[-1] - cs[0], sum(abs(s) for s in steps)
        xm, ym = n / 2, sum(cs) / len(cs)
        sxx = sum((k - xm) ** 2 for k in range(len(cs)))
        slope = sum((k - xm) * (y - ym) for k, y in enumerate(cs)) / sxx
        sst = sum((y - ym) ** 2 for y in cs)
        ssr = sum((y - (ym + slope * (k - xm))) ** 2 for k, y in enumerate(cs))
        r2 = (1 - ssr / sst if sst else 0) * (1 if slope > 0 else -1)
        half = n // 2
        acc = (sum(steps[half:]) / len(steps[half:]) - sum(steps[:half]) / half) / atr[i]
        norm = sum(ref["vol_per_min"][m] for m in idx[i - n + 1:i + 1] if m < 390)
        return {"eff": net / path if path else 0, "r2": r2, "acc": acc,
                "vrel": sum(v[i - n + 1:i + 1]) / norm if norm else None}

    rows = []
    for i in range(len(B)):
        m = min(idx[i], 389)
        r = {"t": t[i], "c": c[i], "vwap": vw[i], "vwap_pct": (c[i] / vw[i] - 1) * 100,
             "noise": abs(c[i] / o - 1) * 100 / max(ref["abs_move_pct"][m], 1e-6),
             "rvol": sum(v[:i + 1]) / max(ref["cum_vol"][m], 1)}
        for n in (5, 15, 30):
            w = win(i, n)
            if w:
                r.update({f"{k}{n}": x for k, x in w.items()})
        fut = c[i + 1:i + 61]
        r["fwd60"] = (max(fut) / c[i] - 1) * 100 if fut else None
        r["to_close"] = (c[-1] / c[i] - 1) * 100
        r["turn"] = bool(r.get("acc15") is not None and r["acc15"] >= 0.5 and r["eff15"] >= 0.6 and r["r215"] >= 0.7
                         and (r["vrel15"] or 0) >= 2 and 0 <= r["vwap_pct"] <= 2)
        r["smooth"] = bool(r.get("r230") is not None and r["r230"] >= 0.8)
        r["climax"] = bool(r.get("acc5") is not None and r["acc5"] >= 1.5 and (r.get("vrel15") or 0) >= 5
                           and r["vwap_pct"] >= 4 and r["noise"] >= 3)
        rows.append(r)
    return rows, o


def firsts(rows, key, gap=15):
    """First minute of each cluster of a signal (clusters separated by `gap` minutes)."""
    out, last = [], -999
    for i, r in enumerate(rows):
        if r[key]:
            if i - last > gap:
                out.append(r)
            last = i
    return out


def chart(day, sym, rows, o):
    x = [datetime.strptime(day + " " + r["t"], "%Y-%m-%d %H:%M") for r in rows]
    g = lambda k: [r.get(k) for r in rows]
    fig, ax = plt.subplots(4, 1, figsize=(11, 8.5), sharex=True, gridspec_kw={"height_ratios": [3, 1.1, 1.1, 1.1]})
    ax[0].plot(x, g("c"), color="#1f3a5f", lw=1, label="price")
    ax[0].plot(x, g("vwap"), color="#c0392b", lw=1, label="VWAP")
    for key, col, lab in (("turn", "green", "TURN"), ("climax", "purple", "CLIMAX")):
        for r in firsts(rows, key):
            tm = datetime.strptime(day + " " + r["t"], "%Y-%m-%d %H:%M")
            ax[0].scatter([tm], [r["c"]], color=col, s=35, zorder=5)
            ax[0].annotate(f"{lab} {r['t']}", (tm, r["c"]), xytext=(4, 6), textcoords="offset points", fontsize=7, color=col)
    sm = [xx for xx, r in zip(x, rows) if r["smooth"]]
    for xx in sm:
        ax[0].axvspan(xx, xx, color="#2a78d6", alpha=0.08)
    ax[0].set_title(f"{sym} {day} — TURN (green), CLIMAX (purple), SMOOTH trend (blue shading)", fontsize=10)
    ax[0].legend(fontsize=7, loc="upper left")
    ax[0].grid(alpha=.3)
    ax[1].plot(x, g("r230"), color="#2a78d6", lw=1)
    ax[1].axhline(0.8, ls=":", color="gray")
    ax[1].set_ylim(-1, 1)
    ax[1].set_ylabel("smoothness\n30-bar R²", fontsize=8)
    ax[2].plot(x, g("acc15"), color="#eb6834", lw=1, label="15-bar")
    ax[2].plot(x, g("acc5"), color="#999", lw=.6, label="5-bar")
    ax[2].axhline(0, color="gray", lw=.5)
    ax[2].set_ylabel("acceleration", fontsize=8)
    ax[2].legend(fontsize=6)
    ax[3].plot(x, g("vrel15"), color="#0ca30c", lw=1, label="15-min vol vs normal")
    ax[3].plot(x, g("noise"), color="#6b3fa0", lw=1, label="move vs normal")
    ax[3].axhline(1, ls=":", color="gray")
    ax[3].legend(fontsize=6)
    for a in ax[1:]:
        a.grid(alpha=.3)
    ax[3].xaxis.set_major_formatter(md.DateFormatter("%H:%M"))
    plt.tight_layout()
    p = f"{OUT}/{sym}_speed_shape_{day}.png"
    plt.savefig(p, dpi=110)
    plt.close()
    return p


if __name__ == "__main__":
    day = sys.argv[1]
    for sym in sys.argv[2:]:
        rows, o = measures(day, sym)
        chart(day, sym, rows, o)
        cl = rows[-1]["c"]
        print(f"\n=== {sym}: open {o:.2f} close {cl:.2f} ({(cl/o-1)*100:+.1f}%)")
        for key in ("turn", "climax"):
            for r in firsts(rows, key):
                print(f"  {key.upper():6} {r['t']} @ {r['c']:.2f}  vsVWAP {r['vwap_pct']:+.1f}%  move/normal {r['noise']:.1f}  "
                      f"vol15 {r.get('vrel15') or 0:.1f}x  -> best next 60m {r['fwd60']:+.1f}%, to close {r['to_close']:+.1f}%")
        sm = [r for r in rows if r["smooth"]]
        ns = [r for r in rows if not r["smooth"] and r.get("r230") is not None]
        avg = lambda L, k: sum(x[k] for x in L if x[k] is not None) / max(1, len([x for x in L if x[k] is not None]))
        print(f"  SMOOTH minutes: {len(sm)} ({', '.join(sorted({r['t'][:2] + ':00' for r in sm}))}); "
              f"avg to-close from a smooth minute {avg(sm,'to_close'):+.2f}% vs {avg(ns,'to_close'):+.2f}% otherwise")
        print(f"  RVOL so far: 9:45 {next(r['rvol'] for r in rows if r['t']>='09:45'):.1f}x, 11:00 {next(r['rvol'] for r in rows if r['t']>='11:00'):.1f}x, close {rows[-1]['rvol']:.1f}x")
