#!/usr/bin/env python3
"""
Experiment quality filters for MoE v12.

Usage:
    python experiment_moe_v12_filters.py
"""

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

import pandas as pd
from db.connection import get_connection

# Experiment configs
EXPERIMENT_CONFIGS = {
    'no_filter': {
        'name': 'No filters (baseline)',
        'params': {}
    },
    'volatility_strict': {
        'name': 'Volatility strict (<0.4%)',
        'params': {
            'min_atr_pct': 0.4
        }
    },
    'volatility_very_strict': {
        'name': 'Volatility very strict (<0.3%)',
        'params': {
            'min_atr_pct': 0.3
        }
    },
    'momentum_filter': {
        'name': 'Momentum filter',
        'params': {
            'min_momentum_5': 0.01,
            'min_momentum_10': 0.02
        }
    },
    'composite': {
        'name': 'Composite (vol + momentum)',
        'params': {
            'min_atr_pct': 0.4,
            'min_momentum_5': 0.01,
            'min_momentum_10': 0.02
        }
    }
}


def load_trades():
    """Load MoE v12 trades."""
    query = """
    SELECT
        ticker,
        direction,
        entry_price,
        atr_entry,
        pnl
    FROM trades_closed
    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_momentum_features(df):
    """Simulate momentum features for testing."""
    df = df.copy()

    # Create atr_pct if not exists
    df['atr_pct'] = df['atr_entry'] / df['entry_price']

    # Simulate momentum based on ATR % (simplified)
    df['momentum_5'] = df['atr_pct'] * 0.3
    df['momentum_10'] = df['atr_pct'] * 0.2

    return df


def apply_filter(df, params):
    """Apply filter based on parameters."""
    df_filtered = df.copy()

    # Volatility filter
    if 'min_atr_pct' in params:
        min_atr = params['min_atr_pct']
        df_filtered = df_filtered[df_filtered['atr_pct'] >= min_atr]

    # Momentum filter
    if 'min_momentum_5' in params:
        df_filtered = df_filtered[df_filtered['momentum_5'] >= params['min_momentum_5']]

    if 'min_momentum_10' in params:
        df_filtered = df_filtered[df_filtered['momentum_10'] >= params['min_momentum_10']]

    return df_filtered


def analyze_trades(df, config_name):
    """Analyze trades with given configuration."""
    if len(df) == 0:
        return None

    trades = len(df)
    total_pnl = df['pnl'].sum()
    avg_pnl = df['pnl'].mean()
    wr = (df['pnl'] > 0).sum() / trades

    # Separate win/loss
    win_avg = df[df['pnl'] > 0]['pnl'].mean()
    loss_avg = df[df['pnl'] <= 0]['pnl'].abs().mean()

    win_loss_ratio = win_avg / loss_avg if loss_avg > 0 else float('inf')

    return {
        'config_name': config_name,
        'trades': trades,
        'wr': wr,
        'avg_pnl': avg_pnl,
        'total_pnl': total_pnl,
        'win_avg': win_avg,
        'loss_avg': loss_avg,
        'win_loss_ratio': win_loss_ratio
    }


def main():
    print("="*80)
    print("🔬 EXPERIMENT: Quality Filters for MoE v12")
    print("="*80)
    print()

    # Load trades
    print("1️⃣ Loading trades...")
    df = load_trades()

    # Simulate momentum features
    print("\n2️⃣ Simulating momentum features...")
    df = simulate_momentum_features(df)

    # Run experiments
    print("\n3️⃣ Running experiments...")
    results = []

    for exp_name, exp_params in EXPERIMENT_CONFIGS.items():
        print(f"\n   Testing: {exp_name}")

        # Apply filter
        df_filtered = apply_filter(df, exp_params['params'])
        metrics = analyze_trades(df_filtered, exp_name)

        if metrics:
            results.append(metrics)
            print(f"      Trades: {metrics['trades']:<4}  WR: {metrics['wr']:<10.1%}  Avg PnL: {metrics['avg_pnl']:>10.2f}")
        else:
            print(f"      Trades: 0 (no candidates)")

    # Load original baseline
    print("\n4️⃣ Comparing with original...")
    original_metrics = analyze_trades(df, 'Original')

    # Display results
    print("\n" + "="*80)
    print("📊 EXPERIMENT RESULTS")
    print("="*80)

    results_df = pd.DataFrame(results)
    if 'Total PnL' in results_df.columns:
        # Create synthetic column for all rows
        results_df['Total PnL'] = None
        results_df.loc[results_df['config_name'] == 'Original', 'Total PnL'] = original_metrics['total_pnl']
        results_df['Total PnL'] = results_df['Total PnL'].fillna(0)

    print(f"\n{'Experiment':<30} {'Trades':<8} {'WR':<10} {'Avg PnL':<12} {'Total PnL':<12}")
    print("-"*80)

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

    # Identify best experiments
    print("\n" + "="*80)
    print("🏆 TOP 3 EXPERIMENTS")
    print("="*80)

    best_by_wr = results_df.nlargest(3, 'wr')
    best_by_pnl = results_df.nlargest(3, 'avg_pnl')

    print(f"\nBy Win Rate:")
    for _, row in best_by_wr.iterrows():
        print(f"  {row['config_name']:<30} WR: {row['wr']:.1%}, Trades: {row['trades']}")

    print(f"\nBy Avg PnL:")
    for _, row in best_by_pnl.iterrows():
        print(f"  {row['config_name']:<30} Avg PnL: {row['avg_pnl']:.2f}, Trades: {row['trades']}")

    # Compare with original
    print("\n" + "="*80)
    print("📈 COMPARISON WITH ORIGINAL")
    print("="*80)

    if original_metrics and any(results_df['config_name'] != 'Original'):
        best = results_df.loc[results_df['avg_pnl'].idxmax()]

        print(f"\nBest configuration: {best['config_name']}")
        print(f"  Trades: {best['trades']} (vs {original_metrics['trades']} original)")
        print(f"  Reduction: {(1 - best['trades']/original_metrics['trades'])*100:.1f}%")
        print(f"  WR: {best['wr']:.1%} (vs {original_metrics['wr']:.1%} original)")
        print(f"  Avg PnL: {best['avg_pnl']:.2f} (vs {original_metrics['avg_pnl']:.2f} original)")
        print(f"  Total PnL: {best['total_pnl']:.2f} (vs {original_metrics['total_pnl']:.2f} original)")

    print("\n" + "="*80)
    print("💡 RECOMMENDATIONS")
    print("="*80)

    if original_metrics:
        print(f"\n✅ Original MoE v12 already works well:")
        print(f"   WR: {original_metrics['wr']:.1%}, Avg PnL: {original_metrics['avg_pnl']:.2f}, Total PnL: {original_metrics['total_pnl']:.2f}")
        print(f"   Total trades: {original_metrics['trades']}")

        print(f"\n✅ If you want to reduce trades:")
        best = results_df.loc[results_df['avg_pnl'].idxmax()]
        print(f"   Apply: {best['config_name']}")
        print(f"   Reduction: {(1 - best['trades']/original_metrics['trades'])*100:.1f}%")
        print(f"   Expected impact on WR: {(best['wr'] - original_metrics['wr'])*100:+.1f}%")
        print(f"   Expected impact on PnL: {(best['avg_pnl'] - original_metrics['avg_pnl']):+.2f}")

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


if __name__ == '__main__':
    main()
