#!/usr/bin/env python3
"""
Test adaptive gating on historical trades.

Usage:
    python test_adaptive_gating.py
"""

import sys
sys.path.insert(0, '/home/ai/projects/AI_Strategy')

import pandas as pd
from db.connection import get_connection
from features.adaptive_gating import AdaptiveGating

# Test configurations
TEST_CONFIGS = {
    'original': {
        'name': 'Original (no adaptive gating)',
        'params': {}
    },
    'adaptive_v1': {
        'name': 'Adaptive Gating v1 (volatility + momentum)',
        'params': {
            'min_volatility_low': 0.3,
            'min_momentum_5': 0.005,
            'min_momentum_10': 0.01,
            'min_adx': 20,
            'max_volatility_high': 1.5
        }
    },
    'adaptive_v2_strict': {
        'name': 'Adaptive Gating v2 (stricter)',
        'params': {
            'min_volatility_low': 0.4,
            'min_momentum_5': 0.01,
            'min_momentum_10': 0.02,
            'min_adx': 25,
            'max_volatility_high': 1.2
        }
    },
    'adaptive_v3_relaxed': {
        'name': 'Adaptive Gating v3 (relaxed)',
        'params': {
            'min_volatility_low': 0.2,
            'min_momentum_5': 0.003,
            'min_momentum_10': 0.005,
            'min_adx': 15,
            'max_volatility_high': 2.0
        }
    }
}


def load_trades():
    """Load historical trades."""
    query = """
    SELECT
        ticker,
        direction,
        entry_price,
        atr_entry,
        pnl
    FROM trades_closed_regression
    WHERE direction = "LONG"
    ORDER BY entry_time
    """

    with get_connection() as conn:
        df = pd.read_sql(query, conn)

    print(f"✅ Loaded {len(df)} trades")
    return df


def simulate_signals_from_trades(trades_df):
    """
    Simulate signals from trades.
    This is a simplified approach - using actual trades as signals.
    """
    signals = trades_df.copy()

    # Simplify features for simulation
    signals['current_price'] = signals['entry_price']
    signals['atr_pct'] = signals['atr_entry'] / signals['entry_price']
    signals['momentum_5'] = signals['atr_pct'] * 0.5  # Placeholder
    signals['momentum_10'] = signals['atr_pct'] * 0.3
    signals['ADX_14'] = 25
    signals['Volume'] = 1000
    signals['volume_avg'] = 1000
    signals['bars_to_target'] = 10

    return signals


def test_gating_config(gating: AdaptiveGating, trades_df, config_name: str):
    """Test a specific gating configuration."""
    print(f"\n{'='*80}")
    print(f"Testing: {config_name}")
    print(f"{'='*80}")

    # Apply adaptive gating
    signals = simulate_signals_from_trades(trades_df)

    gating_results = signals.apply(
        lambda row: gating.check_adaptive_gating(row, max_horizons=[10, 30, 60]),
        axis=1
    )

    signals['gating_passed'] = gating_results.apply(lambda x: x['passed'])
    signals['quality_score'] = gating_results.apply(lambda x: x['passed'] and x['scores']['volatility'])

    # Count trades
    total_trades = len(signals)
    passed_trades = signals['gating_passed'].sum()
    avg_quality_score = signals['quality_score'].mean()

    # Calculate metrics for passed trades
    passed_trades_df = signals[signals['gating_passed']]

    if len(passed_trades_df) > 0:
        avg_pnl = passed_trades_df['pnl'].mean()
        total_pnl = passed_trades_df['pnl'].sum()
        win_rate = (passed_trades_df['pnl'] > 0).sum() / len(passed_trades_df)
    else:
        avg_pnl = 0
        total_pnl = 0
        win_rate = 0

    # Show summary
    print(f"\n📊 Results:")
    print(f"   Total trades: {total_trades}")
    print(f"   Passed: {passed_trades} ({passed_trades/total_trades*100:.1f}%)")
    print(f"   Avg quality score: {avg_quality_score:.3f}")
    print(f"   Win Rate: {win_rate:.1%}")
    print(f"   Avg PnL: {avg_pnl:.2f}")
    print(f"   Total PnL: {total_pnl:.2f}")

    # Show trade distribution by volatility
    print(f"\n📈 Trade Distribution:")
    for vol_regime in ['LOW_VOL', 'NORMAL_VOL', 'HIGH_VOL']:
        regime_trades = signals[signals['atr_pct'] < 0.5] if vol_regime == 'LOW_VOL' else \
                       signals[(signals['atr_pct'] >= 0.5) & (signals['atr_pct'] < 1.5)] if vol_regime == 'NORMAL_VOL' else \
                       signals[signals['atr_pct'] >= 1.5]

        regime_passed = regime_trades['gating_passed'].sum()

        if len(regime_trades) > 0:
            print(f"   {vol_regime:<12} Total: {len(regime_trades):<4}  Passed: {regime_passed}")

    return {
        'config_name': config_name,
        'total_trades': total_trades,
        'passed_trades': passed_trades,
        'pass_rate': passed_trades / total_trades if total_trades > 0 else 0,
        'avg_quality_score': avg_quality_score,
        'win_rate': win_rate,
        'avg_pnl': avg_pnl,
        'total_pnl': total_pnl
    }


def main():
    print("="*80)
    print("🧪 Testing Adaptive Gating on Historical Trades")
    print("="*80)

    # Load trades
    print("\n1️⃣ Loading trades...")
    trades_df = load_trades()

    # Store results
    results = []

    # Test each configuration
    print("\n2️⃣ Testing gating configurations...")

    for config_name, config_params in TEST_CONFIGS.items():
        print(f"\n   Testing: {config_name}")

        # Create gating with specific params
        gating_config = {**AdaptiveGating().get_default_config(), **config_params}
        gating = AdaptiveGating(gating_config)

        # Test configuration
        result = test_gating_config(gating, trades_df, config_name)
        results.append(result)

    # Compare results
    print("\n" + "="*80)
    print("📊 COMPARISON SUMMARY")
    print("="*80)

    results_df = pd.DataFrame(results)

    print(f"\n{'Config':<30} {'Passed':<8} {'Pass Rate':<12} {'WR':<10} {'Avg PnL':<12} {'Total PnL':<12}")
    print("-"*80)

    for _, row in results_df.iterrows():
        print(f"{row['config_name']:<30} {row['passed_trades']:<8} "
              f"{row['pass_rate']:<12.1%} {row['win_rate']:<10.1%} {row['avg_pnl']:>10.2f} "
              f"{row['total_pnl']:>10,.2f}")

    # Identify best configuration
    print("\n" + "="*80)
    print("🏆 BEST CONFIGURATION")
    print("="*80)

    best_by_wr = results_df.loc[results_df['win_rate'].idxmax()]
    best_by_pnl = results_df.loc[results_df['total_pnl'].idxmax()]

    print(f"\nBy Win Rate:")
    print(f"  {best_by_wr['config_name']}")
    print(f"  WR: {best_by_wr['win_rate']:.1%}, Avg PnL: {best_by_wr['avg_pnl']:.2f}")

    print(f"\nBy Total PnL:")
    print(f"  {best_by_pnl['config_name']}")
    print(f"  Total PnL: {best_by_pnl['total_pnl']:.2f}, Avg PnL: {best_by_pnl['avg_pnl']:.2f}")

    # Determine best overall
    print("\n" + "="*80)
    print("🎯 RECOMMENDED CONFIGURATION")
    print("="*80)

    # Best should have: High pass rate, good WR, positive PnL
    # Score = WR * 100 + Avg PnL (normalized)

    def calculate_score(row):
        return row['win_rate'] * 100 + row['avg_pnl']

    results_df['score'] = results_df.apply(calculate_score, axis=1)

    best_overall = results_df.loc[results_df['score'].idxmax()]

    print(f"\nRecommended: {best_overall['config_name']}")
    print(f"   Pass rate: {best_overall['pass_rate']:.1%}")
    print(f"   Win rate: {best_overall['win_rate']:.1%}")
    print(f"   Avg PnL: {best_overall['avg_pnl']:.2f}")
    print(f"   Total PnL: {best_overall['total_pnl']:.2f}")

    print("\n" + "="*80)
    print("✅ Testing complete")
    print("="*80)


if __name__ == '__main__':
    main()
