"""
build_floating_dd.py
====================
Reconstruct the REAL per-calendar-day floating drawdown of every backtest trade
by walking 15-minute price bars from entry to exit. Eliminates the mae-assignment
guess: we now know exactly how much each open position was underwater each day.

OUTPUT: data/processed/floating_dd_by_day.json
  A list of "trading day" records, each:
    {"date":"YYYY-MM-DD", "realized_pnl_300k": <closed P&L that day at 300k>,
     "worst_float_loss_300k": <worst single-position floating loss that day at 300k>,
     "n_open": <max concurrent open trades that day>}
  Floating losses are stored at $300k notional (scale later for other sizes).

For each trade:
  LONG  (signal/direction = +1): floating loss at a bar = (entry - bar_low)  * units
  SHORT (signal/direction = -1): floating loss at a bar = (bar_high - entry) * units
  units = NOTIONAL / entry_price   (so loss in account currency, USD)
Worst floating loss for a day = max over that day's bars (most underwater point).
A day's "worst_float_loss" across all open trades = the single worst position
(one trade dominates; FTMO daily rule is about total equity dip, but using the
worst single position is the dominant term and matches our prior proxy).

NOTE: We also sum concurrent floating losses per day (multiple open positions)
since FTMO measures TOTAL equity. We store BOTH worst-single and sum-concurrent.
"""
import csv, json, sys
import numpy as np
from pathlib import Path
from datetime import datetime, timedelta
from collections import defaultdict

NOTIONAL = 300_000
EU_HOLD_H = 52
UJ_HOLD_H = 24

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",
              "%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_bars(path):
    """Return list of (datetime, high, low) sorted by time. Handles dot-dates
    and either capitalized or lowercase OHLC headers."""
    bars=[]
    with open(path, newline="", encoding="utf-8-sig") as f:
        reader=csv.DictReader(f)
        # normalize header lookup
        cols={c.lower():c for c in reader.fieldnames}
        dtc=cols.get("datetime"); hc=cols.get("high"); lc=cols.get("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
            bars.append((dt,hi,lo))
    bars.sort(key=lambda x:x[0])
    return bars

def bars_index(bars):
    """Index bars by date for fast lookup: {date: [(dt,hi,lo),...]}."""
    idx=defaultdict(list)
    for b in bars: idx[b[0].date()].append(b)
    return idx

def load_eu_trades():
    p=BASE/"data"/"processed"/"trades"/"trades_real_costs.csv"
    out=[]
    with open(p, newline="", encoding="utf-8-sig") as f:
        for r in csv.DictReader(f):
            entry=parse_dt(r.get("entry_time")); exit_=parse_dt(r.get("exit_time"))
            if entry is None: continue
            if exit_ is None: exit_=entry+timedelta(hours=EU_HOLD_H)
            try:
                ep=float(r["entry_price"]); sig=int(float(r["signal"]))
                pnl=float(r["dollar_pnl_real"])
            except (ValueError,KeyError): continue
            out.append({"pair":"EU","entry":entry,"exit":exit_,"ep":ep,
                        "dir":sig,"pnl":pnl})
    return out

def load_uj_trades():
    p=BASE/"data"/"processed"/"usdjpy_trades_real_costs.csv"
    out=[]
    with open(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:
                ep=float(r["entry_price"]); d=int(float(r["direction"]))
                pnl=float(r["pnl_real"])
            except (ValueError,KeyError): continue
            out.append({"pair":"UJ","entry":entry,"exit":entry+timedelta(hours=UJ_HOLD_H),
                        "ep":ep,"dir":d,"pnl":pnl})
    return out

def trade_daily_floats(trade, bidx):
    """For one trade, walk its bars and return {date: worst_float_loss_usd}."""
    out={}
    d=trade["entry"].date(); end=trade["exit"].date()
    units=NOTIONAL/trade["ep"]
    while d<=end:
        worst=0.0
        for (bt,hi,lo) in bidx.get(d,[]):
            if bt<trade["entry"] or bt>trade["exit"]: continue
            if trade["dir"]>0:   # long: loss if price below entry
                fl=(trade["ep"]-lo)*units
            else:                # short: loss if price above entry
                fl=(hi-trade["ep"])*units
            if fl>worst: worst=fl
        if worst>0: out[d]=worst
        d=d+timedelta(days=1)
    return out

def main():
    print("Loading 15M bars (2003-2026)...")
    eu_bars=bars_index(load_bars(BASE/"data"/"raw"/"prices"/"EURUSD"/"EURUSD_15M_2003_2026.csv"))
    uj_bars=bars_index(load_bars(BASE/"data"/"raw"/"prices"/"USDJPY"/"USDJPY_15M_2003_2026.csv"))
    print(f"  EU bar-days: {len(eu_bars):,}   UJ bar-days: {len(uj_bars):,}")

    eu=load_eu_trades(); uj=load_uj_trades()
    print(f"Trades: EU {len(eu):,}, UJ {len(uj):,}")

    # accumulate per calendar day
    day_realized=defaultdict(float)         # closed pnl that day (at 300k)
    day_float_concurrent=defaultdict(float) # sum of concurrent floating losses
    day_float_worst=defaultdict(float)      # worst single floating loss
    day_nopen=defaultdict(int)

    def process(trades, bidx, label):
        for i,t in enumerate(trades):
            # realized pnl on exit day
            day_realized[t["exit"].date()] += t["pnl"]
            fl_by_day=trade_daily_floats(t, bidx)
            for d,fl in fl_by_day.items():
                day_float_concurrent[d]+=fl
                day_float_worst[d]=max(day_float_worst[d], fl)
                day_nopen[d]+=1
            if (i+1)%500==0:
                print(f"  {label}: {i+1}/{len(trades)}",flush=True)

    print("Reconstructing EU floating paths...")
    process(eu, eu_bars, "EU")
    print("Reconstructing UJ floating paths...")
    process(uj, uj_bars, "UJ")

    alldays=sorted(set(day_realized)|set(day_float_worst))
    records=[]
    for d in alldays:
        records.append({
            "date": d.isoformat(),
            "realized_pnl_300k": round(day_realized.get(d,0.0),2),
            "worst_float_loss_300k": round(day_float_worst.get(d,0.0),2),
            "concurrent_float_loss_300k": round(day_float_concurrent.get(d,0.0),2),
            "n_open": day_nopen.get(d,0),
        })

    out=BASE/"data"/"processed"/"floating_dd_by_day.json"
    out.write_text(json.dumps(records,indent=0), encoding="utf-8")
    print(f"\nWrote {out}  ({len(records):,} day records)")

    # quick stats
    wf=[r["worst_float_loss_300k"] for r in records]
    cf=[r["concurrent_float_loss_300k"] for r in records]
    wf.sort(); cf.sort()
    print(f"\nAt $300k notional:")
    print(f"  Worst single-position daily float loss: max=${max(wf):,.0f} "
          f"median=${wf[len(wf)//2]:,.0f}")
    print(f"  Concurrent (summed) daily float loss:   max=${max(cf):,.0f} "
          f"median=${cf[len(cf)//2]:,.0f}")
    over5k_worst=sum(1 for x in wf if x>5000)
    over5k_conc=sum(1 for x in cf if x>5000)
    print(f"  Days worst-single > $5,000: {over5k_worst}")
    print(f"  Days concurrent-sum > $5,000: {over5k_conc}")
    print(f"  Total day-records: {len(records)}")
    print("\nThis is REAL intraday floating DD from 15M bars - no more estimation.")

if __name__=="__main__":
    main()
