"""
sf3_batch.py -- [2026-10-01] STAR FOLLOW v3, the user's rules (approved 2026-10-01):
  references : WINNERS only (data/playbook/winners.json: closed >= +2% over the open with a
               confirmed buy mark), only days before the traded day; marks = confirmed marks
  BUY        : reference buys in the last 5 min >= max(6, 2 x normal for that minute)
  HOLD/EXIT  : every minute count the references that STILL MATCH (whole-day distance <= the
               150th-closest distance at the buy). If fewer than 10 match for 2 minutes in a row:
                 not climbing (close <= close 2 min earlier) -> sell next minute's open
                 climbing -> hold, sell at the next open after the first flat/lower 1-min close
               While >= 10 match: also sell on the references' sell signal (>= max(6, 2 x normal))
               15:55 sell-all. No stop. Cost 0.1% per trade.
    python3 sf3_batch.py 2026-08-27..2026-10-01
"""
import gzip
import json
import sys
from pathlib import Path

import numpy as np

HERE = Path(__file__).resolve().parent
sys.path.insert(0, str(HERE)); sys.path.insert(0, "/var/www/screener/trade")
import reference_trader as RT
import lib_v2 as L2
USE_V2 = "--v2" in sys.argv

WINNERS = Path("/var/www/screener/trade/data/playbook/winners.json")
K, WIN, MIN_N, RATIO, COST, MIN_MATCH, OUT_MIN = 150, 5, 6, 2.0, 0.1, 10, 2
# [2026-10-02] MODE "band": "still matching" = the stock stays at/above the lower BAND_Q quantile of
# the buy-time references' paths SINCE THE BUY (CHPT 10/2: the whole-day count needed 27 min)
MODE, BAND_Q = "count", 0.25
# [2026-10-02] MODE "since" (the user's version): every minute after the buy, search again among the
# POOL winner days that matched best up to the buy for days whose price moved like the stock SINCE THE
# BUY (RMS difference of % change since the buy <= SINCE_RMS); fewer than SINCE_MIN for 2 min -> diverged
POOL, SINCE_RMS, SINCE_MIN = 500, 0.5, 50


def since_count(lib, pool, a, m0, m):
    """How many pool days moved like the stock from minute m0 to m (RMS of % change difference)."""
    s = a["C"][m0:m + 1] / a["C"][m0] - 1
    r = lib.C[pool, m0:m + 1] / lib.C[pool, m0][:, None] - 1
    rms = np.sqrt(((r - s[None, :]) ** 2).mean(1)) * 100
    return int((rms <= SINCE_RMS).sum())


def load():
    W = json.load(open(WINNERS))
    marks = {f"{w['date']}|{w['symbol']}": [(x["entry_m"], x["exit_m"]) for x in (w.get("buy"), w.get("leg2")) if x] for w in W}
    return marks


def normal_rates(marks):
    buy = np.zeros(390); sell = np.zeros(390)
    for mk in marks.values():
        for m in range(390):
            w0 = m - WIN + 1
            if any(w0 <= bm <= m for bm, sm in mk):
                buy[m] += 1
            if any(bm < w0 and w0 <= sm <= m for bm, sm in mk):
                sell[m] += 1
    n = len(marks)
    return buy / n * K, sell / n * K


def run_day(lib, wmask, marks, nb, ns, day, a, can_buy=None):
    ok = np.where((lib.dates < day) & wmask)[0]
    sub = lib.F[ok]
    D = np.zeros(len(ok))
    keys = [lib.keys[i] for i in ok]
    pos, trades = None, []
    for m in range(0, 385):
        D += ((sub[:, m, :] - a["F"][None, m, :]) ** 2).sum(1)
        if m < 5:
            continue
        mean = D / (m + 1)
        order = np.argpartition(mean, K)[:K]
        w0 = m - WIN + 1
        buys = sells = 0
        for i in order:
            for bm, sm in marks.get(keys[i], ()):
                if w0 <= bm <= m:
                    buys += 1
                if bm < w0 and w0 <= sm <= m:
                    sells += 1
        if pos is None:
            if m <= 345 and (can_buy is None or can_buy(m)) and buys >= max(MIN_N, RATIO * nb[m]):
                thr = float(np.sort(mean[order])[-1])
                refs = ok[order]
                pool = ok[np.argpartition(mean, POOL)[:POOL]]
                band = np.quantile(lib.C[refs, m:386] / lib.C[refs, m][:, None], BAND_Q, axis=0)
                pos = {"m0": m, "band": band, "pool": pool, "buy_m": m + 1, "buy": float(a["O"][m + 1]), "thr": thr, "out": 0, "diverged": False,
                       "climb": False, "low": float(a["O"][m + 1]), "min_match": K}
            continue
        pos["low"] = min(pos["low"], float(a["L"][m]))
        if MODE == "since":
            n_match = MIN_MATCH if since_count(lib, pos["pool"], a, pos["m0"], m) >= SINCE_MIN else 0
        elif MODE == "band":
            k = m - pos["m0"]
            n_match = MIN_MATCH if a["C"][m] / a["C"][pos["m0"]] >= pos["band"][min(k, len(pos["band"]) - 1)] else 0
        else:
            n_match = int((mean <= pos["thr"]).sum())
        pos["min_match"] = min(pos["min_match"], n_match)
        why = None
        if not pos["diverged"]:
            pos["out"] = pos["out"] + 1 if n_match < MIN_MATCH else 0
            if pos["out"] >= OUT_MIN:
                pos["diverged"] = True
                if a["C"][m] > a["C"][m - 2]:
                    pos["climb"] = True
                else:
                    why = "no longer matching, not climbing"
            elif sells >= max(MIN_N, RATIO * ns[m]):
                why = "references sell"
        elif pos["climb"] and a["C"][m] < min(a["L"][m - 5:m]):
            # [2026-10-02] the climb ends when a close breaks below the lowest low of the previous
            # 5 minutes (the last higher low) -- a single lower close is only a pause (ONT 10/2)
            why = "no longer matching, climb ended"
        if why:
            trades.append({"buy_m": pos["buy_m"], "buy": pos["buy"], "sell_m": m + 1, "sell": float(a["O"][m + 1]),
                           "why": why, "low": pos["low"], "min_match": pos["min_match"]})
            pos = None
    if pos:
        trades.append({"buy_m": pos["buy_m"], "buy": pos["buy"], "sell_m": 385, "sell": float(a["C"][385]), "why": "15:55",
                       "low": pos["low"], "min_match": pos["min_match"]})
    for t in trades:
        t["pl"] = (t["sell"] / t["buy"] - 1) * 100 - COST
        t["dd"] = (t["low"] / t["buy"] - 1) * 100
        t["held"] = t["sell_m"] - t["buy_m"]
    return trades


def main(rng):
    a_, b_ = rng.split("..")
    lib = L2.Lib() if USE_V2 else RT.Lib()
    marks = load()
    print("features:", "v2 (price, VWAP, volume + SPY, IWM, gap)" if USE_V2 else "v1 (price, VWAP, volume)", flush=True)
    wmask = np.array([k in marks for k in lib.keys])
    nb, ns = normal_rates(marks)
    print(f"winners in the library: {int(wmask.sum())}; normal buys at 9:35/10:00/11:00:",
          [round(nb[m], 1) for m in (5, 30, 90)], flush=True)
    days = sorted({k.split("|")[0] for k in lib.keys if a_ <= k.split("|")[0] <= b_})
    T = []
    for day in days:
        d = json.load(gzip.open(RT.BARS / f"bars_{day}.json.gz", "rt"))
        for s in d["top30_current"]:
            ref = RT.REF.load(day, s)
            if not ref:
                continue
            a = (L2.day_arrays(d["symbols"][s]["bars"], ref, day) if USE_V2 else RT.day_arrays(d["symbols"][s]["bars"], ref))
            if a is None:
                continue
            for t in run_day(lib, wmask, marks, nb, ns, day, a):
                T.append({"day": day, "sym": s, **t})
        RT.REF._CACHE.clear()
        print(day, len(T), flush=True)
    json.dump(T, open(HERE / f"sf3_batch_{a_}_{b_}_m{MIN_MATCH}{'_v2' if USE_V2 else ''}{'_' + MODE if MODE != 'count' else ''}.json", "w"), default=float)
    print(f"MIN_MATCH = {MIN_MATCH}, MODE = {MODE}")
    hm = lambda m: f"{9 + (30 + m) // 60}:{(30 + m) % 60:02d}"
    p = np.array([t["pl"] for t in T]); half = days[len(days) // 2]; h = np.array([t["day"] < half for t in T])
    byd = {}
    for t in T:
        byd[t["day"]] = byd.get(t["day"], 0) + t["pl"]
    print(f"\n{len(days)} days: {len(p)} trades, win {(p > 0).mean():.0%}, avg {p.mean():+.2f}%, median {np.median(p):+.2f}%, "
          f"sum {p.sum():+.1f}% (halves {p[h].sum():+.1f}% / {p[~h].sum():+.1f}%), positive days {sum(v > 0 for v in byd.values())}/{len(byd)}, "
          f"worst trade {p.min():+.1f}%, worst dip {min(t['dd'] for t in T):+.1f}%, median minutes held {int(np.median([t['held'] for t in T]))}")
    for w in sorted({t["why"] for t in T}):
        q = np.array([t["pl"] for t in T if t["why"] == w])
        print(f"   exit '{w}': {len(q)} trades, win {(q > 0).mean():.0%}, avg {q.mean():+.2f}%, sum {q.sum():+.1f}%")
    s = np.sort(p)
    print(f"   without the worst 10 trades: {s[10:].sum():+.1f}%")
    print("   worst:", [(t['day'][5:], t['sym'], hm(t['buy_m']), round(t['pl'], 1), t['why']) for t in sorted(T, key=lambda t: t['pl'])[:6]])


if __name__ == "__main__":
    if "--band" in sys.argv:
        MODE = "band"
    if "--since" in sys.argv:
        MODE = "since"
    if len(sys.argv) > 2 and sys.argv[2].isdigit():
        MIN_MATCH = int(sys.argv[2])          # [2026-10-01] e.g. 80 (user)
    main(sys.argv[1])
