"""
annual_pnl_analysis_v1.py
==========================
Year-by-year P&L analysis of the validated model running from a $100k account.

Uses the real-world costs trade log (trades_real_costs.csv) generated by
real_world_costs_v1.py. If not found, regenerates from the backtest.

Shows:
  - Annual dollar P&L at 0.75% risk ($300k notional)
  - Cumulative account balance growth
  - Annual win rate, trade count, max DD per year
  - Monthly returns heatmap
  - Full equity curve with FTMO challenge lines marked

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

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

import pandas as pd
import numpy as np
import matplotlib
matplotlib.use("Agg")
import matplotlib as mpl
mpl.rcParams["text.usetex"] = False
import matplotlib.pyplot as plt
import matplotlib.gridspec as gridspec
import matplotlib.ticker as mticker
from matplotlib.patches import FancyBboxPatch
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"
CHARTS_DIR = BASE_PATH / "data" / "processed"

# ── Model parameters ──────────────────────────────────────────────────────────
ACCOUNT_START  = 100_000.0
STOP_PCT       = 0.0025
BASE_RISK_PCT  = 0.0075
BASE_NOTIONAL  = (ACCOUNT_START * BASE_RISK_PCT) / STOP_PCT   # $300,000

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

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


# ── Load trade log ────────────────────────────────────────────────────────────
def load_trades() -> pd.DataFrame:
    # Prefer real-costs log; fall back to final validation log
    for fname in ["trades_real_costs.csv", "trades_eurusd_final.csv",
                  "trades_growth_52h.csv"]:
        path = TRADES_DIR / fname
        if path.exists():
            df = pd.read_csv(path, parse_dates=["entry_time", "exit_time"])
            print(f"  Loaded: {fname}  ({len(df):,} trades)")

            # Determine which return column to use
            if "return_real" in df.columns:
                df["ret"] = df["return_real"]
                print(f"  Using real-world cost returns")
            elif "return" in df.columns:
                df["ret"] = df["return"]
                print(f"  Using backtest returns")
            else:
                raise ValueError(f"No return column found in {fname}")

            if "zscore_abs" not in df.columns:
                df["zscore_abs"] = 2.75

            return df.sort_values("exit_time").reset_index(drop=True)

    raise FileNotFoundError(
        "No trade log found. Run real_world_costs_v1.py or "
        "final_model_validation_v1.py first."
    )


# ── Build dollar P&L series ───────────────────────────────────────────────────
def build_dollar_pnl(df: pd.DataFrame) -> pd.DataFrame:
    """Converts returns to dollar P&L using signal-scaled notional."""
    df = df.copy()
    df["multiplier"]   = df["zscore_abs"].apply(get_multiplier)
    df["notional"]     = (BASE_NOTIONAL * df["multiplier"]).clip(
                            upper=BASE_NOTIONAL * 3)
    df["dollar_pnl"]   = df["ret"] * df["notional"]
    df["cum_pnl"]      = df["dollar_pnl"].cumsum()
    df["balance"]      = ACCOUNT_START + df["cum_pnl"]
    df["win"]          = (df["ret"] > 0).astype(int)
    df["year"]         = df["exit_time"].dt.year
    df["month"]        = df["exit_time"].dt.to_period("M")
    return df


# ── Annual summary ────────────────────────────────────────────────────────────
def annual_summary(df: pd.DataFrame) -> pd.DataFrame:
    rows = []
    balance = ACCOUNT_START

    for year, grp in df.groupby("year"):
        grp      = grp.sort_values("exit_time")
        pnl      = grp["dollar_pnl"].sum()
        n        = len(grp)
        wr       = grp["win"].mean() * 100
        avg_ret  = grp["ret"].mean() * 100
        end_bal  = balance + pnl
        pct_ret  = pnl / balance * 100

        # Max DD within the year (on dollar basis from year start)
        grp["cum_from_year_start"] = grp["dollar_pnl"].cumsum()
        grp["peak_year"]           = grp["cum_from_year_start"].cummax()
        grp["dd_year"]             = (grp["cum_from_year_start"] -
                                      grp["peak_year"]) / balance * 100
        max_dd = grp["dd_year"].min()

        er = grp["exit_reason"].value_counts(normalize=True) * 100 \
             if "exit_reason" in grp.columns else {}
        tp_pct   = er.get("tp",               0) if hasattr(er, "get") else 0
        stop_pct = er.get("stop",             0) if hasattr(er, "get") else 0
        z_pct    = er.get("zscore_reversal",  0) if hasattr(er, "get") else 0

        rows.append({
            "year"      : year,
            "trades"    : n,
            "pnl"       : pnl,
            "pct_return": pct_ret,
            "win_rate"  : wr,
            "avg_ret_pct": avg_ret,
            "start_bal" : balance,
            "end_bal"   : end_bal,
            "max_dd"    : max_dd,
            "tp_pct"    : tp_pct,
            "stop_pct"  : stop_pct,
            "z_pct"     : z_pct,
        })
        balance = end_bal

    return pd.DataFrame(rows)


# ── Monthly heatmap data ──────────────────────────────────────────────────────
def monthly_returns_matrix(df: pd.DataFrame) -> pd.DataFrame:
    monthly = df.groupby("month")["dollar_pnl"].sum().reset_index()
    monthly["year"]  = monthly["month"].dt.year
    monthly["month_n"] = monthly["month"].dt.month
    pivot = monthly.pivot(index="year", columns="month_n", values="dollar_pnl")
    pivot.columns = ["Jan","Feb","Mar","Apr","May","Jun",
                     "Jul","Aug","Sep","Oct","Nov","Dec"]
    return pivot


# ── Chart ─────────────────────────────────────────────────────────────────────
def make_chart(df: pd.DataFrame, annual: pd.DataFrame,
               monthly_matrix: pd.DataFrame, save_path: Path):

    # Colour scheme
    BG      = "#0d1117"
    PANEL   = "#161b22"
    GRID    = "#21262d"
    TEXT    = "#e6edf3"
    MUTED   = "#8b949e"
    CYAN    = "#00d4ff"
    GREEN   = "#00ff88"
    RED     = "#ff4444"
    GOLD    = "#ffd700"
    ORANGE  = "#ff8c00"

    fig = plt.figure(figsize=(22, 26), facecolor=BG)
    gs  = gridspec.GridSpec(4, 2, figure=fig,
                            hspace=0.45, wspace=0.28,
                            left=0.06, right=0.97,
                            top=0.93, bottom=0.04)

    def style(ax, title):
        ax.set_facecolor(PANEL)
        ax.set_title(title, color=TEXT, fontsize=11,
                     fontweight="bold", pad=10)
        ax.tick_params(colors=TEXT, labelsize=9)
        for spine in ax.spines.values():
            spine.set_color(GRID)
        ax.grid(True, color=GRID, linewidth=0.5, alpha=0.7)
        ax.yaxis.label.set_color(TEXT)
        ax.xaxis.label.set_color(TEXT)

    # ── Panel 1: Cumulative equity curve (full width) ──────────────────────
    ax1 = fig.add_subplot(gs[0, :])
    style(ax1, "Cumulative Account Balance — 100k Start  |  0.75% Risk  |  $300k Notional  |  Real-World Costs Applied")

    ax1.fill_between(df["exit_time"], df["balance"] / 1000,
                     ACCOUNT_START / 1000,
                     where=df["balance"] >= ACCOUNT_START,
                     alpha=0.15, color=GREEN)
    ax1.fill_between(df["exit_time"], df["balance"] / 1000,
                     ACCOUNT_START / 1000,
                     where=df["balance"] < ACCOUNT_START,
                     alpha=0.15, color=RED)
    ax1.plot(df["exit_time"], df["balance"] / 1000,
             color=CYAN, linewidth=1.4, zorder=5)

    ax1.axhline(y=ACCOUNT_START / 1000, color=MUTED,
                linewidth=0.8, linestyle="--", alpha=0.6)

    # Mark final balance
    final_bal = df["balance"].iloc[-1]
    ax1.annotate(f"{final_bal/1000:.0f}k",
                 xy=(df["exit_time"].iloc[-1], final_bal / 1000),
                 xytext=(-60, 12), textcoords="offset points",
                 color=CYAN, fontsize=10, fontweight="bold",
                 arrowprops=dict(arrowstyle="->", color=CYAN, lw=1.2))

    ax1.yaxis.set_major_formatter(
        mticker.FuncFormatter(lambda x, _: f"{x:.0f}k"))
    ax1.set_ylabel("Account Balance", fontsize=10)

    # Annotate key years
    for _, row in annual.iterrows():
        col = GREEN if row["pct_return"] >= 0 else RED
        ax1.annotate(f"{row['pct_return']:+.0f}%",
                     xy=(pd.Timestamp(f"{int(row['year'])}-07-01"),
                         row["end_bal"] / 1000 + 8),
                     ha="center", fontsize=7.5, color=col, fontweight="bold")

    # ── Panel 2: Annual bar chart ──────────────────────────────────────────
    ax2 = fig.add_subplot(gs[1, 0])
    style(ax2, "Annual P&L ($)")

    colors = [GREEN if v >= 0 else RED for v in annual["pnl"]]
    bars   = ax2.bar(annual["year"], annual["pnl"] / 1000,
                     color=colors, alpha=0.85, width=0.7, zorder=3)
    ax2.axhline(y=0, color=MUTED, linewidth=0.8)

    for bar, (_, row) in zip(bars, annual.iterrows()):
        ypos = bar.get_height()
        ax2.text(bar.get_x() + bar.get_width() / 2,
                 ypos + (0.5 if ypos >= 0 else -1.5),
                 f"{row['pnl']/1000:.0f}k",
                 ha="center", va="bottom" if ypos >= 0 else "top",
                 fontsize=7.5, color=TEXT, fontweight="bold")

    ax2.yaxis.set_major_formatter(
        mticker.FuncFormatter(lambda x, _: f"{x:.0f}k"))
    ax2.set_ylabel("P&L ($000s)", fontsize=9)
    ax2.set_xticks(annual["year"])
    ax2.set_xticklabels(annual["year"], rotation=45, fontsize=8)

    # ── Panel 3: Annual % return bar ──────────────────────────────────────
    ax3 = fig.add_subplot(gs[1, 1])
    style(ax3, "Annual Return % on Starting Balance")

    colors3 = [GREEN if v >= 0 else RED for v in annual["pct_return"]]
    ax3.bar(annual["year"], annual["pct_return"],
            color=colors3, alpha=0.85, width=0.7, zorder=3)
    ax3.axhline(y=0, color=MUTED, linewidth=0.8)
    ax3.yaxis.set_major_formatter(
        mticker.FuncFormatter(lambda x, _: f"{x:.0f}%"))
    ax3.set_ylabel("Annual Return %", fontsize=9)
    ax3.set_xticks(annual["year"])
    ax3.set_xticklabels(annual["year"], rotation=45, fontsize=8)

    # Add avg line
    avg_pct = annual["pct_return"].mean()
    ax3.axhline(y=avg_pct, color=GOLD, linewidth=1.2,
                linestyle="--", alpha=0.8,
                label=f"Avg: {avg_pct:.1f}%/yr")
    ax3.legend(facecolor=PANEL, edgecolor=GRID,
               labelcolor=TEXT, fontsize=8, loc="upper left")

    # ── Panel 4: Monthly heatmap ──────────────────────────────────────────
    ax4 = fig.add_subplot(gs[2, :])
    style(ax4, "Monthly P&L Heatmap ($)  —  Green = Profit  /  Red = Loss")

    mat   = monthly_matrix.values
    vmax  = np.nanpercentile(np.abs(mat[~np.isnan(mat)]), 95)
    vmin  = -vmax

    im = ax4.imshow(mat, aspect="auto", cmap="RdYlGn",
                    vmin=vmin, vmax=vmax, alpha=0.85)

    ax4.set_xticks(range(12))
    ax4.set_xticklabels(["Jan","Feb","Mar","Apr","May","Jun",
                          "Jul","Aug","Sep","Oct","Nov","Dec"],
                         color=TEXT, fontsize=9)
    ax4.set_yticks(range(len(monthly_matrix.index)))
    ax4.set_yticklabels(monthly_matrix.index, color=TEXT, fontsize=8)

    # Annotate each cell
    for i in range(mat.shape[0]):
        for j in range(mat.shape[1]):
            v = mat[i, j]
            if not np.isnan(v):
                txt_color = "black" if abs(v) > vmax * 0.5 else TEXT
                ax4.text(j, i, f"{v/1000:.0f}k",
                         ha="center", va="center",
                         fontsize=6.5, color=txt_color, fontweight="bold")

    plt.colorbar(im, ax=ax4, orientation="vertical",
                 fraction=0.015, pad=0.01,
                 format=mticker.FuncFormatter(lambda x, _: f"{x/1000:.0f}k"))

    # ── Panel 5: Win rate per year ────────────────────────────────────────
    ax5 = fig.add_subplot(gs[3, 0])
    style(ax5, "Win Rate & Trade Count per Year")

    ax5b = ax5.twinx()
    ax5b.set_facecolor(PANEL)

    ax5.plot(annual["year"], annual["win_rate"],
             color=GREEN, linewidth=2, marker="o",
             markersize=5, zorder=5, label="Win Rate %")
    ax5.axhline(y=annual["win_rate"].mean(), color=GREEN,
                linewidth=0.8, linestyle="--", alpha=0.5)
    ax5.set_ylabel("Win Rate %", fontsize=9, color=GREEN)
    ax5.tick_params(axis="y", colors=GREEN)
    ax5.set_ylim(0, 100)
    ax5.set_xticks(annual["year"])
    ax5.set_xticklabels(annual["year"], rotation=45, fontsize=8)

    ax5b.bar(annual["year"], annual["trades"],
             color=CYAN, alpha=0.25, width=0.7, zorder=2)
    ax5b.set_ylabel("Trades", fontsize=9, color=CYAN)
    ax5b.tick_params(axis="y", colors=CYAN)
    for spine in ax5b.spines.values():
        spine.set_color(GRID)

    ax5.legend(facecolor=PANEL, edgecolor=GRID,
               labelcolor=TEXT, fontsize=8, loc="lower right")

    # ── Panel 6: Annual max DD ────────────────────────────────────────────
    ax6 = fig.add_subplot(gs[3, 1])
    style(ax6, "Annual Max Drawdown % (on year-start balance)")

    colors6 = []
    for v in annual["max_dd"]:
        if v > -2:
            colors6.append(GREEN)
        elif v > -5:
            colors6.append(GOLD)
        else:
            colors6.append(RED)

    ax6.bar(annual["year"], annual["max_dd"],
            color=colors6, alpha=0.85, width=0.7, zorder=3)
    ax6.axhline(y=-5,  color=ORANGE, linewidth=1, linestyle="--",
                alpha=0.7, label="Daily DD limit (5%)")
    ax6.axhline(y=-10, color=RED,    linewidth=1, linestyle="--",
                alpha=0.7, label="Overall DD limit (10%)")
    ax6.set_ylabel("Max Drawdown %", fontsize=9)
    ax6.set_xticks(annual["year"])
    ax6.set_xticklabels(annual["year"], rotation=45, fontsize=8)
    ax6.yaxis.set_major_formatter(
        mticker.FuncFormatter(lambda x, _: f"{x:.1f}%"))
    ax6.legend(facecolor=PANEL, edgecolor=GRID,
               labelcolor=TEXT, fontsize=8)

    # ── Main title ────────────────────────────────────────────────────────
    total_pnl   = df["balance"].iloc[-1] - ACCOUNT_START
    total_ret   = total_pnl / ACCOUNT_START * 100
    cagr        = (df["balance"].iloc[-1] / ACCOUNT_START) ** (1/22) - 1

    fig.suptitle(
        f"EURUSD Macro Lead-Lag Model  —  Year-by-Year Performance\n"
        f"100k Starting Capital  |  22 Years (2003-2026)  |  "
        f"Total PnL: {total_pnl/1000:.0f}k  |  "
        f"Total Return: {total_ret:.0f}%  |  "
        f"CAGR: {cagr*100:.1f}%/yr  |  "
        f"Real-World Costs Applied",
        color=TEXT, fontsize=13, fontweight="bold", y=0.965
    )

    plt.savefig(save_path, dpi=150, bbox_inches="tight",
                facecolor=fig.get_facecolor())
    plt.close()
    print(f"\n  Chart saved: {save_path}")


# ── Print annual table ────────────────────────────────────────────────────────
def print_annual_table(annual: pd.DataFrame):
    total_pnl = annual["pnl"].sum()
    avg_wr    = annual["win_rate"].mean()
    best_yr   = annual.loc[annual["pnl"].idxmax()]
    worst_yr  = annual.loc[annual["pnl"].idxmin()]

    print(f"\n{'='*90}")
    print("YEAR-BY-YEAR P&L SUMMARY")
    print(f"  Starting capital: ${ACCOUNT_START:,.0f}  |  "
          f"Notional: ${BASE_NOTIONAL:,.0f} base  |  Real-world costs applied")
    print(f"{'='*90}")
    print(f"\n  {'Year':<6}{'Trades':>7}{'P&L ($)':>12}{'Return%':>9}"
          f"{'WR%':>7}{'Bal (end)':>12}{'MaxDD%':>8}"
          f"{'TP%':>7}{'Stop%':>7}{'Z%':>6}")
    print(f"  {'─'*86}")

    for _, r in annual.iterrows():
        flag = " ◄ BEST" if r["pnl"] == annual["pnl"].max() else \
               " ◄ WORST" if r["pnl"] == annual["pnl"].min() else ""
        print(f"  {int(r['year']):<6}{int(r['trades']):>7}"
              f"  {r['pnl']:>+10,.0f}{r['pct_return']:>8.1f}%"
              f"{r['win_rate']:>7.1f}{r['end_bal']:>12,.0f}"
              f"{r['max_dd']:>8.2f}"
              f"{r['tp_pct']:>7.1f}{r['stop_pct']:>7.1f}"
              f"{r['z_pct']:>6.1f}{flag}")

    print(f"\n  {'─'*86}")
    print(f"  {'TOTAL':<6}{int(annual['trades'].sum()):>7}"
          f"  {total_pnl:>+10,.0f}"
          f"{total_pnl/ACCOUNT_START*100:>8.1f}%"
          f"{avg_wr:>7.1f}")
    print(f"\n  Positive years   : "
          f"{(annual['pnl'] > 0).sum()}/{len(annual)} "
          f"({(annual['pnl'] > 0).mean()*100:.0f}%)")
    print(f"  Best year        : {int(best_yr['year'])}  "
          f"${best_yr['pnl']:>+,.0f}  ({best_yr['pct_return']:+.1f}%)")
    print(f"  Worst year       : {int(worst_yr['year'])}  "
          f"${worst_yr['pnl']:>+,.0f}  ({worst_yr['pct_return']:+.1f}%)")
    print(f"  Avg annual return: {annual['pct_return'].mean():.1f}%")
    print(f"  Avg annual P&L   : ${annual['pnl'].mean():>+,.0f}")
    cagr = (annual["end_bal"].iloc[-1] / ACCOUNT_START) ** (1/len(annual)) - 1
    print(f"  CAGR             : {cagr*100:.2f}%")
    print(f"  Final balance    : ${annual['end_bal'].iloc[-1]:>,.0f}")
    print(f"  Total return     : "
          f"{(annual['end_bal'].iloc[-1] - ACCOUNT_START) / ACCOUNT_START * 100:.0f}%")


# ── Main ──────────────────────────────────────────────────────────────────────
def main():
    print("=" * 70)
    print("ANNUAL P&L ANALYSIS")
    print(f"  Account start: ${ACCOUNT_START:,.0f}")
    print(f"  Base notional: ${BASE_NOTIONAL:,.0f}  (0.75% risk)")
    print(f"  Signal scaling: 1x / 1.5x / 2x")
    print("=" * 70)

    print("\nLoading trade log...")
    df = load_trades()

    print("Building dollar P&L series...")
    df = build_dollar_pnl(df)

    print("Computing annual summaries...")
    annual = annual_summary(df)

    print("Computing monthly matrix...")
    monthly_matrix = monthly_returns_matrix(df)

    print_annual_table(annual)

    print("\nGenerating chart...")
    chart_path = CHARTS_DIR / "annual_pnl_chart.png"
    make_chart(df, annual, monthly_matrix, chart_path)

    # Save annual table to CSV
    annual.to_csv(TRADES_DIR / "annual_pnl_summary.csv", index=False)
    print(f"  Annual table saved: {TRADES_DIR / 'annual_pnl_summary.csv'}")

    print(f"\n{'='*70}")
    print("DONE — open the chart at:")
    print(f"  {chart_path}")


if __name__ == "__main__":
    main()
