from __future__ import annotations

from typing import List, Optional
import json
import numpy as np

import pandas as pd


def calculate_metrics(trades: List[dict], equity_data: List[dict] = None) -> dict:
    if not trades:
        return {
            "total_trades": 0,
            "win_rate": 0.0,
            "profit_factor": 0.0,
            "max_drawdown": 0.0,
            "avg_rr": 0.0,
            "total_pnl": 0.0,
            "sharpe_ratio": 0.0,
            "avg_trade_pnl": 0.0,
            "max_consecutive_wins": 0,
            "max_consecutive_losses": 0,
            "avg_win": 0.0,
            "avg_loss": 0.0,
            "expectancy": 0.0,
        }

    df = pd.DataFrame(trades)
    wins = df[df["pnl"] > 0]
    losses = df[df["pnl"] <= 0]

    win_rate = len(wins) / len(df) * 100 if len(df) > 0 else 0

    total_wins = wins["pnl"].sum() if len(wins) > 0 else 0
    total_losses = abs(losses["pnl"].sum()) if len(losses) > 0 else 1
    profit_factor = total_wins / total_losses if total_losses > 0 else 0

    df_sorted = df.sort_values("exit_time")
    equity = df_sorted["pnl"].cumsum().values
    running_max = pd.Series(equity).cummax()
    drawdown = equity - running_max
    max_drawdown = abs(drawdown.min()) if len(drawdown) > 0 else 0

    # Calculate actual RR from entry/exit vs SL distance
    avg_rr = 0.0
    if "entry_price" in df.columns and "exit_price" in df.columns and "sl_price" in df.columns:
        risk_dist = (df["entry_price"] - df["sl_price"]).abs()
        reward_dist = (df["exit_price"] - df["entry_price"]).abs()
        # For SELL trades, reward is negative of (exit - entry)
        if "direction" in df.columns:
            sell_mask = df["direction"].astype(str).str.contains("SELL", case=False, na=False)
            reward_dist[sell_mask] = (df["entry_price"] - df["exit_price"])[sell_mask].abs()
        valid = risk_dist > 0
        if valid.any():
            rr_series = reward_dist[valid] / risk_dist[valid]
            avg_rr = rr_series.mean()
    elif "rr" in df.columns and len(df) > 0:
        avg_rr = df["rr"].mean()

    total_pnl = df["pnl"].sum()

    # Sharpe Ratio (annualized, assuming ~252 trading days)
    daily_returns = df_sorted["pnl"].values
    if len(daily_returns) > 1 and np.std(daily_returns) > 0:
        sharpe_ratio = np.mean(daily_returns) / np.std(daily_returns) * np.sqrt(252)
    else:
        sharpe_ratio = 0.0

    # Average trade PnL
    avg_trade_pnl = df["pnl"].mean()

    # Average win / loss
    avg_win = wins["pnl"].mean() if len(wins) > 0 else 0.0
    avg_loss = losses["pnl"].mean() if len(losses) > 0 else 0.0

    # Expectancy = WinRate * AvgWin - LossRate * AvgLoss
    wr = win_rate / 100.0
    expectancy = wr * avg_win - (1 - wr) * abs(avg_loss)

    # Max consecutive wins/losses
    is_win = df_sorted["pnl"] > 0
    max_consec_wins = 0
    max_consec_losses = 0
    current_wins = 0
    current_losses = 0
    for w in is_win:
        if w:
            current_wins += 1
            current_losses = 0
            max_consec_wins = max(max_consec_wins, current_wins)
        else:
            current_losses += 1
            current_wins = 0
            max_consec_losses = max(max_consec_losses, current_losses)

    return {
        "total_trades": len(df),
        "winning_trades": len(wins),
        "losing_trades": len(losses),
        "win_rate": round(win_rate, 2),
        "profit_factor": round(profit_factor, 2),
        "max_drawdown": round(max_drawdown, 2),
        "avg_rr": round(avg_rr, 2),
        "total_pnl": round(total_pnl, 2),
        "sharpe_ratio": round(sharpe_ratio, 2),
        "avg_trade_pnl": round(avg_trade_pnl, 2),
        "avg_win": round(avg_win, 2),
        "avg_loss": round(avg_loss, 2),
        "expectancy": round(expectancy, 2),
        "max_consecutive_wins": max_consec_wins,
        "max_consecutive_losses": max_consec_losses,
    }


def print_report(metrics: dict, trades: List[dict] = None, equity_data: List[dict] = None):
    print("\n" + "=" * 40)
    print("BACKTEST RESULTS")
    print("=" * 40)
    print(f"Total Trades:     {metrics['total_trades']}")
    print(f"Winning:          {metrics['winning_trades']}")
    print(f"Losing:           {metrics['losing_trades']}")
    print(f"Win Rate:         {metrics['win_rate']:.2f}%")
    print(f"Profit Factor:    {metrics['profit_factor']:.2f}")
    print(f"Max Drawdown:     {metrics['max_drawdown']:.2f}")
    print(f"Avg RR:           {metrics['avg_rr']:.2f}")
    print(f"Sharpe Ratio:     {metrics.get('sharpe_ratio', 0.0):.2f}")
    print(f"Avg Trade PnL:    {metrics.get('avg_trade_pnl', 0.0):.2f}")
    print(f"Avg Win:          {metrics.get('avg_win', 0.0):.2f}")
    print(f"Avg Loss:         {metrics.get('avg_loss', 0.0):.2f}")
    print(f"Expectancy:       {metrics.get('expectancy', 0.0):.2f}")
    print(f"Max Consec Wins:  {metrics.get('max_consecutive_wins', 0)}")
    print(f"Max Consec Loss:  {metrics.get('max_consecutive_losses', 0)}")
    print(f"Total PnL:        {metrics['total_pnl']:,.2f}")
    print("=" * 40 + "\n")
    
    if trades:
        print("\nDETAILED TRADE LOG (first 10 and last 10):")
        print("-" * 100)
        df_trades = pd.DataFrame(trades)
        for col in ["entry_time", "exit_time"]:
            if df_trades[col].dtype in ("object", "int64", "float64"):
                sample_val = df_trades[col].dropna().iloc[0] if len(df_trades[col].dropna()) > 0 else None
                import numpy as np
                if sample_val is not None and isinstance(sample_val, (int, float, np.integer, np.floating)):
                    df_trades[col] = pd.to_datetime(df_trades[col], unit="s")
        
        display_trades = pd.concat([df_trades.head(10), df_trades.tail(10)])
        for idx, row in display_trades.iterrows():
            status = "WIN" if row["pnl"] > 0 else "LOSS"
            print(f"{row['entry_time'].strftime('%Y-%m-%d %H:%M')} | {row['direction']:4} | "
                  f"Entry: {row['entry_price']:8.2f} | Exit: {row['exit_price']:8.2f} | "
                  f"PnL: {row['pnl']:>10.2f} | {status}")
        print("-" * 100 + "\n")


def plot_equity_curve(equity_data: List[dict], output_path: str = "equity_curve.png"):
    try:
        import matplotlib.pyplot as plt
        import matplotlib.dates as mdates
    except ImportError:
        print("Matplotlib not installed. Install with: pip install matplotlib")
        return
    
    if not equity_data:
        return
    
    df = pd.DataFrame(equity_data)
    df["timestamp"] = pd.to_datetime(df["timestamp"], unit="s")
    df = df.sort_values("timestamp")
    
    fig, ax = plt.subplots(figsize=(14, 7))
    
    ax.plot(df["timestamp"], df["equity"], linewidth=1.5, color="#2E86AB")
    ax.fill_between(df["timestamp"], df["equity"], alpha=0.3, color="#2E86AB")
    
    ax.set_title("Equity Curve", fontsize=14, fontweight="bold")
    ax.set_xlabel("Date", fontsize=12)
    ax.set_ylabel("Equity (RUB)", fontsize=12)
    
    ax.xaxis.set_major_formatter(mdates.DateFormatter("%Y-%m-%d"))
    ax.xaxis.set_major_locator(mdates.MonthLocator())
    plt.xticks(rotation=45)
    
    ax.grid(True, alpha=0.3)
    
    start_equity = df["equity"].iloc[0]
    end_equity = df["equity"].iloc[-1]
    change = end_equity - start_equity
    change_pct = (change / start_equity * 100) if start_equity > 0 else 0
    
    annotation = (
        f"Start: {start_equity:,.0f}\n"
        f"End: {end_equity:,.0f}\n"
        f"Change: {change:+,.0f} ({change_pct:+.1f}%)"
    )
    ax.annotate(
        annotation,
        xy=(0.02, 0.98),
        xycoords="axes fraction",
        fontsize=10,
        verticalalignment="top",
        bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5),
    )
    
    plt.tight_layout()
    plt.savefig(output_path, dpi=150, bbox_inches="tight")
    plt.close()
    print(f"Equity curve saved to {output_path}")


def save_full_report(
    metrics: dict,
    trades: List[dict],
    equity_data: List[dict],
    ticker: str,
    output_dir: str = "reports",
):
    import os
    from datetime import datetime
    
    os.makedirs(output_dir, exist_ok=True)
    timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
    
    trades_df = pd.DataFrame(trades)
    for col in ["entry_time", "exit_time"]:
        if trades_df[col].dtype in ("object", "int64", "float64"):
            non_null = trades_df[col].dropna()
            if len(non_null) > 0:
                import numpy as np
                if isinstance(non_null.iloc[0], (int, float, np.integer, np.floating)):
                    trades_df[col] = pd.to_datetime(trades_df[col], unit="s")
    trades_path = f"{output_dir}/{ticker}_trades_{timestamp}.csv"
    trades_df.to_csv(trades_path, index=False)
    
    equity_df = pd.DataFrame(equity_data)
    equity_df["timestamp"] = pd.to_datetime(equity_df["timestamp"], unit="s")
    equity_path = f"{output_dir}/{ticker}_equity_{timestamp}.csv"
    equity_df.to_csv(equity_path, index=False)
    
    metrics_path = f"{output_dir}/{ticker}_metrics_{timestamp}.json"
    with open(metrics_path, "w") as f:
        json.dump(metrics, f, indent=2)
    
    plot_equity_curve(equity_data, f"{output_dir}/{ticker}_equity_{timestamp}.png")
    
    print(f"\nFull report saved to {output_dir}/")
    print(f"  - Trades: {trades_path}")
    print(f"  - Equity: {equity_path}")
    print(f"  - Metrics: {metrics_path}")