#!/usr/bin/env python3
"""
Step 2: Meal Optimization for 7-Day Vegetarian Indian Meal Plan
==================================================================
Generates a balanced 7-day meal plan ensuring:
- ~2000 kcal/day (±10%): 1800-2200 kcal
- Protein > 50g/day
- Regional diversity (North, South, East, West)
- Variety (no repeated main dishes on consecutive days)
- Inclusion of dairy/accompaniments and staples (rice/chapati)

FIXED: Corrected staple lookup to search across all categories, not just Dairy/Accompaniments
"""

import pandas as pd
import numpy as np
from pathlib import Path
from typing import Dict, List, Tuple
import random

# Set random seed for reproducibility
np.random.seed(42)
random.seed(42)

# Define paths
BASE_DIR = Path("/app/sandbox/session_20251221_093619_3390c5c3bed6")
DATA_FILE = BASE_DIR / "indian_food_nutrition.csv"
OUTPUT_PLAN = BASE_DIR / "results" / "seven_day_meal_plan.csv"
OUTPUT_SUMMARY = BASE_DIR / "results" / "daily_nutrition_summary.csv"

# Nutritional targets
TARGET_CALORIES = 2000
CALORIE_TOLERANCE = 0.10  # ±10%
MIN_CALORIES = TARGET_CALORIES * (1 - CALORIE_TOLERANCE)  # 1800
MAX_CALORIES = TARGET_CALORIES * (1 + CALORIE_TOLERANCE)  # 2200
MIN_PROTEIN = 50  # grams

# Meal structure
MEAL_TYPES = ["Breakfast", "Lunch", "Snack", "Dinner"]
REGIONS = ["North", "South", "East", "West"]
DAYS = 7


def load_and_preprocess_data(filepath: Path) -> pd.DataFrame:
    """Load and preprocess the nutrition dataset."""
    print(f"Loading data from {filepath}...")
    df = pd.read_csv(filepath)

    # Remove any empty rows
    df = df.dropna(subset=['name'])

    print(f"Loaded {len(df)} food items")
    print(f"Regions: {df['region'].unique()}")
    print(f"Categories: {df['category'].unique()}")

    # Validate data
    required_cols = ['name', 'region', 'category', 'calories', 'protein',
                     'carbohydrates', 'fats', 'fiber']
    for col in required_cols:
        if col not in df.columns:
            raise ValueError(f"Missing required column: {col}")

    return df


def group_foods_by_category(df: pd.DataFrame) -> Dict[str, pd.DataFrame]:
    """
    Group foods by meal category.
    Excludes Rice and Chapati from main dish categories (they're used only as accompaniments).
    """
    groups = {}

    # Define staples that should be excluded from main dish selection
    staple_keywords = ['Rice.*White', 'Chapati', 'Roti', 'Phulka']
    staple_pattern = '|'.join(staple_keywords)

    for category in df['category'].unique():
        category_df = df[df['category'] == category].copy()

        # For Lunch and Dinner, exclude staples from main dish selection
        if category in ['Lunch', 'Dinner']:
            # Keep only non-staple items as main dishes
            category_df = category_df[~category_df['name'].str.contains(staple_pattern, case=False, regex=True)]

        groups[category] = category_df

    return groups


def select_meal_item(available_items: pd.DataFrame,
                     used_items: List[str],
                     preferred_region: str = None) -> pd.Series:
    """
    Select a meal item from available options.

    Args:
        available_items: DataFrame of available food items
        used_items: List of item names already used in previous days
        preferred_region: Preferred region for this meal (optional)

    Returns:
        Selected food item as a Series
    """
    # Filter out recently used items (to avoid consecutive repeats)
    candidates = available_items[~available_items['name'].isin(used_items[-2:])]

    # If no candidates after filtering, use all available
    if len(candidates) == 0:
        candidates = available_items

    # Prefer items from the specified region if provided
    if preferred_region:
        region_candidates = candidates[candidates['region'] == preferred_region]
        if len(region_candidates) > 0:
            candidates = region_candidates

    # Randomly select from candidates
    return candidates.sample(n=1, random_state=None).iloc[0]


def add_accompaniments(meal_category: str,
                      all_foods: pd.DataFrame,
                      food_groups: Dict[str, pd.DataFrame],
                      meal_nutrition: Dict[str, float]) -> Tuple[List[str], Dict[str, float]]:
    """
    Add appropriate accompaniments and staples to a meal.

    Args:
        meal_category: Type of meal (Breakfast, Lunch, Dinner)
        all_foods: Complete DataFrame of all food items (for staple lookup)
        food_groups: Dictionary of food groups by category
        meal_nutrition: Nutritional totals for this specific meal

    Returns:
        Tuple of (list of accompaniment names, updated meal nutrition dict)
    """
    added_items = []

    # Get dairy/accompaniments for this function
    accompaniments = food_groups.get("Dairy/Accompaniments", pd.DataFrame())

    if meal_category == "Breakfast":
        # Add yogurt or lassi (single serving)
        if len(accompaniments) > 0:
            yogurt = accompaniments[accompaniments['name'].str.contains('Yogurt|Lassi', case=False)]
            if len(yogurt) > 0:
                item = yogurt.sample(n=1).iloc[0]
                added_items.append(f"{item['name']}")
                for nutrient in ['calories', 'protein', 'carbohydrates', 'fats', 'fiber']:
                    meal_nutrition[nutrient] += item[nutrient]

    elif meal_category == "Lunch":
        # FIXED: Search for rice in ALL foods, not just accompaniments
        rice = all_foods[all_foods['name'].str.contains('Rice.*White', case=False, regex=True)]
        if len(rice) > 0:
            item = rice.iloc[0]  # Use first rice item found
            # Add 2.5 servings of rice for lunch (balanced portion)
            multiplier = 2.5
            added_items.append(f"{item['name']} ({multiplier}x)")
            for nutrient in ['calories', 'protein', 'carbohydrates', 'fats', 'fiber']:
                meal_nutrition[nutrient] += item[nutrient] * multiplier

        # Add raita or yogurt (single serving)
        if len(accompaniments) > 0:
            dairy = accompaniments[accompaniments['name'].str.contains('Raita|Yogurt', case=False)]
            if len(dairy) > 0:
                item = dairy.sample(n=1).iloc[0]
                added_items.append(f"{item['name']}")
                for nutrient in ['calories', 'protein', 'carbohydrates', 'fats', 'fiber']:
                    meal_nutrition[nutrient] += item[nutrient]

    elif meal_category == "Dinner":
        # FIXED: Search for chapati in ALL foods, not just accompaniments
        chapati = all_foods[all_foods['name'].str.contains('Chapati|Roti|Phulka', case=False)]
        if len(chapati) > 0:
            item = chapati.iloc[0]  # Use first chapati/roti item found
            # Add 2.5 servings of chapati for dinner (balanced portion)
            multiplier = 2.5
            added_items.append(f"{item['name']} ({multiplier}x)")
            for nutrient in ['calories', 'protein', 'carbohydrates', 'fats', 'fiber']:
                meal_nutrition[nutrient] += item[nutrient] * multiplier

        # Add raita or chutney (single serving)
        if len(accompaniments) > 0:
            condiment = accompaniments[accompaniments['name'].str.contains('Raita|Chutney', case=False)]
            if len(condiment) > 0:
                item = condiment.sample(n=1).iloc[0]
                added_items.append(f"{item['name']}")
                for nutrient in ['calories', 'protein', 'carbohydrates', 'fats', 'fiber']:
                    meal_nutrition[nutrient] += item[nutrient]

    return added_items, meal_nutrition


def generate_meal_plan(all_foods: pd.DataFrame,
                       food_groups: Dict[str, pd.DataFrame]) -> pd.DataFrame:
    """
    Generate a 7-day balanced meal plan.

    Args:
        all_foods: Complete DataFrame of all food items
        food_groups: Dictionary of food groups by category

    Returns:
        DataFrame with the complete meal plan
    """
    print("\nGenerating 7-day meal plan...")

    meal_plan = []
    used_main_dishes = []  # Track to avoid consecutive repeats

    # Ensure regional diversity - assign regions to days cyclically
    region_schedule = []
    for i in range(DAYS):
        region_schedule.append(REGIONS[i % len(REGIONS)])

    print(f"Regional schedule: {region_schedule}")

    for day in range(1, DAYS + 1):
        print(f"\nPlanning Day {day}...")
        preferred_region = region_schedule[day - 1]
        print(f"  Preferred region: {preferred_region}")

        daily_nutrition = {
            'calories': 0.0,
            'protein': 0.0,
            'carbohydrates': 0.0,
            'fats': 0.0,
            'fiber': 0.0
        }

        # Plan each meal type
        for meal_type in MEAL_TYPES:
            if meal_type in food_groups:
                # Initialize meal-specific nutrition
                meal_nutrition = {
                    'calories': 0.0,
                    'protein': 0.0,
                    'carbohydrates': 0.0,
                    'fats': 0.0,
                    'fiber': 0.0
                }

                # Select main dish
                main_dish = select_meal_item(
                    food_groups[meal_type],
                    used_main_dishes,
                    preferred_region=preferred_region
                )

                meal_items = [main_dish['name']]
                used_main_dishes.append(main_dish['name'])

                # Add main dish nutrition to meal
                # For snacks, use 1.2x portion to boost daily calories
                snack_multiplier = 1.2 if meal_type == "Snack" else 1.0
                for nutrient in ['calories', 'protein', 'carbohydrates', 'fats', 'fiber']:
                    meal_nutrition[nutrient] += main_dish[nutrient] * snack_multiplier

                # Update meal items display if snack is multiplied
                if snack_multiplier > 1.0:
                    meal_items = [f"{main_dish['name']} ({snack_multiplier}x)"]

                # Add accompaniments for main meals
                if meal_type in ["Breakfast", "Lunch", "Dinner"]:
                    accompaniment_items, meal_nutrition = add_accompaniments(
                        meal_type,
                        all_foods,  # Pass complete dataset for staple lookup
                        food_groups,
                        meal_nutrition
                    )
                    meal_items.extend(accompaniment_items)

                # Update daily totals
                for nutrient in ['calories', 'protein', 'carbohydrates', 'fats', 'fiber']:
                    daily_nutrition[nutrient] += meal_nutrition[nutrient]

                # Create meal plan entry with per-meal nutrition
                meal_entry = {
                    'Day': day,
                    'Meal': meal_type,
                    'Dishes': ' + '.join(meal_items),
                    'Main_Region': main_dish['region'],
                    'Meal_Calories': meal_nutrition['calories'],
                    'Meal_Protein': meal_nutrition['protein'],
                    'Meal_Carbs': meal_nutrition['carbohydrates'],
                    'Meal_Fats': meal_nutrition['fats'],
                    'Meal_Fiber': meal_nutrition['fiber']
                }

                meal_plan.append(meal_entry)

                print(f"  {meal_type}: {main_dish['name']} ({main_dish['region']}) - "
                      f"Meal total: {meal_nutrition['calories']:.0f} kcal, {meal_nutrition['protein']:.1f}g protein")

        # Print daily totals
        print(f"  Daily totals: {daily_nutrition['calories']:.0f} kcal, "
              f"{daily_nutrition['protein']:.1f}g protein")

        # Check if within targets
        if MIN_CALORIES <= daily_nutrition['calories'] <= MAX_CALORIES:
            print(f"  ✓ Calories within target range")
        else:
            print(f"  ⚠ Calories outside target range ({MIN_CALORIES}-{MAX_CALORIES})")

        if daily_nutrition['protein'] >= MIN_PROTEIN:
            print(f"  ✓ Protein meets minimum requirement")
        else:
            print(f"  ⚠ Protein below minimum ({MIN_PROTEIN}g)")

    return pd.DataFrame(meal_plan)


def calculate_daily_summary(meal_plan: pd.DataFrame) -> pd.DataFrame:
    """
    Calculate daily nutritional summary from the meal plan.

    Args:
        meal_plan: Complete meal plan DataFrame

    Returns:
        Daily summary DataFrame
    """
    print("\nCalculating daily nutritional summary...")

    # Sum up per-meal nutrition to get daily totals
    daily_summary = meal_plan.groupby('Day').agg({
        'Meal_Calories': 'sum',
        'Meal_Protein': 'sum',
        'Meal_Carbs': 'sum',
        'Meal_Fats': 'sum',
        'Meal_Fiber': 'sum'
    }).reset_index()

    # Rename columns to standard format
    daily_summary.columns = ['Day', 'Calories', 'Protein', 'Carbs', 'Fats', 'Fiber']

    # Add analysis columns
    daily_summary['Within_Calorie_Target'] = daily_summary['Calories'].apply(
        lambda x: "Yes" if MIN_CALORIES <= x <= MAX_CALORIES else "No"
    )
    daily_summary['Meets_Protein_Min'] = daily_summary['Protein'].apply(
        lambda x: "Yes" if x >= MIN_PROTEIN else "No"
    )

    return daily_summary


def validate_meal_plan(meal_plan: pd.DataFrame, daily_summary: pd.DataFrame) -> bool:
    """
    Validate that the meal plan meets all success criteria.

    Returns:
        True if all criteria are met, False otherwise
    """
    print("\n" + "="*60)
    print("VALIDATION RESULTS")
    print("="*60)

    all_valid = True

    # Check 1: 7 days coverage
    days_count = meal_plan['Day'].nunique()
    print(f"\n1. Days coverage: {days_count}/7 days")
    if days_count == 7:
        print("   ✓ PASS")
    else:
        print("   ✗ FAIL")
        all_valid = False

    # Check 2: Average calories within range
    avg_calories = daily_summary['Calories'].mean()
    print(f"\n2. Average daily calories: {avg_calories:.0f} kcal")
    print(f"   Target range: {MIN_CALORIES:.0f}-{MAX_CALORIES:.0f} kcal")
    if MIN_CALORIES <= avg_calories <= MAX_CALORIES:
        print("   ✓ PASS")
    else:
        print("   ✗ FAIL")
        all_valid = False

    # Check 3: Regional diversity
    regions_used = meal_plan['Main_Region'].unique()
    print(f"\n3. Regional diversity: {len(regions_used)}/4 regions")
    print(f"   Regions included: {', '.join(sorted(regions_used))}")
    if len(regions_used) == 4:
        print("   ✓ PASS")
    else:
        print("   ✗ FAIL")
        all_valid = False

    # Check 4: Protein adequacy
    protein_meets_min = (daily_summary['Protein'] >= MIN_PROTEIN).all()
    avg_protein = daily_summary['Protein'].mean()
    print(f"\n4. Protein adequacy: {avg_protein:.1f}g/day average")
    print(f"   Minimum required: {MIN_PROTEIN}g/day")
    if protein_meets_min:
        print("   ✓ PASS - All days meet minimum")
    else:
        days_below = (daily_summary['Protein'] < MIN_PROTEIN).sum()
        print(f"   ⚠ WARNING - {days_below} days below minimum")

    # Check 5: Meal structure completeness
    expected_meals = DAYS * len(MEAL_TYPES)
    actual_meals = len(meal_plan)
    print(f"\n5. Meal structure: {actual_meals}/{expected_meals} meals")
    if actual_meals == expected_meals:
        print("   ✓ PASS")
    else:
        print("   ✗ FAIL")
        all_valid = False

    # Check 6: Staples inclusion (NEW CHECK)
    meals_with_staples = meal_plan[
        (meal_plan['Meal'].isin(['Lunch', 'Dinner'])) &
        (meal_plan['Dishes'].str.contains('Rice|Chapati|Roti|Phulka', case=False))
    ]
    expected_staple_meals = DAYS * 2  # Lunch + Dinner for 7 days
    actual_staple_meals = len(meals_with_staples)
    print(f"\n6. Staples inclusion: {actual_staple_meals}/{expected_staple_meals} Lunch/Dinner meals")
    if actual_staple_meals == expected_staple_meals:
        print("   ✓ PASS - All Lunch/Dinner meals include rice or chapati")
    else:
        print(f"   ⚠ WARNING - {expected_staple_meals - actual_staple_meals} meals missing staples")

    # Check 7: Individual day calorie variance
    print(f"\n7. Daily calorie range: {daily_summary['Calories'].min():.0f} - {daily_summary['Calories'].max():.0f} kcal")
    days_in_range = (daily_summary['Within_Calorie_Target'] == 'Yes').sum()
    print(f"   Days within target: {days_in_range}/{DAYS}")
    if days_in_range >= 5:  # At least 5 out of 7 days
        print("   ✓ PASS")
    else:
        print("   ⚠ WARNING - Less than 5 days within target range")

    print("\n" + "="*60)
    if all_valid:
        print("✓ ALL VALIDATION CHECKS PASSED")
    else:
        print("⚠ SOME VALIDATION CHECKS FAILED")
    print("="*60 + "\n")

    return all_valid


def main():
    """Main execution function."""
    print("="*60)
    print("Step 2: Meal Optimization (FIXED)")
    print("="*60)

    # Load data
    df = load_and_preprocess_data(DATA_FILE)

    # Group foods by category
    food_groups = group_foods_by_category(df)
    print("\nFood groups created:")
    for category, items in food_groups.items():
        print(f"  {category}: {len(items)} items")

    # Verify staples are found
    print("\nVerifying staples location:")
    rice_items = df[df['name'].str.contains('Rice', case=False)]
    chapati_items = df[df['name'].str.contains('Chapati|Roti|Phulka', case=False)]
    print(f"  Rice items found: {len(rice_items)}")
    for _, item in rice_items.iterrows():
        print(f"    - {item['name']} (Category: {item['category']})")
    print(f"  Chapati/Roti items found: {len(chapati_items)}")
    for _, item in chapati_items.iterrows():
        print(f"    - {item['name']} (Category: {item['category']})")

    # Generate meal plan
    meal_plan = generate_meal_plan(df, food_groups)

    # Calculate daily summary
    daily_summary = calculate_daily_summary(meal_plan)

    # Validate meal plan
    validation_passed = validate_meal_plan(meal_plan, daily_summary)

    # Save outputs
    print("\nSaving outputs...")
    OUTPUT_PLAN.parent.mkdir(parents=True, exist_ok=True)

    meal_plan.to_csv(OUTPUT_PLAN, index=False)
    print(f"✓ Saved meal plan to: {OUTPUT_PLAN}")

    daily_summary.to_csv(OUTPUT_SUMMARY, index=False)
    print(f"✓ Saved daily summary to: {OUTPUT_SUMMARY}")

    # Display summary statistics
    print("\n" + "="*60)
    print("SUMMARY STATISTICS")
    print("="*60)
    print(f"\nDaily Nutritional Averages:")
    print(f"  Calories: {daily_summary['Calories'].mean():.0f} ± {daily_summary['Calories'].std():.0f} kcal")
    print(f"  Protein: {daily_summary['Protein'].mean():.1f} ± {daily_summary['Protein'].std():.1f} g")
    print(f"  Carbs: {daily_summary['Carbs'].mean():.1f} ± {daily_summary['Carbs'].std():.1f} g")
    print(f"  Fats: {daily_summary['Fats'].mean():.1f} ± {daily_summary['Fats'].std():.1f} g")
    print(f"  Fiber: {daily_summary['Fiber'].mean():.1f} ± {daily_summary['Fiber'].std():.1f} g")

    print("\nRegional Distribution:")
    region_counts = meal_plan['Main_Region'].value_counts()
    for region in REGIONS:
        count = region_counts.get(region, 0)
        percentage = (count / len(meal_plan)) * 100
        print(f"  {region}: {count} meals ({percentage:.1f}%)")

    print("\n" + "="*60)
    print("✓ Step 2 completed successfully!")
    print("="*60)

    return validation_passed


if __name__ == "__main__":
    success = main()
    exit(0 if success else 1)
