"""
sf_batch.py -- [2026-10-01] STAR FOLLOW (star_follow.py rules, unchanged) over many stock-days.
Per stock-day the per-minute reference signals (reference buys / sells in the last 5 min,
150 closest earlier days re-matched every minute) are computed ONCE, then each variant is
played out on the same signals:
   nostop        no stop, Star Follow controls the exit alone (+ 15:55)
   stop1         1% safety stop
   nostop_skip   no stop, no buys before 9:45 (skips the opening signal)
    python3 sf_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 star_follow as SF

VARIANTS = {"nostop": {"stop": 100.0, "first_buy_m": 5},
            "stop1": {"stop": 1.0, "first_buy_m": 5},
            "nostop_skip": {"stop": 100.0, "first_buy_m": 15}}
K, TRIG, WIN, COST = 150, 10, 5, 0.1


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))
    return out


def play(sig, a, stop, first_buy_m):
    pos, trades = None, []
    for m, buys, sells in sig:
        if pos is None:
            if buys >= TRIG and first_buy_m <= m <= 345:
                e = float(a["O"][m + 1]); pos = {"buy_m": m + 1, "buy": e, "stop": e * (1 - stop / 100), "low": e}
        else:
            pos["low"] = min(pos["low"], float(a["L"][m]))
            if a["L"][m] <= pos["stop"]:
                trades.append({**pos, "sell_m": m, "sell": pos["stop"], "why": "stop"}); pos = None
            elif sells >= TRIG:
                trades.append({**pos, "sell_m": m + 1, "sell": float(a["O"][m + 1]), "why": "exit"}); pos = None
    if pos:
        trades.append({**pos, "sell_m": 385, "sell": float(a["C"][385]), "why": "15:55"})
    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 = SF.load_marks()
    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, c in VARIANTS.items():
                for t in play(sig, a, c["stop"], c["first_buy_m"]):
                    res[v].append({"day": day, "sym": s, **{k: (float(x) if isinstance(x, (np.floating, float)) else x)
                                                           for k, x in t.items()}})
        RT.REF._CACHE.clear()
        print(day, {v: len(r) for v, r in res.items()}, flush=True)
    json.dump(res, open(HERE / f"sf_batch_{a_}_{b_}.json", "w"), default=str)
    half = days[len(days) // 2]
    print(f"\n{len(days)} days ({days[0]}..{days[-1]})")
    print(f"{'variant':>12} {'trades':>6} {'wins':>5} {'win%':>5} {'avg':>7} {'sum':>8} {'1st/2nd half':>16} {'days+':>6} {'worst dip':>9} {'exits/stops/15:55':>18}")
    for v, T in res.items():
        if not T:
            print(f"{v:>12}      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 = {w: sum(t["why"] == w for t in T) for w in ("exit", "stop", "15:55")}
        print(f"{v:>12} {len(p):6} {(p > 0).sum():5} {(p > 0).mean():5.0%} {p.mean():+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}% {why['exit']:>6}/{why['stop']}/{why['15:55']}")


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