"""
sf2_batch.py -- [2026-10-01] STAR FOLLOW v2 over many stock-days:
  - references' marks = the user-approved CONFIRMED marks (day_labels_confirmed.jsonl)
  - signals compared with NORMAL for the minute: normal = 150 x the share of all library days
    with a marked buy (or a sell, with the buy before the window) in that 5-minute window
      BUY  when reference buys  >= max(MIN_N, RATIO x normal buys)
      SELL when reference sells >= max(MIN_N, RATIO x normal sells)   (refs already in their trade)
    (MIN_N = 6, RATIO = 2.0 -- fixed before testing)
  - exit variants (same signals):
      refs      Star Follow exit only (+ 15:55)
      band      + sell when the close is below the references' 25th-percentile band for 3 minutes
      lowstop   + stop 1c under the lowest low of the 15 minutes before the buy
Only days before the traded day are references. Cost 0.1% per trade.
    python3 sf2_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

LABELS = Path("/var/www/screener/trade/data/playbook/day_labels_confirmed.jsonl")
K, WIN, MIN_N, RATIO, COST = 150, 5, 6, 2.0, 0.1
VARIANTS = ("refs", "band", "lowstop")


def load_marks():
    marks = {}
    for l in map(json.loads, open(LABELS)):
        marks[f"{l['date']}|{l['symbol']}"] = [(x["entry_m"], x["exit_m"]) for x in (l.get("buy"), l.get("leg2")) if x]
    return marks


def normal_rates(marks, n_days):
    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
    return buy / n_days * K, sell / n_days * K


def signals(lib, marks, day, a):
    lib.start(day)
    out = []
    for m in range(5, 385):
        idx, _ = lib.match(a["F"], m, day, K)
        w0 = m - WIN + 1
        buys = sells = 0
        for i in idx:
            for bm, sm in marks.get(lib.keys[i], ()):
                if w0 <= bm <= m:
                    buys += 1
                if bm < w0 and w0 <= sm <= m:
                    sells += 1
        out.append((m, buys, sells, idx))
    return out


def play(sig, a, lib, nb, ns, variant):
    pos, trades = None, []
    below = 0
    for m, buys, sells, idx in sig:
        if pos is None:
            if m <= 345 and buys >= max(MIN_N, RATIO * nb[m]):
                e = float(a["O"][m + 1])
                st = float(min(a["L"][max(0, m - 14):m + 1])) - 0.01 if variant == "lowstop" else -1.0
                band = np.quantile(lib.C[idx, m:386] / lib.C[idx, m][:, None], 0.25, axis=0) if variant == "band" else None
                pos = {"buy_m": m + 1, "buy": e, "stop": st, "band": band, "m0": m, "low": e}
                below = 0
            continue
        pos["low"] = min(pos["low"], float(a["L"][m]))
        why = None
        if variant == "lowstop" and a["L"][m] <= pos["stop"]:
            why, px = "stop", pos["stop"]
        elif sells >= max(MIN_N, RATIO * ns[m]):
            why, px = "refs exit", float(a["O"][m + 1])
        elif variant == "band":
            k = m - pos["m0"]
            lo = pos["band"][min(k, len(pos["band"]) - 1)]
            below = below + 1 if a["C"][m] / a["C"][pos["m0"]] < lo else 0
            if below >= 3:
                why, px = "left band", float(a["O"][m + 1])
        if why:
            trades.append({"buy_m": pos["buy_m"], "buy": pos["buy"], "sell_m": m + 1, "sell": px, "why": why, "low": pos["low"]})
            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"]})
    for t in trades:
        t["pl"] = (t["sell"] / t["buy"] - 1) * 100 - COST
        t["dd"] = (t["low"] / t["buy"] - 1) * 100
    return trades


def main(rng):
    a_, b_ = rng.split("..")
    lib = RT.Lib()
    marks = load_marks()
    nb, ns = normal_rates(marks, len(marks))
    print("normal reference buys per 150 at 9:35/10:00/11:00/13:00:",
          [round(nb[m], 1) for m in (5, 30, 90, 210)], "sells:", [round(ns[m], 1) for m in (30, 90, 210, 330)], flush=True)
    days = sorted({k.split("|")[0] for k in lib.keys if a_ <= k.split("|")[0] <= b_})
    res = {v: [] for v in VARIANTS}
    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 = RT.day_arrays(d["symbols"][s]["bars"], ref)
            if a is None:
                continue
            sig = signals(lib, marks, day, a)
            for v in VARIANTS:
                for t in play(sig, a, lib, nb, ns, v):
                    res[v].append({"day": day, "sym": s, **t})
        RT.REF._CACHE.clear()
        print(day, {v: len(r) for v, r in res.items()}, flush=True)
    json.dump(res, open(HERE / f"sf2_batch_{a_}_{b_}.json", "w"), default=float)
    half = days[len(days) // 2]
    print(f"\n{len(days)} days ({days[0]}..{days[-1]})")
    print(f"{'variant':>8} {'trades':>6} {'win%':>5} {'avg':>7} {'median':>7} {'sum':>8} {'1st/2nd half':>17} {'days+':>6} {'worst dip':>9} {'worst trade':>11} {'exits':>30}")
    for v, T in res.items():
        if not T:
            print(f"{v:>8}      0"); continue
        p = np.array([t["pl"] for t in T]); 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"]
        why = {}
        for t in T:
            why[t["why"]] = why.get(t["why"], 0) + 1
        print(f"{v:>8} {len(p):6} {(p > 0).mean():5.0%} {p.mean():+6.2f}% {np.median(p):+6.2f}% {p.sum():+7.1f}% "
              f"{p[h].sum():+7.1f}%/{p[~h].sum():+.1f}% {sum(x > 0 for x in byd.values()):3}/{len(byd)} "
              f"{min(t['dd'] for t in T):+8.1f}% {p.min():+10.1f}% {str(why):>30}")


if __name__ == "__main__":
    main(sys.argv[1])
