#!/usr/bin/env python3
"""
Step 2: Data Preprocessing & Integration
Cleans crop production and soil fertility datasets, calculates yield, and prepares data for analysis.
"""

import pandas as pd
import numpy as np
from pathlib import Path
import warnings
warnings.filterwarnings('ignore')

# Set reproducibility
np.random.seed(42)

# Define paths
BASE_DIR = Path('/app/sandbox/session_20251229_081931_9f2d364070e9')
DATA_DIR = BASE_DIR / 'workflow' / 'data'
CROP_INPUT = DATA_DIR / 'crop_production_district.csv'
SOIL_INPUT = DATA_DIR / 'soil_fertility_data.csv'
CROP_OUTPUT = DATA_DIR / 'clean_crop_production.csv'
SOIL_OUTPUT = DATA_DIR / 'clean_soil_data.csv'
SUMMARY_OUTPUT = DATA_DIR / 'state_crop_yield_summary.csv'

print("=" * 70)
print("Step 2: Data Preprocessing & Integration")
print("=" * 70)

# ============================================================================
# PART 1: Process Crop Production Data
# ============================================================================
print("\n[1/3] Processing Crop Production Data...")
print("-" * 70)

# Load crop production data
print(f"Loading: {CROP_INPUT}")
crop_df = pd.read_csv(CROP_INPUT)
print(f"Initial shape: {crop_df.shape}")
print(f"Columns: {list(crop_df.columns)}")
print(f"Initial data types:\n{crop_df.dtypes}")

# Display initial statistics
print(f"\nInitial statistics:")
print(f"  - Total rows: {len(crop_df):,}")
print(f"  - Unique states: {crop_df['State_Name'].nunique()}")
print(f"  - Unique districts: {crop_df['District_Name'].nunique()}")
print(f"  - Unique crops: {crop_df['Crop'].nunique()}")
print(f"  - Year range: {crop_df['Crop_Year'].min()} - {crop_df['Crop_Year'].max()}")

# Check missing values before cleaning
print(f"\nMissing values before cleaning:")
missing_before = crop_df[['Area', 'Production']].isnull().sum()
print(missing_before)

# Step 1: Drop rows with missing values in 'Area' or 'Production'
initial_count = len(crop_df)
crop_df = crop_df.dropna(subset=['Area', 'Production'])
dropped_missing = initial_count - len(crop_df)
print(f"\nDropped {dropped_missing:,} rows with missing Area or Production")

# Step 2: Standardize 'Season' and 'Crop' columns
# Strip whitespace and convert to title case
print("\nStandardizing text columns...")
print(f"  Before - Unique Seasons: {crop_df['Season'].nunique()}")
print(f"  Sample seasons: {crop_df['Season'].unique()[:5]}")
crop_df['Season'] = crop_df['Season'].str.strip().str.title()
print(f"  After - Unique Seasons: {crop_df['Season'].nunique()}")
print(f"  Sample seasons: {sorted(crop_df['Season'].unique())[:5]}")

print(f"\n  Before - Unique Crops: {crop_df['Crop'].nunique()}")
print(f"  Sample crops: {crop_df['Crop'].unique()[:5]}")
crop_df['Crop'] = crop_df['Crop'].str.strip().str.title()
print(f"  After - Unique Crops: {crop_df['Crop'].nunique()}")
print(f"  Sample crops: {sorted(crop_df['Crop'].unique())[:5]}")

# Step 3: Calculate Yield = Production / Area
print("\nCalculating Yield (Production / Area)...")

# Handle division by zero
zero_area = (crop_df['Area'] == 0).sum()
if zero_area > 0:
    print(f"  Warning: Found {zero_area:,} rows with zero Area")
    # Filter out zero area rows
    crop_df = crop_df[crop_df['Area'] > 0]
    print(f"  Removed {zero_area:,} rows with zero Area")

# Calculate yield
crop_df['Yield'] = crop_df['Production'] / crop_df['Area']

# Check for infinite or NaN values
inf_count = np.isinf(crop_df['Yield']).sum()
nan_count = crop_df['Yield'].isna().sum()
print(f"  Yield statistics:")
print(f"    - Infinite values: {inf_count}")
print(f"    - NaN values: {nan_count}")
print(f"    - Min: {crop_df['Yield'].min():.2f}")
print(f"    - Max: {crop_df['Yield'].max():.2f}")
print(f"    - Mean: {crop_df['Yield'].mean():.2f}")
print(f"    - Median: {crop_df['Yield'].median():.2f}")

# Step 4: Filter out anomalies
print("\nFiltering anomalies...")
initial_count = len(crop_df)

# Filter 1: Zero Area but positive Production (should already be handled)
anomaly1 = ((crop_df['Area'] == 0) & (crop_df['Production'] > 0)).sum()
print(f"  - Rows with 0 Area but positive Production: {anomaly1}")

# Filter 2: Extremely high yields (potential data errors)
# Define threshold as mean + 3*std or 99th percentile, whichever is more conservative
yield_mean = crop_df['Yield'].mean()
yield_std = crop_df['Yield'].std()
yield_99th = crop_df['Yield'].quantile(0.99)
threshold_high = max(yield_mean + 3*yield_std, yield_99th * 2)
print(f"  - High yield threshold: {threshold_high:.2f}")

high_yield_mask = crop_df['Yield'] > threshold_high
high_yield_count = high_yield_mask.sum()
print(f"  - Rows with extremely high yield (>{threshold_high:.2f}): {high_yield_count}")

# Filter 3: Negative or zero Production (if any)
neg_prod_mask = crop_df['Production'] <= 0
neg_prod_count = neg_prod_mask.sum()
print(f"  - Rows with zero or negative Production: {neg_prod_count}")

# Apply filters
crop_df = crop_df[~(high_yield_mask | neg_prod_mask)]
filtered_count = initial_count - len(crop_df)
print(f"  Total filtered: {filtered_count:,} rows")

# Final statistics
print(f"\nFinal crop production statistics:")
print(f"  - Total rows: {len(crop_df):,}")
print(f"  - Unique states: {crop_df['State_Name'].nunique()}")
print(f"  - Unique crops: {crop_df['Crop'].nunique()}")
print(f"  - Yield range: {crop_df['Yield'].min():.2f} - {crop_df['Yield'].max():.2f}")
print(f"  - Mean yield: {crop_df['Yield'].mean():.2f}")

# Save cleaned crop production data
print(f"\nSaving cleaned data to: {CROP_OUTPUT}")
crop_df.to_csv(CROP_OUTPUT, index=False)
print(f"✓ Saved {len(crop_df):,} rows")

# ============================================================================
# PART 2: Process Soil Fertility Data
# ============================================================================
print("\n[2/3] Processing Soil Fertility Data...")
print("-" * 70)

# Load soil fertility data
print(f"Loading: {SOIL_INPUT}")
soil_df = pd.read_csv(SOIL_INPUT)
print(f"Initial shape: {soil_df.shape}")
print(f"Columns: {list(soil_df.columns)}")

# Display data types
print(f"\nData types:")
print(soil_df.dtypes)

# Check missing values
print(f"\nMissing values:")
missing_soil = soil_df.isnull().sum()
print(missing_soil[missing_soil > 0])
if missing_soil.sum() == 0:
    print("  ✓ No missing values found")

# Handle missing values if any
if soil_df.isnull().sum().sum() > 0:
    print("\nHandling missing values...")
    # For numeric columns, we could impute with median
    numeric_cols = soil_df.select_dtypes(include=[np.number]).columns
    for col in numeric_cols:
        if soil_df[col].isnull().sum() > 0:
            median_val = soil_df[col].median()
            soil_df[col].fillna(median_val, inplace=True)
            print(f"  - Filled {col} missing values with median: {median_val:.2f}")

    # For categorical columns, fill with mode
    categorical_cols = soil_df.select_dtypes(include=['object']).columns
    for col in categorical_cols:
        if soil_df[col].isnull().sum() > 0:
            mode_val = soil_df[col].mode()[0]
            soil_df[col].fillna(mode_val, inplace=True)
            print(f"  - Filled {col} missing values with mode: {mode_val}")

# Create "Soil Class" column based on Output
print("\nCreating Soil Class column...")
print(f"  Unique Output values: {soil_df['Output'].unique()}")
print(f"  Output value counts:\n{soil_df['Output'].value_counts()}")

# Standardize Output column (strip whitespace, title case)
soil_df['Output'] = soil_df['Output'].str.strip().str.title()

# Create Soil_Class (same as Output but with consistent naming)
soil_df['Soil_Class'] = soil_df['Output']
print(f"  ✓ Created Soil_Class column based on Output")
print(f"  Soil_Class distribution:\n{soil_df['Soil_Class'].value_counts()}")

# Additional classification based on texture if needed
# Calculate texture-based classification
soil_df['Texture_Class'] = 'Loam'  # Default
soil_df.loc[soil_df['Sand'] > 85, 'Texture_Class'] = 'Sandy'
soil_df.loc[soil_df['Clay'] > 40, 'Texture_Class'] = 'Clayey'
soil_df.loc[(soil_df['Silt'] > 80) & (soil_df['Clay'] < 20), 'Texture_Class'] = 'Silty'

print(f"\n  Texture classification:")
print(f"{soil_df['Texture_Class'].value_counts()}")

# Verify no missing values in final dataset
print(f"\nFinal missing values check:")
final_missing = soil_df.isnull().sum()
if final_missing.sum() == 0:
    print("  ✓ No missing values in cleaned soil data")
else:
    print(final_missing[final_missing > 0])

print(f"\nFinal soil data statistics:")
print(f"  - Total rows: {len(soil_df):,}")
print(f"  - Fertile samples: {(soil_df['Soil_Class'] == 'Fertile').sum()}")
print(f"  - Non-Fertile samples: {(soil_df['Soil_Class'] == 'Non Fertile').sum()}")

# Save cleaned soil data
print(f"\nSaving cleaned data to: {SOIL_OUTPUT}")
soil_df.to_csv(SOIL_OUTPUT, index=False)
print(f"✓ Saved {len(soil_df):,} rows")

# ============================================================================
# PART 3: Generate Integration Summary
# ============================================================================
print("\n[3/3] Generating Integration Summary...")
print("-" * 70)

# Aggregate Crop Yields by State and Crop (mean yield)
print("Creating summary: Mean Yield by State and Crop")
summary_df = crop_df.groupby(['State_Name', 'Crop']).agg({
    'Yield': ['mean', 'std', 'count'],
    'Area': 'sum',
    'Production': 'sum'
}).reset_index()

# Flatten column names
summary_df.columns = ['State', 'Crop', 'Mean_Yield', 'Std_Yield', 'N_Records', 'Total_Area', 'Total_Production']

# Round numeric columns
summary_df['Mean_Yield'] = summary_df['Mean_Yield'].round(2)
summary_df['Std_Yield'] = summary_df['Std_Yield'].round(2)

# Sort by Mean_Yield descending
summary_df = summary_df.sort_values('Mean_Yield', ascending=False)

print(f"\nSummary statistics:")
print(f"  - Total State-Crop combinations: {len(summary_df):,}")
print(f"  - States covered: {summary_df['State'].nunique()}")
print(f"  - Crops covered: {summary_df['Crop'].nunique()}")

print(f"\nTop 10 State-Crop combinations by mean yield:")
print(summary_df.head(10).to_string(index=False))

# Save integration summary
print(f"\nSaving integration summary to: {SUMMARY_OUTPUT}")
summary_df.to_csv(SUMMARY_OUTPUT, index=False)
print(f"✓ Saved {len(summary_df):,} rows")

# ============================================================================
# Summary
# ============================================================================
print("\n" + "=" * 70)
print("PREPROCESSING COMPLETE")
print("=" * 70)
print("\nGenerated files:")
print(f"  1. {CROP_OUTPUT.name}: {len(crop_df):,} rows, {len(crop_df.columns)} columns")
print(f"  2. {SOIL_OUTPUT.name}: {len(soil_df):,} rows, {len(soil_df.columns)} columns")
print(f"  3. {SUMMARY_OUTPUT.name}: {len(summary_df):,} rows, {len(summary_df.columns)} columns")

print("\n✓ All preprocessing steps completed successfully!")
print("=" * 70)
