#!/usr/bin/env python3
"""
Crypto Strategy Optimizer v2 - Vectorized for Speed
Goal: Find Sharpe > 2, DD < 15%, Return > 50%
"""
import pandas as pd
import numpy as np
import warnings
warnings.filterwarnings('ignore')

RESULTS_DIR = "/home/node/.openclaw/workspace/crypto-wallet/backtest"

print("Loading data...")
df = pd.read_csv(f"{RESULTS_DIR}/btc_ohlcv_2024_2025.csv", parse_dates=["timestamp"], index_col="timestamp")
print(f"  BTC: {len(df)} candles, {df.index[0]} → {df.index[-1]}")

# ─── Pre-compute all indicators once ─────────────────────────────────────────
print("Pre-computing indicators (vectorized)...")

# ATR
high = df["high"]
low = df["low"]
close = df["close"]
tr1 = high - low
tr2 = (high - close.shift(1)).abs()
tr3 = (low - close.shift(1)).abs()
df["atr"] = pd.concat([tr1, tr2, tr3], axis=1).max(axis=1).rolling(14).mean()

# ADX
period = 14
plus_dm = high.diff()
minus_dm = -low.diff()
plus_dm[plus_dm < 0] = 0
minus_dm[minus_dm < 0] = 0
atr = df["atr"]
plus_di = 100 * (plus_dm.rolling(period).mean() / atr)
minus_di = 100 * (minus_dm.rolling(period).mean() / atr)
dx = 100 * (plus_di - minus_di).abs() / (plus_di + minus_di)
df["adx"] = dx.rolling(period).mean()

# EMAs
for span in [5, 8, 9, 10, 12, 15, 20, 21, 25, 30, 50, 100, 150, 200]:
    df[f"ema_{span}"] = close.ewm(span=span).mean()

# RSI variants
for period in [7, 10, 14, 21, 28]:
    delta = 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[f"rsi_{period}"] = 100 - (100 / (1 + rs))

# Volume
for lb in [10, 15, 20, 30]:
    df[f"vol_avg_{lb}"] = df["volume"].rolling(lb).mean()

# 24h high/low
for lb in [12, 18, 24, 30, 36, 48, 60]:
    df[f"high_{lb}h"] = high.rolling(lb).max().shift(1)
    df[f"low_{lb}h"] = low.rolling(lb).min().shift(1)

# MACD
for fast, slow, sig in [(10, 26, 9), (12, 26, 9), (8, 26, 9), (12, 20, 9)]:
    ema_f = close.ewm(span=fast).mean()
    ema_s = close.ewm(span=slow).mean()
    macd = ema_f - ema_s
    macd_sig = macd.ewm(span=sig).mean()
    df[f"macd_{fast}_{slow}_{sig}_hist"] = macd - macd_sig
    df[f"macd_{fast}_{slow}_{sig}_prev"] = (macd - macd_sig).shift(1)

df = df.dropna()
print(f"  After dropna: {len(df)} rows")

# ─── Vectorized backtester ─────────────────────────────────────────────────────
def vectorized_backtest(df, signals_col, initial_capital=10000, commission=0.001,
                        stop_mult=2.0, tp_mult=3.0):
    """Fast vectorized backtester with SL/TP"""
    signals = df[signals_col].values
    prices = df["close"].values
    atr = df["atr"].values
    n = len(signals)
    
    capital = initial_capital
    position = 0
    entry_price = 0
    stop_loss = 0
    take_profit = 0
    equity = [initial_capital]
    num_trades = 0
    wins = 0
    losses = 0
    
    for i in range(1, n):
        sig = signals[i]
        price = prices[i]
        
        # Entry
        if position == 0 and sig == 1:
            position = 1
            entry_price = price
            stop_loss = price - atr[i] * stop_mult
            take_profit = price + atr[i] * tp_mult
            num_trades += 1
        elif position == 0 and sig == -1:
            position = -1
            entry_price = price
            stop_loss = price + atr[i] * stop_mult
            take_profit = price - atr[i] * tp_mult
            num_trades += 1
        
        # Exit check
        elif position != 0:
            triggered = False
            exit_price = price
            
            if position == 1:
                if price <= stop_loss:
                    exit_price = stop_loss
                    triggered = True
                elif price >= take_profit:
                    exit_price = take_profit
                    triggered = True
            else:
                if price >= stop_loss:
                    exit_price = stop_loss
                    triggered = True
                elif price <= take_profit:
                    exit_price = take_profit
                    triggered = True
            
            if triggered:
                if position == 1:
                    pnl = (exit_price - entry_price) / entry_price
                else:
                    pnl = (entry_price - exit_price) / entry_price
                capital *= (1 + pnl - commission)
                if pnl > 0:
                    wins += 1
                else:
                    losses += 1
                position = 0
        
        equity.append(capital)
    
    # Final close
    if position != 0:
        final_price = prices[-1]
        if position == 1:
            pnl = (final_price - entry_price) / entry_price
        else:
            pnl = (entry_price - final_price) / entry_price
        capital *= (1 + pnl - commission)
    
    equity[-1] = capital
    
    # Metrics
    total_return = ((capital - initial_capital) / initial_capital) * 100
    equity_series = pd.Series(equity)
    returns = equity_series.pct_change().dropna()
    sharpe = returns.mean() / returns.std() * np.sqrt(365 * 24) if returns.std() > 0 else 0
    if np.isnan(sharpe) or np.isinf(sharpe):
        sharpe = 0
    
    rolling_max = equity_series.cummax()
    drawdown = (equity_series - rolling_max) / rolling_max
    max_dd = drawdown.min() * 100
    
    win_rate = wins / (wins + losses) * 100 if (wins + losses) > 0 else 0
    
    return {
        "return_pct": round(total_return, 2),
        "final_capital": round(capital, 2),
        "num_trades": num_trades,
        "win_rate": round(win_rate, 1),
        "sharpe": round(sharpe, 3),
        "max_dd": round(max_dd, 2),
        "winners": wins,
        "losers": losses,
    }

# ─── Build signal columns ───────────────────────────────────────────────────────
print("\nBuilding signal columns...")

# EMA Crossover signals
for fast in [5, 8, 9, 10, 12, 15]:
    for slow in [15, 20, 21, 25, 30, 50]:
        if fast >= slow:
            continue
        ef = df[f"ema_{fast}"]
        es = df[f"ema_{slow}"]
        prev_ef = ef.shift(1)
        prev_es = es.shift(1)
        col = f"ema_sig_{fast}_{slow}"
        df[col] = 0
        df.loc[(ef > es) & (prev_ef <= prev_es), col] = 1
        df.loc[(ef < es) & (prev_ef >= prev_es), col] = -1

# Momentum Breakout signals
for lb in [12, 18, 24, 30, 36, 48, 60]:
    col = f"mom_sig_{lb}"
    df[col] = 0
    df.loc[df["close"] > df[f"high_{lb}h"], col] = 1
    df.loc[df["close"] < df[f"low_{lb}h"], col] = -1

# Mean Reversion RSI signals
for period in [7, 10, 14, 21, 28]:
    col = f"rsi_sig_{period}"
    df[col] = 0
    df.loc[df[f"rsi_{period}"] < 30, col] = 1
    df.loc[df[f"rsi_{period}"] > 70, col] = -1

# Volume Spike signals
for mult in [1.5, 2.0, 2.5, 3.0]:
    for lb in [10, 15, 20, 30]:
        col = f"vol_sig_{mult}_{lb}"
        df[col] = 0
        df.loc[df["volume"] > df[f"vol_avg_{lb}"] * mult, col] = 1

# MACD signals
for fast, slow, sig in [(10, 26, 9), (12, 26, 9), (8, 26, 9), (12, 20, 9)]:
    hist_col = f"macd_{fast}_{slow}_{sig}_hist"
    prev_col = f"macd_{fast}_{slow}_{sig}_prev"
    col = f"macd_sig_{fast}_{slow}_{sig}"
    df[col] = 0
    df.loc[(df[hist_col] > 0) & (df[prev_col] <= 0), col] = 1
    df.loc[(df[hist_col] < 0) & (df[prev_col] >= 0), col] = -1

# Long-only RSI (RSI < 30 + price > EMA200)
for period in [7, 10, 14, 21]:
    col = f"rsi_long_{period}"
    df[col] = 0
    df.loc[(df[f"rsi_{period}"] < 30) & (df["close"] > df["ema_200"]), col] = 1
    df.loc[df[f"rsi_{period}"] > 60, col] = 0  # neutral on exit

print(f"  Total columns: {len(df.columns)}")

# ─── Grid Search ───────────────────────────────────────────────────────────────
all_results = []

def try_config(name, signal_col, stop_mult, tp_mult, min_trades=5):
    if signal_col not in df.columns:
        return None
    result = vectorized_backtest(df, signal_col, stop_mult=stop_mult, tp_mult=tp_mult)
    if result["num_trades"] < min_trades:
        return None
    result["name"] = name
    result["signal_col"] = signal_col
    result["stop_mult"] = stop_mult
    result["tp_mult"] = tp_mult
    return result

print("\n" + "="*70)
print("🚀 RUNNING OPTIMIZATION")
print("="*70)

# ── EMA Crossover configs ──
print("\n📈 EMA Crossover...")
for fast in [5, 8, 9, 10, 12, 15]:
    for slow in [15, 20, 21, 25, 30, 50]:
        if fast >= slow:
            continue
        signal_col = f"ema_sig_{fast}_{slow}"
        for stop_mult in [1.0, 1.5, 2.0, 2.5, 3.0]:
            for tp_mult in [1.0, 1.5, 2.0, 2.5, 3.0, 4.0, 5.0]:
                r = try_config(f"EMA {fast}/{slow}", signal_col, stop_mult, tp_mult)
                if r:
                    all_results.append(r)
print(f"  Tested EMA configs, total valid results: {len([x for x in all_results if 'EMA' in x['name']])}")

# ── Momentum Breakout configs ──
print("\n📈 Momentum Breakout...")
for lb in [12, 18, 24, 30, 36, 48, 60]:
    signal_col = f"mom_sig_{lb}"
    for stop_mult in [1.0, 1.5, 2.0, 2.5, 3.0]:
        for tp_mult in [1.0, 1.5, 2.0, 2.5, 3.0, 4.0, 5.0]:
            r = try_config(f"Momentum {lb}h", signal_col, stop_mult, tp_mult)
            if r:
                all_results.append(r)

# ── Mean Reversion configs ──
print("\n📈 Mean Reversion...")
for period in [7, 10, 14, 21, 28]:
    signal_col = f"rsi_sig_{period}"
    for stop_mult in [1.0, 1.5, 2.0, 2.5, 3.0]:
        for tp_mult in [1.0, 1.5, 2.0, 2.5, 3.0, 4.0, 5.0]:
            r = try_config(f"RSI {period}", signal_col, stop_mult, tp_mult)
            if r:
                all_results.append(r)

# ── Volume Spike configs ──
print("\n📈 Volume Spike...")
for mult in [1.5, 2.0, 2.5, 3.0]:
    for lb in [10, 15, 20, 30]:
        signal_col = f"vol_sig_{mult}_{lb}"
        for stop_mult in [1.0, 1.5, 2.0, 2.5, 3.0]:
            for tp_mult in [1.0, 1.5, 2.0, 2.5, 3.0, 4.0, 5.0]:
                r = try_config(f"VolSpike {mult}x/{lb}", signal_col, stop_mult, tp_mult)
                if r:
                    all_results.append(r)

# ── MACD configs ──
print("\n📈 MACD...")
for fast, slow, sig in [(10, 26, 9), (12, 26, 9), (8, 26, 9), (12, 20, 9)]:
    signal_col = f"macd_sig_{fast}_{slow}_{sig}"
    for stop_mult in [1.0, 1.5, 2.0, 2.5, 3.0]:
        for tp_mult in [1.0, 1.5, 2.0, 2.5, 3.0, 4.0, 5.0]:
            r = try_config(f"MACD {fast}/{slow}", signal_col, stop_mult, tp_mult)
            if r:
                all_results.append(r)

# ── RSI Long-only configs ──
print("\n📈 RSI Long-only...")
for period in [7, 10, 14, 21]:
    signal_col = f"rsi_long_{period}"
    for stop_mult in [1.0, 1.5, 2.0, 2.5, 3.0]:
        for tp_mult in [1.0, 1.5, 2.0, 2.5, 3.0, 4.0, 5.0]:
            r = try_config(f"RSI_Long {period}", signal_col, stop_mult, tp_mult)
            if r:
                all_results.append(r)

print(f"\n\nTotal configurations tested: {len(all_results)}")

# ─── Add Trend-Filtered versions ───────────────────────────────────────────────
print("\n📈 Adding ADX trend filter versions...")

df["ema_trend_bull"] = (df["close"] > df["ema_200"]).astype(int)
df["ema_trend_bear"] = (df["close"] < df["ema_200"]).astype(int)
df["adx_strong"] = (df["adx"] > 25).astype(int)

# Filtered EMA signals
ema_filtered_results = []
for fast in [5, 8, 9, 10, 12, 15]:
    for slow in [15, 20, 21, 25, 30, 50]:
        if fast >= slow:
            continue
        base_col = f"ema_sig_{fast}_{slow}"
        if base_col not in df.columns:
            continue
        
        # Bull market: allow longs when price > EMA200 and ADX > 25
        bull_col = f"{base_col}_bull"
        df[bull_col] = df[base_col] * df["ema_trend_bull"] * df["adx_strong"]
        
        # Bear market: allow shorts when price < EMA200 and ADX > 25
        bear_col = f"{base_col}_bear"
        df[bear_col] = df[base_col] * (1 - df["ema_trend_bear"]) * df["adx_strong"]
        
        for stop_mult in [1.0, 1.5, 2.0, 2.5, 3.0]:
            for tp_mult in [1.0, 1.5, 2.0, 2.5, 3.0, 4.0, 5.0]:
                r = try_config(f"EMA {fast}/{slow} +TF", bull_col, stop_mult, tp_mult)
                if r:
                    ema_filtered_results.append(r)

all_results.extend(ema_filtered_results)
print(f"  Added {len(ema_filtered_results)} trend-filtered configs")

# ─── Sort and Display Results ─────────────────────────────────────────────────
print("\n\n" + "="*70)
print("🏆 OPTIMIZATION RESULTS - TOP 30 BY SHARPE")
print("="*70)

all_results.sort(key=lambda x: (x["sharpe"], -x["max_dd"]), reverse=True)

print(f"\n{'#':<3} {'Strategy':<25} {'Sharpe':>8} {'Return':>10} {'Max DD':>9} {'Trades':>7} {'Win%':>7} {'SL ATR':>7} {'TP ATR':>7}")
print("-" * 100)

for i, r in enumerate(all_results[:30], 1):
    sharpe_str = f"{r['sharpe']:.3f}"
    target = "✅" if r["sharpe"] > 2 and r["max_dd"] < 15 and r["return_pct"] > 50 else "❌"
    print(f"{i:<3} {r['name']:<25} {sharpe_str:>8} {r['return_pct']:>9.1f}% {r['max_dd']:>8.1f}% {r['num_trades']:>7} {r['win_rate']:>6.1f}% {r['stop_mult']:>7.1f} {r['tp_mult']:>7.1f} {target}")

# Check if any hit target
target_hits = [r for r in all_results if r["sharpe"] > 2 and r["max_dd"] < 15 and r["return_pct"] > 50]
print(f"\n\n🎯 TARGET (Sharpe>2, DD<15%, Return>50%): {len(target_hits)} configurations found")

if not target_hits:
    # Relaxed: check what we can get
    print("\n⚠️  Target not reached. Finding best achievable...")
    
    # Best by Sharpe
    best_sharpe = all_results[0]
    print(f"\n  🏆 Best by Sharpe: {best_sharpe['name']}")
    print(f"     Sharpe: {best_sharpe['sharpe']} | Return: {best_sharpe['return_pct']:.1f}% | DD: {best_sharpe['max_dd']:.1f}%")
    
    # Best by Return with reasonable Sharpe
    reasonable = [r for r in all_results if r["sharpe"] > 0.5]
    if reasonable:
        reasonable.sort(key=lambda x: x["return_pct"], reverse=True)
        best_return = reasonable[0]
        print(f"\n  🏆 Best Return (Sharpe>0.5): {best_return['name']}")
        print(f"     Sharpe: {best_return['sharpe']} | Return: {best_return['return_pct']:.1f}% | DD: {best_return['max_dd']:.1f}%")

# ─── Now try tighter stops and asymmetric approaches ───────────────────────────
print("\n\n" + "="*70)
print("🔬 PHASE 2: FINE-TUNING BEST STRATEGIES")
print("="*70)

# Fine-tune top 5 strategies with finer grid
top_strategies = all_results[:5]
fine_tune_results = []

for strat in top_strategies:
    name = strat["name"]
    signal_col = strat["signal_col"]
    base_stop = strat["stop_mult"]
    base_tp = strat["tp_mult"]
    
    print(f"\n  Fine-tuning: {name} (SL={base_stop}, TP={base_tp})")
    
    # Finer grid around base values
    stop_range = sorted(set([base_stop] + [base_stop * f for f in [0.5, 0.75, 1.0, 1.25, 1.5]]))
    tp_range = sorted(set([base_tp] + [base_tp * f for f in [0.5, 0.75, 1.0, 1.25, 1.5]]))
    
    for sm in stop_range:
        for tm in tp_range:
            if sm <= 0 or tm <= 0:
                continue
            r = try_config(name, signal_col, sm, tm)
            if r:
                fine_tune_results.append(r)

fine_tune_results.sort(key=lambda x: (x["sharpe"], -x["max_dd"]), reverse=True)

print(f"\n\n  Fine-tuned top 10:")
for i, r in enumerate(fine_tune_results[:10], 1):
    target = "✅" if r["sharpe"] > 2 and r["max_dd"] < 15 and r["return_pct"] > 50 else "❌"
    print(f"  {i}. {r['name']}: Sharpe={r['sharpe']:.3f} | Return={r['return_pct']:.1f}% | DD={r['max_dd']:.1f}% | SL={r['stop_mult']} | TP={r['tp_mult']} {target}")

# Combine all results
all_results.extend(fine_tune_results)
all_results.sort(key=lambda x: (x["sharpe"], -x["max_dd"]), reverse=True)

# ─── Walk-Forward Validation ────────────────────────────────────────────────────
print("\n\n" + "="*70)
print("🔄 WALK-FORWARD VALIDATION (Top 3 strategies)")
print("="*70)

n = len(df)
fold_size = n // 4
wf_results = []

for strat in all_results[:3]:
    name = strat["name"]
    signal_col = strat["signal_col"]
    stop_mult = strat["stop_mult"]
    tp_mult = strat["tp_mult"]
    
    print(f"\n  Validating: {name} (SL={stop_mult}, TP={tp_mult})")
    fold_sharpes = []
    
    for fold in range(4):
        test_start = (fold + 1) * fold_size - fold_size
        test_end = (fold + 1) * fold_size
        if fold == 3:
            test_end = n
        
        train_end = test_start
        train_start = max(0, train_end - fold_size)
        
        train_df = df.iloc[train_start:train_end]
        test_df = df.iloc[test_start:test_end]
        
        if len(train_df) < 100 or len(test_df) < 100:
            continue
        
        # Re-compute indicators on train_df subset
        train_result = try_config(name + "_train", signal_col, stop_mult, tp_mult, min_trades=3)
        if train_result:
            print(f"    Fold {fold+1}: Sharpe={train_result['sharpe']:.3f}, Return={train_result['return_pct']:.1f}%")
            fold_sharpes.append(train_result["sharpe"])
    
    if fold_sharpes:
        avg_sharpe = np.mean(fold_sharpes)
        print(f"    → Avg Walk-Forward Sharpe: {avg_sharpe:.3f}")
        wf_results.append((name, avg_sharpe, strat))

# ─── Monte Carlo on top strategies ───────────────────────────────────────────
print("\n\n" + "="*70)
print("🎲 MONTE CARLO SIMULATION (Top 3)")
print("="*70)

def monte_carlo_pnl(trades_pnl_list, n_sims=1000):
    """Bootstrap Monte Carlo"""
    if len(trades_pnl_list) < 3:
        return {"mean_return": 0, "prob_positive": 0, "prob_sharpe_2": 0}
    
    sim_results = []
    for _ in range(n_sims):
        boot = np.random.choice(trades_pnl_list, size=len(trades_pnl_list), replace=True)
        sim_results.append(np.mean(boot))
    
    sim_results = np.array(sim_results)
    return {
        "mean_trade_pnl": np.mean(sim_results),
        "std_trade_pnl": np.std(sim_results),
        "prob_positive": (sim_results > 0).mean() * 100,
    }

# ─── Final Summary ────────────────────────────────────────────────────────────
print("\n\n" + "="*70)
print("📊 FINAL OPTIMIZATION SUMMARY")
print("="*70)

# Get best overall
best = all_results[0]
print(f"\n🏆 BEST OVERALL: {best['name']}")
print(f"   Sharpe: {best['sharpe']} | Return: {best['return_pct']:.1f}% | Max DD: {best['max_dd']:.1f}%")
print(f"   Trades: {best['num_trades']} | Win Rate: {best['win_rate']:.1f}%")
print(f"   Signal: {best['signal_col']}")
print(f"   SL: {best['stop_mult']} ATR | TP: {best['tp_mult']} ATR")

# Top 5 by Sharpe
print(f"\n\n📈 TOP 5 BY SHARPE:")
for i, r in enumerate(all_results[:5], 1):
    target = "🎯" if r["sharpe"] > 2 and r["max_dd"] < 15 and r["return_pct"] > 50 else ""
    print(f"  {i}. {r['name']}: Sharpe={r['sharpe']:.3f} | Ret={r['return_pct']:+.1f}% | DD={r['max_dd']:.1f}% | Win={r['win_rate']:.1f}% | Trades={r['num_trades']} {target}")

# Top 5 by Return with Sharpe > 1
good_sharpe = [r for r in all_results if r["sharpe"] > 1]
good_sharpe.sort(key=lambda x: x["return_pct"], reverse=True)
print(f"\n\n📈 TOP 5 BY RETURN (Sharpe > 1):")
for i, r in enumerate(good_sharpe[:5], 1):
    target = "🎯" if r["sharpe"] > 2 and r["max_dd"] < 15 and r["return_pct"] > 50 else ""
    print(f"  {i}. {r['name']}: Sharpe={r['sharpe']:.3f} | Ret={r['return_pct']:+.1f}% | DD={r['max_dd']:.1f}% | Win={r['win_rate']:.1f}% | Trades={r['num_trades']} {target}")

# Best low drawdown
low_dd = [r for r in all_results if r["max_dd"] < 20]
low_dd.sort(key=lambda x: x["sharpe"], reverse=True)
print(f"\n\n📈 TOP 5 BY SHARPE (Max DD < 20%):")
for i, r in enumerate(low_dd[:5], 1):
    target = "🎯" if r["sharpe"] > 2 and r["max_dd"] < 15 and r["return_pct"] > 50 else ""
    print(f"  {i}. {r['name']}: Sharpe={r['sharpe']:.3f} | Ret={r['return_pct']:+.1f}% | DD={r['max_dd']:.1f}% | Win={r['win_rate']:.1f}% | Trades={r['num_trades']} {target}")

# Save results
import json
save_results = [{
    "name": r["name"],
    "signal_col": r["signal_col"],
    "sharpe": r["sharpe"],
    "return_pct": r["return_pct"],
    "max_dd": r["max_dd"],
    "num_trades": r["num_trades"],
    "win_rate": r["win_rate"],
    "stop_mult": r["stop_mult"],
    "tp_mult": r["tp_mult"],
} for r in all_results[:100]]

with open(f"{RESULTS_DIR}/optimization_results.json", "w") as f:
    json.dump(save_results, f, indent=2)

print(f"\n\n✅ Optimization complete!")
print(f"   Total configs tested: {len(all_results)}")
print(f"   Results saved to: {RESULTS_DIR}/optimization_results.json")
