import os
import json
import numpy as np
import pandas as pd
import geopandas as gpd
import matplotlib.pyplot as plt

TOWNS_PATH = "dataset/ShikokuMetropolitan.geojson"
ROADS_PATH = "dataset/AllSeasonRoads.geojson"
POP_PATH = "dataset/ShikokuPopulation.geojson"
OUT_DIR = "pred_results"
PNG_PATH = os.path.join(OUT_DIR, "accessibility.png")
JSON_PATH = os.path.join(OUT_DIR, "accessibility.json")

os.makedirs(OUT_DIR, exist_ok=True)

CRS = "EPSG:4087"
BUFFER_DIST = 2000.0


def ensure_valid(gdf):
    gdf = gdf.copy()
    try:
        valid_geom = gdf.geometry.make_valid()
        return gdf.set_geometry(valid_geom)
    except Exception:
        try:
            from shapely import make_valid

            geom = gdf.geometry.apply(lambda g: make_valid(g) if not g.is_valid else g)
            return gdf.set_geometry(geom)
        except Exception:
            return gdf


# Load data
towns = gpd.read_file(TOWNS_PATH).to_crs(CRS)
roads = gpd.read_file(ROADS_PATH).to_crs(CRS)
pop = gpd.read_file(POP_PATH).to_crs(CRS)

# Clean geometries
towns = ensure_valid(towns)
roads = ensure_valid(roads)
pop = ensure_valid(pop)

# Filter rural towns
rural = towns[towns["AREATYPE"].astype(str).str.strip().str.upper() == "RURAL"].copy()
if rural.empty:
    raise ValueError("No rural areas found in the towns dataset.")

rural_union = rural.unary_union

# Prepare population data
pop["D0001"] = pd.to_numeric(pop["D0001"], errors="coerce").fillna(0.0)

# Clip population polygons to rural areas
rural_pop = gpd.clip(pop, rural_union, keep_geom_type=False)
rural_pop = rural_pop[rural_pop.geometry.type.isin(["Polygon", "MultiPolygon"])].copy()
rural_pop = ensure_valid(rural_pop)

# Drop zero-area features
rural_pop["area_total"] = rural_pop.geometry.area
rural_pop = rural_pop[rural_pop["area_total"] > 0].copy()

if rural_pop.empty:
    result = {
        "rural_population": 0.0,
        "accessible_population": 0.0,
        "access_percent": 0.0,
        "rural_feature_count": 0,
        "accessible_feature_count": 0,
    }
    with open(JSON_PATH, "w") as f:
        json.dump(result, f, indent=2)

    fig, ax = plt.subplots(figsize=(10, 10))
    ax.set_title("Rural road accessibility (no rural population areas)")
    ax.axis("off")
    plt.savefig(PNG_PATH, dpi=300, bbox_inches="tight")
    plt.close()
else:
    # Dissolved 2 km buffer around all-season roads
    road_buffer_geom = roads.geometry.buffer(BUFFER_DIST).unary_union
    road_buffer_geom = road_buffer_geom.buffer(0)  # clean

    # Area of each rural population polygon within the road buffer
    access_geom = rural_pop.geometry.intersection(road_buffer_geom)
    rural_pop["area_access"] = access_geom.area
    rural_pop["access_ratio"] = np.clip(
        rural_pop["area_access"] / rural_pop["area_total"], 0.0, 1.0
    )
    rural_pop["access_pop"] = rural_pop["D0001"] * rural_pop["access_ratio"]
    rural_pop["access_percent"] = rural_pop["access_ratio"] * 100.0

    # Aggregate statistics
    rural_population = float(rural_pop["D0001"].sum())
    accessible_population = float(rural_pop["access_pop"].sum())
    access_percent = (
        (accessible_population / rural_population * 100.0)
        if rural_population > 0
        else 0.0
    )
    rural_feature_count = int(len(rural_pop))
    accessible_feature_count = int((rural_pop["access_pop"] > 0).sum())

    result = {
        "rural_population": rural_population,
        "accessible_population": accessible_population,
        "access_percent": access_percent,
        "rural_feature_count": rural_feature_count,
        "accessible_feature_count": accessible_feature_count,
    }
    with open(JSON_PATH, "w") as f:
        json.dump(result, f, indent=2)

    # Choropleth map
    fig, ax = plt.subplots(figsize=(12, 12))
    rural_pop.plot(
        column="access_percent",
        cmap="YlGn",
        linewidth=0.5,
        edgecolor="darkgrey",
        legend=True,
        legend_kwds={
            "label": "Percent of rural population within 2 km of all-season roads",
            "shrink": 0.5,
        },
        vmin=0,
        vmax=100,
        ax=ax,
    )
    roads.plot(ax=ax, color="black", linewidth=0.7, alpha=0.6, label="All-season roads")
    ax.set_title("Rural Road Accessibility in Shikoku")
    ax.axis("off")
    plt.tight_layout()
    plt.savefig(PNG_PATH, dpi=300, bbox_inches="tight")
    plt.close()

print(json.dumps(result, indent=2))