#!/usr/bin/env python3
"""Parameter optimization for pump_short_bot V2 - quick grid search on IS data."""
import sys, os, json, glob
import pandas as pd
import numpy as np

sys.path.insert(0, os.path.abspath(os.path.dirname(__file__)))
from config.settings_v2 import V2Params
from strategy.detector_v2 import detect_v2
from strategy.simulator_v2 import simulate_v2
from metrics_v2 import compute_metrics

# Load IS data once
def load_is_data():
    base = 'data/real/IS'
    f15_files = sorted(glob.glob(os.path.join(base, '**', '*_15m.csv'), recursive=True))
    data = {}
    for f15p in f15_files:
        sym = os.path.basename(f15p).replace('_15m.csv', '')
        d = os.path.dirname(f15p)
        f4p = os.path.join(d, f'{sym}_4h.csv')
        if not os.path.exists(f4p):
            continue
        df15 = pd.read_csv(f15p).sort_values('time')
        df4 = pd.read_csv(f4p).sort_values('time')
        k15 = [{'time': int(r.time), 'open': r.open, 'high': r.high, 'low': r.low, 'close': r.close,
                'volume': r.volume, 'quote_volume': r.quote_volume, 'taker_buy_quote': r.taker_buy_quote} 
               for r in df15.itertuples()]
        k4 = [{'time': int(r.time), 'open': r.open, 'high': r.high, 'low': r.low, 'close': r.close,
               'volume': r.volume, 'quote_volume': r.quote_volume, 'taker_buy_quote': r.taker_buy_quote} 
              for r in df4.itertuples()]
        data[sym] = (k4, k15)
    return data

ALL_DATA = load_is_data()
print(f"Loaded {len(ALL_DATA)} symbols for IS")

# Parameter grid - REDUCED for speed
PARAM_GRID = [
    # range_pctile_thr, atr_ratio_thr, volume_ratio_thr, va_pctile_thr, trailing_k
    (20, 0.7, 0.8, 25, 2.5),  # baseline
    (15, 0.6, 0.7, 20, 2.0),  # loose
    (25, 0.8, 0.9, 30, 3.0),  # tight
    (20, 0.7, 0.8, 25, 2.0),  # trailing 2.0
    (20, 0.7, 0.8, 25, 3.0),  # trailing 3.0
]

results = []

for i, (rp, ar, vr, vp, tk) in enumerate(PARAM_GRID):
    p = V2Params()
    p.range_pctile_thr = rp
    p.atr_ratio_thr = ar
    p.volume_ratio_thr = vr
    p.va_pctile_thr = vp
    p.trailing_k = tk
    
    all_trades = []
    total_bars = 0
    all_times = []
    
    for sym, (k4, k15) in ALL_DATA.items():
        total_bars += len(k15)
        if k15:
            all_times.append((k15[0]['time'], k15[-1]['time']))
        signals = detect_v2(k4, k15, p)
        for s in signals:
            t = simulate_v2(sym, k15, s, p, equity=100.0)
            if t and t.r_multiple is not None:
                all_trades.append({
                    'r_multiple': t.r_multiple,
                    'leverage': getattr(t, 'leverage', 1.0),
                    'score': s.get('squeeze_score', 0),
                    'direction': t.direction,
                    'exit_reason': t.exit_reason,
                })
    
    if not all_trades:
        print(f"  [{i+1}/{len(PARAM_GRID)}] rp={rp} ar={ar} vr={vr} vp={vp} tk={tk} -> 0 trades")
        continue
    
    # Equity curve
    equity = 100.0
    curve = [equity]
    for tr in all_trades:
        risk_pct = p.risk_score4 if tr['score'] >= 4 else p.risk_score3
        risk_usd = equity * risk_pct
        pnl_usd = tr['r_multiple'] * risk_usd * tr['leverage']
        cost = risk_usd * tr['leverage'] * (p.fee_pct * 2 + p.slippage_pct * 2)
        equity += pnl_usd - cost
        curve.append(equity)
    
    period_days = 0
    if all_times:
        t0 = min(a[0] for a in all_times)
        t1 = max(a[1] for a in all_times)
        period_days = (t1 - t0) / (1000 * 3600 * 24)
    if period_days == 0:
        period_days = 60  # IS is ~60 days
    
    m = compute_metrics(all_trades, curve, period_days, total_bars=total_bars)
    
    r = {
        'params': {'range_pctile_thr': rp, 'atr_ratio_thr': ar, 'volume_ratio_thr': vr, 
                   'va_pctile_thr': vp, 'trailing_k': tk},
        'trades': len(all_trades),
        'win_rate': m.get('win_rate', 0),
        'profit_factor': m.get('profit_factor', 0),
        'sharpe': m.get('sharpe', 0),
        'sortino': m.get('sortino', 0),
        'max_dd_pct': m.get('max_drawdown_pct', 0),
        'avg_dd_pct': m.get('avg_drawdown_pct', 0),
        'total_return_pct': m.get('total_return_pct', 0),
        'cagr_pct': m.get('cagr_pct', 0),
        'calmar': m.get('calmar', 0),
        'sterling': m.get('sterling', 0),
        'avg_r': m.get('avg_r', 0),
        'final_equity': m.get('final_equity', 100),
    }
    results.append(r)
    
    print(f"  [{i+1}/{len(PARAM_GRID)}] rp={rp} ar={ar} vr={vr} vp={vp} tk={tk} -> "
          f"n={r['trades']} wr={r['win_rate']:.1f} pf={r['profit_factor']:.2f} "
          f"sh={r['sharpe']:.2f} dd={r['max_dd_pct']:.1f}% ret={r['total_return_pct']:.1f}%")

# Save results
with open('opt_results.json', 'w') as f:
    json.dump(results, f, indent=2)

# Print best
if results:
    best = max(results, key=lambda x: x['sharpe'])
    print("\n=== BEST (by Sharpe) ===")
    print(json.dumps(best, indent=2))
    
    best_ret = max(results, key=lambda x: x['total_return_pct'])
    print("\n=== BEST (by Return) ===")
    print(json.dumps(best_ret, indent=2))
    
    best_dd = min(results, key=lambda x: x['max_dd_pct'])
    print("\n=== BEST (by Min DD) ===")
    print(json.dumps(best_dd, indent=2))

print("\nDone. Results saved to opt_results.json")