#!/usr/bin/env python3
"""
Step 3: Visualization Generation
Generate professional visualizations for the Indian meal plan analysis.
"""

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

# Set style for professional plots
plt.rcParams['font.family'] = 'sans-serif'
plt.rcParams['font.size'] = 10
plt.rcParams['axes.linewidth'] = 1.0
plt.rcParams['figure.dpi'] = 300
sns.set_palette("husl")

# Define paths
BASE_DIR = Path('/app/sandbox/session_20251221_093619_3390c5c3bed6')
RESULTS_DIR = BASE_DIR / 'results'
FIGURES_DIR = BASE_DIR / 'figures'

# Load data
print("Loading data files...")
meal_plan = pd.read_csv(RESULTS_DIR / 'seven_day_meal_plan.csv')
daily_summary = pd.read_csv(RESULTS_DIR / 'daily_nutrition_summary.csv')
print(f"✓ Loaded meal plan: {meal_plan.shape}")
print(f"✓ Loaded daily summary: {daily_summary.shape}")

# ============================================================================
# Figure 1: Macronutrient Distribution (Pie Chart)
# Average distribution of Calories from Protein, Carbs, and Fats
# ============================================================================
print("\nGenerating Figure 1: Macronutrient Distribution (Pie Chart)...")

# Calculate average macros across the week
avg_protein = daily_summary['Protein'].mean()  # grams
avg_carbs = daily_summary['Carbs'].mean()  # grams
avg_fats = daily_summary['Fats'].mean()  # grams

# Convert to calories (Protein: 4 kcal/g, Carbs: 4 kcal/g, Fats: 9 kcal/g)
protein_kcal = avg_protein * 4
carbs_kcal = avg_carbs * 4
fats_kcal = avg_fats * 9

# Create pie chart
fig, ax = plt.subplots(figsize=(8, 6))
sizes = [protein_kcal, carbs_kcal, fats_kcal]
labels = [f'Protein\n{protein_kcal:.0f} kcal\n({100*protein_kcal/sum(sizes):.1f}%)',
          f'Carbohydrates\n{carbs_kcal:.0f} kcal\n({100*carbs_kcal/sum(sizes):.1f}%)',
          f'Fats\n{fats_kcal:.0f} kcal\n({100*fats_kcal/sum(sizes):.1f}%)']
colors = ['#ff9999', '#66b3ff', '#99ff99']
explode = (0.05, 0.05, 0.05)

ax.pie(sizes, labels=labels, colors=colors, explode=explode, autopct='',
       shadow=True, startangle=90)
ax.set_title('Average Weekly Macronutrient Distribution\n(Caloric Contribution)',
             fontsize=14, fontweight='bold', pad=20)
plt.tight_layout()
plt.savefig(FIGURES_DIR / '01_macro_distribution.png', dpi=300, bbox_inches='tight')
plt.close()
print("✓ Saved: figures/01_macro_distribution.png")

# ============================================================================
# Figure 2: Caloric Breakdown per Meal (Stacked Bar)
# Daily calories broken down by Breakfast, Lunch, Snack, Dinner
# ============================================================================
print("\nGenerating Figure 2: Caloric Breakdown per Meal (Stacked Bar)...")

# Pivot data to get meals as columns
meal_calories = meal_plan.pivot_table(index='Day', columns='Meal',
                                       values='Meal_Calories', aggfunc='sum')
# Ensure correct meal order
meal_order = ['Breakfast', 'Lunch', 'Snack', 'Dinner']
meal_calories = meal_calories[meal_order]

fig, ax = plt.subplots(figsize=(10, 6))
meal_calories.plot(kind='bar', stacked=True, ax=ax,
                   color=['#FFD700', '#FF6347', '#87CEEB', '#9370DB'],
                   edgecolor='black', linewidth=0.5)
ax.set_xlabel('Day', fontsize=12, fontweight='bold')
ax.set_ylabel('Calories (kcal)', fontsize=12, fontweight='bold')
ax.set_title('Daily Caloric Breakdown by Meal Type', fontsize=14, fontweight='bold', pad=15)
ax.legend(title='Meal Type', bbox_to_anchor=(1.05, 1), loc='upper left')
ax.set_xticklabels([f'Day {i}' for i in range(1, 8)], rotation=0)
ax.axhline(y=2000, color='red', linestyle='--', linewidth=2, label='Target: 2000 kcal')
ax.grid(axis='y', alpha=0.3)
plt.tight_layout()
plt.savefig(FIGURES_DIR / '02_calories_per_meal.png', dpi=300, bbox_inches='tight')
plt.close()
print("✓ Saved: figures/02_calories_per_meal.png")

# ============================================================================
# Figure 3: Daily Macronutrient Composition (Stacked Bar)
# Daily grams of Protein, Carbs, and Fats
# ============================================================================
print("\nGenerating Figure 3: Daily Macronutrient Composition (Stacked Bar)...")

fig, ax = plt.subplots(figsize=(10, 6))
x = np.arange(len(daily_summary))
width = 0.6

p1 = ax.bar(x, daily_summary['Protein'], width, label='Protein (g)',
            color='#ff9999', edgecolor='black', linewidth=0.5)
p2 = ax.bar(x, daily_summary['Carbs'], width, bottom=daily_summary['Protein'],
            label='Carbohydrates (g)', color='#66b3ff', edgecolor='black', linewidth=0.5)
p3 = ax.bar(x, daily_summary['Fats'], width,
            bottom=daily_summary['Protein'] + daily_summary['Carbs'],
            label='Fats (g)', color='#99ff99', edgecolor='black', linewidth=0.5)

ax.set_xlabel('Day', fontsize=12, fontweight='bold')
ax.set_ylabel('Macronutrients (grams)', fontsize=12, fontweight='bold')
ax.set_title('Daily Macronutrient Composition (Stacked)', fontsize=14, fontweight='bold', pad=15)
ax.set_xticks(x)
ax.set_xticklabels([f'Day {i}' for i in range(1, 8)])
ax.legend(loc='upper right')
ax.grid(axis='y', alpha=0.3)
plt.tight_layout()
plt.savefig(FIGURES_DIR / '03_daily_macros.png', dpi=300, bbox_inches='tight')
plt.close()
print("✓ Saved: figures/03_daily_macros.png")

# ============================================================================
# Figure 4: Regional Diversity (Bar Chart)
# Count of meals per region (North, South, East, West)
# ============================================================================
print("\nGenerating Figure 4: Regional Diversity (Bar Chart)...")

# Count meals by region
region_counts = meal_plan['Main_Region'].value_counts().sort_index()

fig, ax = plt.subplots(figsize=(8, 6))
bars = ax.bar(region_counts.index, region_counts.values,
              color=['#FF6B6B', '#4ECDC4', '#45B7D1', '#FFA07A'],
              edgecolor='black', linewidth=1.5)

# Add value labels on bars
for bar in bars:
    height = bar.get_height()
    ax.text(bar.get_x() + bar.get_width()/2., height,
            f'{int(height)} meals',
            ha='center', va='bottom', fontweight='bold', fontsize=11)

ax.set_xlabel('Region', fontsize=12, fontweight='bold')
ax.set_ylabel('Number of Meals', fontsize=12, fontweight='bold')
ax.set_title('Regional Diversity in 7-Day Meal Plan', fontsize=14, fontweight='bold', pad=15)
ax.set_ylim(0, max(region_counts.values) * 1.15)
ax.grid(axis='y', alpha=0.3)
plt.tight_layout()
plt.savefig(FIGURES_DIR / '04_regional_diversity.png', dpi=300, bbox_inches='tight')
plt.close()
print("✓ Saved: figures/04_regional_diversity.png")

# ============================================================================
# Figure 5: RDA Compliance (Radar Chart)
# Compare average daily values against standard targets
# ============================================================================
print("\nGenerating Figure 5: RDA Compliance (Radar Chart)...")

# Define RDA targets
rda_targets = {
    'Calories': 2000,
    'Protein': 50,
    'Carbs': 260,
    'Fats': 70,
    'Fiber': 30
}

# Calculate actual averages
actual_values = {
    'Calories': daily_summary['Calories'].mean(),
    'Protein': daily_summary['Protein'].mean(),
    'Carbs': daily_summary['Carbs'].mean(),
    'Fats': daily_summary['Fats'].mean(),
    'Fiber': daily_summary['Fiber'].mean()
}

# Calculate percentages of RDA
categories = list(rda_targets.keys())
rda_percentages = [100 * actual_values[cat] / rda_targets[cat] for cat in categories]

# Number of variables
N = len(categories)
angles = np.linspace(0, 2 * np.pi, N, endpoint=False).tolist()
rda_percentages += rda_percentages[:1]  # Complete the circle
angles += angles[:1]

# Create radar chart
fig, ax = plt.subplots(figsize=(8, 8), subplot_kw=dict(projection='polar'))
ax.plot(angles, rda_percentages, 'o-', linewidth=2, color='#4169E1', label='Actual')
ax.fill(angles, rda_percentages, alpha=0.25, color='#4169E1')

# Add 100% reference line
reference = [100] * (N + 1)
ax.plot(angles, reference, '--', linewidth=2, color='red', label='RDA Target (100%)')

# Customize
ax.set_theta_offset(np.pi / 2)
ax.set_theta_direction(-1)
ax.set_xticks(angles[:-1])
ax.set_xticklabels(categories, fontsize=11, fontweight='bold')
ax.set_ylim(0, 150)
ax.set_yticks([50, 100, 150])
ax.set_yticklabels(['50%', '100%', '150%'], fontsize=9)
ax.set_title('RDA Compliance: Meal Plan vs. Recommended Daily Allowance\n',
             fontsize=14, fontweight='bold', y=1.08)
ax.legend(loc='upper right', bbox_to_anchor=(1.3, 1.1), fontsize=10)
ax.grid(True, alpha=0.3)

# Add text annotations for actual percentages
for angle, percentage, category in zip(angles[:-1], rda_percentages[:-1], categories):
    ax.text(angle, percentage + 10, f'{percentage:.1f}%',
            ha='center', va='center', fontsize=9, fontweight='bold',
            bbox=dict(boxstyle='round,pad=0.3', facecolor='white', alpha=0.8))

plt.tight_layout()
plt.savefig(FIGURES_DIR / '05_rda_radar_chart.png', dpi=300, bbox_inches='tight')
plt.close()
print("✓ Saved: figures/05_rda_radar_chart.png")

# ============================================================================
# Summary
# ============================================================================
print("\n" + "="*70)
print("VISUALIZATION GENERATION COMPLETE")
print("="*70)
print(f"\n✓ All 5 figures generated successfully in: {FIGURES_DIR}")
print("\nGenerated files:")
print("  1. figures/01_macro_distribution.png - Macronutrient distribution pie chart")
print("  2. figures/02_calories_per_meal.png - Daily caloric breakdown by meal type")
print("  3. figures/03_daily_macros.png - Daily macronutrient composition (stacked)")
print("  4. figures/04_regional_diversity.png - Regional diversity bar chart")
print("  5. figures/05_rda_radar_chart.png - RDA compliance radar chart")
print("\nAll visualizations saved at 300 DPI with proper labels and legends.")
print("="*70)
