import requests
import gzip
import io
from netCDF4 import Dataset
import matplotlib.pyplot as plt
import matplotlib.colors as mcolors
import matplotlib.colors as Normalize
import cartopy.crs as ccrs
import cartopy.feature as cfeature
from cartopy.mpl.ticker import LongitudeFormatter, LatitudeFormatter
import matplotlib.ticker as mticker
import matplotlib as mpl
import xarray as xr
import netCDF4
import numpy as np
import pandas as pd
import glob
import dask
import os
from matplotlib.colors import LinearSegmentedColormap, TwoSlopeNorm
from metpy.plots import colortables, USCOUNTIES
from datetime import datetime, timedelta

def download_and_extract_gz(url):
    """Downloads and extracts a .gz file from a given URL."""
    response = requests.get(url)
    if response.status_code == 200:
        with gzip.GzipFile(fileobj=io.BytesIO(response.content)) as gz_file:
            return gz_file.read()
    else:
        raise Exception(f"Failed to download file. HTTP status code: {response.status_code}")

def read_grib_data(file_content):
    """Reads GRIB data from the decompressed file content."""
    temp_filename = 'temporary.grib'
    with open(temp_filename, 'wb') as temp_file:
        temp_file.write(file_content)
    return xr.load_dataset(temp_filename, engine="cfgrib")

def select_region(bounding_box, current_ds, nexrad=True):
        min_lon,min_lat,max_lon,max_lat = bounding_box[0], bounding_box[1], bounding_box[2], bounding_box[3]

        # Access the latitude and longitude arrays (XLAT, XLONG)
        if nexrad==False:
            lats = current_ds['XLAT']
            lons = current_ds['XLONG']
            # Create a boolean mask for the region of interest
            region_mask = (lats >= min_lat) & (lats <= max_lat) & (lons >= min_lon) & (lons <= max_lon)
            #print(region_mask)
            # Subset the data using the bounding box
            current_ds = current_ds.where(region_mask, drop=True)

        else:
            current_ds['longitude'] = ((current_ds['longitude'] + 180) % 360) - 180
            current_ds = current_ds.sortby('longitude')
            current_ds = current_ds.sortby('latitude')
            current_ds = current_ds.sel(longitude=slice(bounding_box[0], bounding_box[2]), latitude=slice(bounding_box[1], bounding_box[3]))


        return current_ds

bounding_box=[-112.99227905,26.33052444,-82.00772095,46.40573883] # to select only the region we want

def download_mrms_data(start_date, end_date):
    """Downloads MRMS precip data for all days within the specified range."""
    current_date = start_date
    hours = np.arange(0, 24, 1)
    
    datasets=[]
    while current_date <= end_date:
        for hour in hours:
            url = (f"https://mtarchive.geol.iastate.edu/{current_date:%Y/%m/%d}/mrms/ncep/"
                   f"GaugeCorr_QPE_01H/GaugeCorr_QPE_01H_00.00_"
                   f"{current_date:%Y%m%d}-{hour:02}0000.grib2.gz")
            
            try:
                #print(f"Downloading: {url}")
                file_content = download_and_extract_gz(url)
                ds = read_grib_data(file_content)
                ds = select_region(bounding_box, ds, nexrad=True)
                #print(ds)
                ds.to_netcdf(f'/pscratch/sd/d/dbrooks/acc2017_analysis/model_evaluation/mrms_data/0{start_date.month}/mrms_gauge_corrected_{current_date:%Y%m%d}-{hour:02}_.nc')
                #datasets.append(ds)
                print(f"Successfully downloaded and read data for {current_date} hour {hour}")
            except Exception as e:
                print(f"Error processing {current_date} hour {hour}: {e}")
        
        current_date += timedelta(days=1)
        print(current_date)

    #combined_ds = xr.concat(datasets, dim='time')
    #combined_ds.to_netcdf(f'/pscratch/sd/d/dbrooks/acc2017_analysis/model_evaluation/mrms_gauge_corr_month{start_date.month}.nc')
    print('done')

"""
# Example usage
start_date = datetime(2017, 4, 1)
end_date = datetime(2017, 4, 30)
download_mrms_data(start_date, end_date)
"""

start_date = datetime(2017, 5, 1)
end_date = datetime(2017, 5, 31)
download_mrms_data(start_date, end_date)

start_date = datetime(2017, 6, 1)
end_date = datetime(2017, 6, 30)
download_mrms_data(start_date, end_date)