"""TR-Keltner W7-TEST: vectorbtpro run on Dukascopy H1 Bid (EURUSD, EURJPY).

Run (canonical): py -3.11 test.py   (from this folder; vectorbtpro required)
Reads working copies in ../_data (gitignored), writes results/ CSVs.
Base rule: fully-closed candle outside EMA20 +/- 2.0xATR20 Keltner band, SL at
the opposite band capped to a 1.0xATR fixed stop when the opposite band is more
than 5.0xATR away (see RULES.md A6), fixed 1.5R target, next-bar open entry,
one position per symbol, stop-first on ambiguous H1 bars.
Costs: fixed assumed spread 1.0 pip EURUSD / 1.2 pip EURJPY + $7/lot round
trip ($0.07 at 0.01 lot); EURJPY commission at fixed USDJPY=140.
See RULES.md for frozen assumptions; RESULTS.md for measurements.
"""

import os
import sys

import vectorbtpro  # noqa: F401  (hard requirement: real engine, no fallback)

sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import w7_vbt as W

STRATEGY = "tr-keltner"


def compute_all(df, p):
    """Causal indicators: EMA20, ATR20 (channel), ATR14 (walk needs the key)."""
    ind = {}
    ind["ema20"] = W.ema(df["close"], p["kc_ema"])
    ind["atr20"] = W.atr_wilder(df, p["kc_atr"])
    ind["atr14"] = W.atr_wilder(df, 14)
    return ind


STOP_CAP_ATR = 5.0  # "too wide" threshold for the opposite band (RULES.md A6)


def _scan(o, h, low, c, ema20, atr20, p, start):
    """Outside-band scan with capped opposite-band stop (see RULES.md A6)."""
    sigs = []
    n = len(c)
    k = p["kc_mult"]
    for i in range(start, n - 1):  # need bar i+1 to enter
        e, a = ema20[i], atr20[i]
        if not (e == e and a == a and a > 0):
            continue
        up, lo = e + k * a, e - k * a
        if o[i] > up and c[i] > up:  # fully closed candle above (strict)
            d_opp = c[i] - lo
            sl = lo if d_opp <= STOP_CAP_ATR * a else c[i] - 1.0 * a
            sigs.append((i, 1, float(sl)))
        elif o[i] < lo and c[i] < lo:  # fully closed candle below (strict)
            d_opp = up - c[i]
            sl = up if d_opp <= STOP_CAP_ATR * a else c[i] + 1.0 * a
            sigs.append((i, -1, float(sl)))
    return sigs


def build_signals(strategy, df, ind, p):
    if strategy != STRATEGY:
        raise ValueError(strategy)
    return _scan(df["open"].to_numpy(), df["high"].to_numpy(),
                 df["low"].to_numpy(), df["close"].to_numpy(),
                 ind["ema20"].to_numpy(), ind["atr20"].to_numpy(),
                 p, max(W.WARMUP_BARS, 1))


def _unit_tests():
    """Positive/negative/boundary checks on the real _scan (hand-built bars)."""
    import numpy as np

    n = 30
    o = np.full(n, 1.10)
    h = np.full(n, 1.101)
    low = np.full(n, 1.099)
    c = np.full(n, 1.10)
    ema = np.full(n, 1.10)
    atr = np.full(n, 0.002)  # band +/-0.004 at k=2
    p = {"kc_mult": 2.0}
    o[15] = 1.1045
    c[15] = 1.1050  # both above upper 1.104 -> long; opposite 4.5 ATR -> band stop
    o[16] = 1.1060
    c[16] = 1.1070  # extended breakout; opposite 5.5 ATR -> capped 1 ATR stop
    sigs = _scan(o, h, low, c, ema, atr, p, 1)
    assert [(i, d) for (i, d, _) in sigs] == [(15, 1), (16, 1)], sigs
    assert abs(sigs[0][2] - 1.0960) < 1e-12, sigs  # opposite (lower) band
    assert abs(sigs[1][2] - (1.1070 - 0.002)) < 1e-12, sigs  # capped 1 ATR stop
    print("unit keltner positive long (band + capped stop): PASS (bars 15-16)")

    o2, c2 = o.copy(), c.copy()
    o2[15] = 1.1030  # open inside the band -> no signal
    o2[16] = c2[16] = 1.10  # neutralize the bar-16 candle
    assert _scan(o2, h, low, c2, ema, atr, p, 1) == [], "inside-open leak"
    print("unit keltner negative (open inside band): PASS (no signal)")

    o3, c3 = o.copy(), c.copy()
    o3[15] = 1.104  # open exactly on the band is not above (boundary)
    c3[15] = 1.1050
    o3[16] = c3[16] = 1.10  # neutralize the bar-16 candle
    assert _scan(o3, h, low, c3, ema, atr, p, 1) == [], "boundary leak"
    print("unit keltner boundary (open on band): PASS (no signal)")

    # short mirror with tight bands -> opposite-band stop applies
    ema4 = np.full(n, 1.10)
    atr4 = np.full(n, 0.0004)  # band +/-0.0008 at k=2
    o4 = np.full(n, 1.10)
    c4 = np.full(n, 1.10)
    h4 = np.full(n, 1.101)
    low4 = np.full(n, 1.099)
    o4[15] = 1.0990
    c4[15] = 1.0989  # both below lower 1.0992; opposite 1.1008 within cap
    sigs4 = _scan(o4, h4, low4, c4, ema4, atr4, p, 1)
    assert [(i, d) for (i, d, _) in sigs4] == [(15, -1)], sigs4
    assert abs(sigs4[0][2] - (1.10 + 2.0 * 0.0004)) < 1e-12, sigs4
    print("unit keltner positive short (band stop): PASS (signal bar 15)")


W.compute_all = compute_all
W.build_signals = build_signals

VARIANTS = [
    dict(id="base", kc_ema=20, kc_atr=20, kc_mult=2.0),
    dict(id="k15", kc_ema=20, kc_atr=20, kc_mult=1.5),
]

if __name__ == "__main__":
    _unit_tests()
    outdir = os.path.join(os.path.dirname(os.path.abspath(__file__)), "results")
    rows = W.run_strategy(STRATEGY, VARIANTS, outdir, use_vbt=True)
    for s in rows:
        print("%s %-8s n=%d TP=%d SL=%d END=%d wr=%.2f%% be=%.2f%% effR=%s ddR=%s IS=%d/%.0f%% OOS=%d/%.0f%% both=%d" % (
            s["scope"], s["variant"], s["n_trades"], s["n_tp"], s["n_sl"], s["n_end"],
            100 * s["win_rate"], 100 * s["breakeven_win_rate"], s["effective_R"], s["maxdd_R"],
            s["is_n"], 100 * s["is_winrate"], s["oos_n"], 100 * s["oos_winrate"], s["both_hit_trades"]))
    print("python", sys.version.split()[0], "vbt", rows[0]["vbt_version"],
          "reconciled", rows[0]["vbt_reconciled"])
