"""
risk_scaling_v1.py
===================
Tests higher base risk levels given the validated model's very low DD profile.

Current state
--------------
  Base risk  : 0.30% ($120k notional)
  95th pct DD: 1.927%  (FTMO limit = 10%)
  Pass rate  : 100%
  Headroom   : 8.07% of unused DD budget

Objective
----------
Find the highest base risk level where:
  1. 95th pct DD stays below 9.0% (1% safety buffer)
  2. Daily DD breach rate stays at 0%
  3. Pass rate stays above 99%

This maximises dollar P&L per challenge attempt without
approaching the FTMO hard limits.

Signal scaling continues to apply on top of base risk:
  |z| 2.75-3.50 → 1.0x base
  |z| 3.50-4.50 → 1.5x base
  |z| 4.50+     → 2.0x base

Risk levels tested: 0.30%, 0.50%, 0.75%, 1.00%, 1.25%, 1.50%

Place this file in:
  C:\\Users\\paul_\\OneDrive\\fx_macro_intraday\\src\\research\\risk_scaling_v1.py

Run from project root:
  python src/research/risk_scaling_v1.py
"""

import pandas as pd
import numpy as np
from pathlib import Path
import sys

BASE_PATH = Path(__file__).resolve().parents[2]
SRC_PATH  = BASE_PATH / "src"
if str(SRC_PATH) not in sys.path:
    sys.path.append(str(SRC_PATH))

TRADES_DIR = BASE_PATH / "data" / "processed" / "trades"

# ── FTMO parameters ───────────────────────────────────────────────────────────
ACCOUNT_START  = 100_000.0
PROFIT_TARGET  = ACCOUNT_START * 0.10   # Phase 1
PROFIT_P2      = ACCOUNT_START * 0.05   # Phase 2
MAX_OVERALL    = ACCOUNT_START * 0.10
MAX_DAILY      = ACCOUNT_START * 0.05   # $5,000 per day — KEY constraint

STOP_PCT       = 0.0025

ZSCORE_BANDS = [
    (2.75, 3.50, 1.0),
    (3.50, 4.50, 1.5),
    (4.50, 99.0, 2.0),
]

# Risk levels to test
RISK_LEVELS = [0.003, 0.005, 0.0075, 0.010, 0.0125, 0.015]

N_SIMS      = 2000
RANDOM_SEED = 42
YEARS       = 22.0

# Safety thresholds
DD95_LIMIT        = 9.0    # stay 1% below FTMO 10% hard limit
DAILY_BREACH_LIMIT= 0.0    # must be zero daily DD breaches
PASS_RATE_MIN     = 99.0   # must stay above 99%


def get_multiplier(z: float) -> float:
    for lo, hi, mult in ZSCORE_BANDS:
        if lo <= z < hi:
            return mult
    return ZSCORE_BANDS[-1][2]


# ── FTMO Monte Carlo ──────────────────────────────────────────────────────────
def run_mc(
    returns      : np.ndarray,
    zscores      : np.ndarray,
    exits        : np.ndarray,
    base_risk_pct: float,
    profit_pct   : float = 0.10,
    n_sims       : int   = N_SIMS,
) -> dict:
    rng           = np.random.default_rng(RANDOM_SEED)
    base_notional = (ACCOUNT_START * base_risk_pct) / STOP_PCT
    max_notional  = base_notional * 3.0
    profit_target = ACCOUNT_START * profit_pct
    per_yr        = len(returns) / YEARS

    results = []
    for _ in range(n_sims):
        idx    = rng.permutation(len(returns))
        sh_ret = returns[idx]
        sh_mult= np.array([get_multiplier(z) for z in zscores[idx]])

        balance = ACCOUNT_START
        peak    = ACCOUNT_START
        max_dd  = 0.0
        dpnl    = {}
        outcome = "INCOMPLETE"
        ttrades = 0

        for ret, mult, ex_dt in zip(sh_ret, sh_mult, exits):
            dk = str(pd.Timestamp(ex_dt).date())
            if dk not in dpnl:
                dpnl[dk] = 0.0
            notional  = min(base_notional * mult, max_notional)
            pnl       = ret * notional
            balance  += pnl
            ttrades  += 1
            dpnl[dk] += pnl

            if balance > peak:
                peak = balance
            dd = (peak - balance) / ACCOUNT_START
            if dd > max_dd:
                max_dd = dd

            if dpnl[dk] < -MAX_DAILY:
                outcome = "BREACH_DAILY"; break
            if balance <= ACCOUNT_START - MAX_OVERALL:
                outcome = "BREACH_OVERALL"; break
            if balance >= ACCOUNT_START + profit_target:
                outcome = "PASS"; break

        if outcome == "INCOMPLETE":
            outcome = ("PASS" if balance >= ACCOUNT_START + profit_target
                       else "BREACH_OVERALL"
                       if (peak - balance) / ACCOUNT_START >= 0.10
                       else "INCOMPLETE")

        results.append({
            "outcome": outcome,
            "max_dd" : max_dd * 100,
            "trades" : ttrades,
            "balance": balance,
        })

    mc = pd.DataFrame(results)
    passing = mc[mc["outcome"] == "PASS"]

    return {
        "pass_rate"    : (mc["outcome"] == "PASS").mean() * 100,
        "breach_daily" : (mc["outcome"] == "BREACH_DAILY").mean() * 100,
        "breach_overall": (mc["outcome"] == "BREACH_OVERALL").mean() * 100,
        "dd_50"        : mc["max_dd"].quantile(0.50),
        "dd_75"        : mc["max_dd"].quantile(0.75),
        "dd_90"        : mc["max_dd"].quantile(0.90),
        "dd_95"        : mc["max_dd"].quantile(0.95),
        "dd_99"        : mc["max_dd"].quantile(0.99),
        "p10_months"   : (passing["trades"].quantile(0.10) / per_yr * 12
                          if len(passing) > 0 else np.nan),
        "p25_months"   : (passing["trades"].quantile(0.25) / per_yr * 12
                          if len(passing) > 0 else np.nan),
        "med_months"   : (passing["trades"].median() / per_yr * 12
                          if len(passing) > 0 else np.nan),
        "avg_pnl_pass" : (passing["balance"] - ACCOUNT_START).mean()
                          if len(passing) > 0 else np.nan,
        "base_notional": base_notional,
        "per_yr"       : per_yr,
    }


# ── Main ──────────────────────────────────────────────────────────────────────
def main():
    print("=" * 70)
    print("RISK SCALING V1")
    print(f"  FTMO limits: overall DD 10%, daily DD 5%")
    print(f"  Safety target: 95th pct DD < {DD95_LIMIT}%,  daily breach = 0%")
    print(f"  Signal scaling: 1x / 1.5x / 2x on top of base")
    print("=" * 70)

    # Load validated trade log
    trade_path = TRADES_DIR / "trades_eurusd_final.csv"
    if not trade_path.exists():
        print(f"\n[ERROR] {trade_path} not found.")
        print("  Run final_model_validation_v1.py first.")
        return

    trades = pd.read_csv(trade_path, parse_dates=["entry_time", "exit_time"])
    trades = trades.sort_values("entry_time").reset_index(drop=True)

    returns = trades["return"].values
    zscores = trades["zscore_abs"].values
    exits   = trades["exit_time"].values

    n       = len(trades)
    per_yr  = n / YEARS

    print(f"\n  Loaded: {n:,} trades  ({per_yr:.1f}/year)")
    print(f"  Win rate  : {(trades['return']>0).mean():.2%}")
    print(f"  Avg return: {trades['return'].mean():.6f}")
    print(f"  Date range: {trades['entry_time'].min().date()} "
          f"to {trades['exit_time'].max().date()}")

    # ── Run MC for each risk level ────────────────────────────────────────────
    print(f"\nRunning {len(RISK_LEVELS)} risk levels × 2 phases × {N_SIMS:,} sims...")
    print(f"  (This may take 5-10 minutes)\n")

    rows_p1 = []
    rows_p2 = []

    for risk_pct in RISK_LEVELS:
        notional = (ACCOUNT_START * risk_pct) / STOP_PCT
        label    = f"{risk_pct*100:.2f}%"

        mc_p1 = run_mc(returns, zscores, exits, risk_pct, 0.10)
        mc_p2 = run_mc(returns, zscores, exits, risk_pct, 0.05)

        valid_p1 = (mc_p1["dd_95"] < DD95_LIMIT and
                    mc_p1["breach_daily"] <= DAILY_BREACH_LIMIT and
                    mc_p1["pass_rate"] >= PASS_RATE_MIN)

        # Dollar P&L at median passing trade count
        med_trades_p1 = mc_p1["med_months"] * per_yr / 12
        est_pnl_p1    = med_trades_p1 * trades["return"].mean() * notional

        print(f"  Risk {label:<7} Notional ${notional:>8,.0f}  "
              f"P1: pass={mc_p1['pass_rate']:.1f}%  "
              f"DD95={mc_p1['dd_95']:.2f}%  "
              f"DailyBreach={mc_p1['breach_daily']:.2f}%  "
              f"Med={mc_p1['med_months']:.1f}m  "
              f"{'✓' if valid_p1 else '✗'}")

        rows_p1.append({
            "risk_pct"     : risk_pct,
            "risk_label"   : label,
            "notional"     : notional,
            "pass_rate"    : mc_p1["pass_rate"],
            "breach_daily" : mc_p1["breach_daily"],
            "breach_overall": mc_p1["breach_overall"],
            "dd_50"        : mc_p1["dd_50"],
            "dd_75"        : mc_p1["dd_75"],
            "dd_90"        : mc_p1["dd_90"],
            "dd_95"        : mc_p1["dd_95"],
            "dd_99"        : mc_p1["dd_99"],
            "p10_months"   : mc_p1["p10_months"],
            "p25_months"   : mc_p1["p25_months"],
            "med_months"   : mc_p1["med_months"],
            "est_pnl"      : est_pnl_p1,
            "valid"        : valid_p1,
        })

        rows_p2.append({
            "risk_pct"  : risk_pct,
            "risk_label": label,
            "pass_rate" : mc_p2["pass_rate"],
            "dd_95"     : mc_p2["dd_95"],
            "med_months": mc_p2["med_months"],
            "p25_months": mc_p2["p25_months"],
        })

    # ── Full results table ────────────────────────────────────────────────────
    print(f"\n{'='*80}")
    print("PHASE 1 RESULTS TABLE (10% profit target)")
    print(f"{'='*80}")
    print(f"\n  {'Risk%':<8}{'Notional':>10}  {'Pass%':>7}  {'Daily%':>7}  "
          f"{'DD50':>6}  {'DD90':>6}  {'DD95':>6}  {'DD99':>6}  "
          f"{'P25m':>6}  {'Medm':>6}  {'EstPnL':>9}  {'OK?':>5}")
    print(f"  {'─'*94}")

    best_risk = 0
    best_label = ""
    for r in rows_p1:
        flag = "✓" if r["valid"] else "✗"
        curr = " ◄ CURRENT" if abs(r["risk_pct"] - 0.003) < 0.0001 else ""
        if r["valid"] and r["risk_pct"] > best_risk:
            best_risk  = r["risk_pct"]
            best_label = r["risk_label"]

        print(f"  {r['risk_label']:<8}{r['notional']:>10,.0f}  "
              f"{r['pass_rate']:>7.2f}  {r['breach_daily']:>7.2f}  "
              f"{r['dd_50']:>6.3f}  {r['dd_90']:>6.3f}  "
              f"{r['dd_95']:>6.3f}  {r['dd_99']:>6.3f}  "
              f"{r['p25_months']:>6.1f}  {r['med_months']:>6.1f}  "
              f"${r['est_pnl']:>8,.0f}  {flag}{curr}")

    # ── Phase 2 table ─────────────────────────────────────────────────────────
    print(f"\n{'='*80}")
    print("PHASE 2 RESULTS TABLE (5% profit target)")
    print(f"{'='*80}")
    print(f"\n  {'Risk%':<8}{'Pass%':>8}{'DD95':>8}{'P25m':>8}{'Medm':>8}")
    print(f"  {'─'*40}")
    for r in rows_p2:
        print(f"  {r['risk_label']:<8}{r['pass_rate']:>8.2f}"
              f"{r['dd_95']:>8.3f}{r['p25_months']:>8.1f}{r['med_months']:>8.1f}")

    # ── Combined P1 + P2 summary ──────────────────────────────────────────────
    print(f"\n{'='*80}")
    print("COMBINED PHASE 1 + PHASE 2 TIMELINE")
    print(f"{'='*80}")
    print(f"\n  {'Risk%':<8}{'Notional':>10}  {'P25 total':>10}  "
          f"{'Med total':>10}  {'DD95 P1':>9}  {'Valid':>6}")
    print(f"  {'─'*58}")

    for r1, r2 in zip(rows_p1, rows_p2):
        combined_p25 = r1["p25_months"] + r2["p25_months"]
        combined_med = r1["med_months"]  + r2["med_months"]
        flag         = "✓" if r1["valid"] else "✗"
        curr         = " ◄ CURRENT" if abs(r1["risk_pct"] - 0.003) < 0.0001 else ""
        print(f"  {r1['risk_label']:<8}{r1['notional']:>10,.0f}  "
              f"{combined_p25:>10.1f}  {combined_med:>10.1f}  "
              f"{r1['dd_95']:>9.3f}  {flag}{curr}")

    # ── Recommendation ────────────────────────────────────────────────────────
    print(f"\n{'='*80}")
    print("RECOMMENDATION")
    print(f"{'='*80}")

    valid_rows = [r for r in rows_p1 if r["valid"]]
    current    = next(r for r in rows_p1 if abs(r["risk_pct"]-0.003) < 0.0001)

    if valid_rows:
        best = max(valid_rows, key=lambda r: r["risk_pct"])

        # Find matching P2 row
        best_p2 = next(r for r in rows_p2
                       if abs(r["risk_pct"] - best["risk_pct"]) < 0.0001)

        pnl_uplift = (best["est_pnl"] / current["est_pnl"] - 1) * 100 \
                     if current["est_pnl"] > 0 else 0

        print(f"""
  Highest valid risk level: {best['risk_label']}
  Notional per trade      : ${best['notional']:,.0f}
  (with 2x signal scaling): ${best['notional']*2:,.0f} max

  vs current 0.30% ($120k):
    DD95 change   : {current['dd_95']:.3f}% → {best['dd_95']:.3f}%
    Pass rate     : {current['pass_rate']:.2f}% → {best['pass_rate']:.2f}%
    Daily breach  : {current['breach_daily']:.2f}% → {best['breach_daily']:.2f}%
    P25 months P1 : {current['p25_months']:.1f} → {best['p25_months']:.1f}
    Med months P1 : {current['med_months']:.1f} → {best['med_months']:.1f}
    Est P&L uplift: +{pnl_uplift:.0f}%

  Combined challenge (P1 + P2) at {best['risk_label']}:
    Best 25% of attempts : {best['p25_months'] + best_p2['p25_months']:.1f} months
    Median               : {best['med_months']  + best_p2['med_months']:.1f} months

  Daily DD note:
    At {best['risk_label']} risk, stop = ${best['notional']*0.0025:,.0f}/trade
    Worst realistic day (3 stops): ${best['notional']*0.0025*3:,.0f}
    vs FTMO daily limit          : $5,000
    Buffer remaining             : ${5000 - best['notional']*0.0025*3:,.0f}

  IMPORTANT: The walk-forward validation was run at 0.30% risk.
  Higher risk scales P&L proportionally but does NOT change the
  signal quality or win rate. The edge itself is unchanged.
  Only the dollar magnitude of wins and losses scales up.
""")

    # Save results
    pd.DataFrame(rows_p1).to_csv(
        TRADES_DIR / "risk_scaling_results.csv", index=False)
    print(f"  Results saved: {TRADES_DIR / 'risk_scaling_results.csv'}")


if __name__ == "__main__":
    main()
