#!/usr/bin/env python3
"""
Visualize architecture comparison: MoE v12 vs MoERegression.

Usage:
    python visualize_architecture_comparison.py
"""

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

import pandas as pd
import matplotlib.pyplot as plt
import matplotlib
matplotlib.use('Agg')  # Non-interactive backend
import numpy as np

# Style
plt.style.use('seaborn-v0_8-whitegrid')
colors = ['steelblue', 'forestgreen', 'lightblue', 'lightgreen', 'orange', 'red']

def load_data():
    """Load trade data from database."""
    from db.connection import get_connection

    # MoE v12
    query_moe = """
    SELECT
        ticker, direction, close_reason,
        pnl, pnl_pct,
        entry_time
    FROM trades_closed
    ORDER BY entry_time DESC
    """
    with get_connection() as conn:
        df_moe = pd.read_sql(query_moe, conn)

    # MoERegression
    query_reg = """
    SELECT
        ticker, direction, close_reason,
        pnl, pnl_pct,
        entry_time
    FROM trades_closed_regression
    ORDER BY entry_time DESC
    """
    with get_connection() as conn:
        df_reg = pd.read_sql(query_reg, conn)

    return df_moe, df_reg


def plot_pnl_distribution(df_moe, df_reg, colors):
    """Plot PnL distribution comparison."""
    fig, axes = plt.subplots(2, 2, figsize=(16, 12))

    # MoE v12 PnL
    ax1 = axes[0, 0]
    ax1.hist(df_moe['pnl'].dropna(), bins=20, alpha=0.7, color=colors[0], edgecolor='black')
    ax1.axvline(df_moe['pnl'].mean(), color='red', linestyle='--', linewidth=2, label=f'Mean: {df_moe["pnl"].mean():.2f}')
    ax1.set_title('MoE v12: PnL Distribution', fontsize=14, fontweight='bold')
    ax1.set_xlabel('PnL (RUB)')
    ax1.set_ylabel('Count')
    ax1.legend()
    ax1.grid(True, alpha=0.3)

    # MoERegression PnL
    ax2 = axes[0, 1]
    ax2.hist(df_reg['pnl'].dropna(), bins=20, alpha=0.7, color=colors[1], edgecolor='black')
    ax2.axvline(df_reg['pnl'].mean(), color='red', linestyle='--', linewidth=2, label=f'Mean: {df_reg["pnl"].mean():.2f}')
    ax2.set_title('MoERegression: PnL Distribution', fontsize=14, fontweight='bold')
    ax2.set_xlabel('PnL (RUB)')
    ax2.set_ylabel('Count')
    ax2.legend()
    ax2.grid(True, alpha=0.3)

    # Box plots
    ax3 = axes[1, 0]
    data = [df_moe['pnl'], df_reg['pnl']]
    bp = ax3.boxplot(data, labels=['MoE v12', 'MoERegression'], patch_artist=True)
    colors = ['lightblue', 'lightgreen']
    for patch, color in zip(bp['boxes'], colors):
        patch.set_facecolor(color)
    ax3.set_title('PnL Box Plot Comparison', fontsize=14, fontweight='bold')
    ax3.set_ylabel('PnL (RUB)')
    ax3.grid(True, alpha=0.3, axis='y')

    # Win/Loss ratio
    ax4 = axes[1, 1]
    win_moe = (df_moe['pnl'] > 0).sum()
    loss_moe = (df_moe['pnl'] <= 0).sum()
    win_reg = (df_reg['pnl'] > 0).sum()
    loss_reg = (df_reg['pnl'] <= 0).sum()

    categories = ['MoE v12', 'MoERegression']
    wins = [win_moe, win_reg]
    losses = [loss_moe, loss_reg]

    x = range(len(categories))
    width = 0.35

    ax4.bar([i - width/2 for i in x], wins, width, label='Wins', color='green', alpha=0.7)
    ax4.bar([i + width/2 for i in x], losses, width, label='Losses', color='red', alpha=0.7)
    ax4.set_title('Win/Loss Ratio', fontsize=14, fontweight='bold')
    ax4.set_ylabel('Count')
    ax4.set_xticks(x)
    ax4.set_xticklabels(categories)
    ax4.legend()
    ax4.grid(True, alpha=0.3, axis='y')

    plt.tight_layout()
    plt.savefig('/home/ai/projects/AI_Strategy/figures/architecture_comparison_1.png', dpi=150)
    print("✅ Saved: figures/architecture_comparison_1.png")


def plot_ticker_performance(df_moe, df_reg, palette):
    """Plot per-ticker performance comparison."""
    fig, axes = plt.subplots(2, 2, figsize=(18, 12))

    # MoE v12 Ticker PnL
    ax1 = axes[0, 0]
    ticker_pnl_moe = df_moe.groupby('ticker')['pnl'].sum().sort_values()
    ticker_pnl_moe.plot(kind='barh', alpha=0.7, ax=ax1, color=colors[0])
    ax1.set_title('MoE v12: Total PnL per Ticker', fontsize=14, fontweight='bold')
    ax1.set_xlabel('Total PnL (RUB)')
    ax1.grid(True, alpha=0.3, axis='x')

    # MoERegression Ticker PnL
    ax2 = axes[0, 1]
    ticker_pnl_reg = df_reg.groupby('ticker')['pnl'].sum().sort_values()
    ticker_pnl_reg.plot(kind='barh', alpha=0.7, ax=ax2, color=colors[1])
    ax2.set_title('MoERegression: Total PnL per Ticker', fontsize=14, fontweight='bold')
    ax2.set_xlabel('Total PnL (RUB)')
    ax2.grid(True, alpha=0.3, axis='x')

    # Cumulative PnL over time
    ax3 = axes[1, 0]

    # MoE v12
    df_moe_sorted = df_moe.sort_values('entry_time')
    df_moe_sorted['cumulative_pnl'] = df_moe_sorted['pnl'].cumsum()
    df_moe_sorted['entry_date'] = pd.to_datetime(df_moe_sorted['entry_time'], unit='s')
    df_moe_sorted.set_index('entry_date')['cumulative_pnl'].plot(
        label='MoE v12', linewidth=2, ax=ax3
    )

    # MoERegression
    df_reg_sorted = df_reg.sort_values('entry_time')
    df_reg_sorted['cumulative_pnl'] = df_reg_sorted['pnl'].cumsum()
    df_reg_sorted['entry_date'] = pd.to_datetime(df_reg_sorted['entry_time'], unit='s')
    df_reg_sorted.set_index('entry_date')['cumulative_pnl'].plot(
        label='MoERegression', linewidth=2, ax=ax3, linestyle='--'
    )

    ax3.set_title('Cumulative PnL over Time', fontsize=14, fontweight='bold')
    ax3.set_xlabel('Date')
    ax3.set_ylabel('Cumulative PnL (RUB)')
    ax3.legend()
    ax3.grid(True, alpha=0.3)

    # Close reasons comparison
    ax4 = axes[1, 1]

    close_moe = df_moe['close_reason'].value_counts()
    close_reg = df_reg['close_reason'].value_counts()

    close_moe.plot(kind='barh', alpha=0.7, ax=ax4, color=palette[0], label='MoE v12')
    close_reg.plot(kind='barh', alpha=0.7, ax=ax4, color=palette[1], label='MoERegression')

    ax4.set_title('Close Reasons Comparison', fontsize=14, fontweight='bold')
    ax4.set_xlabel('Count')
    ax4.legend()
    ax4.grid(True, alpha=0.3, axis='x')

    plt.tight_layout()
    plt.savefig('/home/ai/projects/AI_Strategy/figures/architecture_comparison_2.png', dpi=150)
    print("✅ Saved: figures/architecture_comparison_2.png")


def plot_directional_performance(df_moe, df_reg, palette):
    """Plot directional (LONG/SHORT) performance."""
    fig, axes = plt.subplots(1, 3, figsize=(18, 5))

    # MoE v12 directional
    ax1 = axes[0]
    moe_long = df_moe[df_moe['direction'] == 'LONG']
    moe_short = df_moe[df_moe['direction'] == 'SHORT']

    if len(moe_long) > 0:
        moe_long_wr = (moe_long['close_reason'] == 'TP').sum() / len(moe_long)
    else:
        moe_long_wr = 0

    if len(moe_short) > 0:
        moe_short_wr = (moe_short['close_reason'] == 'TP').sum() / len(moe_short)
    else:
        moe_short_wr = 0

    directions = ['LONG', 'SHORT']
    win_rates = [moe_long_wr, moe_short_wr]

    # Simplified plotting without complex color management
    ax1.plot(directions, [wr * 100 for wr in win_rates], 'o-', linewidth=2, markersize=10, color='blue')
    ax1.set_xticks(range(len(directions)))
    ax1.set_xticklabels(directions)
    ax1.set_title('MoE v12: Win Rate by Direction', fontsize=12, fontweight='bold')
    ax1.set_ylabel('Win Rate (%)')
    ax1.set_ylim(0, 50)
    ax1.grid(True, alpha=0.3, axis='y')

    # Add text labels
    for i, (dir, wr) in enumerate(zip(directions, win_rates)):
        ax1.text(i, wr * 100, f'{wr:.1%}', ha='center', va='bottom', fontweight='bold', fontsize=12)

    # MoERegression directional
    ax2 = axes[1]
    reg_long = df_reg[df_reg['direction'] == 'LONG']
    reg_short = df_reg[df_reg['direction'] == 'SHORT']

    if len(reg_long) > 0:
        reg_long_wr = (reg_long['close_reason'] == 'TP').sum() / len(reg_long)
    else:
        reg_long_wr = 0

    if len(reg_short) > 0:
        reg_short_wr = (reg_short['close_reason'] == 'TP').sum() / len(reg_short)
    else:
        reg_short_wr = 0

    directions = ['LONG', 'SHORT']
    win_rates = [reg_long_wr, reg_short_wr]

    # Simplified plotting
    ax2.plot(directions, [wr * 100 for wr in win_rates], 'o-', linewidth=2, markersize=10, color='green')
    ax2.set_xticks(range(len(directions)))
    ax2.set_xticklabels(directions)
    ax2.set_title('MoERegression: Win Rate by Direction', fontsize=12, fontweight='bold')
    ax2.set_ylabel('Win Rate (%)')
    ax2.set_ylim(0, 50)
    ax2.grid(True, alpha=0.3, axis='y')

    # Add text labels
    for i, (dir, wr) in enumerate(zip(directions, win_rates)):
        ax2.text(i, wr * 100, f'{wr:.1%}', ha='center', va='bottom', fontweight='bold', fontsize=12)

    # PnL by direction
    ax3 = axes[2]

    moe_long_pnl = moe_long['pnl'].mean()
    moe_short_pnl = moe_short['pnl'].mean()
    reg_long_pnl = reg_long['pnl'].mean()
    reg_short_pnl = reg_short['pnl'].mean()

    directions = ['MoE\nLONG', 'MoE\nSHORT', 'Reg\nLONG', 'Reg\nSHORT']
    pnl = [moe_long_pnl, moe_short_pnl, reg_long_pnl, reg_short_pnl]
    pnl_colors = ['green' if pnl_i > 0 else 'red' for pnl_i in pnl]

    bars = ax3.bar(directions, pnl, color=pnl_colors, alpha=0.7)
    ax3.axhline(y=0, color='black', linewidth=2)
    ax3.set_title('Avg PnL per Direction', fontsize=12, fontweight='bold')
    ax3.set_ylabel('Avg PnL (RUB)')
    ax3.grid(True, alpha=0.3, axis='y')

    for bar, pnl_i in zip(bars, pnl):
        height = bar.get_height()
        ax3.text(bar.get_x() + bar.get_width()/2., height,
                f'{pnl_i:.1f}', ha='center', va='bottom' if pnl_i >= 0 else 'top', fontweight='bold')

    plt.tight_layout()
    plt.savefig('/home/ai/projects/AI_Strategy/figures/architecture_comparison_3.png', dpi=150)
    print("✅ Saved: figures/architecture_comparison_3.png")


def plot_win_loss_balance(df_moe, df_reg, palette):
    """Plot win/loss balance analysis."""
    fig, axes = plt.subplots(1, 2, figsize=(16, 6))

    # MoE v12 Win/Loss
    ax1 = axes[0]

    moe_win = df_moe[df_moe['pnl'] > 0]
    moe_loss = df_moe[df_moe['pnl'] <= 0]

    if len(moe_win) > 0:
        moe_win_avg = moe_win['pnl'].mean()
        moe_win_sum = moe_win['pnl'].sum()
    else:
        moe_win_avg = 0
        moe_win_sum = 0

    if len(moe_loss) > 0:
        moe_loss_avg = moe_loss['pnl'].mean()
        moe_loss_sum = moe_loss['pnl'].sum()
    else:
        moe_loss_avg = 0
        moe_loss_sum = 0

    labels = ['Wins', 'Losses']
    win_avg = [moe_win_avg, moe_loss_avg]
    win_sum = [moe_win_sum, moe_loss_sum]
    x = range(len(labels))
    width = 0.35

    bars1 = ax1.bar([i - width/2 for i in x], win_avg, width, label='Avg PnL', color='steelblue')
    bars2 = ax1.bar([i + width/2 for i in x], win_sum, width, label='Total PnL', color='orange')

    ax1.set_title('MoE v12: Win/Loss Balance', fontsize=14, fontweight='bold')
    ax1.set_ylabel('PnL (RUB)')
    ax1.set_xticks(x)
    ax1.set_xticklabels(labels)
    ax1.legend()
    ax1.grid(True, alpha=0.3, axis='y')

    # MoERegression Win/Loss
    ax2 = axes[1]

    reg_win = df_reg[df_reg['pnl'] > 0]
    reg_loss = df_reg[df_reg['pnl'] <= 0]

    if len(reg_win) > 0:
        reg_win_avg = reg_win['pnl'].mean()
        reg_win_sum = reg_win['pnl'].sum()
    else:
        reg_win_avg = 0
        reg_win_sum = 0

    if len(reg_loss) > 0:
        reg_loss_avg = reg_loss['pnl'].mean()
        reg_loss_sum = reg_loss['pnl'].sum()
    else:
        reg_loss_avg = 0
        reg_loss_sum = 0

    labels = ['Wins', 'Losses']
    win_avg = [reg_win_avg, reg_loss_avg]
    win_sum = [reg_win_sum, reg_loss_sum]
    x = range(len(labels))
    width = 0.35

    bars1 = ax2.bar([i - width/2 for i in x], win_avg, width, label='Avg PnL', color='forestgreen')
    bars2 = ax2.bar([i + width/2 for i in x], win_sum, width, label='Total PnL', color='orange')

    ax2.set_title('MoERegression: Win/Loss Balance', fontsize=14, fontweight='bold')
    ax2.set_ylabel('PnL (RUB)')
    ax2.set_xticks(x)
    ax2.set_xticklabels(labels)
    ax2.legend()
    ax2.grid(True, alpha=0.3, axis='y')

    plt.tight_layout()
    plt.savefig('/home/ai/projects/AI_Strategy/figures/architecture_comparison_4.png', dpi=150)
    print("✅ Saved: figures/architecture_comparison_4.png")


def main():
    print("📊 Visualizing Architecture Comparison")
    print("=" * 60)

    # Load data
    print("\n1️⃣ Loading data from database...")
    df_moe, df_reg = load_data()
    print(f"   MoE v12: {len(df_moe)} trades")
    print(f"   MoERegression: {len(df_reg)} trades")

    # Create figures directory
    import os
    os.makedirs('/home/ai/projects/AI_Strategy/figures', exist_ok=True)

    # Create visualizations
    print("\n2️⃣ Creating visualizations...")
    plot_pnl_distribution(df_moe, df_reg, colors)
    plot_ticker_performance(df_moe, df_reg, colors)
    plot_directional_performance(df_moe, df_reg, colors)
    plot_win_loss_balance(df_moe, df_reg, colors)

    print("\n" + "=" * 60)
    print("✅ All visualizations saved to: figures/")
    print("=" * 60)


if __name__ == '__main__':
    main()
