#!/usr/bin/env python3
"""
Data Acquisition Script for Air Pollution Data (OpenAQ v3 API)
Fetches historical pollution data from OpenAQ v3 and World Bank APIs
for New Delhi, Mumbai, Bengaluru, and Hyderabad.
"""

import requests
import pandas as pd
import json
import time
from datetime import datetime, timedelta
from pathlib import Path
import sys

# Configuration
API_KEY = "5e9335f14620aa0dcb851537e1a57a6d201210fa06448a3dd0186e7816d5515a"
OUTPUT_DIR = Path("/app/sandbox/session_20251224_200424_4a6cc33ed0fd/workflow/raw_data")
RESULTS_DIR = Path("/app/sandbox/session_20251224_200424_4a6cc33ed0fd/results")

# Target cities - using city names as recognized by OpenAQ
CITIES = ["New Delhi", "Mumbai", "Bengaluru", "Hyderabad"]

# Alternative: search by country code
COUNTRY_CODE = "IN"

# Target parameters (OpenAQ v3 uses IDs)
# Common parameter IDs: 2 = PM2.5, 1 = PM10, 3 = NO2
PARAMETERS = {
    "pm25": 2,
    "pm10": 1,
    "no2": 3
}

# World Bank country code for India
WB_COUNTRY_CODE = "IND"

# World Bank indicators for pollution data
WB_INDICATORS = {
    "EN.ATM.PM25.MC.M3": "PM2.5 air pollution, mean annual exposure (micrograms per cubic meter)",
    "EN.ATM.PM25.MC.ZS": "PM2.5 air pollution, population exposed to levels exceeding WHO guideline value (% of total)"
}


def search_locations(city_name, api_key):
    """
    Search for location IDs in OpenAQ v3 API.

    Args:
        city_name: Name of the city to search
        api_key: OpenAQ API key

    Returns:
        List of location dictionaries
    """
    base_url = "https://api.openaq.org/v3/locations"
    headers = {"X-API-Key": api_key}

    params = {
        "country": "IN",
        "limit": 100
    }

    try:
        response = requests.get(base_url, headers=headers, params=params, timeout=30)
        response.raise_for_status()
        data = response.json()

        if "results" in data:
            # Filter locations by city name (case-insensitive match)
            city_lower = city_name.lower()
            matching_locations = [
                loc for loc in data["results"]
                if city_lower in loc.get("name", "").lower() or
                   city_lower in loc.get("city", "").lower() or
                   city_lower in loc.get("locality", "").lower()
            ]
            return matching_locations
        return []
    except Exception as e:
        print(f"  Error searching locations for {city_name}: {e}")
        return []


def fetch_openaq_v3_measurements(location_ids, parameter_ids, api_key, limit_per_location=10000):
    """
    Fetch measurements from OpenAQ v3 API for specific locations and parameters.

    Args:
        location_ids: List of location IDs
        parameter_ids: Dictionary of parameter names to IDs
        api_key: OpenAQ API key
        limit_per_location: Maximum measurements to fetch per location

    Returns:
        DataFrame with measurements
    """
    base_url = "https://api.openaq.org/v3/measurements"
    headers = {"X-API-Key": api_key}

    all_data = []

    for location_id in location_ids:
        print(f"  Fetching data for location ID: {location_id}")

        for param_name, param_id in parameter_ids.items():
            page = 1
            total_fetched = 0

            while total_fetched < limit_per_location:
                params = {
                    "locations_id": location_id,
                    "parameters_id": param_id,
                    "limit": 1000,
                    "page": page,
                    "date_from": "2015-01-01",
                    "order_by": "datetime"
                }

                try:
                    response = requests.get(base_url, headers=headers, params=params, timeout=30)
                    response.raise_for_status()
                    data = response.json()

                    if "results" not in data or len(data["results"]) == 0:
                        break

                    results = data["results"]
                    total_fetched += len(results)

                    for measurement in results:
                        record = {
                            "location_id": location_id,
                            "parameter": param_name,
                            "value": measurement.get("value"),
                            "date": measurement.get("datetime"),
                            "coordinates": f"{measurement.get('coordinates', {}).get('latitude')},{measurement.get('coordinates', {}).get('longitude')}"
                        }
                        all_data.append(record)

                    if page % 5 == 0:
                        print(f"    Progress: {total_fetched} measurements for {param_name} (page {page})")

                    page += 1
                    time.sleep(0.5)

                except requests.exceptions.RequestException as e:
                    print(f"    Error fetching {param_name}, page {page}: {e}")
                    break
                except Exception as e:
                    print(f"    Unexpected error: {e}")
                    break

            if total_fetched > 0:
                print(f"    ✓ Fetched {total_fetched} measurements for {param_name}")

    return pd.DataFrame(all_data) if all_data else pd.DataFrame()


def fetch_openaq_data_all_india(api_key):
    """
    Fetch air quality data from OpenAQ v3 API for Indian cities.

    Args:
        api_key: OpenAQ API key

    Returns:
        DataFrame with the fetched data
    """
    print(f"\n{'='*60}")
    print("Fetching OpenAQ v3 data for Indian cities...")
    print(f"{'='*60}")

    # Method 1: Try to get recent measurements by country
    base_url = "https://api.openaq.org/v3/measurements"
    headers = {"X-API-Key": api_key}

    all_data = []

    # Fetch latest measurements for India
    print("\nFetching latest measurements for India...")

    for param_name, param_id in PARAMETERS.items():
        print(f"\nFetching {param_name.upper()} data...")
        page = 1
        total_fetched = 0
        max_pages = 50  # Limit to avoid excessive API calls

        while page <= max_pages:
            params = {
                "countries_id": "IN",
                "parameters_id": param_id,
                "limit": 1000,
                "page": page,
                "date_from": "2020-01-01",  # Get data from 2020 onwards
                "order_by": "datetime"
            }

            try:
                response = requests.get(base_url, headers=headers, params=params, timeout=30)

                # Check if endpoint exists
                if response.status_code == 404:
                    print(f"  Endpoint not found. Trying alternative approach...")
                    break

                response.raise_for_status()
                data = response.json()

                if "results" not in data or len(data["results"]) == 0:
                    print(f"  No more data for {param_name} (page {page})")
                    break

                results = data["results"]
                total_fetched += len(results)

                for measurement in results:
                    # Extract location information
                    location_name = measurement.get("location", {}).get("name", "Unknown")
                    city = measurement.get("location", {}).get("city", "Unknown")

                    # Filter for target cities
                    city_lower = city.lower() if city else ""
                    location_lower = location_name.lower() if location_name else ""

                    is_target_city = any(
                        target.lower() in city_lower or target.lower() in location_lower
                        for target in CITIES
                    )

                    if is_target_city:
                        record = {
                            "city": city,
                            "location": location_name,
                            "parameter": param_name,
                            "value": measurement.get("value"),
                            "unit": measurement.get("parameter", {}).get("units"),
                            "date": measurement.get("datetime"),
                            "coordinates": f"{measurement.get('coordinates', {}).get('latitude')},{measurement.get('coordinates', {}).get('longitude')}"
                        }
                        all_data.append(record)

                if page % 5 == 0:
                    print(f"  Progress: Processed {total_fetched} measurements (page {page}), found {len(all_data)} for target cities")

                page += 1
                time.sleep(0.5)

            except requests.exceptions.RequestException as e:
                print(f"  Error fetching data: {e}")
                break
            except Exception as e:
                print(f"  Unexpected error: {e}")
                break

        if total_fetched > 0:
            print(f"  ✓ Processed {total_fetched} total measurements for {param_name}")
            print(f"  ✓ Found {len([r for r in all_data if r['parameter'] == param_name])} measurements for target cities")

    if all_data:
        df = pd.DataFrame(all_data)
        print(f"\n✓ Total measurements for target cities: {len(df)}")
        return df
    else:
        print(f"\n⚠ No data fetched from OpenAQ v3 API")
        return pd.DataFrame()


def fetch_world_bank_data(country_code, indicators):
    """
    Fetch historical pollution data from World Bank API.

    Args:
        country_code: World Bank country code (e.g., 'IND' for India)
        indicators: Dictionary of indicator codes and descriptions

    Returns:
        DataFrame with the fetched data
    """
    print(f"\n{'='*60}")
    print(f"Fetching World Bank data for {country_code}...")
    print(f"{'='*60}")

    all_data = []

    for indicator_code, description in indicators.items():
        print(f"\nFetching: {description}")

        url = f"https://api.worldbank.org/v2/country/{country_code}/indicator/{indicator_code}"
        params = {
            "format": "json",
            "per_page": 1000,
            "date": "1990:2024"
        }

        try:
            response = requests.get(url, params=params, timeout=30)
            response.raise_for_status()
            data = response.json()

            if len(data) > 1 and data[1]:
                results = data[1]
                print(f"  Fetched {len(results)} yearly records")

                for record in results:
                    entry = {
                        "country": record.get("country", {}).get("value"),
                        "country_code": country_code,
                        "indicator": description,
                        "indicator_code": indicator_code,
                        "year": record.get("date"),
                        "value": record.get("value")
                    }
                    all_data.append(entry)
            else:
                print(f"  No data available for {indicator_code}")

        except Exception as e:
            print(f"  Error: {e}")

        time.sleep(0.3)

    if all_data:
        df = pd.DataFrame(all_data)
        print(f"\n✓ Total World Bank records fetched: {len(df)}")
        return df
    else:
        print(f"\n⚠ No World Bank data fetched")
        return pd.DataFrame()


def generate_summary(openaq_df, wb_df, output_file):
    """Generate a summary of the acquired data."""
    print(f"\n{'='*60}")
    print("Generating Data Acquisition Summary...")
    print(f"{'='*60}")

    summary_lines = []
    summary_lines.append("=" * 80)
    summary_lines.append("DATA ACQUISITION SUMMARY")
    summary_lines.append(f"Generated: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}")
    summary_lines.append("=" * 80)

    # OpenAQ Summary
    summary_lines.append("\n1. OPENAQ DATA")
    summary_lines.append("-" * 80)

    if not openaq_df.empty:
        summary_lines.append(f"Total Records: {len(openaq_df):,}")
        summary_lines.append(f"\nDate Range: {openaq_df['date'].min()} to {openaq_df['date'].max()}")

        summary_lines.append("\nRecords by City:")
        for city in openaq_df['city'].unique():
            city_data = openaq_df[openaq_df['city'] == city]
            summary_lines.append(f"  - {city}: {len(city_data):,} measurements")
            if not city_data.empty:
                summary_lines.append(f"    Date range: {city_data['date'].min()} to {city_data['date'].max()}")

        summary_lines.append("\nRecords by Parameter:")
        for param in openaq_df['parameter'].unique():
            param_data = openaq_df[openaq_df['parameter'] == param]
            summary_lines.append(f"  - {param.upper()}: {len(param_data):,} measurements")
    else:
        summary_lines.append("No OpenAQ data acquired.")

    # World Bank Summary
    summary_lines.append("\n\n2. WORLD BANK DATA")
    summary_lines.append("-" * 80)

    if not wb_df.empty:
        summary_lines.append(f"Total Records: {len(wb_df):,}")

        wb_df_clean = wb_df[wb_df['value'].notna()]
        if not wb_df_clean.empty:
            summary_lines.append(f"Year Range: {wb_df_clean['year'].min()} to {wb_df_clean['year'].max()}")

        summary_lines.append("\nRecords by Indicator:")
        for indicator in wb_df['indicator'].unique():
            ind_data = wb_df[wb_df['indicator'] == indicator]
            ind_data_clean = ind_data[ind_data['value'].notna()]
            summary_lines.append(f"  - {indicator}")
            summary_lines.append(f"    Total: {len(ind_data_clean):,} yearly records")
            if not ind_data_clean.empty:
                summary_lines.append(f"    Years: {ind_data_clean['year'].min()} - {ind_data_clean['year'].max()}")
    else:
        summary_lines.append("No World Bank data acquired.")

    # Overall assessment
    summary_lines.append("\n\n3. DATA COVERAGE ASSESSMENT")
    summary_lines.append("-" * 80)

    if not openaq_df.empty:
        openaq_start = pd.to_datetime(openaq_df['date']).min().year
        summary_lines.append(f"OpenAQ earliest data: {openaq_start}")

    if not wb_df.empty:
        wb_df_clean = wb_df[wb_df['value'].notna()]
        if not wb_df_clean.empty:
            wb_start = int(wb_df_clean['year'].min())
            summary_lines.append(f"World Bank earliest data: {wb_start}")

    summary_lines.append("\n" + "=" * 80)

    # Save summary
    summary_text = "\n".join(summary_lines)
    with open(output_file, 'w') as f:
        f.write(summary_text)

    print(summary_text)
    print(f"\n✓ Summary saved to: {output_file}")


def main():
    """Main execution function."""
    print("\n" + "=" * 80)
    print("AIR POLLUTION DATA ACQUISITION (OpenAQ v3 + World Bank)")
    print("=" * 80)
    print(f"Start time: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}")

    # Ensure output directories exist
    OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
    RESULTS_DIR.mkdir(parents=True, exist_ok=True)

    # Fetch OpenAQ data
    openaq_df = fetch_openaq_data_all_india(API_KEY)

    if not openaq_df.empty:
        openaq_file = OUTPUT_DIR / "openaq_data.csv"
        openaq_df.to_csv(openaq_file, index=False)
        print(f"\n✓ OpenAQ data saved to: {openaq_file}")
        print(f"  Total records: {len(openaq_df):,}")
    else:
        print("\n⚠ No OpenAQ data to save")
        openaq_df = pd.DataFrame()

    # Fetch World Bank data
    wb_df = fetch_world_bank_data(WB_COUNTRY_CODE, WB_INDICATORS)
    if not wb_df.empty:
        wb_file = OUTPUT_DIR / "historical_macro_data.csv"
        wb_df.to_csv(wb_file, index=False)
        print(f"\n✓ World Bank data saved to: {wb_file}")
        print(f"  Total records: {len(wb_df):,}")
    else:
        print("\n⚠ No World Bank data to save")

    # Generate summary
    summary_file = RESULTS_DIR / "data_acquisition_summary.txt"
    generate_summary(openaq_df, wb_df, summary_file)

    print("\n" + "=" * 80)
    print("DATA ACQUISITION COMPLETE")
    print(f"End time: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}")
    print("=" * 80)


if __name__ == "__main__":
    main()
