"""
Per-tier edge analysis: is the 2.0× tier worth keeping?

We're testing whether higher-z trades have *higher per-unit-of-risk edge*,
not just higher win rate. WR can be high while expected value per dollar of
risk is low — what matters is Sharpe-per-trade and expected $ per dollar at
risk.

Three perspectives are computed for each tier:
  1. Win rate, profit factor, expected $ per trade
  2. Per-trade Sharpe-equivalent: mean / std of trade returns
  3. Risk-adjusted edge: expected $ / max-loss-magnitude per trade

If the 2.0x tier has LOWER Sharpe-per-trade than the 1.0x tier, the dynamic
sizing logic is amplifying weaker per-unit edge with bigger size. That would
hurt overall Sharpe even if total $ goes up.
"""
import pandas as pd
import numpy as np

eu = pd.read_csv("data/processed/trades/trades_real_costs.csv")
uj = pd.read_csv("data/processed/usdjpy_trades_real_costs.csv")

eu_pnl_col = "dollar_pnl_real"
uj_pnl_col = "pnl_real"


def analyze(df, pnl_col, tiers, pair):
    print(f"\n{'='*72}")
    print(f"  {pair} — Per-tier edge analysis")
    print(f"{'='*72}")
    print(f"{'Tier':<22}{'n':>6}{'WR%':>7}{'Avg $':>10}"
          f"{'Std $':>10}{'Sharpe-T':>10}{'PF':>8}{'EV/Risk':>10}")
    print("-" * 84)

    rows = []
    for label, lo, hi, _mult in tiers:
        sub = df[(df.zscore_abs >= lo) & (df.zscore_abs < hi)]
        if len(sub) == 0:
            continue
        pnl = sub[pnl_col].values
        wr = (pnl > 0).mean() * 100
        mean_pnl = pnl.mean()
        std_pnl = pnl.std(ddof=1)

        # Per-trade Sharpe (NOT annualized — comparable across tiers within pair)
        sharpe_t = mean_pnl / std_pnl if std_pnl > 0 else 0

        # Profit factor
        gross_win = pnl[pnl > 0].sum()
        gross_loss = abs(pnl[pnl < 0].sum())
        pf = gross_win / gross_loss if gross_loss > 0 else float("inf")

        # Expected $ per dollar of "risk" (using avg loss as risk proxy)
        avg_loss = abs(pnl[pnl < 0].mean()) if (pnl < 0).any() else 1.0
        ev_per_risk = mean_pnl / avg_loss

        print(f"  {label:<20}{len(sub):>6}{wr:>6.1f}%{mean_pnl:>10.2f}"
              f"{std_pnl:>10.2f}{sharpe_t:>10.3f}{pf:>8.2f}{ev_per_risk:>10.3f}")
        rows.append({
            "tier": label, "n": len(sub), "wr": wr,
            "mean_pnl": mean_pnl, "std_pnl": std_pnl,
            "sharpe_t": sharpe_t, "pf": pf, "ev_per_risk": ev_per_risk,
        })
    return pd.DataFrame(rows)


eu_tiers = [
    ("1.0x (2.75-3.5)", 2.75, 3.5,  1.0),
    ("1.5x (3.5-4.5)",  3.5,  4.5,  1.5),
    ("2.0x (4.5+)",     4.5,  999,  2.0),
]
uj_tiers = [
    ("1.0x (2.0-2.5)",  2.0,  2.5,  1.0),
    ("1.5x (2.5-3.5)",  2.5,  3.5,  1.5),
    ("2.0x (3.5+)",     3.5,  999,  2.0),
]

eu_results = analyze(eu, eu_pnl_col, eu_tiers, "EURUSD")
uj_results = analyze(uj, uj_pnl_col, uj_tiers, "USDJPY")


# ── Hypothesis test: is the 2.0x tier per-trade Sharpe HIGHER than the 1.0x tier? ──
print(f"\n{'='*72}")
print("  THE QUESTION: Does 2.0x tier have stronger per-trade edge than 1.0x?")
print(f"{'='*72}")

for pair, results in [("EURUSD", eu_results), ("USDJPY", uj_results)]:
    if len(results) < 3:
        print(f"\n  {pair}: insufficient tiers")
        continue
    s1, s2, s3 = results.iloc[0], results.iloc[1], results.iloc[2]
    print(f"\n  {pair}:")
    print(f"    1.0x tier Sharpe-per-trade: {s1.sharpe_t:.3f}")
    print(f"    1.5x tier Sharpe-per-trade: {s2.sharpe_t:.3f}  "
          f"({'BETTER' if s2.sharpe_t > s1.sharpe_t else 'WORSE'} than 1.0x)")
    print(f"    2.0x tier Sharpe-per-trade: {s3.sharpe_t:.3f}  "
          f"({'BETTER' if s3.sharpe_t > s1.sharpe_t else 'WORSE'} than 1.0x)")

    # If 2.0x has LOWER Sharpe-per-trade than 1.0x, scaling that tier up
    # disproportionately is dragging overall portfolio Sharpe.
    if s3.sharpe_t < s1.sharpe_t:
        print(f"    ⚠  Scaling 2.0x tier up by 2x amplifies weaker per-unit edge")
    elif s3.sharpe_t > s1.sharpe_t * 1.2:
        print(f"    ✓  2.0x tier has substantially better edge — scaling justified")
    else:
        print(f"    ~  2.0x tier has similar edge — sizing is neutral")


# ── Counterfactual: what if we capped at 1.5x? ──
print(f"\n{'='*72}")
print("  COUNTERFACTUAL: capping at 1.5x (no 2.0x tier)")
print(f"{'='*72}")


def tier_mult(z, tiers):
    for lo, hi, m in tiers:
        if lo <= z < hi:
            return m
    return 0.0


# Currrent dynamic sizing
eu_mult_current = eu.zscore_abs.apply(lambda z: tier_mult(z, [(2.75,3.5,1.0),(3.5,4.5,1.5),(4.5,999,2.0)]))
uj_mult_current = uj.zscore_abs.apply(lambda z: tier_mult(z, [(2.0,2.5,1.0),(2.5,3.5,1.5),(3.5,999,2.0)]))
# Capped at 1.5x
eu_mult_capped = eu.zscore_abs.apply(lambda z: tier_mult(z, [(2.75,3.5,1.0),(3.5,999,1.5)]))
uj_mult_capped = uj.zscore_abs.apply(lambda z: tier_mult(z, [(2.0,2.5,1.0),(2.5,999,1.5)]))


def portfolio_stats(eu_pnl_scaled, uj_pnl_scaled, label):
    """Combine pair P&L by date and compute portfolio Sharpe + total."""
    eu_t = pd.DataFrame({
        "date": pd.to_datetime(eu.exit_time).dt.normalize(),
        "pnl": eu_pnl_scaled,
    })
    uj_t = pd.DataFrame({
        "date": pd.to_datetime(uj.entry_time).dt.normalize(),
        "pnl": uj_pnl_scaled,
    })
    daily = (pd.concat([eu_t, uj_t])
             .groupby("date")["pnl"].sum()
             .rename("daily_pnl").reset_index())
    daily = daily.sort_values("date").reset_index(drop=True)

    # Reindex to full date range so idle days count as 0 (matches main backtest)
    full_range = pd.date_range(daily.date.min(), daily.date.max(), freq="D")
    daily = daily.set_index("date").reindex(full_range, fill_value=0).reset_index()

    daily_ret = daily["daily_pnl"].values / 100_000  # vs $100k starting balance
    sharpe = (daily_ret.mean() / daily_ret.std() * np.sqrt(252)
              if daily_ret.std() > 0 else 0)

    # Worst calendar year DD (FTMO fresh-account view)
    daily["year"] = daily["index"].dt.year
    worst_yr_dd = 0
    worst_yr = None
    for yr, sub in daily.groupby("year"):
        bal = 100_000 + sub["daily_pnl"].cumsum()
        peak = bal.cummax()
        dd = (bal - peak).min() / 100_000 * 100
        if dd < worst_yr_dd:
            worst_yr_dd, worst_yr = dd, yr

    total_pnl = (eu_pnl_scaled.sum() + uj_pnl_scaled.sum())
    print(f"\n  {label}:")
    print(f"    Total P&L:        ${total_pnl:>13,.0f}")
    print(f"    Sharpe:           {sharpe:.3f}")
    print(f"    Worst year DD:    {worst_yr_dd:.2f}%  ({worst_yr})")
    return sharpe, total_pnl, worst_yr_dd


s_curr, p_curr, w_curr = portfolio_stats(
    eu[eu_pnl_col] * eu_mult_current,
    uj[uj_pnl_col] * uj_mult_current,
    "Current dynamic (1.0/1.5/2.0)"
)
s_cap, p_cap, w_cap = portfolio_stats(
    eu[eu_pnl_col] * eu_mult_capped,
    uj[uj_pnl_col] * uj_mult_capped,
    "Capped at 1.5x (no 2.0x tier)"
)
s_flat, p_flat, w_flat = portfolio_stats(
    eu[eu_pnl_col] * 1.0,
    uj[uj_pnl_col] * 1.0,
    "Flat 1.0x baseline"
)


print(f"\n{'='*72}")
print("  VERDICT")
print(f"{'='*72}")
print(f"  Capped 1.5x vs Current 2.0x: ", end="")
if s_cap > s_curr:
    edge = (s_cap - s_curr) / s_curr * 100
    pnl_lost = p_curr - p_cap
    dd_better = w_cap - w_curr  # less negative is better
    print(f"BETTER Sharpe (+{edge:.1f}%), gives up ${pnl_lost:,.0f} P&L, "
          f"DD {dd_better:+.2f}pp")
    print(f"  → 2.0x tier is dragging risk-adjusted performance")
else:
    edge = (s_curr - s_cap) / s_cap * 100
    pnl_kept = p_curr - p_cap
    print(f"WORSE Sharpe (-{edge:.1f}%), kept ${pnl_kept:,.0f} extra P&L")
    print(f"  → 2.0x tier is genuinely accretive — keep it")
