# -*- coding: utf-8 -*-
"""그리드봇 8.7년 실측 — 영상 재현 코드 (AlgoLab)

영상: 「코인 그리드봇 8년, 1000만원이 254만원이 됐습니다」

- 기본 실행: 영상에 나온 헤드라인 수치를 전부 재현합니다(몇 분).
    python grid_backtest.py
- 전체 재현: 설정 162개 조합 × 자산 2종 = 324칸 전수 스윕(수십 분)
    python grid_backtest.py --full

데이터: 업비트 원화(KRW-BTC/ETH) · 바이낸스 달러(BTCUSDT/ETHUSDT) 4시간봉, 공개 API.
처음 실행하면 내려받아 CSV 로 저장하고, 다음부터는 재사용합니다.

⚠ ASOF: 영상 제작 시점까지로 데이터를 잘라야 영상과 같은 숫자가 나옵니다.
   오늘까지 보려면 ASOF = None 으로 바꾸세요(숫자는 영상과 달라집니다).
⚠ 시뮬레이션입니다. 4시간봉 기준이고 호가 미끄러짐(슬리피지)은 반영하지 않았습니다.
   과거 성적은 미래 수익을 보장하지 않습니다.
"""
import argparse, csv, datetime, json, os, statistics, sys, time, urllib.parse, urllib.request

ASOF = "2026-08-31T12:00:00"   # 영상 제작 기준 마지막 봉(UTC). None 이면 오늘까지.
START = "2017-12-25"           # 실측 창 시작 — 그리드가 90일 되돌아보기를 채울 수 있는 첫날
COST = 0.001                   # 편도 0.1%
NGRID = 20                     # 격자 수
BLOCK = LOOK = 540             # 4시간봉 540개 = 90일 (블록 길이 · 되돌아보기)
HERE = os.path.dirname(os.path.abspath(__file__))


# ── 데이터 수집 (공개 API · 캐시) ────────────────────────────────────────────
def _get(url):
    req = urllib.request.Request(url, headers={"User-Agent": "Mozilla/5.0",
                                               "Accept": "application/json"})
    return json.loads(urllib.request.urlopen(req, timeout=30).read().decode())


def fetch_upbit(market, out):
    if os.path.exists(out):
        return
    print("  업비트 %s 4시간봉 내려받는 중 …" % market)
    rows, to = [], ""
    while True:
        url = ("https://api.upbit.com/v1/candles/minutes/240?market=%s&count=200" % market
               + ("&to=%s" % urllib.parse.quote(to) if to else ""))
        try:
            batch = _get(url)
        except Exception as e:
            print("   재시도(%s)" % e); time.sleep(2); continue
        if not batch:
            break
        for b in batch:
            rows.append((b["candle_date_time_utc"], b["opening_price"], b["high_price"],
                         b["low_price"], b["trade_price"]))
        to = batch[-1]["candle_date_time_utc"]
        time.sleep(0.12)
    rows.reverse()
    _write(out, rows)


def fetch_binance(symbol, out):
    if os.path.exists(out):
        return
    print("  바이낸스 %s 4시간봉 내려받는 중 …" % symbol)
    rows, start = [], int(datetime.datetime(2017, 8, 1).timestamp() * 1000)
    while True:
        url = ("https://api.binance.com/api/v3/klines?symbol=%s&interval=4h&limit=1000"
               "&startTime=%d" % (symbol, start))
        try:
            batch = _get(url)
        except Exception as e:
            print("   재시도(%s)" % e); time.sleep(2); continue
        if not batch:
            break
        for b in batch:
            t = datetime.datetime.utcfromtimestamp(b[0] / 1000).strftime("%Y-%m-%dT%H:%M:%S")
            rows.append((t, b[1], b[2], b[3], b[4]))
        if len(batch) < 1000:
            break
        start = batch[-1][0] + 1
        time.sleep(0.15)
    _write(out, rows)


def _write(out, rows):
    with open(out, "w", newline="", encoding="utf-8") as f:
        w = csv.writer(f); w.writerow(["date", "open", "high", "low", "close"]); w.writerows(rows)
    print("  ✓ %s봉 → %s" % (format(len(rows), ","), os.path.basename(out)))


def load(fname):
    ds, o, h, l, c = [], [], [], [], []
    with open(os.path.join(HERE, fname), encoding="utf-8") as f:
        for r in csv.DictReader(f):
            if ASOF and r["date"] > ASOF:      # ⚠ 캐시를 쓰든 새로 받든 똑같이 자른다
                continue
            ds.append(r["date"]); o.append(float(r["open"])); h.append(float(r["high"]))
            l.append(float(r["low"])); c.append(float(r["close"]))
    return ds, o, h, l, c


# ── 그리드봇 ────────────────────────────────────────────────────────────────
def grid_lines(lo, hi, n):
    """기하 간격 격자. 9년에 가격이 몇십 배 움직이므로 등차로 나누면 위쪽만 촘촘해진다."""
    r = (hi / lo) ** (1.0 / n)
    return [lo * (r ** i) for i in range(n + 1)]


def run_grid(path, lines, fee=COST, top_hold=False, bot_stop=False):
    """현물 롱온리 그리드 — 업비트·바이낸스 봇과 같은 동작.

    · 칸마다 자본 1/n 배정. 시작가 위 칸은 팔 물건이 필요하므로 시작가에 미리 산다.
    · 매수선을 아래로 지나면 매수, 그 위 매도선을 위로 지나면 매도.
    · 범위 위로 뚫리면 전량 매도되어 현금으로 남고(상승을 놓친다),
      아래로 뚫리면 전량 매수되어 물린 채로 보유한다.
    · top_hold: 위를 뚫으면 팔지 않고 현금까지 실어 그대로 들고 간다
    · bot_stop: 아래를 뚫으면 전량 팔고 나온다
    """
    n = len(lines) - 1
    unit = 1.0 / n
    top, bot = lines[-1], lines[0]
    p0 = path[0]
    qty = [0.0] * n
    cash = 1.0
    for i in range(n):
        if lines[i] >= p0:
            qty[i] = unit * (1 - fee) / p0
            cash -= unit

    state, held_all, prev = "grid", 0.0, p0
    for px in path[1:]:
        if state == "grid":
            if top_hold and px > top:
                held_all = sum(qty) + cash * (1 - fee) / px
                cash = 0.0; qty = [0.0] * n; state = "hold"; prev = px; continue
            if bot_stop and px < bot:
                cash += sum(qty) * px * (1 - fee)
                qty = [0.0] * n; state = "out"; prev = px; continue
            if px < prev:
                for i in range(n - 1, -1, -1):
                    if qty[i] == 0.0 and prev > lines[i] >= px and cash >= unit - 1e-12:
                        qty[i] = unit * (1 - fee) / lines[i]
                        cash -= unit
            elif px > prev:
                for i in range(n):
                    if qty[i] > 0.0 and prev < lines[i + 1] <= px:
                        cash += qty[i] * lines[i + 1] * (1 - fee)
                        qty[i] = 0.0
        prev = px
    last = path[-1]
    return held_all * last if state == "hold" else cash + sum(qty) * last


def path_close(o, h, l, c, a, b):
    """보수 — 종가만 격자를 지난 것으로 친다(체결을 적게 잡는다)."""
    return c[a:b]


def path_ohlc(o, h, l, c, a, b):
    """낙관 — 봉 안을 단조 경로로 편다. 양봉이면 저→고, 음봉이면 고→저."""
    out = [o[a]]
    for i in range(a, b):
        out += ([o[i], l[i], h[i], c[i]] if c[i] >= o[i] else [o[i], h[i], l[i], c[i]])
    return out


def blocks(ds, o, h, l, c, mode="close", block=BLOCK, look=LOOK, ngrid=NGRID,
           widen=1.0, fee=COST, **kw):
    """START 부터 90일 블록으로 계속 돌린다. 범위는 **직전 90일만** 보고 정한다(룩어헤드 금지)."""
    fn = path_close if mode == "close" else path_ohlc
    i0 = next(i for i, d in enumerate(ds) if d[:10] >= START)
    eq, curve, segs, a = 1.0, [1.0], [], i0
    while a < len(c) - 1:
        b = min(a + block, len(c))
        j = max(0, a - look)
        lo, hi = min(l[j:a]), max(h[j:a])
        if widen != 1.0:
            mid = (lo * hi) ** 0.5
            lo, hi = mid * (lo / mid) ** widen, mid * (hi / mid) ** widen
        v = run_grid(fn(o, h, l, c, a, b), grid_lines(lo, hi, ngrid), fee=fee, **kw)
        eq *= v
        curve.append(eq)
        segs.append({"시작": ds[a][:10], "그리드": v - 1.0, "보유": c[b - 1] / c[a] - 1.0,
                     "온전": b - a == block})   # 짧은 마지막 블록은 장세 집계에서 뺀다
        a = b
    return i0, eq, curve, segs


# ── 잣대를 통과한 전략 (지난 편 생존 4칸 중 하나) ─────────────────────────────
def sma(x, n):
    out = [None] * len(x)
    s = 0.0
    for i, v in enumerate(x):
        s += v
        if i >= n:
            s -= x[i - n]
        if i >= n - 1:
            out[i] = s / n
    return out


def atr(h, l, c, n):
    tr = [h[0] - l[0]]
    for i in range(1, len(c)):
        tr.append(max(h[i] - l[i], abs(h[i] - c[i - 1]), abs(l[i] - c[i - 1])))
    out = [None] * len(c)
    if len(tr) >= n:
        out[n - 1] = sum(tr[:n]) / n
        for i in range(n, len(tr)):
            out[i] = out[i - 1] + (tr[i] - out[i - 1]) / n
    return out


def ut_bot(h, l, c, key=2.0, atr_n=21):
    """UT Bot Alerts (QuantNomad) — 변동성(ATR)만큼 떨어진 손절선을 따라 올리는 추세 지표."""
    a = atr(h, l, c, atr_n)
    pos = [0.0] * len(c)
    st = None
    for i in range(len(c)):
        if a[i] is None:
            continue
        loss = key * a[i]
        if st is None:
            st = c[i] - loss
        prev = st
        if c[i] > prev and c[i - 1] > prev:
            st = max(prev, c[i] - loss)
        elif c[i] < prev and c[i - 1] < prev:
            st = min(prev, c[i] + loss)
        else:
            st = c[i] - loss if c[i] > prev else c[i] + loss
        pos[i] = 1.0 if c[i] > st else 0.0
    return pos


def run_signal(ds, o, h, l, c, sig, filt_days=50, fee=COST):
    """0/1 신호를 그대로 따라간다. 진입·청산마다 편도 비용을 뗀다. 신호는 **다음 봉**에 반영."""
    f = sma(c, filt_days * 6)
    s = [sig[i] * (1.0 if (f[i] is not None and c[i] > f[i]) else 0.0) for i in range(len(c))]
    i0 = next(i for i, d in enumerate(ds) if d[:10] >= START)
    eq, curve, prev = 1.0, [1.0], 0.0
    for i in range(1, len(c)):
        pos = s[i - 1]
        g = (1 + pos * (c[i] / c[i - 1] - 1)) * (1 - fee * abs(pos - prev))
        prev = pos
        if i >= i0:
            eq *= g
            curve.append(eq)
    return eq, curve


def mdd(curve):
    peak, worst = curve[0], 0.0
    for e in curve:
        peak = max(peak, e)
        worst = min(worst, e / peak - 1.0)
    return worst


MK = [("업비트", "BTC", "KRW-BTC", "grid_btc_krw_4h.csv"),
      ("업비트", "ETH", "KRW-ETH", "grid_eth_krw_4h.csv"),
      ("바이낸스", "BTC", "BTCUSDT", "grid_btcusdt_4h.csv"),
      ("바이낸스", "ETH", "ETHUSDT", "grid_ethusdt_4h.csv")]


def ensure():
    print("데이터 확인")
    for mk, sym, code, fname in MK:
        p = os.path.join(HERE, fname)
        (fetch_upbit if mk == "업비트" else fetch_binance)(code, p)
    print()


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--full", action="store_true", help="162개 조합 324칸 전수 스윕")
    args = ap.parse_args()
    ensure()
    D = {}
    for mk, sym, code, fname in MK:
        D["%s_%s" % (mk, sym)] = load(fname)

    ds, o, h, l, c = D["업비트_BTC"]
    i0, geq, gcur, segs = blocks(ds, o, h, l, c)
    _, teq, tcur, _ = blocks(ds, o, h, l, c, top_hold=True, bot_stop=True)
    _, oeq, _, _ = blocks(ds, o, h, l, c, mode="ohlc")
    hold = [x / c[i0] for x in c[i0:]]
    ueq, ucur = run_signal(ds, o, h, l, c, ut_bot(h, l, c, 2.0, 21), 50)
    yrs = (datetime.date.fromisoformat(ds[-1][:10])
           - datetime.date.fromisoformat(ds[i0][:10])).days / 365.25

    print("창 %s ~ %s (%.1f년) · 업비트 BTC · 1,000만원 기준 · 편도 %.1f%%"
          % (ds[i0][:10], ds[-1][:10], yrs, COST * 100))
    print("  그리드봇 기본            %8.0f만원  (%.2f배 · MDD %.1f%%)"
          % (geq * 1000, geq, mdd(gcur) * 100))
    print("  그리드봇 기본(낙관 가정)  %8.0f만원  (%.2f배)" % (oeq * 1000, oeq))
    print("  그리드봇 + 이탈처리 두 줄 %8.0f만원  (%.2f배 · MDD %.1f%%)"
          % (teq * 1000, teq, mdd(tcur) * 100))
    print("  그냥 보유                %8.0f만원  (%.2f배 · MDD %.1f%%)"
          % (hold[-1] * 1000, hold[-1], mdd(hold) * 100))
    print("  잣대 통과 전략(UT Bot)   %8.0f만원  (%.2f배 · MDD %.1f%%)"
          % (ueq * 1000, ueq, mdd(ucur) * 100))

    print("\n수수료를 낮추면 (자가반박 — 진 게 비용 탓인가)")
    for fee in (0.001, 0.0005, 0.0):
        _, e, _, _ = blocks(ds, o, h, l, c, fee=fee)
        print("  편도 %.2f%% → %6.0f만원" % (fee * 100, e * 1000))

    print("\n장세별 (네 시장 합산 · 90일 구간 · 보수 가정)")
    agg = {}
    for k in D:
        _, _, _, s = blocks(*D[k])
        for r in s:
            if not r["온전"]:
                continue
            lab = "상승" if r["보유"] >= 0.20 else ("하락" if r["보유"] <= -0.20 else "횡보")
            a = agg.setdefault(lab, {"n": 0, "win": 0, "g": [], "h": []})
            a["n"] += 1; a["win"] += r["그리드"] > r["보유"]
            a["g"].append(r["그리드"]); a["h"].append(r["보유"])
    for lab in ("하락", "횡보", "상승"):
        a = agg[lab]
        print("  %s 구간 %2d개 · 그리드 승 %2d · 그리드 평균 %+5.1f%% · 보유 평균 %+5.1f%%"
              % (lab, a["n"], a["win"], statistics.mean(a["g"]) * 100,
                 statistics.mean(a["h"]) * 100))

    if not args.full:
        print("\n(전수 스윕은 --full · 162개 조합 324칸)")
        return

    print("\n전수 스윕 — 두 시장 동시 잣대(연 30% 이상 · 최대낙폭 -35% 이내)")
    passed = total = 0
    for ng in (10, 20, 40):
        for lb in (60, 90, 180):
            for wd in (0.7, 1.0, 1.4):
                for pr in (30, 90, 180):
                    for name, kw in (("기본", {}), ("이탈처리", dict(top_hold=True, bot_stop=True))):
                        for sym in ("BTC", "ETH"):
                            stat = []
                            for mk in ("업비트", "바이낸스"):
                                d = D["%s_%s" % (mk, sym)]
                                i, e, cur, _ = blocks(*d, block=pr * 6, look=lb * 6,
                                                      ngrid=ng, widen=wd, **kw)
                                y = (datetime.date.fromisoformat(d[0][-1][:10])
                                     - datetime.date.fromisoformat(d[0][i][:10])).days / 365.25
                                stat.append((e ** (1 / y) - 1, mdd(cur)))
                            total += 1
                            if min(s[0] for s in stat) >= 0.30 and min(s[1] for s in stat) >= -0.35:
                                passed += 1
                                print("  통과: 격자%d 과거%d일 폭%.1f 주기%d %s %s"
                                      % (ng, lb, wd, pr, name, sym))
    print("  %d칸 중 통과 %d칸" % (total, passed))


if __name__ == "__main__":
    main()
