"""TR-SMA-cross 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: SMA9/SMA21 cross + signal candle at/touching the averages + SMA200
filter, SL at the 5-bar swing +/- 1 pip, 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-sma-cross"


def compute_all(df, p):
    """Causal indicators: SMA9/21/200, ATR14 (walk needs the key)."""
    ind = {}
    ind["sma9"] = df["close"].rolling(9, min_periods=9).mean()
    ind["sma21"] = df["close"].rolling(21, min_periods=21).mean()
    ind["sma200"] = df["close"].rolling(200, min_periods=200).mean()
    ind["atr14"] = W.atr_wilder(df, 14)
    return ind


def _scan(o, h, low, c, s9, s21, s200, p, pip, start):
    """Cross on closed bars; swing SL over the last `swing` closed bars."""
    sigs = []
    n = len(c)
    k = p["swing"]
    for i in range(max(start, k), n - 1):  # need bar i+1 to enter
        a9, a9p, a21, a21p, s = s9[i], s9[i - 1], s21[i], s21[i - 1], s200[i]
        if not (a9 == a9 and a9p == a9p and a21 == a21 and a21p == a21p and s == s):
            continue
        up = a9p <= a21p and a9 > a21  # equality before counts as a cross
        dn = a9p >= a21p and a9 < a21
        if up and low[i] >= min(a9, a21) and c[i] > s:
            sl = float(low[i - k + 1:i + 1].min() - pip)
            sigs.append((i, 1, sl))
        elif dn and h[i] <= max(a9, a21) and c[i] < s:
            sl = float(h[i - k + 1:i + 1].max() + pip)
            sigs.append((i, -1, sl))
    return sigs


def build_signals(strategy, df, ind, p):
    if strategy != STRATEGY:
        raise ValueError(strategy)
    # 1 pip in price units, inferred from price magnitude (documented freeze).
    pip = 0.01 if float(df["close"].median()) > 10.0 else 0.0001
    return _scan(df["open"].to_numpy(), df["high"].to_numpy(),
                 df["low"].to_numpy(), df["close"].to_numpy(),
                 ind["sma9"].to_numpy(), ind["sma21"].to_numpy(),
                 ind["sma200"].to_numpy(), p, pip, 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.10)
    c = np.full(n, 1.10)
    s9 = np.full(n, 1.099)
    s21 = np.full(n, 1.0995)
    s200 = np.full(n, 1.05)
    s9[14] = 1.0990  # at/below s21 -> cross up at bar 15
    s9[15] = 1.1000
    low[15] = 1.0998  # whole candle at/above both averages
    h[15] = 1.1010
    c[15] = 1.1005
    p = {"swing": 5}
    pip = 0.0001
    sigs = _scan(o, h, low, c, s9, s21, s200, p, pip, 5)
    assert [(i, d) for (i, d, _) in sigs] == [(15, 1)], sigs
    assert abs(sigs[0][2] - (low[11:16].min() - pip)) < 1e-12, sigs
    print("unit smacross positive long: PASS (signal bar 15, 5-bar swing -1pip)")

    low2 = low.copy()
    low2[15] = 1.0980  # candle dips under the averages -> blocked
    assert _scan(o, h, low2, c, s9, s21, s200, p, pip, 5) == [], "location leak"
    print("unit smacross negative (candle under averages): PASS (no signal)")

    low3 = low.copy()
    low3[15] = min(s9[15], s21[15])  # exact touch counts as above (boundary)
    assert [(i, d) for (i, d, _) in _scan(o, h, low3, c, s9, s21, s200, p, pip, 5)] == [(15, 1)]
    print("unit smacross boundary (touching averages): PASS (signal fires)")

    # short mirror with 10-bar swing
    s9b = np.full(n, 1.101)
    s21b = np.full(n, 1.1005)
    s200b = np.full(n, 1.15)
    s9b[14] = 1.1010
    s9b[15] = 1.1000  # cross down at bar 15
    hb = np.full(n, 1.101)
    cb = np.full(n, 1.10)
    hb[15] = 1.1005
    cb[15] = 1.1002
    p10 = {"swing": 10}
    sigs4 = _scan(o, hb, low, cb, s9b, s21b, s200b, p10, pip, 10)
    assert [(i, d) for (i, d, _) in sigs4] == [(15, -1)], sigs4
    assert abs(sigs4[0][2] - (hb[6:16].max() + pip)) < 1e-12, sigs4
    print("unit smacross positive short mirror (10-bar): PASS (signal bar 15)")


W.compute_all = compute_all
W.build_signals = build_signals

VARIANTS = [
    dict(id="base", swing=5),
    dict(id="swing10", swing=10),
]

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"])
