"""
02_prepare_graph.py
===================
Prepara una sola volta il grafo metrico + lo snap degli edifici.

Cosa fa:
  1. Carica osm_walk_parma.graphml (output di 01).
  2. Proietta il grafo in CRS metrico (EPSG:32632) e RIPROIETTA anche le
     geometrie degli edge (osmnx 2.x non lo fa di suo, bug noto: vedi note
     in dash/script/03_compute_paths.py della pipeline storica).
  3. Per ogni edificio (44k) calcola il baricentro, trova l'arco OSM piu'
     vicino (`ox.distance.nearest_edges`) e inserisce un NODO VIRTUALE
     "virt_b_<i>" nel punto proiettato sull'arco. L'arco originale viene
     spezzato in due sotto-archi che conservano la geometria reale.
  4. Salva in cache:
       _cache/graph_prepared.pkl         (MultiDiGraph con nodi virtuali)
       _cache/graph_simple.pkl           (DiGraph "ridotto" per Dijkstra veloce)
       _cache/buildings_meta.pkl         (DataFrame: fid, virt_node, snap_dist,
                                          entry_cost, centroid_geom_metric, ...)

Tempo stimato: 3-8 minuti (la chiamata pesante e' nearest_edges su 44k punti).
Questa preparazione vale per TUTTE le categorie POI -> ammortizziamo il costo.
"""
import os
import time
import pickle

import numpy as np
import pandas as pd
import geopandas as gpd
import networkx as nx
import osmnx as ox
from pyproj import Transformer
from shapely.geometry import LineString, Point
from shapely.ops import substring

import config as C


PREPARED_GRAPH_PKL = os.path.join(C.CACHE_DIR, "graph_prepared.pkl")
SIMPLE_GRAPH_PKL   = os.path.join(C.CACHE_DIR, "graph_simple.pkl")
BUILDINGS_META_PKL = os.path.join(C.CACHE_DIR, "buildings_meta.pkl")


# -----------------------------------------------------------------------------
# Helper geometrici (riusati anche in 03)
# -----------------------------------------------------------------------------
def project_point_on_edge(G, pt_xy, u, v, k):
    """Proietta pt_xy sull'edge (u,v,k) orientato u->v.
    Restituisce (proj_xy, s_da_u, total_len, edge_geom_oriented).
    """
    data = G.get_edge_data(u, v, k)
    geom = data.get("geometry")
    if geom is None:
        nu = G.nodes[u]; nv = G.nodes[v]
        geom = LineString([(nu["x"], nu["y"]), (nv["x"], nv["y"])])
    coords = list(geom.coords)
    nu = G.nodes[u]; nv = G.nodes[v]
    d_first_to_u = (coords[0][0] - nu["x"])**2 + (coords[0][1] - nu["y"])**2
    d_first_to_v = (coords[0][0] - nv["x"])**2 + (coords[0][1] - nv["y"])**2
    if d_first_to_v < d_first_to_u:
        geom = LineString(coords[::-1])
    pt = Point(pt_xy)
    s = geom.project(pt)
    proj = geom.interpolate(s)
    return (proj.x, proj.y), s, geom.length, geom


def slice_edge_geometry(edge_geom, s_start, s_end):
    """Estrae sotto-LineString da s_start a s_end lungo edge_geom."""
    try:
        sub = substring(edge_geom, s_start, s_end)
        if isinstance(sub, Point):
            return LineString([(sub.x, sub.y), (sub.x, sub.y)])
        return sub
    except Exception:
        p1 = edge_geom.interpolate(s_start)
        p2 = edge_geom.interpolate(s_end)
        return LineString([(p1.x, p1.y), (p2.x, p2.y)])


def reproject_edge_geometries(G, src_crs="EPSG:4326", dst_crs=None):
    """ox.project_graph() proietta x/y dei nodi ma NON gli attributi
    geometry degli edge. Senza questo fix il risultato finale ha
    LineString in WGS84 mescolate con metri UTM. -> riproietto a mano."""
    if dst_crs is None:
        dst_crs = C.METRIC_CRS
    transformer = Transformer.from_crs(src_crs, dst_crs, always_xy=True)
    n_fixed = 0
    for u, v, k, data in G.edges(keys=True, data=True):
        geom = data.get("geometry")
        if geom is None:
            nu = G.nodes[u]; nv = G.nodes[v]
            data["geometry"] = LineString([(nu["x"], nu["y"]),
                                           (nv["x"], nv["y"])])
            continue
        coords = list(geom.coords)
        if not coords:
            continue
        x0, y0 = coords[0]
        # se sembra WGS84 (lon/lat) -> riproietto
        if -180 <= x0 <= 180 and -90 <= y0 <= 90:
            xs, ys = zip(*coords)
            xr, yr = transformer.transform(xs, ys)
            data["geometry"] = LineString(list(zip(xr, yr)))
            n_fixed += 1
    return n_fixed


def main():
    t0 = time.time()
    C.log("=== 02 - Preparazione grafo + snap edifici ===", "==")

    # ------------------------------------------------------------------
    # Skip se gia' in cache
    # ------------------------------------------------------------------
    if (os.path.exists(PREPARED_GRAPH_PKL)
            and os.path.exists(SIMPLE_GRAPH_PKL)
            and os.path.exists(BUILDINGS_META_PKL)):
        C.log("Cache gia' presente:")
        C.log(f"  {PREPARED_GRAPH_PKL}")
        C.log(f"  {SIMPLE_GRAPH_PKL}")
        C.log(f"  {BUILDINGS_META_PKL}")
        C.log("Skippo. Elimina i file per forzare la ricostruzione.")
        return

    # ------------------------------------------------------------------
    # 1) Carico grafo OSM e proietto in metrico
    # ------------------------------------------------------------------
    C.log(f"Carico grafo: {C.OSM_GRAPHML_FILE}")
    G = ox.load_graphml(C.OSM_GRAPHML_FILE)
    C.log(f"  nodi={G.number_of_nodes()}  archi={G.number_of_edges()}")

    C.log(f"Proietto grafo in {C.METRIC_CRS}")
    G = ox.project_graph(G, to_crs=C.METRIC_CRS)

    C.log("Fix riproiezione geometrie edge (WGS84 -> UTM)...")
    n_fixed = reproject_edge_geometries(G, "EPSG:4326", C.METRIC_CRS)
    C.log(f"  riproiettate {n_fixed} geometrie edge")

    # ------------------------------------------------------------------
    # 2) Carico edifici e calcolo baricentri
    # ------------------------------------------------------------------
    C.log(f"Carico edifici: layer '{C.FOOTPRINTS_LAYER}'")
    fp = gpd.read_file(C.SOURCE_GPKG, layer=C.FOOTPRINTS_LAYER).to_crs(C.METRIC_CRS)

    # ------------------------------------------------------------------
    # FILTRO QUARTIERE (San Leonardo)
    # Se QUARTIERE_GPKG e' configurato, tengo solo gli edifici il cui
    # CENTROIDE cade dentro il poligono del quartiere. Stesso criterio
    # geometrico che usiamo per il baricentro -> coerente.
    # ------------------------------------------------------------------
    quart_poly = C.load_quartiere_polygon()
    if quart_poly is not None:
        n_before = len(fp)
        centroids_pre = fp.geometry.centroid
        mask = centroids_pre.within(quart_poly)
        fp = fp[mask].reset_index(drop=True)
        C.log(f"  filtro quartiere edifici: {n_before} -> {len(fp)} "
              f"(scartati {n_before - len(fp)} fuori dal quartiere)")
        if len(fp) == 0:
            raise RuntimeError("Nessun edificio dentro il quartiere: "
                               "verifica QUARTIERE_GPKG / QUARTIERE_LAYER.")

    # campo id: preferisco 'fid' (presente nella maggior parte dei GPKG)
    if "fid" in fp.columns:
        id_field = "fid"
    elif "id" in fp.columns:
        id_field = "id"
    else:
        fp["_autoid"] = range(len(fp))
        id_field = "_autoid"
    C.log(f"  campo id edifici: '{id_field}'")

    centroids = fp.geometry.centroid
    fids = fp[id_field].tolist()
    n_b = len(fp)
    C.log(f"  edifici: {n_b}")

    # ------------------------------------------------------------------
    # 3) nearest_edges in batch (una sola chiamata vettoriale)
    # ------------------------------------------------------------------
    C.log("nearest_edges: trovo l'arco piu' vicino per ogni edificio...")
    xs = np.array([p.x for p in centroids])
    ys = np.array([p.y for p in centroids])
    t1 = time.time()
    near = ox.distance.nearest_edges(G, X=xs, Y=ys)
    C.log(f"  fatto in {time.time()-t1:.1f} s")

    # ------------------------------------------------------------------
    # 4) Inserisco i nodi virtuali "virt_b_<i>" nel grafo
    # ------------------------------------------------------------------
    C.log("Inserisco nodi virtuali sugli archi (uno per edificio)...")
    virt_nodes = []
    snap_dists = []
    proj_points = []  # tuple (x,y) del punto proiettato sull'arco
    n_warn = 0
    for i in range(n_b):
        u, v, k = near[i]
        proj_xy, s, total_len, edge_geom = project_point_on_edge(
            G, (xs[i], ys[i]), u, v, k
        )
        snap_d = ((proj_xy[0] - xs[i])**2 + (proj_xy[1] - ys[i])**2) ** 0.5
        if snap_d > C.SNAP_WARN_DIST_M:
            n_warn += 1

        vname = f"virt_b_{i}"
        # PATCH-LUCA most-used-fix (06/2026): salvo come attributo del nodo
        # virtuale i due nodi OSM reali (u, v) dell'arco di "ospitalita'", cosi'
        # in 03 posso ricondurre gli archi (u, virt) e (virt, v) all'arco OSM
        # originale (u, v) e contare correttamente fid_count su TUTTE le strade
        # percorse, anche quelle dove il nodo virtuale spezza un arco.
        G.add_node(vname, x=proj_xy[0], y=proj_xy[1],
                   host_u=u, host_v=v)

        geom_u2virt = slice_edge_geometry(edge_geom, 0.0, s)
        geom_virt2v = slice_edge_geometry(edge_geom, s, total_len)
        # archi bidirezionali (pedonale)
        G.add_edge(u, vname, length=float(s),
                   geometry=geom_u2virt)
        G.add_edge(vname, u, length=float(s),
                   geometry=LineString(list(geom_u2virt.coords)[::-1]))
        G.add_edge(vname, v, length=float(total_len - s),
                   geometry=geom_virt2v)
        G.add_edge(v, vname, length=float(total_len - s),
                   geometry=LineString(list(geom_virt2v.coords)[::-1]))

        virt_nodes.append(vname)
        snap_dists.append(float(snap_d))
        proj_points.append(proj_xy)

    C.log(f"  inseriti {len(virt_nodes)} nodi virtuali edifici")
    C.log(f"  snap_dist medio: {np.mean(snap_dists):.1f} m  "
          f"max: {np.max(snap_dists):.1f} m")
    if n_warn:
        C.log(f"  ATTENZIONE: {n_warn} edifici distano > {C.SNAP_WARN_DIST_M:.0f} m"
              f" dall'arco OSM piu' vicino")

    # ------------------------------------------------------------------
    # 5) DataFrame meta edifici
    # ------------------------------------------------------------------
    meta = pd.DataFrame({
        "i":           range(n_b),
        "fid":         fids,
        "virt_node":   virt_nodes,
        "entry_cost":  snap_dists,
        "centroid_x":  xs,
        "centroid_y":  ys,
        "proj_x":      [p[0] for p in proj_points],
        "proj_y":      [p[1] for p in proj_points],
    })

    # ------------------------------------------------------------------
    # 6) Grafo "ridotto" DiGraph per Dijkstra rapido
    #    (un edge per (u,v), il piu' corto in length)
    # ------------------------------------------------------------------
    C.log("Costruisco DiGraph ridotto per Dijkstra veloce...")
    H = nx.DiGraph()
    for u, v, k, d in G.edges(keys=True, data=True):
        w = d.get("length", None)
        if w is None:
            continue
        if H.has_edge(u, v):
            if w < H[u][v]["weight"]:
                H[u][v]["weight"] = w
                H[u][v]["key"]    = k
        else:
            H.add_edge(u, v, weight=float(w), key=k)
    for n in G.nodes:
        if n not in H:
            H.add_node(n)
    C.log(f"  DiGraph ridotto: {H.number_of_nodes()} nodi, "
          f"{H.number_of_edges()} archi")

    # ------------------------------------------------------------------
    # 7) Salvataggio cache
    # ------------------------------------------------------------------
    C.log("Salvo cache...")
    with open(PREPARED_GRAPH_PKL, "wb") as f:
        pickle.dump(G, f, protocol=pickle.HIGHEST_PROTOCOL)
    with open(SIMPLE_GRAPH_PKL, "wb") as f:
        pickle.dump(H, f, protocol=pickle.HIGHEST_PROTOCOL)
    meta.to_pickle(BUILDINGS_META_PKL)
    C.log(f"  {PREPARED_GRAPH_PKL}")
    C.log(f"  {SIMPLE_GRAPH_PKL}")
    C.log(f"  {BUILDINGS_META_PKL}")

    C.log(f"FATTO in {time.time()-t0:.1f} s", "==")


if __name__ == "__main__":
    main()
