"""
ev_trader.py -- [2026-10-01] Step 1 after the calibration check: score each minute by the
EXPECTED RESULT of the references, not by their "up" share (which mostly measured how much a
stock was about to move, in either direction).

For every test stock-day, every minute 9:31-15:15 (whole day so far matched to the 150
closest EARLIER library days, as reference_trader.py):
  target   = target_frac x the references' median best gain from this minute
  ref EV   = mean over the 150 references of the trade "buy at this minute's close, sell at
             +target, or at -stop_pct, or at 15:55" minus the round-trip cost
Then for the test stock itself, the same trade from the next bar's open (stop / target /
15:55) gives the real result. Reported:
  (a) calibration: ref EV buckets vs the real result of the same trade
  (b) trading: per stock-day, buy the first minute ref EV >= threshold (then again after the
      exit, max 2 trades), for several thresholds and target fractions
    python3 ev_trader.py 2026-08-27..2026-09-30
"""
import gzip
import json
import sys
from pathlib import Path

import numpy as np

sys.path.insert(0, str(Path(__file__).resolve().parent))
import reference_trader as RT

STOP = 1.0
COST = 0.1
FRACS = (1.0, 0.5)
THRS = (0.0, 0.1, 0.2, 0.3, 0.5)


def ref_ev(lib, idx, m, target_pct):
    """Mean result (%) of buy-at-close m / +target / -STOP / 15:55 over references idx."""
    e = lib.C[idx, m][:, None]
    Hs, Ls = lib.H[idx, m + 1:386], lib.L[idx, m + 1:386]
    if Hs.shape[1] == 0:
        return -COST
    up = Hs >= e * (1 + target_pct / 100)
    dn = Ls <= e * (1 - STOP / 100)
    BIG = 10 ** 6
    fu = np.where(up.any(1), up.argmax(1), BIG)
    fd = np.where(dn.any(1), dn.argmax(1), BIG)
    res = np.where((fd <= fu) & (fd < BIG), -STOP,
                   np.where(fu < BIG, target_pct, (lib.C[idx, 385] / e[:, 0] - 1) * 100))
    return float(res.mean()) - COST


def real_trade(a, m, target_pct):
    """The test stock: buy at the next bar's open, same exits. Returns (result %, exit minute)."""
    e = a["O"][m + 1]
    for k in range(m + 1, 386):
        if a["L"][k] <= e * (1 - STOP / 100):
            return -STOP - COST, k
        if a["H"][k] >= e * (1 + target_pct / 100):
            return target_pct - COST, k
    return (a["C"][385] / e - 1) * 100 - COST, 385


def main(rng):
    a_, b_ = rng.split("..")
    lib = RT.Lib()
    days = sorted({k.split("|")[0] for k in lib.keys if a_ <= k.split("|")[0] <= b_})
    series = []   # per stock-day: list of (m, {frac: (ev, target)})
    cal = {f: [] for f in FRACS}
    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
            lib.start(day)
            ser = []
            for m in range(1, 346):
                idx, _ = lib.match(a["F"], m, day, 150)
                outc = lib.outcomes(idx, m)
                med = float(np.median([o[2] for o in outc]))
                row = {}
                for f in FRACS:
                    tgt = max(med * f, 0.3)
                    ev = ref_ev(lib, idx, m, tgt)
                    row[f] = (ev, tgt)
                    if m % 5 == 0:
                        cal[f].append((day, ev, real_trade(a, m, tgt)[0]))
                ser.append((m, row))
            series.append((day, s, a, ser))
        RT.REF._CACHE.clear()
        print(day, len(series), flush=True)

    h_split = days[len(days) // 2]
    for f in FRACS:
        C = cal[f]
        ev = np.array([c[1] for c in C]); real = np.array([c[2] for c in C]); h1 = np.array([c[0] < h_split for c in C])
        print(f"\n=== target = {f:.0%} of the references' median best gain ===")
        print(f"(a) calibration, {len(C)} forecasts (every 5 min): correlation forecast vs real {np.corrcoef(ev, real)[0, 1]:.3f}")
        print(f"   {'ref EV':>14} {'n':>6} {'real avg':>9} {'real win%':>9} {'1st half':>9} {'2nd half':>9}")
        for lo, hi in ((-9, -0.5), (-0.5, -0.3), (-0.3, -0.1), (-0.1, 0), (0, 0.1), (0.1, 0.2), (0.2, 0.3), (0.3, 0.5), (0.5, 9)):
            k = (ev >= lo) & (ev < hi)
            if k.sum() < 15:
                print(f"   {lo:+5.1f}..{hi:+5.1f}% {k.sum():6}")
                continue
            print(f"   {lo:+5.1f}..{hi:+5.1f}% {k.sum():6} {real[k].mean():+8.2f}% {(real[k] > 0).mean():8.0%} "
                  f"{real[k & h1].mean() if (k & h1).any() else float('nan'):+8.2f}% {real[k & ~h1].mean() if (k & ~h1).any() else float('nan'):+8.2f}%")
        print(f"(b) trading (first minute ref EV >= threshold, max 2 trades per stock-day, every signal)")
        print(f"   {'threshold':>9} {'trades':>7} {'wins':>5} {'avg':>7} {'sum':>8} {'1st/2nd half sum':>18} {'days+':>6}")
        for thr in THRS:
            trades = []
            for day, s, a, ser in series:
                n_tr, m_free = 0, 1
                for m, row in ser:
                    if m < m_free or n_tr >= 2:
                        continue
                    ev, tgt = row[f]
                    if ev >= thr:
                        res, k_exit = real_trade(a, m, tgt)
                        trades.append((day, s, m, res))
                        n_tr += 1
                        m_free = k_exit + 1
            if not trades:
                print(f"   {thr:+8.1f}% {0:7}")
                continue
            r = np.array([t[3] for t in trades]); h = np.array([t[0] < h_split for t in trades])
            byday = {}
            for t in trades:
                byday[t[0]] = byday.get(t[0], 0) + t[3]
            print(f"   {thr:+8.1f}% {len(r):7} {(r > 0).sum():5} {r.mean():+6.2f}% {r.sum():+7.1f}% "
                  f"{r[h].sum():+8.1f}% / {r[~h].sum():+.1f}% {sum(v > 0 for v in byday.values()):3}/{len(byday)}")


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