#!/usr/bin/env python3
"""
Test adaptive gating on volatility-based analysis.

Key insight from experiment: All 39 trades have ATR < 0.5%
"""

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

import pandas as pd
from db.connection import get_connection


def analyze_volatility_distribution():
    """Analyze volatility distribution of 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)

    # Calculate ATR % of price
    df['atr_pct'] = df['atr_entry'] / df['entry_price']

    # Define volatility regimes
    def get_volatility_regime(atr_pct):
        if atr_pct < 0.5:
            return 'LOW_VOL (<0.5%)'
        elif atr_pct < 1.5:
            return 'NORMAL (0.5-1.5%)'
        else:
            return 'HIGH_VOL (>1.5%)'

    df['vol_regime'] = df['atr_pct'].apply(get_volatility_regime)

    # Show distribution
    print("="*80)
    print("📊 Volatility Distribution of Historical Trades")
    print("="*80)
    print()

    summary = df.groupby('vol_regime').agg({
        'pnl': ['count', 'sum', 'mean', lambda x: (x > 0).sum() / len(x)]
    }).round(2)

    summary.columns = ['Trades', 'Total PnL', 'Avg PnL', 'Win Rate']
    summary = summary.sort_values('Trades', ascending=False)

    print(summary.to_string())
    print()

    # Calculate pass rates for different ATR thresholds
    print("="*80)
    print("📈 Pass Rate by ATR Threshold")
    print("="*80)
    print()

    thresholds = [0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9, 1.0]
    results = []

    for thresh in thresholds:
        pass_trades = df[df['atr_pct'] <= thresh]
        pass_rate = len(pass_trades) / len(df) if len(df) > 0 else 0

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

        results.append({
            'ATR Threshold': f'<= {thresh}',
            'Passed': len(pass_trades),
            'Pass Rate': pass_rate,
            'Win Rate': wr,
            'Avg PnL': avg_pnl,
            'Total PnL': total_pnl
        })

    pass_rates_df = pd.DataFrame(results)
    print(pass_rates_df.to_string(index=False))
    print()

    # Find optimal threshold
    print("="*80)
    print("🏆 Optimal ATR Threshold")
    print("="*80)
    print()

    # Score = Win Rate * 100 + Avg PnL
    pass_rates_df['score'] = pass_rates_df['Win Rate'] * 100 + pass_rates_df['Avg PnL']

    best = pass_rates_df.loc[pass_rates_df['score'].idxmax()]

    print(f"Optimal Threshold: <= {best['ATR Threshold']}")
    print(f"  Pass Rate: {best['Pass Rate']:.1%}")
    print(f"  Win Rate: {best['Win Rate']:.1%}")
    print(f"  Avg PnL: {best['Avg PnL']:.2f}")
    print(f"  Total PnL: {best['Total PnL']:.2f}")
    print(f"  Trades Passed: {best['Passed']}")

    print()
    print("="*80)

    return df, pass_rates_df


def simulate_adaptive_gating(df, atr_threshold):
    """Simulate adaptive gating with specific ATR threshold."""
    # Apply threshold filter
    passed_df = df[df['atr_pct'] <= atr_threshold].copy()

    # Calculate metrics
    total_trades = len(df)
    passed_trades = len(passed_df)
    pass_rate = passed_trades / total_trades if total_trades > 0 else 0

    if passed_trades > 0:
        avg_pnl = passed_df['pnl'].mean()
        total_pnl = passed_df['pnl'].sum()
        wr = (passed_df['pnl'] > 0).sum() / passed_trades
    else:
        avg_pnl = 0
        total_pnl = 0
        wr = 0

    return {
        'atr_threshold': atr_threshold,
        'total_trades': total_trades,
        'passed_trades': passed_trades,
        'pass_rate': pass_rate,
        'win_rate': wr,
        'avg_pnl': avg_pnl,
        'total_pnl': total_pnl,
        'original_trades': total_trades,
        'reduced': total_trades - passed_trades,
        'reduction_pct': (total_trades - passed_trades) / total_trades * 100 if total_trades > 0 else 0
    }


def main():
    print("🧪 Volatility-Based Adaptive Gating Test")
    print("="*80)
    print()

    # Analyze volatility distribution
    df, pass_rates_df = analyze_volatility_distribution()

    # Show original vs adaptive
    print("="*80)
    print("📊 Original vs Adaptive Comparison")
    print("="*80)
    print()

    original_metrics = simulate_adaptive_gating(df, 1.0)  # No filter
    optimal_metrics = simulate_adaptive_gating(df, 0.5)  # Optimal threshold

    print(f"{'Metric':<20} {'Original':<20} {'Adaptive':<20} {'Improvement'}")
    print("-"*80)
    print(f"{'Trades':<20} {original_metrics['total_trades']:<20} {optimal_metrics['passed_trades']:<20} "
          f"{original_metrics['total_trades'] - optimal_metrics['passed_trades']} ({optimal_metrics['reduction_pct']:.0f}%)")
    print(f"{'Pass Rate':<20} {original_metrics['pass_rate']:<20} {optimal_metrics['pass_rate']:<20} -")
    print(f"{'Win Rate':<20} {original_metrics['win_rate']:<20} {optimal_metrics['win_rate']:<20} "
          f"{(optimal_metrics['win_rate'] - original_metrics['win_rate'])*100:+.1f}%")
    print(f"{'Avg PnL':<20} {original_metrics['avg_pnl']:<20} {optimal_metrics['avg_pnl']:<20} "
          f"{(optimal_metrics['avg_pnl'] - original_metrics['avg_pnl']):+.2f}")
    print(f"{'Total PnL':<20} {original_metrics['total_pnl']:<20} {optimal_metrics['total_pnl']:<20} "
          f"{(optimal_metrics['total_pnl'] - original_metrics['total_pnl']):+.2f}")

    print()
    print("="*80)
    print("💡 Recommendation")
    print("="*80)
    print()
    print(f"Implement adaptive gating with ATR threshold = 0.5%")
    print(f"Expected impact:")
    print(f"  - Reduce trades by {(optimal_metrics['reduction_pct']):.0f}%")
    print(f"  - Improve WR by {(optimal_metrics['win_rate'] - original_metrics['win_rate'])*100:+.1f}%")
    print(f"  - Improve Avg PnL by {optimal_metrics['avg_pnl'] - original_metrics['avg_pnl']:+.2f}")
    print(f"  - Improve Total PnL by {optimal_metrics['total_pnl'] - original_metrics['total_pnl']:+.2f}")
    print()
    print("="*80)


if __name__ == '__main__':
    main()
