"""
combined_ftmo_mc_v2.py  (corrected mae scaling)
===============================================
HONEST FTMO Monte Carlo for EURUSD + USDJPY. Fixes v1's four flaws AND the
mae-scaling bug found in the first v2 run.

mae UNITS (verified):
  UJ 'mae_pct' = PERCENT (e.g. 0.76 = 0.76%)  -> float_loss$ = mae/100 * notional
  EU 'mae'     = NEGATIVE DECIMAL (e.g. -0.0159 = -1.59%) -> float_loss$ = |mae| * notional

Intraday DD proxy uses the WORST SINGLE TRADE's floating loss in a day
(not the sum) — realistic, since one trade's adverse extreme dominates.

Two DD modes, both reported:
  REALIZED  : close-to-close P&L only (we know this gives ~0 daily breaches)
  INTRADAY  : also apply worst single-trade floating loss within the day

Methods A (sequential replay), B (block bootstrap), C (day-level MC),
plus D (intervention stress). $300k notional/pair, pnl already includes tier mult.
"""
import csv, sys
import numpy as np
from pathlib import Path
from datetime import datetime, timedelta
from collections import defaultdict

ACCOUNT=100_000; TARGET_USD=10_000; DAILY_DD_USD=5_000; TOTAL_DD_USD=10_000
N_SIMS=5_000; SEED=42; MAX_MONTHS=12
EU_HOLD_H=52; UJ_HOLD_H=24
EU_NOTIONAL=300_000; UJ_NOTIONAL=300_000
rng=np.random.default_rng(SEED)

def resolve_base():
    for b in [Path(r"C:\Users\Administrator\OneDrive\fx_macro_intraday"),
              Path(r"C:\Users\paul_\OneDrive\fx_macro_intraday"),
              Path(__file__).resolve().parents[2]]:
        if b.exists(): return b
    sys.exit("base path not found")
BASE=resolve_base()

def parse_dt(s):
    s=(s or "").strip()
    for f in ("%Y-%m-%d %H:%M:%S","%Y-%m-%d %H:%M","%d/%m/%Y %H:%M"):
        try: return datetime.strptime(s[:19], f)
        except ValueError: continue
    return None

def load_trades():
    eu_p=BASE/"data"/"processed"/"trades"/"trades_real_costs.csv"
    uj_p=BASE/"data"/"processed"/"usdjpy_trades_real_costs.csv"
    trades=[]
    # EU: mae is NEGATIVE DECIMAL fraction
    with open(eu_p, newline="", encoding="utf-8-sig") as f:
        for r in csv.DictReader(f):
            entry=parse_dt(r.get("entry_time"))
            if entry is None: continue
            try:
                pnl=float(r["dollar_pnl_real"]); mae=float(r.get("mae",0) or 0)
            except (ValueError,KeyError): continue
            float_loss=abs(mae)*EU_NOTIONAL     # |decimal| * notional
            trades.append({"pair":"EU","entry":entry,"pnl":pnl,"float_loss":float_loss})
    # UJ: mae_pct is PERCENT
    with open(uj_p, newline="", encoding="utf-8-sig") as f:
        for r in csv.DictReader(f):
            entry=parse_dt(r.get("entry_time"))
            if entry is None: continue
            try:
                pnl=float(r["pnl_real"]); mae=float(r.get("mae_pct",0) or 0)
            except (ValueError,KeyError): continue
            float_loss=(mae/100.0)*UJ_NOTIONAL  # percent/100 * notional
            trades.append({"pair":"UJ","entry":entry,"pnl":pnl,"float_loss":float_loss})
    trades.sort(key=lambda t:t["entry"])
    return trades

def group_by_day(trades):
    """Each day: (date, day_pnl, worst_single_float_loss, n)."""
    by=defaultdict(lambda:{"pnl":0.0,"wfl":0.0,"n":0})
    for t in trades:
        d=t["entry"].date()
        by[d]["pnl"]+=t["pnl"]
        by[d]["wfl"]=max(by[d]["wfl"], t["float_loss"])  # MAX single-trade
        by[d]["n"]+=1
    return [(d,by[d]["pnl"],by[d]["wfl"],by[d]["n"]) for d in sorted(by)]

def evaluate(day_seq, intraday=True):
    eq=0.0; peak=0.0; first=None
    for (d,day_pnl,wfl,n) in day_seq:
        if first is None: first=d
        # intraday low = equity + (realized day move if negative) - worst floating
        fl = wfl if intraday else 0.0
        intraday_low = eq + min(0.0, day_pnl) - fl
        # Daily DD: worst within-day point measured from day's OPEN equity (eq)
        day_drawdown = min(day_pnl, -fl, day_pnl - fl)  # most negative excursion
        if day_drawdown <= -DAILY_DD_USD:
            return ("daily",(d-first).days)
        if (intraday_low - peak) <= -TOTAL_DD_USD:
            return ("total",(d-first).days)
        eq+=day_pnl; peak=max(peak,eq)
        if (eq-peak) <= -TOTAL_DD_USD:
            return ("total",(d-first).days)
        if eq>=TARGET_USD:
            return ("pass",(d-first).days)
        if (d-first).days > MAX_MONTHS*30:
            return ("timeout",(d-first).days)
    return ("timeout",(day_seq[-1][0]-first).days if first else 0)

def redate(seq):
    out=[]; d=datetime(2000,1,3).date()
    for (_od,pnl,wfl,n) in seq:
        out.append((d,pnl,wfl,n))
        d+=timedelta(days=1)
        while d.weekday()>=5: d+=timedelta(days=1)
    return out

def summarize(label, outs, passdays, n):
    P=sum(1 for o in outs if o=="pass"); DA=sum(1 for o in outs if o=="daily")
    TO=sum(1 for o in outs if o=="total"); TM=sum(1 for o in outs if o=="timeout")
    r={"label":label,"n":n,
       "pass_rate":100*P/n if n else 0,"daily":100*DA/n if n else 0,
       "total":100*TO/n if n else 0,"timeout":100*TM/n if n else 0}
    if passdays:
        a=np.array(passdays)/30.0
        r.update({"med":float(np.median(a)),"p25":float(np.percentile(a,25)),
                  "p75":float(np.percentile(a,75)),"min":float(np.min(a))})
    else:
        r.update({"med":0,"p25":0,"p75":0,"min":0})
    return r

def m_sequential(days, intraday):
    outs=[]; pd=[]; n=len(days)
    for s in range(n):
        if n-s<5: break
        res,el=evaluate(days[s:], intraday)
        outs.append(res)
        if res=="pass": pd.append(el)
    return summarize("A. Sequential replay", outs, pd, len(outs))

def m_block(days, intraday, n_sims=N_SIMS, blk=21):
    outs=[]; pd=[]; n=len(days); tl=min(n,13*21)
    for _ in range(n_sims):
        seq=[]
        while len(seq)<tl:
            st=rng.integers(0,max(1,n-blk)); seq.extend(days[st:st+blk])
        seq=redate(seq); res,el=evaluate(seq,intraday); outs.append(res)
        if res=="pass": pd.append(el)
    return summarize("B. Block bootstrap", outs, pd, n_sims)

def m_daylevel(days, intraday, n_sims=N_SIMS):
    outs=[]; pd=[]; n=len(days); tl=min(n,13*21); pool=np.arange(n)
    for _ in range(n_sims):
        pick=rng.choice(pool,size=tl,replace=True)
        seq=redate([days[i] for i in pick]); res,el=evaluate(seq,intraday); outs.append(res)
        if res=="pass": pd.append(el)
    return summarize("C. Day-level MC", outs, pd, n_sims)

def m_intervention(days, intraday, n_sims=2000, n_events=3, extra=2):
    outs=[]; pd=[]; n=len(days); tl=min(n,13*21)
    dp=sorted(d[1] for d in days); med=dp[len(dp)//2]
    bad=[d for d in days if d[1]<med]
    for _ in range(n_sims):
        pick=rng.choice(np.arange(n),size=tl,replace=True)
        seq=[list(days[i]) for i in pick]
        for _e in range(n_events):
            pos=rng.integers(0,len(seq)); spnl=seq[pos][1]; swfl=seq[pos][2]; sn=seq[pos][3]
            for _c in range(extra):
                bd=bad[rng.integers(0,len(bad))]
                spnl+=bd[1]; swfl=max(swfl,bd[2]); sn+=bd[3]
            seq[pos][1]=spnl; seq[pos][2]=swfl; seq[pos][3]=sn
        seq=redate([tuple(s) for s in seq]); res,el=evaluate(seq,intraday); outs.append(res)
        if res=="pass": pd.append(el)
    return summarize(f"D. Intervention stress ({n_events} clustered days)", outs, pd, n_sims)

def show(r):
    print(f"\n  {'-'*64}\n  {r['label']}   (n={r['n']:,})\n  {'-'*64}")
    print(f"    Pass rate        : {r['pass_rate']:6.2f}%")
    print(f"    Daily DD breach  : {r['daily']:6.2f}%")
    print(f"    Total DD breach  : {r['total']:6.2f}%")
    print(f"    Timeout (>12mo)  : {r['timeout']:6.2f}%")
    print(f"    Median pass time : {r['med']:.1f} months")
    print(f"    P25 / P75        : {r['p25']:.1f} / {r['p75']:.1f} months")
    print(f"    Fastest          : {r['min']:.1f} months")

def run_suite(days, intraday, tag):
    print(f"\n{'='*72}\nMODE: {tag}\n{'='*72}")
    rA=m_sequential(days,intraday); print("  A done",flush=True)
    rB=m_block(days,intraday); print("  B done",flush=True)
    rC=m_daylevel(days,intraday); print("  C done",flush=True)
    rD=m_intervention(days,intraday); print("  D done",flush=True)
    for r in (rA,rB,rC,rD): show(r)
    meds=[rA['med'],rB['med'],rC['med']]; prs=[rA['pass_rate'],rB['pass_rate'],rC['pass_rate']]
    print(f"\n  Convergence ({tag}):")
    print(f"    Median A/B/C: {meds[0]:.1f}/{meds[1]:.1f}/{meds[2]:.1f} mo  "
          f"(spread {max(meds)-min(meds):.1f})")
    print(f"    Pass%  A/B/C: {prs[0]:.1f}/{prs[1]:.1f}/{prs[2]:.1f}  "
          f"| Intervention {rD['pass_rate']:.1f}%")
    return rA,rB,rC,rD

def main():
    print("="*72)
    print("COMBINED FTMO MC v2 (corrected) | +10% | 5% daily | 10% total | $300k/pair")
    print("="*72)
    trades=load_trades()
    days=group_by_day(trades)
    print(f"Trades {len(trades):,} (EU {sum(1 for t in trades if t['pair']=='EU'):,}, "
          f"UJ {sum(1 for t in trades if t['pair']=='UJ'):,}) | days {len(days):,}")
    worst=min(days,key=lambda d:d[1])
    print(f"Worst realized day ${worst[1]:,.0f} ({worst[3]} trades) | "
          f"Total P&L ${sum(d[1] for d in days):,.0f}")
    # sanity: max single-trade floating loss
    mfl=max(d[2] for d in days)
    print(f"Max single-trade floating loss (intraday proxy): ${mfl:,.0f}")

    # REALIZED mode (trustworthy baseline; matches our known 0 daily breaches)
    run_suite(days, intraday=False, tag="REALIZED (close-to-close)")
    # INTRADAY mode (adds floating-loss proxy, correctly scaled)
    run_suite(days, intraday=True,  tag="INTRADAY (with mae floating-loss proxy)")

    print(f"\n{'='*72}")
    print("vs v1 claim: ~1.7 months / ~100% pass (flat-time approx + per-trade DD)")
    print("Read REALIZED as the trustworthy baseline; INTRADAY as the stress view.")
    print("="*72)

if __name__=="__main__":
    main()
