#!/usr/bin/env python3
"""
Crypto Backtest Engine
Strategies: Momentum Breakout, Mean Reversion, EMA Crossover, Volume Spike
Data: BTC/USDT 1h from Binance via CCXT (2024-2025)
"""

import ccxt
import pandas as pd
import numpy as np
import matplotlib.pyplot as plt
import matplotlib.dates as mdates
from datetime import datetime, timedelta
import json
import warnings
warnings.filterwarnings('ignore')

# ─── CONFIG ────────────────────────────────────────────────────────────────
SYMBOL = "BTC/USDT"
TIMEFRAME = "1h"
START_DATE = "2024-01-01"
END_DATE = "2025-03-01"
EXCHANGE = "binance"
INITIAL_CAPITAL = 10_000
RESULTS_DIR = "/home/node/.openclaw/workspace/crypto-wallet/backtest"
# ─────────────────────────────────────────────────────────────────────────────

print("=" * 60)
print("🚀 CRYPTO BACKTEST ENGINE - Gandalf's Trading Lab")
print("=" * 60)

# ─── FETCH DATA ──────────────────────────────────────────────────────────────
def fetch_ohlcv(symbol, timeframe, start, end, exchange_id="binance"):
    """Fetch OHLCV data from exchange"""
    print(f"\n📡 Fetching {symbol} {timeframe} from {exchange_id}...")
    exchange = getattr(ccxt, exchange_id)({"enableRateLimit": True})
    
    # Convert dates to milliseconds
    since = int(pd.Timestamp(start).timestamp() * 1000)
    end_ts = int(pd.Timestamp(end).timestamp() * 1000)
    
    all_data = []
    while since < end_ts:
        try:
            ohlcv = exchange.fetch_ohlcv(symbol, timeframe, since, limit=1000)
            if not ohlcv:
                break
            all_data.extend(ohlcv)
            since = ohlcv[-1][0] + 1
            print(f"  ✓ Got {len(all_data)} candles...", end="\r")
        except Exception as e:
            print(f"\n  ⚠ Error: {e}")
            break
    
    df = pd.DataFrame(all_data, columns=["timestamp", "open", "high", "low", "close", "volume"])
    df["timestamp"] = pd.to_datetime(df["timestamp"], unit="ms")
    df.set_index("timestamp", inplace=True)
    
    # Filter date range
    df = df[(df.index >= start) & (df.index <= end)]
    print(f"\n  ✅ Loaded {len(df)} candles ({start} → {end})")
    return df

# ─── STRATEGIES ──────────────────────────────────────────────────────────────

def strategy_momentum_breakout(df, lookback=24):
    """Momentum Breakout: Price breaks 24h high → LONG, breaks 24h low → SHORT"""
    print("\n📈 Strategy: Momentum Breakout (24h lookback)")
    
    df = df.copy()
    df["high_24h"] = df["high"].rolling(lookback).max().shift(1)
    df["low_24h"] = df["low"].rolling(lookback).min().shift(1)
    
    df["signal"] = 0
    df.loc[df["close"] > df["high_24h"], "signal"] = 1    # LONG
    df.loc[df["close"] < df["low_24h"], "signal"] = -1   # SHORT
    
    return df

def strategy_mean_reversion(df, period=14):
    """Mean Reversion: RSI < 30 → LONG, RSI > 70 → SHORT"""
    print("📈 Strategy: Mean Reversion (RSI)")
    
    df = df.copy()
    delta = df["close"].diff()
    gain = delta.where(delta > 0, 0).rolling(period).mean()
    loss = (-delta.where(delta < 0, 0)).rolling(period).mean()
    rs = gain / loss.replace(0, np.inf)
    df["rsi"] = 100 - (100 / (1 + rs))
    
    df["signal"] = 0
    df.loc[df["rsi"] < 30, "signal"] = 1    # LONG (oversold)
    df.loc[df["rsi"] > 70, "signal"] = -1   # SHORT (overbought)
    
    return df

def strategy_ema_crossover(df, fast=9, slow=21):
    """EMA Crossover: EMA 9 crosses EMA 21 → signal"""
    print("📈 Strategy: EMA Crossover (9/21)")
    
    df = df.copy()
    df["ema_fast"] = df["close"].ewm(span=fast).mean()
    df["ema_slow"] = df["close"].ewm(span=slow).mean()
    
    df["prev_fast"] = df["ema_fast"].shift(1)
    df["prev_slow"] = df["ema_slow"].shift(1)
    
    # Crossover detection
    df["signal"] = 0
    # Golden cross (fast crosses above slow) → LONG
    df.loc[(df["ema_fast"] > df["ema_slow"]) & (df["prev_fast"] <= df["prev_slow"]), "signal"] = 1
    # Death cross (fast crosses below slow) → SHORT
    df.loc[(df["ema_fast"] < df["ema_slow"]) & (df["prev_fast"] >= df["prev_slow"]), "signal"] = -1
    
    return df

def strategy_volume_spike(df, multiplier=2, lookback=20):
    """Volume Spike: Volume > 2x average → entry"""
    print("📈 Strategy: Volume Spike (2x avg)")
    
    df = df.copy()
    df["vol_avg"] = df["volume"].rolling(lookback).mean()
    
    df["signal"] = 0
    df.loc[df["volume"] > (df["vol_avg"] * multiplier), "signal"] = 1
    
    return df

# ─── BACKTESTER ──────────────────────────────────────────────────────────────

def backtest(df, strategy_name, initial_capital=INITIAL_CAPITAL, commission=0.001):
    """Run backtest on a strategy"""
    print(f"\n🔄 Running backtest...")
    
    df = df.copy().dropna()
    
    capital = initial_capital
    position = 0  # 0=none, 1=long, -1=short
    entry_price = 0
    trades = []
    equity_curve = [initial_capital]
    
    for i in range(1, len(df)):
        row = df.iloc[i]
        prev_row = df.iloc[i-1]
        signal = row["signal"]
        price = row["close"]
        
        # Entry logic
        if position == 0 and signal == 1:  # LONG signal
            position = 1
            entry_price = price
            trades.append({"type": "LONG", "entry": price, "time": df.index[i]})
        elif position == 0 and signal == -1:  # SHORT signal
            position = -1
            entry_price = price
            trades.append({"type": "SHORT", "entry": price, "time": df.index[i]})
        
        # Exit logic (simple: exit on opposite signal)
        elif position == 1 and signal == -1:
            pnl = (price - entry_price) / entry_price
            capital *= (1 + pnl - commission)
            trades[-1].update({"exit": price, "pnl_pct": round(pnl*100, 2)})
            position = -1
            entry_price = price
            trades.append({"type": "SHORT", "entry": price, "time": df.index[i]})
        elif position == -1 and signal == 1:
            pnl = (entry_price - price) / entry_price
            capital *= (1 + pnl - commission)
            trades[-1].update({"exit": price, "pnl_pct": round(pnl*100, 2)})
            position = 1
            entry_price = price
            trades.append({"type": "LONG", "entry": price, "time": df.index[i]})
        
        equity_curve.append(capital)
    
    # Close final position
    if position != 0:
        final_price = df.iloc[-1]["close"]
        if position == 1:
            pnl = (final_price - entry_price) / entry_price
        else:
            pnl = (entry_price - final_price) / entry_price
        capital *= (1 + pnl - commission)
        trades[-1].update({"exit": final_price, "pnl_pct": round(pnl*100, 2)})
    
    equity_curve[-1] = capital
    
    # Calculate metrics
    equity_series = pd.Series(equity_curve, index=df.index[:len(equity_curve)])
    returns = equity_series.pct_change().dropna()
    
    total_return = ((capital - initial_capital) / initial_capital) * 100
    winning_trades = [t for t in trades if t.get("pnl_pct", 0) > 0]
    losing_trades = [t for t in trades if t.get("pnl_pct", 0) < 0]
    win_rate = len(winning_trades) / len(trades) * 100 if trades else 0
    
    # Sharpe ratio (simplified)
    sharpe = returns.mean() / returns.std() * np.sqrt(365 * 24) if returns.std() > 0 else 0
    
    # Max drawdown
    rolling_max = equity_series.cummax()
    drawdown = (equity_series - rolling_max) / rolling_max
    max_dd = drawdown.min() * 100
    
    metrics = {
        "strategy": strategy_name,
        "total_return_pct": round(total_return, 2),
        "final_capital": round(capital, 2),
        "num_trades": len(trades),
        "win_rate": round(win_rate, 1),
        "sharpe_ratio": round(sharpe, 2),
        "max_drawdown_pct": round(max_dd, 2),
        "winning_trades": len(winning_trades),
        "losing_trades": len(losing_trades),
    }
    
    print(f"  📊 Total Return: {total_return:.2f}%")
    print(f"  📊 Final Capital: ${capital:.2f}")
    print(f"  📊 Trades: {len(trades)} (Win: {len(winning_trades)}, Lose: {len(losing_trades)})")
    print(f"  📊 Win Rate: {win_rate:.1f}%")
    print(f"  📊 Sharpe: {sharpe:.2f}")
    print(f"  📊 Max Drawdown: {max_dd:.2f}%")
    
    return metrics, equity_curve, trades, df

# ─── PLOT RESULTS ────────────────────────────────────────────────────────────

def plot_results(all_results, df, equity_curves, strategies):
    """Generate comparison charts"""
    fig, axes = plt.subplots(2, 2, figsize=(14, 10))
    fig.suptitle("📊 Crypto Backtest Results 2024-2025\nBTC/USDT | Gandalf's Trading Lab", fontsize=14, fontweight='bold')
    
    colors = ["#2ecc71", "#3498db", "#e74c3c", "#9b59b6"]
    
    # Plot 1: Equity curves
    ax1 = axes[0, 0]
    for i, (name, curve) in enumerate(equity_curves.items()):
        ax1.plot(curve, label=name, color=colors[i % len(colors)], linewidth=1.5)
    ax1.set_title("Equity Curves")
    ax1.set_ylabel("Capital ($)")
    ax1.legend()
    ax1.grid(True, alpha=0.3)
    
    # Plot 2: Bar chart of returns
    ax2 = axes[0, 1]
    names = [r["strategy"] for r in all_results]
    returns = [r["total_return_pct"] for r in all_results]
    bars = ax2.bar(names, returns, color=colors[:len(names)])
    ax2.set_title("Total Return (%)")
    ax2.set_ylabel("Return (%)")
    ax2.axhline(0, color='black', linewidth=0.5)
    for bar, ret in zip(bars, returns):
        ax3_height = bar.get_height()
        ax2.text(bar.get_x() + bar.get_width()/2., ax3_height, f'{ret:.1f}%', ha='center', va='bottom', fontsize=9)
    
    # Plot 3: Metrics table
    ax3 = axes[1, 0]
    ax3.axis('off')
    table_data = [
        [r["strategy"], f"{r['total_return_pct']:.1f}%", f"{r['win_rate']:.1f}%", f"{r['sharpe_ratio']:.2f}", f"{r['max_drawdown_pct']:.1f}%", str(r['num_trades'])]
        for r in all_results
    ]
    table = ax3.table(
        cellText=table_data,
        colLabels=["Strategy", "Return", "Win Rate", "Sharpe", "Max DD", "Trades"],
        loc='center',
        cellLoc='center'
    )
    table.auto_set_font_size(False)
    table.set_fontsize(9)
    table.scale(1.2, 1.5)
    ax3.set_title("Performance Metrics", pad=20)
    
    # Plot 4: BTC price with volume
    ax4 = axes[1, 1]
    ax4_twin = ax4.twinx()
    ax4.plot(df.index, df["close"], color="#f39c12", linewidth=1, label="BTC Price")
    ax4_twin.bar(df.index, df["volume"], alpha=0.3, color="#3498db", width=0.8, label="Volume")
    ax4.set_title("BTC/USDT Price Chart")
    ax4.set_ylabel("Price ($)")
    ax4_twin.set_ylabel("Volume")
    ax4.grid(True, alpha=0.3)
    
    plt.tight_layout()
    plt.savefig(f"{RESULTS_DIR}/backtest_results.png", dpi=150, bbox_inches='tight')
    print(f"\n  💾 Chart saved: {RESULTS_DIR}/backtest_results.png")
    plt.close()

# ─── MAIN ────────────────────────────────────────────────────────────────────

def main():
    # Fetch data
    df = fetch_ohlcv(SYMBOL, TIMEFRAME, START_DATE, END_DATE, EXCHANGE)
    
    if len(df) < 100:
        print("❌ Not enough data fetched!")
        return
    
    # Save raw data info
    df.to_csv(f"{RESULTS_DIR}/btc_ohlcv_2024_2025.csv")
    print(f"  💾 OHLCV data saved: {RESULTS_DIR}/btc_ohlcv_2024_2025.csv")
    
    strategies = {
        "Momentum Breakout": strategy_momentum_breakout,
        "Mean Reversion": strategy_mean_reversion,
        "EMA Crossover": strategy_ema_crossover,
        "Volume Spike": strategy_volume_spike,
    }
    
    all_results = []
    equity_curves = {}
    
    print("\n" + "=" * 60)
    print("⚔️  RUNNING BACKTESTS")
    print("=" * 60)
    
    for name, func in strategies.items():
        df_strat = func(df.copy())
        metrics, equity, trades, df_clean = backtest(df_strat, name)
        all_results.append(metrics)
        equity_curves[name] = equity
        
        # Save trades
        if trades:
            trades_df = pd.DataFrame(trades)
            trades_df.to_csv(f"{RESULTS_DIR}/trades_{name.replace(' ', '_').lower()}.csv", index=False)
    
    # Plot
    print("\n" + "=" * 60)
    print("📊 GENERATING CHARTS")
    print("=" * 60)
    plot_results(all_results, df, equity_curves, strategies)
    
    # Save summary
    summary_df = pd.DataFrame(all_results)
    summary_df.to_csv(f"{RESULTS_DIR}/backtest_summary.csv", index=False)
    
    # Final comparison
    print("\n" + "=" * 60)
    print("🏆 FINAL RANKING - BEST STRATEGIES")
    print("=" * 60)
    sorted_results = sorted(all_results, key=lambda x: x["total_return_pct"], reverse=True)
    for i, r in enumerate(sorted_results, 1):
        print(f"  {i}. {r['strategy']}: {r['total_return_pct']:+.1f}% | Sharpe: {r['sharpe_ratio']} | Win Rate: {r['win_rate']:.0f}%")
    
    print("\n✅ BACKTEST COMPLETE!")
    print(f"   Results saved to: {RESULTS_DIR}/")
    
    return all_results

if __name__ == "__main__":
    main()
