"""
concurrent_trade_audit_v1.py
==============================
Investigates whether the original backtest allowed concurrent trades
and what the correct trade frequency should be for the live system.

The gap:
  Original backtest : 4,019 trades (3.5/week)
  Entry method test : 1,607 trades (1.4/week)

Hypothesis: original backtest allowed new signals to fire while
a trade was already open, treating each signal independently.

Tests:
  1. Concurrent trades allowed (no position limit)
  2. One trade at a time (current live system)
  3. Max 2 concurrent trades

Also checks:
  - How often do signals overlap with open positions?
  - What is the signal autocorrelation (do signals cluster)?
  - What does the trade log show about concurrent trade timing?

Place in:
  C:\\Users\\paul_\\OneDrive\\fx_macro_intraday\\src\\research\\concurrent_trade_audit_v1.py

Run from project root:
  python src/research/concurrent_trade_audit_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))

from research.combined_candidate_matrix_v1 import (
    build_frozen_signals,
    find_entry_pullback,
    get_first_m15_idx_at_or_after,
)
from ingestion.price_loader_15m import load_eurusd_15m

TRADES_DIR    = BASE_PATH / "data" / "processed" / "trades"
THRESHOLD     = 2.75
FIB           = 0.786
STOP          = 0.0025
TP            = 0.0020
ZSCORE_EXIT   = 1.5
SPREAD_COST   = 0.0001
ALLOWED_HOURS = set(range(7, 17))
ACCOUNT_START = 100_000.0
BASE_NOTIONAL = 300_000.0
YEARS         = 22.0
HOLD_HOURS    = 52

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

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


def simulate_trade(m15, signal_df, entry_time, entry_price, signal_dir):
    entry_idx = get_first_m15_idx_at_or_after(m15, entry_time)
    if entry_idx is None: return None, None, None

    hold_bars = HOLD_HOURS * 4
    exit_idx  = min(entry_idx + hold_bars, len(m15)-1)
    path      = m15.iloc[entry_idx:exit_idx+1]
    if path.empty: return None, None, None

    sig_window = signal_df[
        (signal_df["datetime"] >= entry_time) &
        (signal_df["datetime"] <= path.iloc[-1]["datetime"])
    ]

    exit_price  = float(path.iloc[-1]["close"])
    exit_reason = "time"
    exit_time   = path.iloc[-1]["datetime"]

    for _, bar in path.iterrows():
        bt = bar["datetime"]
        if signal_dir == 1:
            tp_hit   = (float(bar["high"])-entry_price)/entry_price >= TP
            stop_hit = (float(bar["low"])-entry_price)/entry_price  <= -STOP
        else:
            tp_hit   = (entry_price-float(bar["low"]))/entry_price  >= TP
            stop_hit = (entry_price-float(bar["high"]))/entry_price <= -STOP

        if tp_hit:
            exit_price  = entry_price*(1+TP) if signal_dir==1 else entry_price*(1-TP)
            exit_reason = "tp"; exit_time = bt; break
        if stop_hit:
            exit_price  = entry_price*(1-STOP) if signal_dir==1 else entry_price*(1+STOP)
            exit_reason = "stop"; exit_time = bt; break

        z_bars = sig_window[sig_window["datetime"]<=bt]
        if not z_bars.empty:
            cz = float(z_bars.iloc[-1]["lag_zscore_24h_v3"])
            if (signal_dir==1 and cz<=-ZSCORE_EXIT) or (signal_dir==-1 and cz>=ZSCORE_EXIT):
                exit_price  = float(bar["close"])
                exit_reason = "zexit"; exit_time = bt; break

    raw_ret = (exit_price-entry_price)/entry_price * signal_dir
    return raw_ret-SPREAD_COST, exit_reason, exit_time


def calc_metrics(trades_list, label):
    if not trades_list: return None
    df = pd.DataFrame(trades_list)
    r  = df["ret"].values
    d  = df["dollar_pnl"].values
    n  = len(df)
    equity = ACCOUNT_START + d.cumsum()
    peak   = np.maximum.accumulate(equity)
    dd     = (equity-peak)/peak*100
    per_yr = n/YEARS
    rf_pt  = (1+0.04)**(1/per_yr)-1
    excess = d/ACCOUNT_START-rf_pt
    sharpe = (excess.mean()/excess.std()*np.sqrt(per_yr) if excess.std()>0 else 0)
    gp = d[d>0].sum(); gl = abs(d[d<0].sum())
    pf = gp/gl if gl>0 else 999
    monthly_pnl = d.sum()/YEARS/12
    return {
        "label"      : label,
        "n"          : n,
        "per_week"   : n/YEARS/52,
        "win_rate"   : df["win"].mean()*100,
        "avg_ret"    : r.mean()*100,
        "sharpe"     : sharpe,
        "pf"         : pf,
        "max_dd"     : dd.min(),
        "net_pnl"    : d.sum(),
        "cagr"       : (equity[-1]/ACCOUNT_START)**(1/YEARS)*100-100,
        "monthly_pnl": monthly_pnl,
        "p1_months"  : 10000/monthly_pnl if monthly_pnl>0 else 99,
        "tp_pct"     : df["exit_reason"].value_counts(normalize=True).get("tp",0)*100,
        "stop_pct"   : df["exit_reason"].value_counts(normalize=True).get("stop",0)*100,
    }


def run_with_concurrent_limit(signals, m15, signal_df, max_concurrent):
    """
    Runs backtest allowing up to max_concurrent open trades simultaneously.
    Each signal is independent — new trades can open while others are running.
    This matches how the original backtest likely worked.
    """
    trades      = []
    open_trades = []  # list of exit_times for currently open trades

    for _, sig in signals.iterrows():
        sig_time   = sig["datetime"]
        signal_dir = int(sig["signal"])
        zscore_abs = abs(float(sig["lag_zscore_24h_v3"]))
        mult       = get_multiplier(zscore_abs)
        notional   = min(BASE_NOTIONAL*mult, BASE_NOTIONAL*3)

        # Remove trades that have closed by now
        open_trades = [et for et in open_trades if et > sig_time]

        # Check concurrent limit
        if len(open_trades) >= max_concurrent:
            continue

        # Find Fibonacci entry
        entry = find_entry_pullback(m15=m15, signal_row=sig,
                                     fib=FIB, wait_hours=6)
        if entry is None:
            continue

        entry_time  = entry["entry_time"]
        entry_price = entry["entry_price"]

        ret, reason, exit_time = simulate_trade(
            m15, signal_df, entry_time, entry_price, signal_dir
        )
        if ret is None:
            continue

        trades.append({
            "entry_time"  : entry_time,
            "exit_time"   : exit_time,
            "ret"         : ret,
            "exit_reason" : reason,
            "win"         : int(ret>0),
            "dollar_pnl"  : ret*notional,
            "zscore_abs"  : zscore_abs,
        })

        if exit_time is not None:
            open_trades.append(exit_time)

    return trades


def main():
    print("="*70)
    print("CONCURRENT TRADE AUDIT  —  Resolving the 4,019 vs 1,607 Gap")
    print("="*70)

    print("\nLoading data...")
    m15     = load_eurusd_15m().sort_values("datetime").reset_index(drop=True)
    signals = build_frozen_signals(threshold=THRESHOLD, allowed_hours=ALLOWED_HOURS)

    try:
        from features.spot_lag_v3 import get_model_ready_spot_lag_v3
        signal_df = get_model_ready_spot_lag_v3().copy()
        signal_df["datetime"] = pd.to_datetime(signal_df["datetime"])
        signal_df = signal_df.sort_values("datetime").reset_index(drop=True)
    except Exception as e:
        print(f"  Warning: {e}")
        signal_df = pd.DataFrame(columns=["datetime","lag_zscore_24h_v3"])

    # ── PART 1: Signal overlap analysis ──────────────────────────────────────
    print(f"\n{'─'*70}")
    print("PART 1: HOW OFTEN DO SIGNALS OVERLAP WITH OPEN POSITIONS?")
    print(f"{'─'*70}")

    overlap_count = 0
    total_checked = 0
    prev_exit = None

    for _, sig in signals.iterrows():
        sig_time = sig["datetime"]
        total_checked += 1
        if prev_exit is not None and sig_time < prev_exit:
            overlap_count += 1
        entry = find_entry_pullback(m15=m15, signal_row=sig,
                                     fib=FIB, wait_hours=6)
        if entry:
            prev_exit = entry["entry_time"] + pd.Timedelta(hours=HOLD_HOURS)

    print(f"\n  Total signals              : {total_checked:,}")
    print(f"  Signals during open trade  : {overlap_count:,} ({overlap_count/total_checked*100:.1f}%)")
    print(f"  This explains the gap —")
    print(f"  if original backtest ignored open positions, it got more trades")

    # ── PART 2: Check original trade log for overlaps ─────────────────────────
    print(f"\n{'─'*70}")
    print("PART 2: DID THE ORIGINAL BACKTEST ALLOW CONCURRENT TRADES?")
    print(f"{'─'*70}")

    for fname in ["trades_eurusd_final.csv", "trades_real_costs.csv",
                  "trades_growth_52h.csv"]:
        path = TRADES_DIR / fname
        if path.exists():
            df_trades = pd.read_csv(path, parse_dates=["entry_time","exit_time"])
            df_trades = df_trades.sort_values("entry_time").reset_index(drop=True)

            # Check for overlapping entry/exit times
            overlaps = 0
            for i in range(1, len(df_trades)):
                if df_trades.iloc[i]["entry_time"] < df_trades.iloc[i-1]["exit_time"]:
                    overlaps += 1

            overlap_pct = overlaps / len(df_trades) * 100
            print(f"\n  File: {fname}")
            print(f"  Total trades    : {len(df_trades):,}")
            print(f"  Overlapping pairs: {overlaps:,} ({overlap_pct:.1f}%)")
            if overlap_pct > 10:
                print(f"  -> CONFIRMED: Original backtest ALLOWED concurrent trades")
                print(f"  -> This explains the 4,019 vs 1,607 gap")
            else:
                print(f"  -> Original backtest was sequential (one trade at a time)")
            break

    # ── PART 3: Test concurrent limits ───────────────────────────────────────
    print(f"\n{'─'*70}")
    print("PART 3: RESULTS WITH DIFFERENT CONCURRENT TRADE LIMITS")
    print(f"{'─'*70}\n")

    results = []
    for max_c in [1, 2, 3, 5, 999]:
        label = f"Max {max_c} concurrent" if max_c < 999 else "Unlimited concurrent"
        print(f"  Testing {label}...", end="", flush=True)
        trades = run_with_concurrent_limit(signals, m15, signal_df, max_c)
        m = calc_metrics(trades, label)
        if m:
            results.append(m)
            print(f"  n={m['n']:,}  WR={m['win_rate']:.1f}%  "
                  f"Sharpe={m['sharpe']:.2f}  {m['per_week']:.1f}/wk  "
                  f"P1={m['p1_months']:.1f}mo")

    # ── Print results table ───────────────────────────────────────────────────
    print(f"\n{'='*85}")
    print("CONCURRENT TRADE LIMIT COMPARISON")
    print(f"{'='*85}")
    print(f"\n  {'Limit':<25}{'Trades':>7}{'WR%':>7}{'Sharpe':>8}{'PF':>7}"
          f"{'CAGR%':>7}{'MaxDD%':>8}{'P1 Mo':>7}{'Per Wk':>8}{'Monthly PnL':>13}")
    print(f"  {'─'*85}")

    for r in results:
        flag = " <- matches backtest" if 3800 <= r["n"] <= 4200 else ""
        print(f"  {r['label']:<25}{r['n']:>7}{r['win_rate']:>7.1f}"
              f"{r['sharpe']:>8.2f}{r['pf']:>7.2f}{r['cagr']:>7.2f}"
              f"{r['max_dd']:>8.2f}{r['p1_months']:>7.1f}"
              f"{r['per_week']:>8.1f}  {r['monthly_pnl']:>10,.0f}{flag}")

    print(f"\n  Reference: Original backtest = 4,019 trades | 3.5/week | Phase 1: 3.4 months")

    # ── Final recommendation ──────────────────────────────────────────────────
    best = max(results, key=lambda r: r["sharpe"])
    match = next((r for r in results if 3800 <= r["n"] <= 4200), None)

    print(f"\n{'='*70}")
    print("CONCLUSION")
    print(f"{'='*70}")

    if match:
        print(f"""
  The original backtest used {match['label']}.
  This produces {match['n']:,} trades matching the validated 4,019.

  However — should the LIVE system also allow concurrent trades?

  ARGUMENT FOR concurrent trades:
    - Matches the validated backtest exactly
    - Higher trade frequency = faster Phase 1
    - Each signal is independent from a macro perspective
    - FTMO allows multiple open positions on same symbol

  ARGUMENT AGAINST concurrent trades:
    - Multiple EURUSD positions in same direction = higher correlation risk
    - If signal is wrong, both trades lose simultaneously
    - Max drawdown increases significantly
    - Harder to manage manually if automation fails

  Best Sharpe: {best['label']} (Sharpe={best['sharpe']:.2f})
  Best frequency match: {match['label']} ({match['n']:,} trades)
""")
    else:
        print(f"""
  Could not find a concurrent limit matching 4,019 trades exactly.
  The gap may also involve different signal logic in the original backtest.
  Best available: {best['label']} (Sharpe={best['sharpe']:.2f}, {best['n']:,} trades)
""")

    pd.DataFrame(results).to_csv(
        TRADES_DIR/"concurrent_trade_results.csv", index=False)
    print(f"  Results saved: concurrent_trade_results.csv")


if __name__ == "__main__":
    main()
