"""
eu_boost_mc_v3.py — TRUE intraday equity path (point-in-time)
==============================================================
Run: python src\\research\\eu_boost_mc_v3.py   (takes a few minutes)

WHY v3: v1/v2 summed each trade's worst daily float as if concurrent.
EU is single-position (overlap check: max concurrent = 1), so sequential
same-day floats NEVER coexist — that summing manufactured phantom daily
drawdown, worst on multi-trade days, multiplied by any boost. Both prior
tables are void (Bennett's catch, twice).

v3 walks the merged 15M bar timeline chronologically. At each bar:
  equity_delta(t) = realized P&L closed so far today
                  + signed float of positions open RIGHT NOW
                    (worst-in-bar per open trade — mildly conservative)
Day low = min over the day. FTMO checks use that true low.

Rule under test unchanged: EU x BOOST only when no UJ open at EU entry.
Sweep 1.0 / 1.5 / 2.0, FULL vs QUIET years.
"""
import csv, sys
import numpy as np
from pathlib import Path
from datetime import datetime, timedelta
from collections import defaultdict

ACCOUNT=100_000; P1_TARGET=10_000; P2_TARGET=5_000
DAILY_LIMIT=5_000; MAXLOSS_FLOOR=90_000
N_SIMS=5_000; SEED=42
BASE_NOTIONAL=300_000
EU_HOLD_H=52; UJ_HOLD_H=24
BOOSTS=[1.0, 1.5, 2.0]
QUIET={2005,2011,2017,2019,2020,2021,2024,2025}
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 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"):
        try: return datetime.strptime(s[:19], f)
        except ValueError: continue
    return None

def load_bars(path):
    by_day=defaultdict(list)
    with open(path, newline="", encoding="utf-8-sig") as f:
        reader=csv.DictReader(f); cols={c.lower().strip():c for c in reader.fieldnames}
        dtc=next(cols[k] for k in cols if "time" in k or "date" in k)
        hc=cols["high"]; lc=cols["low"]
        for r in reader:
            dt=parse_dt(r.get(dtc))
            if dt is None: continue
            try: hi=float(r[hc]); lo=float(r[lc])
            except (ValueError,KeyError): continue
            by_day[dt.date()].append((dt,hi,lo))
    for d in by_day: by_day[d].sort()
    return by_day

def load_trades():
    eu_p=BASE/"data"/"processed"/"trades"/"trades_real_costs.csv"
    uj_p=BASE/"data"/"processed"/"usdjpy_trades_real_costs.csv"
    T=[]
    with open(eu_p, newline="", encoding="utf-8-sig") as f:
        for r in csv.DictReader(f):
            e=parse_dt(r.get("entry_time")); x=parse_dt(r.get("exit_time"))
            if e is None: continue
            if x is None: x=e+timedelta(hours=EU_HOLD_H)
            try:
                ep=float(r["entry_price"]); sig=int(float(r["signal"]))
                pnl=float(r["dollar_pnl_real"]); za=float(r.get("zscore_abs",2.0))
            except (ValueError,KeyError): continue
            tier=2.0 if za>=3.5 else (1.5 if za>=3.0 else 1.0)
            T.append({"pair":"EU","entry":e,"exit":x,"ep":ep,"dir":sig,
                      "pnl_1x":pnl/tier,"tier":tier})
    miss=0
    with open(uj_p, newline="", encoding="utf-8-sig") as f:
        for r in csv.DictReader(f):
            e=parse_dt(r.get("entry_time"))
            if e is None: continue
            x=parse_dt(r.get("exit_time"))
            if x is None: x=e+timedelta(hours=UJ_HOLD_H); miss+=1
            try:
                ep=float(r["entry_price"]); d=int(float(r["direction"]))
                pnl=float(r["pnl_real"]); mult=float(r.get("mult",1.0))
            except (ValueError,KeyError): continue
            T.append({"pair":"UJ","entry":e,"exit":x,
                      "ep":ep,"dir":d,"pnl_1x":pnl/mult,"tier":mult})
    if miss: print(f"  WARNING: {miss} UJ rows missing exit_time (used +24h)")
    T.sort(key=lambda t:t["entry"])
    return T

def build_days(trades, bars, boost):
    """True point-in-time daily records: (date, realized_pnl, day_low, n)."""
    # 1) assign notionals under the rule (chronological, entry-time info only)
    open_uj_exit=None
    for t in trades:
        if t["pair"]=="EU":
            uj_open = open_uj_exit is not None and open_uj_exit > t["entry"]
            t["gmult"] = 1.0 if uj_open else boost
        else:
            t["gmult"]=1.0
            open_uj_exit=t["exit"]
        t["notional"]=BASE_NOTIONAL*t["tier"]*t["gmult"]
        t["pnl"]=t["pnl_1x"]*t["tier"]*t["gmult"]
    # 2) index trades by active day
    active=defaultdict(list); exits=defaultdict(list)
    for t in trades:
        d=t["entry"].date()
        while d<=t["exit"].date():
            active[d].append(t); d+=timedelta(days=1)
        exits[t["exit"].date()].append(t)
    # 3) walk each day's merged bar timeline
    days=[]
    all_dates=sorted(set(active)|set(exits))
    for d in all_dates:
        acts=active.get(d,[])
        # merged events: bars of both pairs + exit events
        events=[]
        pairs_needed={t["pair"] for t in acts}
        for p in pairs_needed:
            for (bt,hi,lo) in bars[p].get(d,[]):
                events.append((bt,"bar",p,hi,lo))
        for t in exits.get(d,[]):
            events.append((t["exit"],"exit",t,None,None))
        events.sort(key=lambda e:(e[0], 0 if e[1]=="exit" else 1))
        realized=0.0; low=0.0
        floats={}   # trade id -> current signed float
        n_today=sum(1 for t in acts if t["entry"].date()==d)
        for ev in events:
            if ev[1]=="exit":
                t=ev[2]; realized+=t["pnl"]; floats.pop(id(t),None)
                low=min(low, realized+sum(floats.values()))
                continue
            bt,_,p,hi,lo_=ev
            for t in acts:
                if t["pair"]!=p: continue
                if t["entry"]<=bt<t["exit"]:
                    px = lo_ if t["dir"]>0 else hi   # worst-in-bar
                    fl = (px-t["ep"])/t["ep"]*t["notional"]*t["dir"]
                    floats[id(t)]=fl
            low=min(low, realized+sum(floats.values()))
        days.append((d, realized, low, n_today))
    return days

# ---- FTMO challenge on true day-lows ----
def run_phase(seq, target, start=0):
    eq=ACCOUNT; td=0; first=None; i=start; n=len(seq)
    while i<n:
        d,dp,dlow,ntr=seq[i]
        if first is None: first=d
        do=eq
        intraday_low=do+dlow
        if intraday_low<=MAXLOSS_FLOOR: return ("maxloss",i,(d-first).days)
        if intraday_low<=do-DAILY_LIMIT: return ("daily",i,(d-first).days)
        eq+=dp
        if ntr>0: td+=1
        if eq<=MAXLOSS_FLOOR: return ("maxloss",i,(d-first).days)
        if eq>=ACCOUNT+target and td>=4: return ("pass",i,(d-first).days)
        i+=1
    return ("incomplete",n-1,(seq[-1][0]-first).days if first else 0)

def run_challenge(seq):
    r1,i1,e1=run_phase(seq,P1_TARGET,0)
    if r1!="pass": return ("fail_p1_"+r1,e1)
    r2,i2,e2=run_phase(seq,P2_TARGET,i1+1)
    if r2!="pass": return ("fail_p2_"+r2,e1+e2)
    return ("pass",e1+e2)

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

def summ(outs,pd_,n):
    p=sum(1 for o in outs if o=="pass")
    f1d=sum(1 for o in outs if o.startswith("fail_p1_daily"))
    f1m=sum(1 for o in outs if o.startswith("fail_p1_maxloss"))
    f2=sum(1 for o in outs if o.startswith("fail_p2"))
    med=float(np.median(np.array(pd_)/30.0)) if pd_ else 0
    return {"pass":100*p/n if n else 0,"daily":100*f1d/n if n else 0,
            "maxloss":100*f1m/n if n else 0,"p2":100*f2/n if n else 0,"med":med}

def m_seq(days):
    outs=[];pd_=[];n=len(days)
    for s in range(n):
        if n-s<8: break
        o,el=run_challenge(days[s:]); outs.append(o)
        if o=="pass": pd_.append(el)
    return summ(outs,pd_,len(outs))

def m_block(days,n_sims=N_SIMS,blk=21):
    outs=[];pd_=[];n=len(days);tl=min(n,24*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); o,el=run_challenge(seq); outs.append(o)
        if o=="pass": pd_.append(el)
    return summ(outs,pd_,n_sims)

def main():
    print("="*74)
    print("EU CONDITIONAL BOOST v3 — TRUE INTRADAY EQUITY PATH | $300k=1x")
    print("(v1/v2 void: summed sequential floats as concurrent — Bennett)")
    print("="*74)
    print("Loading bars + trades...")
    bars={"EU":load_bars(next(BASE.rglob("EURUSD_15M*.csv"))),
          "UJ":load_bars(next(BASE.rglob("USDJPY_15M*.csv")))}
    trades=load_trades()
    print(f"Trades: {len(trades):,}")
    header=(f"  {'Run':>6}{'Boost':>7}{'Meth':>8}{'Pass%':>8}{'DailyF%':>9}"
            f"{'MaxL%':>7}{'P2F%':>7}{'MedMo':>7}")
    for boost in BOOSTS:
        print(f"\n{'─'*74}\nBOOST {boost}x"
              + ("  (baseline — also the CORRECTED v4 $300k figure)" if boost==1.0 else ""))
        days=build_days(trades,bars,boost)
        lows=[x[2] for x in days]
        yrs=len({x[0].year for x in days})
        breach=sum(1 for x in days if x[2]<=-DAILY_LIMIT)
        print(f"  Worst true intraday day-low: ${min(lows):,.0f} | "
              f"99th pctile ${np.percentile(lows,1):,.0f} | "
              f"days below -$5k: {breach} (~{breach/yrs:.1f}/yr)")
        print(header)
        dq=[x for x in days if x[0].year in QUIET]
        for run,dd in (("FULL",days),("QUIET",dq)):
            rS=m_seq(dd); rB=m_block(dd)
            print(f"  {run:>6}{boost:>7.1f}{'A-seq':>8}{rS['pass']:>8.1f}{rS['daily']:>9.1f}"
                  f"{rS['maxloss']:>7.1f}{rS['p2']:>7.1f}{rS['med']:>7.1f}")
            print(f"  {'':>6}{'':>7}{'B-block':>8}{rB['pass']:>8.1f}{rB['daily']:>9.1f}"
                  f"{rB['maxloss']:>7.1f}{rB['p2']:>7.1f}{rB['med']:>7.1f}")
    print(f"\n{'='*74}")
    print("Same honest read: boost row vs same-run baseline; weight QUIET;")
    print("failed challenge = fee + restart, slow pass = patience only.")
    print("="*74)

if __name__=="__main__":
    main()
