"""
reference_trader.py -- [2026-10-01] The user's playbook-reference trading idea, as a replay:
every minute, match the stock's WHOLE day so far (from 9:30) against library stock-days at
the same minute, follow the closest group, buy where they say up, hold toward their target,
sell when the stock leaves their band (version A) or the re-match turns down (version B).

Library: every stock-day in reports/open_window/bars (~2 years, 9:28/9:30 lists) with a
14-day reference file. Only days BEFORE the test day are used as references (no look-ahead).
Per stock-day, per minute (forward-filled), scaled by the stock's daily ATR% (shape, not size):
    P = price vs open / ATR%,  W = price vs VWAP / ATR%,  X = 0.5*log(1 + volume vs normal)
Distance at minute m = mean over minutes 0..m of (dP^2 + dW^2 + dX^2)  (whole day so far)

Rules (defaults; --set key=value):
  k              150 closest references
  start_m        1   (9:31) first possible buy;  last_buy_m 345 (15:15)
  min_p_up       0.40  buy if share of references that went +2% before -1% from here >= this
  min_edge       0.10  ... and that share exceeds the -1%-first share by at least this much
  band_q         (0.25, 0.75) the references' price band after the buy (relative to minute m)
  band_minutes   3    version A: sell when below the band's lower edge this many minutes in a row
  target         median of the references' best gain from the buy minute to 15:55
  stop_pct       1.0  hard stop
  version        "A" (leave band -> sell) or "B" (leave band -> re-match; sell only if the new
                 group's -1%-first share > +2%-first share)
  max_trades     2 per stock-day
Outcomes use the test day's own 1-min bars: buy at the next bar's open, sell at the next
bar's open (stop: at the stop price), 0.1% round-trip cost.

    python3 reference_trader.py --days 2026-10-01 --symbols SVCO,ORN,LWLG,...   (test set)
    python3 reference_trader.py --days 2026-08-27..2026-09-30                   (whole lists)
"""
import argparse
import gzip
import json
import math
import sys
from datetime import datetime
from pathlib import Path
from zoneinfo import ZoneInfo

import numpy as np

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

ET = ZoneInfo("America/New_York")
BARS = TRADE / "reports" / "open_window" / "bars"
LIB_CACHE = TRADE / "data" / "playbook" / "library_paths.npz"
N = 390
DEF = {"k": 150, "start_m": 1, "last_buy_m": 345, "min_p_up": 0.40, "min_edge": 0.10,
       "band_lo": 0.25, "band_hi": 0.75, "band_minutes": 3, "stop_pct": 1.0, "version": "A",
       "max_trades": 2, "cost": 0.1}


def day_arrays(rows, ref):
    """rows [[ts,o,h,l,c,v]] -> dict of per-minute arrays (forward-filled) + features."""
    by = {}
    for b in rows:
        t = datetime.fromtimestamp(b[0], ET) if not isinstance(b[0], datetime) else b[0].astimezone(ET)
        m = (t.hour - 9) * 60 + t.minute - 30
        if 0 <= m < N:
            by[m] = b
    if len(by) < 100:
        return None
    o = by[min(by)][1]
    atrp = ref["atr14"] / ref["prev_close"] * 100 if ref.get("atr14") and ref.get("prev_close") else 5.0
    C, H, L, O = np.zeros(N), np.zeros(N), np.zeros(N), np.zeros(N)
    P, W, X = np.zeros(N, np.float32), np.zeros(N, np.float32), np.zeros(N, np.float32)
    last, pv, cv = o, 0.0, 0.0
    for m in range(N):
        b = by.get(m)
        v = 0.0
        if b:
            O[m], H[m], L[m], last, v = b[1], b[2], b[3], b[4], b[5]
            pv += (b[2] + b[3] + b[4]) / 3 * v
            cv += v
        else:
            O[m] = H[m] = L[m] = last
        C[m] = last
        vw = pv / cv if cv else last
        P[m] = (last / o - 1) * 100 / atrp
        W[m] = (last / vw - 1) * 100 / atrp
        X[m] = 0.5 * math.log1p(v / max(ref["vol_per_min"][m], 1))
    return {"o": o, "C": C, "H": H, "L": L, "O": O, "F": np.stack([P, W, X], 1)}


def outcomes_from(C, H, L, m):
    """For a reference day at minute m: +2% first? -1% first? best gain to 15:55, path rel to C[m]."""
    e = C[m]
    up = dn = False
    for k in range(m + 1, 386):
        if L[k] <= e * 0.99:
            dn = True
            break
        if H[k] >= e * 1.02:
            up = True
            break
    best = (H[m + 1:386].max() / e - 1) * 100 if m + 1 < 386 else 0.0
    return up, dn, best


def build_library():
    """All stock-days -> F (n, 390, 3) float32, C/H/L (n, 390) float32, keys."""
    if LIB_CACHE.exists():
        z = np.load(LIB_CACHE, allow_pickle=True)
        return z["F"], z["C"], z["H"], z["L"], list(z["keys"])
    F, C, H, L, keys = [], [], [], [], []
    for p in sorted(BARS.glob("bars_20*.json.gz")):
        d = json.load(gzip.open(p, "rt"))
        for s in d["top30_current"]:
            ref = REF.load(d["date"], s)
            if not ref:
                continue
            a = day_arrays(d["symbols"][s]["bars"], ref)
            if a is None:
                continue
            F.append(a["F"]); C.append(a["C"]); H.append(a["H"]); L.append(a["L"])
            keys.append(f"{d['date']}|{s}")
        REF._CACHE.clear()
        print("library", p.name[5:15], len(keys), flush=True)
    F = np.array(F, np.float32); C = np.array(C, np.float32); H = np.array(H, np.float32); L = np.array(L, np.float32)
    np.savez(LIB_CACHE, F=F, C=C, H=H, L=L, keys=np.array(keys))
    return F, C, H, L, keys


class Lib:
    def __init__(self):
        self.F, self.C, self.H, self.L, self.keys = build_library()
        self.dates = np.array([k.split("|")[0] for k in self.keys])
        # forward outcome tables are computed lazily per minute
        self._out = {}

    def outcomes(self, idx, m):
        """Vectorized: per reference (up first?, down first?, best gain %) from minute m."""
        e = self.C[idx, m][:, None]
        Hs, Ls = self.H[idx, m + 1:386], self.L[idx, m + 1:386]
        if Hs.shape[1] == 0:
            return [(False, False, 0.0)] * len(idx)
        up = Hs >= e * 1.02
        dn = Ls <= e * 0.99
        BIG = 10 ** 6
        fu = np.where(up.any(1), up.argmax(1), BIG)
        fd = np.where(dn.any(1), dn.argmax(1), BIG)
        best = (Hs.max(1) / e[:, 0] - 1) * 100
        return list(zip((fu < fd) & (fu < BIG), (fd <= fu) & (fd < BIG), best))

    def start(self, before_date):
        """Begin matching a test day: running whole-day distance to every earlier day."""
        self.ok = np.where(self.dates < before_date)[0]
        self.sub = self.F[self.ok]
        self.D = np.zeros(len(self.ok), np.float64)
        self.m_done = -1

    def match(self, f, m, before_date, k):
        """f: (390,3) test-day features; distance = mean over minutes 0..m (whole day so far).
        Running sum: each new minute adds one column."""
        while self.m_done < m:
            self.m_done += 1
            j = self.m_done
            self.D += ((self.sub[:, j, :] - f[None, j, :]) ** 2).sum(1)
        order = np.argpartition(self.D, k)[:k]
        order = order[np.argsort(self.D[order])]          # closest first
        return self.ok[order], float(self.D[order].max() / (m + 1))


def trade_day(lib, day, sym, a, c):
    """Replay one test stock-day minute by minute; returns trades."""
    trades = []
    pos = None
    below = 0
    lib.start(day)
    for m in range(c["start_m"], 385):
        f = a["F"]
        if pos is None:
            if m > c["last_buy_m"] or len(trades) >= c["max_trades"]:
                continue
            idx, _ = lib.match(f, m, day, c["k"])
            outc = lib.outcomes(idx, m)
            p_up = sum(o[0] for o in outc) / len(outc)
            p_dn = sum(o[1] for o in outc) / len(outc)
            if p_up >= c["min_p_up"] and p_up - p_dn >= c["min_edge"]:
                e = a["O"][m + 1]
                paths = np.array([lib.C[i][m:386] / lib.C[i][m] for i in idx])  # relative paths
                band_lo = np.quantile(paths, c["band_lo"], axis=0)
                target = float(np.median([o[2] for o in outc]))
                pos = {"m": m + 1, "entry": e, "stop": e * (1 - c["stop_pct"] / 100), "band_lo": band_lo,
                       "target": e * (1 + target / 100), "p_up": round(p_up, 2), "p_dn": round(p_dn, 2),
                       "refs": [lib.keys[i] for i in idx[:5]], "target_pct": round(target, 2)}
                below = 0
            continue
        # holding
        k = m - (pos["m"] - 1)
        exit_px, why = None, None
        if a["L"][m] <= pos["stop"]:
            exit_px, why = pos["stop"], "stop"
        elif a["H"][m] >= pos["target"]:
            exit_px, why = pos["target"], "target"
        else:
            rel = a["C"][m] / pos["entry"]
            lo = pos["band_lo"][min(k, len(pos["band_lo"]) - 1)]
            below = below + 1 if rel < lo else 0
            if below >= c["band_minutes"]:
                if c["version"] == "A":
                    exit_px, why = a["O"][m + 1], "left the reference band"
                else:
                    idx, _ = lib.match(a["F"], m, day, c["k"])
                    outc = lib.outcomes(idx, m)
                    p_up = sum(o[0] for o in outc) / len(outc)
                    p_dn = sum(o[1] for o in outc) / len(outc)
                    if p_dn > p_up:
                        exit_px, why = a["O"][m + 1], f"re-match says down ({p_dn:.2f} vs {p_up:.2f})"
                    else:
                        paths = np.array([lib.C[i][m:386] / lib.C[i][m] for i in idx])
                        pos["band_lo"] = np.concatenate([np.zeros(k), np.quantile(paths, c["band_lo"], axis=0)
                                                         * a["C"][m] / pos["entry"]])
                        pos["target"] = a["C"][m] * (1 + float(np.median([o[2] for o in outc])) / 100)
                        below = 0
        if exit_px is None and m >= 384:
            exit_px, why = a["C"][m], "end of day"
        if exit_px is not None:
            pl_pct = (exit_px / pos["entry"] - 1) * 100 - c["cost"]
            trades.append({"date": day, "symbol": sym, "buy_m": pos["m"], "buy": round(float(pos["entry"]), 4),
                           "sell_m": m + 1, "sell": round(float(exit_px), 4), "pl_pct": round(float(pl_pct), 2),
                           "why": why, "p_up": pos["p_up"], "p_dn": pos["p_dn"], "target_pct": pos["target_pct"],
                           "refs": pos["refs"]})
            pos = None
    return trades


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


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--days", required=True, help="YYYY-MM-DD or A..B")
    ap.add_argument("--symbols", default="", help="comma list (default: the day's list)")
    ap.add_argument("--set", action="append", default=[])
    ap.add_argument("--out", default="")
    args = ap.parse_args()
    c = {**DEF, **{kv.split("=", 1)[0]: json.loads(kv.split("=", 1)[1]) for kv in args.set}}
    lib = Lib()
    if ".." in args.days:
        a_, b_ = args.days.split("..")
        days = sorted({k.split("|")[0] for k in lib.keys if a_ <= k.split("|")[0] <= b_})
    else:
        days = [args.days]
    from alpaca_client import get_client
    client = None
    all_trades = []
    for day in days:
        syms = [s for s in args.symbols.split(",") if s]
        rowsets = {}
        p = BARS / f"bars_{day}.json.gz"
        dd = json.load(gzip.open(p, "rt")) if p.exists() else {"top30_current": [], "symbols": {}}
        if not syms:
            syms = dd["top30_current"]
        for s in syms:
            if s in dd["symbols"]:
                rowsets[s] = dd["symbols"][s]["bars"]
            else:
                client = client or get_client()
                d0 = datetime.fromisoformat(day).replace(tzinfo=ET)
                raw = client.get_minute_bars(s, start=d0.replace(hour=9, minute=30), end=d0.replace(hour=16), limit=1000)
                rowsets[s] = [[b.timestamp, float(b.open), float(b.high), float(b.low), float(b.close), float(b.volume)]
                              for b in raw]
        for s, rows in rowsets.items():
            ref = REF.load(day, s)
            if not ref:
                try:
                    REF.build(datetime.fromisoformat(day).date(), [s], verbose=False)
                    REF._CACHE.clear()
                    ref = REF.load(day, s)
                except Exception:
                    ref = None
            if not ref:
                print(day, s, "no reference"); continue
            a = day_arrays(rows, ref)
            if a is None:
                continue
            tr = trade_day(lib, day, s, a, c)
            all_trades += tr
            for t in tr:
                print(f"{day} {s:5} buy {hm(t['buy_m'])} {t['buy']:.2f} (p_up {t['p_up']}, p_dn {t['p_dn']}, target +{t['target_pct']}%) "
                      f"-> sell {hm(t['sell_m'])} {t['sell']:.2f} {t['pl_pct']:+.2f}% [{t['why']}]  refs {', '.join(t['refs'][:3])}", flush=True)
            if not tr:
                print(f"{day} {s:5} no trade", flush=True)
    n = len(all_trades)
    if n:
        w = sum(t["pl_pct"] > 0 for t in all_trades)
        print(f"\nTOTAL {n} trades, {w} wins, avg {np.mean([t['pl_pct'] for t in all_trades]):+.2f}% per trade, "
              f"sum {sum(t['pl_pct'] for t in all_trades):+.1f}% (after {c['cost']}% cost) | version {c['version']}")
    if args.out:
        json.dump(all_trades, open(args.out, "w"), indent=1, default=str)


if __name__ == "__main__":
    main()
