import requests
import pandas as pd
from datetime import datetime
from io import StringIO
from datetime import timedelta

################## Keep datapoints only at 3 hour time intervals (to match with wrf)
# Calculate difference from nearest 3-hour UTC time
def nearest_three_hour_anchor(ts):
    floored = ts.replace(minute=0, second=0, microsecond=0)
    hour_mod = floored.hour % 3
    before = floored - timedelta(hours=hour_mod)
    after = before + timedelta(hours=3)

    # pick whichever is closer
    return before if abs(ts - before) <= abs(ts - after) else after

# Filter to get mean wind speed for station within time window
def filter_and_mean(df, buffer_minutes=10):
    df = df.copy()
    df["valid"] = pd.to_datetime(df["valid"])

    # Find anchor and filter
    df["anchor_time"] = df["valid"].apply(nearest_three_hour_anchor)
    mask = (
        (abs(df["valid"] - df["anchor_time"]) <= pd.Timedelta(minutes=buffer_minutes))
    )
    df = df[mask]

    # Keep mean wind per station per anchor time
    df = (
        df.groupby(["station", "anchor_time"], as_index=False)
        .agg({"lon": "first", "lat": "first", "sped": "mean"})
    )

    return df.reset_index(drop=True)

############### Function to request and retrieve asos data for 24 hr period ####################
def get_asos_wind_data(start: str, end: str, variable: str, states: list) -> pd.DataFrame:
    """
    Fetch wind speed data (or other variables) from the ASOS network for selected states.

    Args:
        start (str): Start datetime in 'YYYY-MM-DD HH:MM' format.
        end (str): End datetime in 'YYYY-MM-DD HH:MM' format.
        variable (str): Variable to fetch, e.g., 'sped' for wind speed.
        states (list): List of U.S. state abbreviations, e.g., ['IA', 'MN'].

    Returns:
        pd.DataFrame: DataFrame with columns: station, valid, lon, lat, elevation, sped (in m/s).
    """
    base_url = "https://mesonet.agron.iastate.edu/cgi-bin/request/asos.py"

    # Convert start/end to datetime objects
    start_dt = datetime.strptime(start, "%Y-%m-%d %H:%M")
    end_dt = datetime.strptime(end, "%Y-%m-%d %H:%M")

    # Build the request payload
    payload = {
        "data": variable,
        "tz": "Etc/UTC",
        "format": "comma",
        "latlon": "yes",
        "missing": "empty",
        "trace": "empty",
        "direct": "no",
        "report_type": "3",
        "year1": start_dt.year,
        "month1": start_dt.month,
        "day1": start_dt.day,
        "hour1": start_dt.hour,
        "minute1": start_dt.minute,
        "year2": end_dt.year,
        "month2": end_dt.month,
        "day2": end_dt.day,
        "hour2": end_dt.hour,
        "minute2": end_dt.minute,
    }

    # Add one line per state network
    for state in states:
        payload.setdefault("state", []).append(state.upper())

    # Send the request
    response = requests.post(base_url, data=payload)
    if response.status_code != 200:
        raise Exception(f"Failed to fetch data: {response.status_code}")

    
    # Parse CSV, skipping debug lines
    df = pd.read_csv(StringIO(response.text), comment="#")

    # Keep only required columns
    df = df[["station", "valid", "lon", "lat", variable]]

    # Rename for consistency
    df.columns = ["station", "valid", "lon", "lat", "sped_mph"]

    # Convert wind speed from mph to m/s
    df["sped"] = df["sped_mph"] * 0.44704

    # Drop original mph column
    df = df.drop(columns=["sped_mph"])

    # Drop rows with missing wind speed
    df = df.dropna(subset=["sped"])

    # Convert valid column to datetime
    df["valid"] = pd.to_datetime(df["valid"])

    # Apply filtering
    df["valid"] = pd.to_datetime(df["valid"])
    print(df)
    df = filter_and_mean(df)
    print(df)

    # Only keep rows with wind speed > 2 m/s
    df = df[df["sped"] > 2.0]

    return df

############### Wrapper function to get data for time periods longer than 24 hrs ###################
def get_asos_wind_data_range(start: str, end: str, variable: str, states: list) -> pd.DataFrame:
    """
    Wrapper to fetch ASOS wind data in 1-day increments over a longer date range.

    Args:
        start (str): Start datetime in 'YYYY-MM-DD HH:MM' format.
        end (str): End datetime in 'YYYY-MM-DD HH:MM' format.
        variable (str): Variable to fetch, e.g., 'sped'.
        states (list): List of U.S. state abbreviations.

    Returns:
        pd.DataFrame: Concatenated and filtered data over the full date range.
    """
    start_dt = datetime.strptime(start, "%Y-%m-%d %H:%M")
    end_dt = datetime.strptime(end, "%Y-%m-%d %H:%M")
    current = start_dt
    results = []

    while current < end_dt:
        print(f"Retrieving data for: {current}")
        next_day = min(current + timedelta(days=1), end_dt)
        try:
            df = get_asos_wind_data( # call first function
                start=current.strftime("%Y-%m-%d %H:%M"),
                end=next_day.strftime("%Y-%m-%d %H:%M"),
                variable=variable,
                states=states
            )
            if not df.empty:
                results.append(df)
        except Exception as e:
            print(f"Failed to fetch data for {current} → {next_day}: {e}")
        current = next_day

    # Combine all DataFrames
    return pd.concat(results, ignore_index=True) if results else pd.DataFrame()


states1=['TX','LA','MS','AL','FL','GA','TN','AR','OK','NM','AZ','CO','WY','UT','MT','KS','NE','SD','ND','WI','MN','IA','IL','MO','MI','IN','OH','KY']

for state in states1:
    df = get_asos_wind_data_range(
        start="2017-04-01 00:00",
        end="2017-04-30 23:59",
        variable="sped",
        states=state
    )

    df.to_csv(f'/pscratch/sd/d/dbrooks/acc2017_analysis/model_evaluation/asos_data/04/{state}_wind_mean_04.csv', index=True)
    print(f'Done: {state}')
print('Done April')

for state in states1:
    df = get_asos_wind_data_range(
        start="2017-05-01 00:00",
        end="2017-05-31 23:59",
        variable="sped",
        states=state
    )

    df.to_csv(f'/pscratch/sd/d/dbrooks/acc2017_analysis/model_evaluation/asos_data/05/{state}_wind_mean_05.csv', index=True)
    print(f'Done: {state}')
print('Done May')

for state in states1:
    df = get_asos_wind_data_range(
        start="2017-06-01 00:00",
        end="2017-06-30 23:59",
        variable="sped",
        states=state
    )

    df.to_csv(f'/pscratch/sd/d/dbrooks/acc2017_analysis/model_evaluation/asos_data/06/{state}_wind_mean_06.csv', index=True)
    print(f'Done: {state}')
print('Done June')