"""
improvement_testing_v1.py
==========================
Tests two legitimate improvements to reduce trades-to-pass while keeping
95th pct max DD below the 10% FTMO hard limit.

Improvement 1: Compounding Sizing
----------------------------------
Currently the simulator uses FIXED notional ($120k forever).
With compounding, notional = current_balance * risk_pct / stop_pct
As the account grows toward $110k target, each winner slightly increases
the next position size. This accelerates the path to the profit target
without changing any signal parameters.

Improvement 2: Signal-Strength Scaling
----------------------------------------
Z-score of 4.5 has more macro conviction than z-score of 2.75.
Scale risk proportionally within defined bands:

  Band 1: |z| 2.75 - 3.50  →  base risk (0.30%)  →  $120k notional
  Band 2: |z| 3.50 - 4.50  →  1.5x risk (0.45%)  →  $180k notional
  Band 3: |z| 4.50+         →  2.0x risk (0.60%)  →  $240k notional

  But: each band is capped so the maximum single-trade loss never
  pushes daily DD past $4,000 (leaving $1,000 buffer to the $5,000 limit)

Improvement 3: Both Combined
------------------------------
Compounding + signal-strength scaling together.

Validation
----------
Each approach runs 2,000 Monte Carlo simulations.
Must keep 95th pct max DD < 10.0% to be accepted.
Primary metric: avg trades to pass (lower = better challenge pacing).

Output
------
  Full comparison table
  Recommendation
  CSVs saved to data/processed/trades/

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

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

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

# ── Path setup ────────────────────────────────────────────────────────────────
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"
OUTPUT_DIR = TRADES_DIR
OUTPUT_DIR.mkdir(parents=True, exist_ok=True)

# ── FTMO Rules ────────────────────────────────────────────────────────────────
ACCOUNT_START    = 100_000.0
PROFIT_TARGET    = ACCOUNT_START * 0.10     # $110,000
MAX_OVERALL_LOSS = ACCOUNT_START * 0.10     # floor $90,000
MAX_DAILY_LOSS   = ACCOUNT_START * 0.05     # $5,000 per day

# ── Validated frozen parameters ───────────────────────────────────────────────
STOP_PCT         = 0.0025                   # 0.25% — validated optimal
BASE_RISK_PCT    = 0.0030                   # 0.30% — validated from calibration
BASE_NOTIONAL    = (ACCOUNT_START * BASE_RISK_PCT) / STOP_PCT  # $120,000

# ── Signal-strength scaling bands ────────────────────────────────────────────
# Multipliers applied to base notional when |z-score| is in each band
# Band edges chosen to have meaningful population in each bucket
ZSCORE_BANDS = [
    # (min_z,  max_z,  risk_multiplier,  label)
    (2.75, 3.50, 1.0,  "base  2.75-3.50"),
    (3.50, 4.50, 1.5,  "med   3.50-4.50"),
    (4.50, 99.0, 2.0,  "high  4.50+    "),
]

# Single-trade max dollar loss cap — protects daily DD limit
# Max 2 trades can close on same day. $4k / 2 = $2k max per trade.
# But to be conservative we cap at $3,500 per trade loss
MAX_TRADE_LOSS_DOLLAR = 3_500.0
MAX_NOTIONAL_CAP      = MAX_TRADE_LOSS_DOLLAR / STOP_PCT  # $1,400,000 (never binding)
# More practically cap at 3x base
NOTIONAL_HARD_CAP     = BASE_NOTIONAL * 3.0               # $360,000

N_SIMS      = 2000
RANDOM_SEED = 42
CANDIDATE   = "growth_52h"   # validated winner from calibration


# ── Load the original trade log with z-scores ─────────────────────────────────
def load_trades_with_zscore() -> pd.DataFrame:
    """
    Loads the trade log and attaches z-scores from the original signal data.
    The z-score at each trade entry is needed for signal-strength scaling.
    We rebuild the signal data and merge on entry datetime.
    """
    # Load base trade log
    csv_path = TRADES_DIR / f"trades_{CANDIDATE}.csv"
    if not csv_path.exists():
        raise FileNotFoundError(f"{csv_path} not found. Run export_trade_logs_v1.py first.")

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

    # Load signal data to get z-scores
    from features.spot_lag_v3 import get_model_ready_spot_lag_v3
    signals_df = get_model_ready_spot_lag_v3().copy()
    signals_df["datetime"] = pd.to_datetime(signals_df["datetime"])
    signals_df = signals_df[["datetime", "lag_zscore_24h_v3"]].copy()

    # The entry_time is on 15M bars (pullback entry), but the signal fires on 1H
    # We need to match each trade back to its originating signal z-score
    # Strategy: for each trade, find the signal bar closest to and before entry_time
    # within a 6-hour window (the pullback wait window)
    trades_sorted  = trades.sort_values("entry_time").reset_index(drop=True)
    signals_sorted = signals_df.sort_values("datetime").reset_index(drop=True)

    # Merge backward: for each entry_time, find the most recent signal datetime
    merged = pd.merge_asof(
        trades_sorted.rename(columns={"entry_time": "datetime"}),
        signals_sorted,
        on="datetime",
        direction="backward",
        tolerance=pd.Timedelta("7h"),  # signal must be within 7h before entry
    )
    merged = merged.rename(columns={"datetime": "entry_time"})
    merged["zscore_abs"] = merged["lag_zscore_24h_v3"].abs()

    missing = merged["zscore_abs"].isna().sum()
    if missing > 0:
        print(f"  [NOTE] {missing} trades missing z-score match — "
              f"filling with median {merged['zscore_abs'].median():.3f}")
        merged["zscore_abs"] = merged["zscore_abs"].fillna(merged["zscore_abs"].median())

    return merged


# ── Assign signal-strength multiplier ────────────────────────────────────────
def get_signal_multiplier(zscore_abs: float) -> float:
    for min_z, max_z, mult, _ in ZSCORE_BANDS:
        if min_z <= zscore_abs < max_z:
            return mult
    return ZSCORE_BANDS[-1][2]  # fallback to last band


# ── Build dollar P&L arrays for each approach ─────────────────────────────────
def build_pnl_arrays(trades: pd.DataFrame) -> dict:
    """
    Returns a dict of {approach_name: dollar_pnl_array} for Monte Carlo.
    """
    returns    = trades["return"].values
    zscores    = trades["zscore_abs"].values

    # 1. Fixed baseline (reference — already validated)
    pnl_fixed = returns * BASE_NOTIONAL

    # 2. Signal-strength scaling (fixed account base)
    multipliers = np.array([get_signal_multiplier(z) for z in zscores])
    notional_ss = np.clip(BASE_NOTIONAL * multipliers, BASE_NOTIONAL, NOTIONAL_HARD_CAP)
    pnl_ss      = returns * notional_ss

    # Return dict — compound approaches need per-sim calculation
    return {
        "returns"       : returns,
        "zscores"       : zscores,
        "pnl_fixed"     : pnl_fixed,
        "pnl_ss"        : pnl_ss,
        "multipliers"   : multipliers,
        "notional_ss"   : notional_ss,
    }


# ── FTMO sim (supports compounding) ──────────────────────────────────────────
def simulate_ftmo(
    returns      : np.ndarray,
    exit_dates   : pd.Series,
    compound     : bool  = False,
    multipliers  : np.ndarray | None = None,  # signal-strength multipliers
) -> dict:
    """
    Runs one FTMO challenge simulation.

    compound    : if True, notional = current_balance * BASE_RISK_PCT / STOP_PCT
    multipliers : per-trade signal-strength multiplier array (optional)
                  applied on top of base or compound notional
    """
    balance          = ACCOUNT_START
    peak_balance     = ACCOUNT_START
    max_dd_reached   = 0.0
    daily_pnl        = {}
    outcome          = "INCOMPLETE"
    trading_days_set = set()
    trades_taken     = 0

    for i, (ret, exit_dt) in enumerate(zip(returns, exit_dates)):
        date_key = str(pd.Timestamp(exit_dt).date())

        if date_key not in daily_pnl:
            daily_pnl[date_key] = 0.0

        # Compute notional for this trade
        if compound:
            base_n = (balance * BASE_RISK_PCT) / STOP_PCT
        else:
            base_n = BASE_NOTIONAL

        # Apply signal-strength multiplier if provided
        mult = float(multipliers[i]) if multipliers is not None else 1.0
        notional = min(base_n * mult, NOTIONAL_HARD_CAP)

        dollar_pnl = ret * notional

        balance              += dollar_pnl
        trades_taken         += 1
        trading_days_set.add(date_key)
        daily_pnl[date_key]  += dollar_pnl

        if balance > peak_balance:
            peak_balance = balance

        current_dd = (peak_balance - balance) / ACCOUNT_START
        if current_dd > max_dd_reached:
            max_dd_reached = current_dd

        # FTMO breach checks
        if daily_pnl[date_key] < -MAX_DAILY_LOSS:
            outcome = "BREACH_DAILY_DD"
            break

        if balance <= (ACCOUNT_START - MAX_OVERALL_LOSS):
            outcome = "BREACH_OVERALL_DD"
            break

        if balance >= (ACCOUNT_START + PROFIT_TARGET):
            outcome = "PASS"
            break

    if outcome == "INCOMPLETE":
        if balance >= (ACCOUNT_START + PROFIT_TARGET):
            outcome = "PASS"
        elif (peak_balance - balance) / ACCOUNT_START >= 0.10:
            outcome = "BREACH_OVERALL_DD"

    return {
        "outcome"             : outcome,
        "final_balance"       : balance,
        "max_dd_reached_pct"  : max_dd_reached * 100,
        "total_pnl_pct"       : (balance / ACCOUNT_START - 1) * 100,
        "trades_taken"        : trades_taken,
        "trading_days"        : len(trading_days_set),
    }


# ── Monte Carlo runner ────────────────────────────────────────────────────────
def run_mc(
    returns     : np.ndarray,
    exit_dates  : pd.Series,
    compound    : bool             = False,
    multipliers : np.ndarray | None = None,
    n_sims      : int              = N_SIMS,
) -> pd.DataFrame:
    rng     = np.random.default_rng(RANDOM_SEED)
    results = []

    for sim_i in range(n_sims):
        # Shuffle returns — keep multipliers aligned with shuffled returns
        idx      = rng.permutation(len(returns))
        sh_ret   = returns[idx]
        sh_mult  = multipliers[idx] if multipliers is not None else None

        r          = simulate_ftmo(sh_ret, exit_dates, compound=compound,
                                   multipliers=sh_mult)
        r["sim"]   = sim_i
        results.append(r)

    return pd.DataFrame(results)


# ── Summarise MC results ──────────────────────────────────────────────────────
def summarise_mc(mc: pd.DataFrame) -> dict:
    n          = len(mc)
    passing    = mc[mc["outcome"] == "PASS"]
    return {
        "pass_rate"         : (mc["outcome"] == "PASS").mean() * 100,
        "breach_rate"       : mc["outcome"].str.startswith("BREACH").mean() * 100,
        "daily_breach_rate" : (mc["outcome"] == "BREACH_DAILY_DD").mean() * 100,
        "dd_50"             : mc["max_dd_reached_pct"].quantile(0.50),
        "dd_75"             : mc["max_dd_reached_pct"].quantile(0.75),
        "dd_90"             : mc["max_dd_reached_pct"].quantile(0.90),
        "dd_95"             : mc["max_dd_reached_pct"].quantile(0.95),
        "danger_pct"        : (mc["max_dd_reached_pct"] > 8.0).mean() * 100,
        "avg_trades_pass"   : passing["trades_taken"].mean() if len(passing) > 0 else np.nan,
        "med_trades_pass"   : passing["trades_taken"].median() if len(passing) > 0 else np.nan,
        "avg_pnl_pct"       : mc["total_pnl_pct"].mean(),
    }


# ── Print comparison ──────────────────────────────────────────────────────────
def print_comparison(results: dict):
    dd_limit = 10.0

    print(f"\n{'='*80}")
    print("IMPROVEMENT COMPARISON — growth_52h")
    print(f"{'='*80}")
    print(f"\n  {'Approach':<35}{'Pass%':>7}  {'95pctDD':>8}  "
          f"{'DailyBreach':>12}  {'Danger%':>8}  "
          f"{'AvgTrades':>10}  {'MedTrades':>10}  {'Status':>8}")
    print(f"  {'─'*105}")

    for name, s in results.items():
        status = "✓ VALID" if s["dd_95"] < dd_limit else "✗ BREACH"
        current = " ◄ CURRENT" if name == "1_fixed_baseline" else ""
        print(f"  {name:<35}{s['pass_rate']:>7.2f}  "
              f"{s['dd_95']:>8.3f}  "
              f"{s['daily_breach_rate']:>12.2f}  "
              f"{s['danger_pct']:>8.1f}  "
              f"{s['avg_trades_pass']:>10.0f}  "
              f"{s['med_trades_pass']:>10.0f}  "
              f"{status}{current}")


# ── Z-score band analysis ─────────────────────────────────────────────────────
def analyse_zscore_bands(trades: pd.DataFrame):
    print(f"\n{'─'*60}")
    print("Z-SCORE BAND ANALYSIS")
    print(f"{'─'*60}")
    print(f"\n  {'Band':<20}{'Trades':>8}{'%Total':>8}{'WinRate':>9}"
          f"{'AvgRet':>10}{'AvgRet/$':>10}")
    print(f"  {'─'*65}")

    total = len(trades)
    for min_z, max_z, mult, label in ZSCORE_BANDS:
        mask   = (trades["zscore_abs"] >= min_z) & (trades["zscore_abs"] < max_z)
        subset = trades[mask]
        if len(subset) == 0:
            continue
        wr     = (subset["return"] > 0).mean() * 100
        avgret = subset["return"].mean()
        # Dollar expectancy at this band's notional
        notional = BASE_NOTIONAL * mult
        dollar_exp = avgret * notional
        print(f"  {label:<20}{len(subset):>8}{len(subset)/total*100:>7.1f}%"
              f"{wr:>8.2f}%{avgret:>10.6f}{dollar_exp:>+10.2f}")

    print(f"\n  Key: if higher z-score bands have better win rate AND avg return,")
    print(f"  signal-strength scaling is justified by the data.")


# ── Main ──────────────────────────────────────────────────────────────────────
def main():
    print("=" * 70)
    print("IMPROVEMENT TESTING V1")
    print(f"  Candidate     : {CANDIDATE}")
    print(f"  Base risk     : {BASE_RISK_PCT:.2%}  (${BASE_NOTIONAL:,.0f} notional)")
    print(f"  Stop          : {STOP_PCT:.2%}")
    print(f"  MC sims       : {N_SIMS:,}")
    print(f"  DD hard limit : 10%  (95th pct)")
    print("=" * 70)

    # Load trades with z-scores
    print("\nLoading trades and attaching z-scores...")
    trades = load_trades_with_zscore()
    print(f"  Trades loaded : {len(trades)}")
    print(f"  Z-score range : {trades['zscore_abs'].min():.3f} — "
          f"{trades['zscore_abs'].max():.3f}")
    print(f"  Z-score median: {trades['zscore_abs'].median():.3f}")

    # Analyse z-score bands
    analyse_zscore_bands(trades)

    # Build arrays
    arrays = build_pnl_arrays(trades)
    returns     = arrays["returns"]
    multipliers = arrays["multipliers"]
    exit_dates  = trades["exit_time"]

    ones = np.ones(len(returns))  # neutral multiplier array

    # ── Run all four approaches ───────────────────────────────────────────────
    approaches = {
        "1_fixed_baseline"     : (False, ones),
        "2_compounding"        : (True,  ones),
        "3_signal_scaling"     : (False, multipliers),
        "4_compound_AND_scale" : (True,  multipliers),
    }

    mc_results = {}
    summaries  = {}

    for approach_name, (compound, mults) in approaches.items():
        print(f"\n  Running MC: {approach_name}...")
        mc = run_mc(returns, exit_dates, compound=compound, multipliers=mults)
        mc_results[approach_name] = mc
        summaries[approach_name]  = summarise_mc(mc)

    # Print comparison
    print_comparison(summaries)

    # ── Detailed breakdown of valid approaches ────────────────────────────────
    print(f"\n{'='*70}")
    print("DETAILED BREAKDOWN — VALID APPROACHES ONLY (95th pct DD < 10%)")
    print(f"{'='*70}")

    valid_approaches = {k: v for k, v in summaries.items() if v["dd_95"] < 10.0}

    if not valid_approaches:
        print("\n  [WARNING] No approaches keep 95th pct DD below 10%")
        print("  Showing best available (lowest 95th pct DD):")
        best_k = min(summaries, key=lambda k: summaries[k]["dd_95"])
        valid_approaches = {best_k: summaries[best_k]}

    for name, s in valid_approaches.items():
        mc = mc_results[name]
        passing = mc[mc["outcome"] == "PASS"]
        print(f"\n  {name}:")
        print(f"    Pass rate         : {s['pass_rate']:.2f}%")
        print(f"    Breach rate       : {s['breach_rate']:.2f}%")
        print(f"    Daily DD breach   : {s['daily_breach_rate']:.2f}%")
        print(f"    Median max DD     : {s['dd_50']:.3f}%")
        print(f"    75th pct max DD   : {s['dd_75']:.3f}%")
        print(f"    95th pct max DD   : {s['dd_95']:.3f}%  (limit 10%)")
        print(f"    Sims >8% DD       : {s['danger_pct']:.1f}%")
        print(f"    Avg trades to pass: {s['avg_trades_pass']:.0f}")
        print(f"    Med trades to pass: {s['med_trades_pass']:.0f}")

        if len(passing) > 0:
            trades_p5  = passing["trades_taken"].quantile(0.05)
            trades_p25 = passing["trades_taken"].quantile(0.25)
            trades_p75 = passing["trades_taken"].quantile(0.75)
            trades_p95 = passing["trades_taken"].quantile(0.95)
            print(f"\n    Trades-to-pass distribution (passing sims only):")
            print(f"      5th  pct : {trades_p5:.0f}")
            print(f"      25th pct : {trades_p25:.0f}")
            print(f"      Median   : {s['med_trades_pass']:.0f}")
            print(f"      75th pct : {trades_p75:.0f}")
            print(f"      95th pct : {trades_p95:.0f}")

            # Annualised estimate
            trades_per_year = len(trades) / 22.0  # 2003-2026 ≈ 22 years
            avg_months = (s["avg_trades_pass"] / trades_per_year) * 12
            med_months = (s["med_trades_pass"] / trades_per_year) * 12
            p25_months = (trades_p25 / trades_per_year) * 12
            print(f"\n    Estimated challenge duration (at ~{trades_per_year:.0f} trades/year):")
            print(f"      Best quarter  : ~{p25_months:.1f} months")
            print(f"      Median        : ~{med_months:.1f} months")
            print(f"      Average       : ~{avg_months:.1f} months")

    # ── Sequential validation of best valid approach ──────────────────────────
    print(f"\n{'='*70}")
    print("SEQUENTIAL VALIDATION (historical order)")
    print(f"{'='*70}")

    for approach_name, (compound, mults) in approaches.items():
        s = summaries[approach_name]
        if s["dd_95"] >= 10.0:
            continue

        result = simulate_ftmo(returns, exit_dates, compound=compound,
                               multipliers=mults)
        print(f"\n  {approach_name}:")
        print(f"    Outcome      : {result['outcome']}")
        print(f"    Final bal    : ${result['final_balance']:>12,.2f}")
        print(f"    P&L          : {result['total_pnl_pct']:>+.2f}%")
        print(f"    Max DD       : {result['max_dd_reached_pct']:.3f}%")
        print(f"    Trades taken : {result['trades_taken']}")

    # ── Save results ──────────────────────────────────────────────────────────
    summary_rows = []
    for name, s in summaries.items():
        row = {"approach": name}
        row.update(s)
        summary_rows.append(row)

    summary_df = pd.DataFrame(summary_rows)
    out_path   = OUTPUT_DIR / f"improvement_results_{CANDIDATE}.csv"
    summary_df.to_csv(out_path, index=False)
    print(f"\n  Results saved: {out_path}")

    # ── Final recommendation ──────────────────────────────────────────────────
    print(f"\n{'='*70}")
    print("RECOMMENDATION")
    print(f"{'='*70}")

    if valid_approaches:
        # Best = lowest avg_trades_pass among valid
        best_name = min(valid_approaches,
                        key=lambda k: valid_approaches[k]["avg_trades_pass"])
        best      = valid_approaches[best_name]
        baseline  = summaries["1_fixed_baseline"]

        trades_improvement = (
            (baseline["avg_trades_pass"] - best["avg_trades_pass"])
            / baseline["avg_trades_pass"] * 100
        )
        trades_per_year = len(trades) / 22.0
        months_saved    = (
            (baseline["avg_trades_pass"] - best["avg_trades_pass"])
            / trades_per_year * 12
        )

        print(f"""
  Best approach: {best_name}

  vs fixed baseline:
    Pass rate     : {baseline['pass_rate']:.2f}%  →  {best['pass_rate']:.2f}%
    95th pct DD   : {baseline['dd_95']:.3f}%  →  {best['dd_95']:.3f}%
    Avg trades    : {baseline['avg_trades_pass']:.0f}  →  {best['avg_trades_pass']:.0f}
    Improvement   : {trades_improvement:.1f}% fewer trades to pass
    Time saved    : ~{months_saved:.1f} months off average challenge duration

  If compounding alone is the winner:
    → Zero new signal research needed. Just size on current balance
      instead of fixed initial. Safe, validated, easy to implement.

  If signal scaling adds genuine improvement:
    → Validate that z-score bands genuinely differ in quality
      (check the z-score band analysis table above — look for
      higher win rate and avg return in stronger z-score bands).
    → If the data supports it, scaling is legitimate.
    → If win rates are similar across bands, don't scale — 
      the macro signal conviction doesn't differentiate within
      the trade population at that resolution.

  Next: freeze whichever approach wins here and move to the
  regime filter (Layer 2) which will further reduce losing
  clusters and improve the DD profile.
""")


if __name__ == "__main__":
    main()
