"""
group_charts.py -- [2026-10-01] The user: sort the marked charts (22-day list stocks +
sip_bot stocks) into groups with a similar pattern; how many groups, how many per group.

Each stock-day -> its whole day as a vector, every 5 minutes 9:30-15:55 (78 points each):
   price vs the open %, price vs VWAP %, 0.5 x log(1 + volume vs its 14-day normal)
k-means (numpy, k-means++ start, 10 restarts) for k = 3..14; the elbow of the within-group
spread picks k. Per group: size, average curve, open->close, best holdable gain and where
its buy/sell sit (from day_labels / sip_day_labels), and the bots' P/L in the group.
Output: reports/playbook/groups/ (one chart per group + overview) and groups.json.
"""
import gzip
import json
import math
import sys
from collections import Counter
from datetime import datetime
from pathlib import Path
from zoneinfo import ZoneInfo

import numpy as np
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt

HERE = Path(__file__).resolve().parent
TRADE = Path("/var/www/screener/trade")
sys.path.insert(0, str(TRADE))
sys.path.insert(0, str(HERE))
import reference as REF
import label_days as LD
from alpaca_client import get_client

ET = ZoneInfo("America/New_York")
OUT = HERE / "groups"
STEP = 5


def vec(B, ref=None):
    """[shape version] price paths divided by the stock's normal daily move (ATR14 % of
    the previous close), so groups follow the SHAPE of the day, not its size."""
    by = {x["m"]: x for x in B}
    o = B[0]["o"]
    atrp = (ref["atr14"] / ref["prev_close"] * 100) if ref and ref.get("atr14") and ref.get("prev_close") else 5.0
    P, W, X = [], [], []
    last = B[0]
    for m in range(390):
        last = by.get(m, last)
        if m % STEP == 0:
            P.append((last["c"] / o - 1) * 100 / atrp)
            W.append((last["c"] / last["vwap"] - 1) * 100 / atrp)
            win = [by[k]["volx"] for k in range(m, m + STEP) if k in by]
            X.append(0.5 * math.log1p(sum(win) / len(win)) if win else 0.0)
    return np.array(P + W + X)


def kmeans(X, k, restarts=10, iters=60, seed=1):
    rng = np.random.default_rng(seed)
    best = None
    for _ in range(restarts):
        C = [X[rng.integers(len(X))]]
        for _ in range(1, k):
            d = np.min(((X[:, None, :] - np.array(C)[None]) ** 2).sum(-1), axis=1)
            C.append(X[rng.choice(len(X), p=d / d.sum())])
        C = np.array(C)
        for _ in range(iters):
            lab = ((X[:, None, :] - C[None]) ** 2).sum(-1).argmin(1)
            newC = np.array([X[lab == j].mean(0) if (lab == j).any() else C[j] for j in range(k)])
            if np.allclose(newC, C):
                break
            C = newC
        inertia = ((X - C[lab]) ** 2).sum()
        if best is None or inertia < best[0]:
            best = (inertia, lab, C)
    return best


def collect():
    items = []
    labels = {(l["date"], l["symbol"]): l for l in map(json.loads, open(LD.OUT)) if l["date"] >= "2026-08-27"}
    trades = {}
    for p in sorted((TRADE / "data" / "trades").glob("*_trades.jsonl")):
        for l in open(p):
            t = json.loads(l)
            trades.setdefault((p.name[:10], t["symbol"]), []).append(t.get("pl_dollars") or 0)
    for p in sorted(LD.BARS.glob("bars_2026-*.json.gz")):
        d = json.load(gzip.open(p, "rt"))
        if d["date"] < "2026-08-27":
            continue
        for s in d["top30_current"]:
            ref = REF.load(d["date"], s)
            lab = labels.get((d["date"], s))
            if not ref or not lab:
                continue
            B = LD.build(d["symbols"][s]["bars"], ref)
            items.append({"date": d["date"], "symbol": s, "set": "22d", "lab": lab, "v": vec(B, ref),
                          "bot_pl": round(sum(trades.get((d["date"], s), [])), 2), "bot_traded": (d["date"], s) in trades})
        REF._CACHE.clear()
    have = {(i["date"], i["symbol"]) for i in items}
    client = get_client()
    for lab in map(json.loads, open(TRADE / "data" / "playbook" / "sip_day_labels.jsonl")):
        if (lab["date"], lab["symbol"]) in have:
            for i in items:
                if (i["date"], i["symbol"]) == (lab["date"], lab["symbol"]):
                    i["sip_pl"] = lab["bot_pl"]
            continue
        ref = REF.load(lab["date"], lab["symbol"])
        if not ref:
            continue
        dd = datetime.fromisoformat(lab["date"]).replace(tzinfo=ET)
        raw = client.get_minute_bars(lab["symbol"], start=dd.replace(hour=9, minute=30), end=dd.replace(hour=16), limit=1000)
        rows = [[int(b.timestamp.timestamp()), float(b.open), float(b.high), float(b.low), float(b.close), float(b.volume)]
                for b in raw]
        B = LD.build(rows, ref)
        if len(B) < 100:
            continue
        items.append({"date": lab["date"], "symbol": lab["symbol"], "set": "sip", "lab": lab, "v": vec(B, ref),
                      "bot_pl": 0.0, "bot_traded": False, "sip_pl": lab["bot_pl"]})
    return items


def hm(m):
    return f"{9 + (30 + m) // 60}:{(30 + m) % 60:02d}"


def main():
    OUT.mkdir(exist_ok=True)
    items = collect()
    X = np.array([i["v"] for i in items])
    print(f"{len(items)} stock-days ({sum(i['set'] == '22d' for i in items)} list, {sum(i['set'] == 'sip' for i in items)} sip-only)")
    curve = {}
    for k in range(5, 12):
        curve[k] = kmeans(X, k)
        print(f"   k={k:2}: within-group spread {curve[k][0]:.0f}", flush=True)
    # elbow: largest drop of the spread's improvement
    ks = sorted(curve)
    gains = {k: curve[k - 1][0] - curve[k][0] for k in ks[1:]}
    k_best = 8   # [fixed] spread drops ~10% per extra group up to 8, less after; small groups = rare patterns
    inertia, lab, C = curve[k_best]
    print(f"chosen k = {k_best}")
    order = sorted(range(k_best), key=lambda j: -(lab == j).sum())
    groups = []
    t = np.arange(0, 390, STEP)
    n = len(t)
    fig, axs = plt.subplots(math.ceil(k_best / 3), 3, figsize=(13, 3.1 * math.ceil(k_best / 3)), sharey=True)
    axs = axs.flatten()
    for g, j in enumerate(order, 1):
        mem = [items[i] for i in np.where(lab == j)[0]]
        P = np.array([m["v"][:n] for m in mem])
        bh = [m["lab"]["best_held"] for m in mem if m["lab"]["best_held"]]
        oc = [m["lab"]["open_to_close_pct"] for m in mem]
        traded = [m for m in mem if m["bot_traded"] or "sip_pl" in m]
        info = {"group": g, "n": len(mem), "list_days": sum(m["set"] == "22d" for m in mem),
                "avg_open_to_close": round(float(np.mean(oc)), 2), "up_days_pct": round(sum(x > 0 for x in oc) / len(oc) * 100),
                "best_held_median": round(float(np.median([b["gain_pct"] for b in bh])), 2) if bh else None,
                "best_buy_time_median": hm(int(np.median([b["entry_m"] for b in bh]))) if bh else None,
                "best_sell_time_median": hm(int(np.median([b["exit_m"] for b in bh]))) if bh else None,
                "buy_tags": Counter(tg.split(" (")[0] for b in bh for tg in b["entry"]["tags"]).most_common(3),
                "sell_tags": Counter(tg.split(" (")[0].split(" 1")[0].split(" 2")[0].split(" 3")[0].split(" 4")[0]
                                     for b in bh for tg in b["exit"]["tags"]).most_common(3),
                "bots_traded": len(traded),
                "trade_bot_pl": round(sum(m["bot_pl"] for m in mem), 2),
                "sip_bot_pl": round(sum(m.get("sip_pl", 0) for m in mem), 2),
                "examples": [f"{m['symbol']} {m['date'][5:]}" for m in sorted(mem, key=lambda m: ((m['v'] - C[j]) ** 2).sum())[:6]]}
        groups.append(info)
        ax = axs[g - 1]
        for row in P[:80]:
            ax.plot(t, row, color="#9bb7d4", lw=.5, alpha=.5)
        ax.plot(t, P.mean(0), color="black", lw=2)
        ax.axhline(0, color="grey", lw=.6)
        ax.set_title(f"G{g}: {len(mem)} days, avg {info['avg_open_to_close']:+.1f}% o->c, best {info['best_held_median']}%", fontsize=8.5)
        ax.set_ylabel("price vs open, in daily ATRs", fontsize=7)
        ax.set_xticks([0, 60, 150, 270, 385])
        ax.set_xticklabels(["9:30", "10:30", "12:00", "14:00", "16:00"], fontsize=7)
        ax.set_ylim(-4, 4)
        ax.grid(alpha=.3)
        # one chart per group
        f2, a2 = plt.subplots(figsize=(9, 3.4))
        for row in P:
            a2.plot(t, row, color="#9bb7d4", lw=.5, alpha=.45)
        a2.plot(t, P.mean(0), color="black", lw=2.2, label="group average")
        a2.axhline(0, color="grey", lw=.6)
        a2.set_title(f"Group {g}: {len(mem)} stock-days -- price vs open, in the stock's normal daily moves (ATR)", fontsize=9)
        a2.set_xticks([0, 30, 90, 150, 210, 270, 330, 385])
        a2.set_xticklabels(["9:30", "10:00", "11:00", "12:00", "13:00", "14:00", "15:00", "16:00"], fontsize=7)
        a2.grid(alpha=.3)
        plt.tight_layout()
        f2.savefig(OUT / f"group_{g}.png", dpi=105)
        plt.close(f2)
    for ax in axs[k_best:]:
        ax.axis("off")
    plt.tight_layout()
    fig.savefig(OUT / "groups_overview.png", dpi=105)
    json.dump({"k": k_best, "spread": {k: round(float(curve[k][0])) for k in ks}, "groups": groups,
               "members": [{"date": i["date"], "symbol": i["symbol"], "group": order.index(int(lab[n_])) + 1}
                           for n_, i in enumerate(items)]},
              open(OUT / "groups.json", "w"), indent=1, default=str)
    for gi in groups:
        print(json.dumps(gi, default=str))


if __name__ == "__main__":
    main()
