import os
import json
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
import geopandas as gpd
from shapely.geometry import Polygon, MultiPolygon, GeometryCollection

os.makedirs('pred_results', exist_ok=True)

towns = gpd.read_file('dataset/ShikokuMetropolitan.geojson')
roads = gpd.read_file('dataset/AllSeasonRoads.geojson')
pop = gpd.read_file('dataset/ShikokuPopulation.geojson')

rural_mask = towns['AREATYPE'].astype(str).str.strip().str.upper() == 'RURAL'
rural_towns = towns[rural_mask].copy()
rural_union = rural_towns.unary_union

def to_polygonal(geom):
    if geom is None or geom.is_empty:
        return Polygon()
    if isinstance(geom, (Polygon, MultiPolygon)):
        return geom
    if isinstance(geom, GeometryCollection):
        polys = [g for g in geom.geoms if isinstance(g, (Polygon, MultiPolygon))]
        if not polys:
            return Polygon()
        if len(polys) == 1:
            return polys[0]
        return MultiPolygon(polys)
    return Polygon()

if rural_union is None or rural_union.is_empty:
    pop_rural = gpd.GeoDataFrame(columns=pop.columns, crs=pop.crs)
else:
    pop_rural = pop[pop.intersects(rural_union)].copy()
    pop_rural['geometry'] = pop_rural.intersection(rural_union).apply(to_polygonal)
    pop_rural = pop_rural[~pop_rural.is_empty].copy()
    pop_rural = pop_rural[pop_rural.geom_type.isin(['Polygon', 'MultiPolygon'])].copy()

pop_rural['D0001'] = pd.to_numeric(pop_rural['D0001'], errors='coerce').fillna(0.0)
pop_rural['area_total'] = pop_rural.area

road_buffer = roads.buffer(2000).unary_union

if road_buffer is None or road_buffer.is_empty:
    pop_rural['area_access'] = 0.0
else:
    pop_rural['area_access'] = pop_rural.intersection(road_buffer).area

with np.errstate(divide='ignore', invalid='ignore'):
    pop_rural['access_fraction'] = np.where(
        pop_rural['area_total'] > 0,
        pop_rural['area_access'] / pop_rural['area_total'],
        0.0
    )
pop_rural['accessible_pop'] = pop_rural['D0001'] * pop_rural['access_fraction']

rural_population = float(pop_rural['D0001'].sum())
accessible_population = float(pop_rural['accessible_pop'].sum())
access_percent = float((accessible_population / rural_population) * 100.0) if rural_population > 0 else 0.0
rural_feature_count = int(len(pop_rural))
accessible_feature_count = int((pop_rural['access_fraction'] > 0).sum())

with open('pred_results/accessibility.json', 'w') as f:
    json.dump({
        'rural_population': rural_population,
        'accessible_population': accessible_population,
        'access_percent': access_percent,
        'rural_feature_count': rural_feature_count,
        'accessible_feature_count': accessible_feature_count
    }, f, indent=2)

fig, ax = plt.subplots(figsize=(12, 10))
if not pop_rural.empty:
    pop_rural.plot(
        column='access_fraction',
        cmap='YlGn',
        linewidth=0.5,
        ax=ax,
        edgecolor='0.8',
        legend=True,
        legend_kwds={'label': 'Access Fraction', 'orientation': 'horizontal', 'pad': 0.02}
    )
if not roads.empty:
    roads.plot(ax=ax, color='darkred', linewidth=0.6, alpha=0.7)

ax.set_title('Rural Area Accessibility to All-Season Roads (2 km buffer)')
ax.set_axis_off()
plt.tight_layout()
plt.savefig('pred_results/accessibility.png', dpi=300, bbox_inches='tight')
plt.close()