# -*- coding: utf-8 -*-
"""
AI旺财 EA 参数搜索：解析筛选 -> 蒙特卡洛验证
本金 60000 USD | XAUUSD 4320 | 1手=100oz | 1point=0.01USD | 杠杆1:2000
"""
import random
import math

PRICE0 = 4320.0
CONTRACT = 100.0
POINT = 0.01
EQUITY = 60000.0
LEV = 2000.0
STEP = 0.01
SWAP_PER_LOT_DAY = -20.0      # 隔夜利息(USD/手/天), 双向收费保守估计
HOURS_PER_YEAR = 6240         # 52周 x 120交易小时


def ladder(init, mult, n, maxlot=99.0):
    out, v = [], init
    for _ in range(n):
        v = min(v, maxlot)
        q = round(round(v / STEP) * STEP, 2)
        out.append(q if q >= STEP else STEP)
        v = v * mult
    return out


def grid_stats(init, mult, n, dist_pts, maxlot=99.0):
    L = ladder(init, mult, n, maxlot)
    d = dist_pts * POINT
    S = sum(L)
    dd = sum(L[i] * CONTRACT * (i * d) for i in range(n))
    return dict(L=L, S=S, dd=dd, d=d, span=(n - 1) * d,
                margin=S * CONTRACT * PRICE0 / LEV,
                be=(dd / (S * CONTRACT)) if S else 0.0)


def solve_init(mult, n, dist_pts, target_dd, maxlot=99.0):
    lo, hi, best = 0.001, 50.0, 0.001
    for _ in range(80):
        mid = (lo + hi) / 2
        if grid_stats(mid, mult, n, dist_pts, maxlot)['dd'] <= target_dd:
            best = mid
            lo = mid
        else:
            hi = mid
    return max(0.01, round(best / STEP) * STEP)


def unbias_ratio(mult, n):
    s0 = sum(mult ** i for i in range(n))
    s1 = sum((mult ** i) * i for i in range(n))
    return (s1 / s0) / (n - 1)


# ------------------------------------------------------------------ 蒙特卡洛
def simulate(init, mult, n, dist_pts, maxlot, tp, sl,
             n_sim=200, sigma=0.28, seed=7, entry_p=0.015):
    rng = random.Random(seed)
    L = ladder(init, mult, n, maxlot)
    d = dist_pts * POINT
    dt = 1.0 / HOURS_PER_YEAR
    sq = math.sqrt(dt)
    drift = -0.5 * sigma * sigma * dt
    vol = sigma * sq
    swap_h = SWAP_PER_LOT_DAY / 24.0

    rets, dds, dead_n, rounds_all = [], [], 0, []
    for _ in range(n_sim):
        eq = EQUITY
        peak = EQUITY
        maxdd = 0.0
        price = PRICE0
        t = 0
        pos = []
        nxt = 0
        direction = 1
        rounds = 0
        dead = False

        while t < HOURS_PER_YEAR:
            price *= math.exp(drift + vol * rng.gauss(0, 1))
            t += 1

            if not pos:
                if rng.random() < entry_p:
                    direction = 1 if rng.random() < 0.5 else -1
                    pos.append((L[0], price))
                    nxt = 1
                    rounds += 1
                continue

            if nxt < n:
                last = pos[-1][1]
                if direction * (price - last) <= -d:
                    pos.append((L[nxt], price))
                    nxt += 1

            lots = sum(lot for lot, _ in pos)
            pnl = sum(lot * CONTRACT * direction * (price - e) for lot, e in pos)
            pnl += lots * swap_h * (t - (t - 1))     # 每小时计息
            equity = eq + pnl

            if equity > peak:
                peak = equity
            if peak - equity > maxdd:
                maxdd = peak - equity

            margin = lots * CONTRACT * price / LEV
            if equity <= 0 or equity - margin <= 0:
                dead = True
                break

            if pnl >= tp:
                eq += pnl
                pos = []
                nxt = 0
                continue
            if sl > 0 and pnl <= -sl:
                eq += pnl
                pos = []
                nxt = 0
                continue

        if pos:
            pnl = sum(lot * CONTRACT * direction * (price - e) for lot, e in pos)
            eq += pnl
        rets.append((eq - EQUITY) / EQUITY)
        dds.append(maxdd / EQUITY)
        rounds_all.append(rounds)
        if dead:
            dead_n += 1

    rets.sort()
    dds.sort()
    k = len(rets)
    return dict(
        dead_rate=dead_n / k,
        ret_med=rets[k // 2],
        ret_p10=rets[int(k * 0.10)],
        ret_p90=rets[int(k * 0.90)],
        dd_med=dds[k // 2],
        dd_p90=dds[int(k * 0.90)],
        rounds=sum(rounds_all) / k,
    )


def score(r):
    """风险调整后期望: 爆仓按-100%计入"""
    exp = (1 - r['dead_rate']) * r['ret_med'] + r['dead_rate'] * (-1.0)
    return exp


if __name__ == '__main__':
    # ---------------- 阶段1: 解析筛选 ----------------
    print('=' * 104)
    print('[阶段1] 解析筛选  约束: 满仓浮亏<=25%本金 / 覆盖>=9%金价 / 解套比<=72%')
    print('=' * 104)
    print('%-28s %6s %6s %9s %7s %8s %7s %8s %7s' % (
        '配置', '起手', '满仓', '浮亏', '回撤%', '覆盖USD', '覆盖%', '保本反弹', '解套比'))
    print('-' * 104)

    cands = []
    for mult in (1.15, 1.20, 1.25, 1.30):
        for n in (10, 12, 15):
            for dpt in (2500, 3500, 4500):
                g0 = grid_stats(0.01, mult, n, dpt)
                span_pct = g0['span'] / PRICE0
                ur = unbias_ratio(mult, n)
                if span_pct < 0.09 or ur > 0.72:
                    continue
                init = solve_init(mult, n, dpt, EQUITY * 0.25)
                g = grid_stats(init, mult, n, dpt)
                if g['S'] < 0.25:          # 满仓太小则无意义
                    continue
                cands.append((mult, n, dpt, init, g))
                print('%-28s %6.2f %6.2f %9.0f %6.1f%% %8.0f %6.2f%% %8.1f %6.1f%%' % (
                    'mult%.2f N%d 间距%dpt' % (mult, n, dpt), init, g['S'], g['dd'],
                    g['dd'] / EQUITY * 100, g['span'], g['span'] / PRICE0 * 100,
                    g['be'], ur * 100))

    print('\n候选数: %d' % len(cands))

    # ---------------- 阶段2: 蒙特卡洛 ----------------
    print()
    print('=' * 104)
    print('[阶段2] 蒙特卡洛验证 (含隔夜利息 -20USD/手/天, 年化波动28%, 200次)')
    print('=' * 104)
    rows = []
    for mult, n, dpt, init, g in cands:
        for tp, sl in ((600.0, 0.0), (600.0, 9000.0), (600.0, 15000.0),
                       (400.0, 9000.0), (900.0, 12000.0)):
            r = simulate(init, mult, n, dpt, 99.0, tp, sl, n_sim=150)
            rows.append((mult, n, dpt, init, tp, sl, r, g))

    rows.sort(key=lambda x: -score(x[6]))
    print('%-22s %6s %7s %7s %8s %8s %9s %9s %9s %6s' % (
        '配置', '起手', '止盈', '止损', '爆仓率', '收益中位', '期望(含爆)', '回撤中位', '回撤P90', '轮次'))
    print('-' * 104)
    for mult, n, dpt, init, tp, sl, r, g in rows[:22]:
        print('%-22s %6.2f %7.0f %7.0f %7.1f%% %8.1f%% %9.1f%% %9.1f%% %8.1f%% %6.0f' % (
            'm%.2f N%d d%d' % (mult, n, dpt), init, tp, sl,
            r['dead_rate'] * 100, r['ret_med'] * 100, score(r) * 100,
            r['dd_med'] * 100, r['dd_p90'] * 100, r['rounds']))
