"""
03_compute_routes.py
====================
CUORE DELLA PIPELINE.

Per ogni categoria POI esegue:
  1. Carica i POI (da layer GPKG o da CSV)
  2. Snap dei POI sui rispettivi archi OSM con nodi virtuali "virt_p_<i>"
  3. MULTI-SOURCE DIJKSTRA: parte da TUTTI i POI contemporaneamente.
     In una sola chiamata, per ogni nodo del grafo (= ogni edificio
     virtuale "virt_b_<i>") otteniamo:
       - distanza pedonale al POI piu' vicino
       - lista dei nodi del cammino
       - implicitamente quale POI ha "vinto"
  4. Conta gli archi piu' usati (fid_count) sommando tutti i cammini.
  5. Calcola per ogni POI: quanti edifici servono, dist media, dist max.
  6. Scrive output/routes_<categoria>.gpkg con 5 layer:
       - shortest_path                   (1 LineString per edificio)
       - most_used_path                  (archi OSM con fid_count > 0)
       - footprints_with_shortest_path   (poligoni edifici + costi joinati)
       - pois                            (i POI della categoria + statistiche)
       - origin_points                   (baricentri edifici + assigned_poi_id)

Complessita': O((V+E) log V) per categoria, indipendente dal numero di
edifici. Tipico: 1-3 minuti per categoria.
"""
import os
import sys
import time
import pickle
from collections import defaultdict, Counter

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

import config as C
# Riuso gli helper di 02
sys.path.insert(0, C.PIPELINE_DIR)
prep = __import__("02_prepare_graph")
PREPARED_GRAPH_PKL = prep.PREPARED_GRAPH_PKL
SIMPLE_GRAPH_PKL   = prep.SIMPLE_GRAPH_PKL
BUILDINGS_META_PKL = prep.BUILDINGS_META_PKL
project_point_on_edge = prep.project_point_on_edge
slice_edge_geometry   = prep.slice_edge_geometry


# -----------------------------------------------------------------------------
# Caricamento POI di una categoria
# -----------------------------------------------------------------------------
def load_pois(category):
    """Restituisce un GeoDataFrame in CRS metrico con almeno [geometry, NOME]
    + colonna 'poi_id' interno (0..N-1)."""
    src = category["source"]
    name = category["name"]

    if src == "gpkg_layer":
        gdf = gpd.read_file(C.SOURCE_GPKG, layer=category["layer"])
    elif src == "csv":
        path = os.path.join(C.DATA_DIR, category["csv"])
        df = pd.read_csv(path)
        if "lon" not in df.columns or "lat" not in df.columns:
            raise RuntimeError(f"CSV {path} senza colonne lon/lat")
        gdf = gpd.GeoDataFrame(
            df,
            geometry=gpd.points_from_xy(df["lon"], df["lat"]),
            crs="EPSG:4326",
        )
    else:
        raise RuntimeError(f"source non supportato: {src}")

    gdf = gdf.to_crs(C.METRIC_CRS).reset_index(drop=True)

    # ------------------------------------------------------------------
    # NIENTE FILTRO QUARTIERE SUI POI (rev. 06/2026)
    # ------------------------------------------------------------------
    # Per richiesta utente: vogliamo che gli edifici di San Leonardo possano
    # raggiungere ANCHE i POI fuori dal quartiere (es. la fontanella subito
    # oltre il confine). Quindi NON filtriamo piu' i POI per poligono.
    # Il filtro quartiere resta SOLO sugli edifici (vedi 02_prepare_graph.py)
    # e sul layer footprints_with_shortest_path (sotto, nel main).
    # Risultato: il Dijkstra multi-source parte da TUTTI i POI di Parma e
    # ogni edificio di San Leonardo si aggancia al POI piu' vicino in
    # assoluto, anche se questo si trova in un altro quartiere.
    C.log(f"  POI usati per la categoria '{name}': {len(gdf)} "
          f"(tutti i POI di Parma, nessun filtro quartiere)")
    # FINE rimozione filtro POI ----------------------------------------

    gdf["poi_id"] = range(len(gdf))

    # Cerco una colonna "nome" utile per i log/output
    name_candidates = ["NOME", "Name", "nome", "name", "TIPO", "tipo",
                       "INSEGNA", "TitoloITA", "PARCOUNICO"]
    name_col = next((c for c in name_candidates if c in gdf.columns), None)
    if name_col is None:
        gdf["poi_name"] = [f"poi_{i}" for i in range(len(gdf))]
    else:
        gdf["poi_name"] = gdf[name_col].astype(str)

    return gdf


# -----------------------------------------------------------------------------
# Snap POI: stessa logica edifici ma nodi "virt_p_<poi_id>"
# -----------------------------------------------------------------------------
def snap_pois_to_graph(G, pois_gdf, cat_name):
    """Modifica G inserendo un nodo virtuale per ogni POI. Restituisce
    DataFrame con (poi_id, virt_node, exit_cost)."""
    xs = pois_gdf.geometry.x.to_numpy()
    ys = pois_gdf.geometry.y.to_numpy()
    near = ox.distance.nearest_edges(G, X=xs, Y=ys)

    rows = []
    n_warn = 0
    for i in range(len(pois_gdf)):
        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_p_{cat_name}_{i}"
        # PATCH-LUCA most-used-fix (06/2026): host_u/host_v = i due nodi OSM
        # reali dell'arco di provenienza, identici alla logica di 02 per gli
        # edifici. Serve a 03 per ricondurre gli archi (u, virt) / (virt, v)
        # all'arco OSM originale (u, v) nel conteggio fid_count.
        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)
        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]))

        rows.append({
            "poi_id": i,
            "virt_node": vname,
            "exit_cost": float(snap_d),
            "proj_x": proj_xy[0],
            "proj_y": proj_xy[1],
        })

    df = pd.DataFrame(rows)
    if n_warn:
        C.log(f"    ATTENZIONE: {n_warn} POI distano > "
              f"{C.SNAP_WARN_DIST_M:.0f} m dall'arco piu' vicino")
    return df


# -----------------------------------------------------------------------------
# Aggiunge i nodi virtuali POI anche al DiGraph "ridotto" H usato da Dijkstra
# -----------------------------------------------------------------------------
def add_virt_pois_to_simple_graph(H, G, virt_p_nodes):
    """Per ogni virt_p_X aggiungo gli archi (u<->virt) e (virt<->v) in H."""
    for vp in virt_p_nodes:
        for u, v in G.in_edges(vp):
            # in_edges restituisce edge entranti: u -> vp
            data = G.get_edge_data(u, vp)
            # multidigraph: prendi il piu' corto
            best = min(data.values(), key=lambda d: d.get("length", float("inf")))
            w = float(best["length"])
            if H.has_edge(u, vp):
                if w < H[u][vp]["weight"]:
                    H[u][vp]["weight"] = w
            else:
                H.add_edge(u, vp, weight=w)
        for u, v in G.out_edges(vp):
            data = G.get_edge_data(vp, v)
            best = min(data.values(), key=lambda d: d.get("length", float("inf")))
            w = float(best["length"])
            if H.has_edge(vp, v):
                if w < H[vp][v]["weight"]:
                    H[vp][v]["weight"] = w
            else:
                H.add_edge(vp, v, weight=w)


# -----------------------------------------------------------------------------
# Costruzione geometria di un cammino (LineString in metri)
# -----------------------------------------------------------------------------
def path_to_linestring(G, path):
    """Costruisce LineString concatenando le geometrie reali degli archi."""
    if len(path) < 2:
        return None
    coords = []
    for u, v in zip(path[:-1], path[1:]):
        data = G.get_edge_data(u, v)
        if not data:
            # cerca senso inverso
            data = G.get_edge_data(v, u)
            if not data:
                continue
            best = min(data.values(), key=lambda d: d.get("length", float("inf")))
            geom = best.get("geometry")
            if geom is not None:
                pts = list(geom.coords)[::-1]
            else:
                nu = G.nodes[u]; nv = G.nodes[v]
                pts = [(nu["x"], nu["y"]), (nv["x"], nv["y"])]
        else:
            best = min(data.values(), key=lambda d: d.get("length", float("inf")))
            geom = best.get("geometry")
            if geom is not None:
                pts = list(geom.coords)
            else:
                nu = G.nodes[u]; nv = G.nodes[v]
                pts = [(nu["x"], nu["y"]), (nv["x"], nv["y"])]

        if not coords:
            coords.extend(pts)
        else:
            # evito di duplicare il punto di giunzione
            coords.extend(pts[1:])

    if len(coords) < 2:
        return None
    return LineString(coords)


# -----------------------------------------------------------------------------
# Pipeline per UNA categoria
# -----------------------------------------------------------------------------
def process_category(category, G_template, H_template, b_meta, fp_gdf):
    """Esegue tutta la pipeline per una categoria. Scrive routes_<name>.gpkg."""
    name = category["name"]
    out_gpkg = os.path.join(C.OUTPUT_DIR, f"routes_{name}.gpkg")
    C.log(f"\n>>> Categoria: {name}", "##")
    t0 = time.time()

    # ------------------------------------------------------------------
    # 1) Carico POI
    # ------------------------------------------------------------------
    pois_gdf = load_pois(category)
    C.log(f"  POI caricati: {len(pois_gdf)}")
    if len(pois_gdf) == 0:
        C.log(f"  Nessun POI per la categoria '{name}', salto.")
        return

    # ------------------------------------------------------------------
    # 2) Clono i grafi (per non sporcare H_template / G_template fra categorie)
    # ------------------------------------------------------------------
    C.log("  Clono grafo per la categoria...")
    G = G_template.copy()
    H = H_template.copy()

    # ------------------------------------------------------------------
    # 3) Snap dei POI sul grafo
    # ------------------------------------------------------------------
    C.log("  Snap POI sugli archi OSM + nodi virtuali...")
    t1 = time.time()
    poi_meta = snap_pois_to_graph(G, pois_gdf, name)
    virt_p_nodes = poi_meta["virt_node"].tolist()
    add_virt_pois_to_simple_graph(H, G, virt_p_nodes)
    C.log(f"    fatto in {time.time()-t1:.1f} s")

    # ------------------------------------------------------------------
    # 4) MULTI-SOURCE DIJKSTRA  <<< IL TRUCCO MAGICO >>>
    # ------------------------------------------------------------------
    C.log("  MULTI-SOURCE DIJKSTRA (parto da tutti i POI insieme)...")
    t1 = time.time()
    sources = set(virt_p_nodes)
    if C.DIJKSTRA_CUTOFF_M is not None:
        dist_map, path_map = nx.multi_source_dijkstra(
            H, sources=sources, weight="weight",
            cutoff=C.DIJKSTRA_CUTOFF_M,
        )
    else:
        dist_map, path_map = nx.multi_source_dijkstra(
            H, sources=sources, weight="weight",
        )
    C.log(f"    fatto in {time.time()-t1:.1f} s "
          f"(nodi raggiunti: {len(dist_map)})")

    # ------------------------------------------------------------------
    # 5) Per ogni edificio: estraggo destinazione + costo + path
    # ------------------------------------------------------------------
    C.log("  Estraggo cammini per ogni edificio...")
    poi_id_by_virt = dict(zip(poi_meta["virt_node"], poi_meta["poi_id"]))
    exit_cost_by_pid = dict(zip(poi_meta["poi_id"], poi_meta["exit_cost"]))
    # FIX gate-geometry (06/2026): mappa poi_id -> coordinate REALI del POI
    # (in CRS metrico), serve per "appendere" alla linea il tratto di uscita
    # ultimo-nodo-cammino -> POI. Senza questo, la linea finisce sulla strada
    # e in dashboard non si vede il collegamento alla destinazione.
    poi_xy_by_id = {
        int(pois_gdf.loc[i, "poi_id"]): (
            float(pois_gdf.geometry.iloc[i].x),
            float(pois_gdf.geometry.iloc[i].y),
        )
        for i in range(len(pois_gdf))
    }

    rows_paths    = []     # per layer shortest_path
    edge_counter  = Counter()  # frozenset({u,v}) -> count (chiave su nodi reali OSM, per stats)
    # PATCH-LUCA most-used-fix-v2 (06/2026): conto direttamente i SEGMENTI
    # percorsi (anche quelli che toccano nodi virtuali). Chiave: la coppia
    # ORDINATA (a, b) di nodi del grafo G (interi o stringhe 'virt_*').
    # Cosi' ogni step del path contribuisce a uno e un solo segmento, e
    # la geometria del segmento e' SEMPRE recuperabile da G.get_edge_data.
    # Non perdiamo MAI un pezzo di strada percorso.
    seg_counter   = Counter()  # (a, b) -> count  (a e b nell'ordine usato nel path)
    rows_origins  = []
    unreachable   = 0

    # PATCH-LUCA most-used-fix-v1 (mantenuto per stats): helper per
    # "normalizzare" uno step del cammino (a, b) restituendo la coppia di
    # nodi OSM REALI dell'arco da contare. Usato SOLO per edge_counter,
    # che ora ha solo scopo diagnostico.
    def _resolve_osm_endpoint(node, other):
        if not (isinstance(node, str) and node.startswith("virt_")):
            return node
        ndata = G.nodes.get(node, {})
        hu = ndata.get("host_u")
        hv = ndata.get("host_v")
        if hu is None or hv is None:
            return None
        if other == hu:
            return hv
        if other == hv:
            return hu
        try:
            if G.has_edge(hu, other) or G.has_edge(other, hu):
                return hu
            if G.has_edge(hv, other) or G.has_edge(other, hv):
                return hv
        except Exception:
            pass
        return hu

    def _path_to_osm_edges(path):
        """SOLO per stats: collassa gli step virtuali sugli host OSM."""
        out = []
        for a, b in zip(path[:-1], path[1:]):
            ua = _resolve_osm_endpoint(a, b)
            ub = _resolve_osm_endpoint(b, a)
            if ua is None or ub is None:
                continue
            if ua == ub:
                for n in (a, b):
                    if isinstance(n, str) and n.startswith("virt_"):
                        nd = G.nodes.get(n, {})
                        hu = nd.get("host_u"); hv = nd.get("host_v")
                        if hu is not None and hv is not None and hu != hv:
                            out.append(frozenset({hu, hv}))
                            break
                continue
            out.append(frozenset({ua, ub}))
        return out

    for _, b in b_meta.iterrows():
        vb = b["virt_node"]
        if vb not in path_map:
            unreachable += 1
            rows_origins.append({
                "fid": b["fid"],
                "assigned_poi_id": None,
                "total_cost": None,
                "geometry": Point(b["centroid_x"], b["centroid_y"]),
            })
            continue

        path = path_map[vb]
        # multi_source_dijkstra restituisce path[0] = source (POI virtuale),
        # path[-1] = target (edificio virtuale). Quindi il POI e' path[0].
        poi_virt = path[0]
        poi_id   = poi_id_by_virt[poi_virt]

        # costi
        network_cost = dist_map[vb] - exit_cost_by_pid[poi_id]
        # Nota: multi_source_dijkstra mette il source a costo 0; il costo finale
        # include i tratti dal POI al primo nodo reale (exit_cost) ma NON il
        # tratto edificio->virt_b (che e' "fuori dal grafo"). Quindi:
        #   total_cost (uscita-rete-ingresso) = entry_cost + dist_map[vb]
        # dove dist_map[vb] include gia' exit_cost. Riscompongo:
        entry_cost   = float(b["entry_cost"])
        exit_cost    = float(exit_cost_by_pid[poi_id])
        total_cost   = entry_cost + float(dist_map[vb])
        net_only     = float(dist_map[vb]) - exit_cost

        # PATCH-LUCA most-used-fix-v2 (06/2026): conteggio dei SEGMENTI percorsi.
        # Per ogni step (a, b) del path, normalizzo l'ordine (per evitare di
        # contare a->b e b->a come segmenti diversi) e incremento il counter.
        # Cosi' OGNI step contribuisce esattamente a un segmento, qualunque
        # sia la natura (reale/virtuale) di a e b.
        for a, b_ in zip(path[:-1], path[1:]):
            # normalizzo per usare la stessa chiave indipendentemente dal verso
            key = (a, b_) if str(a) <= str(b_) else (b_, a)
            seg_counter[key] += 1

        # vecchio counter (mantenuto solo per backwards compat / stats)
        for ek in _path_to_osm_edges(path):
            edge_counter[ek] += 1

        ls = path_to_linestring(G, path)
        if ls is not None:
            # ----------------------------------------------------------
            # FIX gate-geometry (06/2026): replico il comportamento della
            # vecchia pipeline (dash/script/03_compute_paths.py righe ~485):
            # alla LineString del cammino OSM appendo in TESTA il tratto
            # baricentro_edificio -> primo_nodo_path (gate di ACCESSO) e
            # in CODA il tratto ultimo_nodo_path -> POI_reale (gate di
            # USCITA). Cosi' in dashboard cliccando un edificio si vede:
            #   1) la "stanghetta" dal centro dell'edificio alla strada
            #   2) il cammino lungo la rete OSM
            #   3) la "stanghetta" dalla strada al POI di destinazione
            # ----------------------------------------------------------
            try:
                bx = float(b["centroid_x"])
                by = float(b["centroid_y"])
                # primo punto della linea OSM (= nodo virtuale POI, sulla strada)
                first_xy = ls.coords[0]
                # ultimo punto della linea OSM (= nodo virtuale edificio, sulla strada)
                last_xy  = ls.coords[-1]
                # coordinate reali del POI (in metri, CRS metrico)
                px, py = poi_xy_by_id.get(int(poi_id), (None, None))

                full_coords = []
                # 1) gate di ACCESSO: baricentro edificio -> ultimo_xy.
                #    NOTA: nella nostra pipeline path[0]=POI e path[-1]=edificio,
                #    quindi "ultimo_xy" e' il punto sulla strada vicino all'edificio
                #    e "first_xy" e' il punto sulla strada vicino al POI.
                #    Ma per il rendering in dashboard scriviamo la linea in ordine
                #    edificio -> ... -> POI, quindi:
                #       - inizio: baricentro edificio
                #       - poi: la linea OSM PERCORSA AL CONTRARIO (last_xy -> first_xy)
                #       - fine: coordinate del POI
                osm_coords_rev = list(ls.coords)[::-1]  # da edificio-lato-strada a POI-lato-strada
                full_coords.append((bx, by))
                full_coords.extend(osm_coords_rev)
                if px is not None and py is not None:
                    full_coords.append((px, py))

                # Rimuovo eventuali duplicati consecutivi (utile se baricentro
                # coincide gia' col primo nodo, cosa rara ma possibile)
                dedup = [full_coords[0]]
                for c in full_coords[1:]:
                    if (abs(c[0] - dedup[-1][0]) > 1e-6 or
                        abs(c[1] - dedup[-1][1]) > 1e-6):
                        dedup.append(c)
                if len(dedup) >= 2:
                    ls_full = LineString(dedup)
                else:
                    ls_full = ls
            except Exception as _e:
                # Se qualcosa va storto, ripiego sulla linea originale OSM
                ls_full = ls

            rows_paths.append({
                "origin_id":      b["fid"],
                "destination_id": int(poi_id),
                "entry_cost":     entry_cost,
                "network_cost":   net_only if net_only >= 0 else 0.0,
                "exit_cost":      exit_cost,
                "total_cost":     total_cost,
                "geometry":       ls_full,
            })

        rows_origins.append({
            "fid": b["fid"],
            "assigned_poi_id": int(poi_id),
            "total_cost": total_cost,
            "geometry": Point(b["centroid_x"], b["centroid_y"]),
        })

    C.log(f"    edifici con percorso: {len(rows_paths)}")
    C.log(f"    edifici irraggiungibili: {unreachable}")

    # ------------------------------------------------------------------
    # 6) Layer shortest_path
    # ------------------------------------------------------------------
    sp_gdf = gpd.GeoDataFrame(rows_paths, geometry="geometry", crs=C.METRIC_CRS)
    sp_gdf = sp_gdf.to_crs(C.OUTPUT_CRS)

    # ------------------------------------------------------------------
    # 7) Layer most_used_path: SEGMENTI PERCORSI con count
    # PATCH-LUCA road-aggr-v3 (2026-06-24): RIPRISTINO della logica della
    # vecchia pipeline (dash/script/04_count_usage.py PATCH-LUCA road-aggr).
    # ------------------------------------------------------------------
    # PROBLEMA che risolviamo:
    #   La precedente most-used-fix-v2 (06/2026) aggregava i sotto-segmenti
    #   per `groupby("osmid")` e fondeva la geometria con unary_union +
    #   linemerge. Effetti collaterali:
    #     1) Si perdevano i sotto-segmenti come record distinti -> la
    #        dashboard mostrava una sola MultiLineString per via OSM, non
    #        piu' i micro-segmenti reali.
    #     2) Il `fid_count` finale era il MAX dei soli sotto-segmenti
    #        VERAMENTE percorsi: i pezzetti della stessa "strada logica"
    #        che il Dijkstra NON aveva attraversato sparivano del tutto
    #        (perche' avevano fid_count=0 e non finivano nemmeno in
    #        seg_counter).
    #     3) `embed_percorsi.py` cerca il campo `fid_count_road` per il
    #        filtro: nella v2 non veniva piu' scritto -> filtro caduto su
    #        `fid_count` per-segmento -> con soglia 6 sparivano moltissime
    #        strade.
    #
    # COSA FA v3 (identica alla pipeline vecchia "04_count_usage.py"):
    #   a) Tiene un RECORD per ogni sotto-segmento percorso (geometria
    #      granulare, niente unary_union).
    #   b) Costruisce un grafo OSM non-orientato UG escludendo i nodi
    #      virtuali (virt_b_*, virt_p_*), e calcola il grado topologico
    #      di ogni nodo reale.
    #   c) Fa una "passeggiata" lungo le catene di nodi di grado 2 partendo
    #      dai sotto-segmenti percorsi, e assegna un `road_id` condiviso a
    #      tutta la catena.
    #   d) `fid_count_road` = MAX(fid_count) tra i sotto-segmenti della
    #      stessa catena. Cosi' se anche UN solo pezzo di una via e' stato
    #      percorso N volte, tutti i suoi pezzetti ereditano N.
    #   e) Il filtro in embed_percorsi.py (soglia 1, vedi PATCH show-all-used)
    #      applicato su `fid_count_road` -> appaiono TUTTE le strade
    #      effettivamente percorse, complete, anche quelle con un solo
    #      passaggio.
    # ------------------------------------------------------------------
    C.log("  Costruisco most_used_path (PATCH road-aggr-v3)...")

    # ---- 7a) Raccolgo un record per sotto-segmento percorso --------------
    # Chiave canonica frozenset({a, b_}) per il lookup (la rete walk e' di
    # fatto bidirezionale).
    seg_records = []     # lista di dict (uno per sotto-segmento usato)
    n_skipped = 0
    for (a, b_), count in seg_counter.items():
        data = G.get_edge_data(a, b_) or G.get_edge_data(b_, a)
        if not data:
            n_skipped += 1
            continue
        best = min(data.values(), key=lambda d: d.get("length", float("inf")))
        geom = best.get("geometry")
        if geom is None:
            na = G.nodes.get(a, {}); nb = G.nodes.get(b_, {})
            if "x" in na and "y" in na and "x" in nb and "y" in nb:
                geom = LineString([(na["x"], na["y"]), (nb["x"], nb["y"])])
            else:
                n_skipped += 1
                continue
        osmid_str = str(best.get("osmid", ""))
        name_str  = str(best.get("name", ""))
        hwy_str   = str(best.get("highway", ""))
        seg_records.append({
            "_a":        a,
            "_b":        b_,
            "osm_id":    osmid_str,
            "name":      name_str,
            "highway":   hwy_str,
            "length":    float(best.get("length", geom.length)),
            "fid_count": int(count),
            "geometry":  geom,
        })
    if n_skipped:
        C.log(f"    segmenti scartati per geom mancante: {n_skipped}")
    C.log(f"    sotto-segmenti percorsi: {len(seg_records)}")

    # ---- 7b) Grafo OSM non-orientato per il grado topologico -------------
    # Escludo i nodi virtuali: non sono intersezioni reali della rete OSM,
    # altrimenti la "passeggiata" verrebbe interrotta nei punti dove un
    # POI o un edificio si e' agganciato a meta' strada.
    UG = nx.Graph()
    for u_e, v_e, _k_e, _d_e in G.edges(keys=True, data=True):
        if isinstance(u_e, str) and u_e.startswith(("virt_b_", "virt_p_")):
            continue
        if isinstance(v_e, str) and v_e.startswith(("virt_b_", "virt_p_")):
            continue
        if u_e == v_e:
            continue  # self-loop: ignoro
        UG.add_edge(u_e, v_e)
    node_degree = dict(UG.degree())

    # ---- 7c) Insieme degli archi usati (chiave canonica frozenset) -------
    # NB: se un sotto-segmento ha un endpoint virtuale, lo "promuovo" sul
    # nodo OSM ospite (host_u / host_v) cosi' la chain-walking sui nodi di
    # grado 2 funziona anche per quei pezzi. Senza questa promozione le
    # strade in prossimita' di POI/edifici si frammenterebbero.
    def _real_endpoint(n):
        if isinstance(n, str) and n.startswith(("virt_b_", "virt_p_")):
            nd = G.nodes.get(n, {})
            hu = nd.get("host_u"); hv = nd.get("host_v")
            # arbitrariamente uso host_u: tanto serve solo per identificare
            # la "catena" cui questo pezzo appartiene topologicamente.
            return hu if hu is not None else (hv if hv is not None else n)
        return n

    # Mappo ogni sotto-segmento alla sua chiave canonica "reale" per la
    # walk; mantengo anche la mappa inversa per assegnare road_id ai record.
    seg_real_key = []                # parallelo a seg_records
    used_edges_set = set()           # frozenset di nodi reali
    neigh_used = defaultdict(set)    # adiacenze fra nodi reali via archi usati

    for rec in seg_records:
        ra = _real_endpoint(rec["_a"])
        rb = _real_endpoint(rec["_b"])
        if ra == rb:
            # pezzo interno a uno stesso arco "host" spezzato da due virtuali:
            # uso comunque la coppia originaria host_u/host_v come catena.
            for n in (rec["_a"], rec["_b"]):
                if isinstance(n, str) and n.startswith(("virt_b_", "virt_p_")):
                    nd = G.nodes.get(n, {})
                    hu = nd.get("host_u"); hv = nd.get("host_v")
                    if hu is not None and hv is not None and hu != hv:
                        ra, rb = hu, hv
                        break
        key = frozenset((ra, rb))
        seg_real_key.append(key)
        if ra != rb:
            used_edges_set.add(key)
            neigh_used[ra].add(rb)
            neigh_used[rb].add(ra)

    # ---- 7d) Walk: assegno road_id seguendo nodi di grado 2 sul grafo OSM
    def _is_intersection(n):
        # Grado != 2 sul grafo OSM completo (no nodi virtuali) = intersezione.
        return node_degree.get(n, 0) != 2

    edge_to_road = {}    # frozenset({u,v}) -> road_id
    road_id_counter = 0

    for start_edge in used_edges_set:
        if start_edge in edge_to_road:
            continue
        road_id_counter += 1
        rid = road_id_counter
        edge_to_road[start_edge] = rid

        a0, b0 = tuple(start_edge)
        frontier = [(a0, start_edge), (b0, start_edge)]
        while frontier:
            node, came_from = frontier.pop()
            if _is_intersection(node):
                continue  # fine strada
            # nodo di grado 2: avanzo verso l'altro vicino USATO
            for nb in neigh_used[node]:
                cand = frozenset((node, nb))
                if cand == came_from or cand in edge_to_road:
                    continue
                edge_to_road[cand] = rid
                frontier.append((nb, cand))

    # ---- 7e) fid_count_road = MAX(fid_count) per road_id -----------------
    # Devo aggregare i fid_count dei record che condividono la stessa
    # chiave canonica reale (puo' capitare che piu' sotto-segmenti virtuali
    # mappino alla stessa coppia (ra, rb)).
    fid_per_real_key = defaultdict(int)
    for rec, key in zip(seg_records, seg_real_key):
        c = int(rec["fid_count"])
        if c > fid_per_real_key[key]:
            fid_per_real_key[key] = c

    road_max = defaultdict(int)
    for key, rid in edge_to_road.items():
        c = fid_per_real_key.get(key, 0)
        if c > road_max[rid]:
            road_max[rid] = c

    # ---- 7f) Scrivo road_id + fid_count_road nei record ------------------
    n_orphans = 0
    for rec, key in zip(seg_records, seg_real_key):
        rid = edge_to_road.get(key)
        if rid is None:
            # fallback: il sotto-segmento non e' entrato in nessuna catena
            # (caso limite, p.es. ra == rb). Lo tratto come strada a se'.
            road_id_counter += 1
            rid = road_id_counter
            edge_to_road[key] = rid
            road_max[rid] = int(rec["fid_count"])
            n_orphans += 1
        rec["road_id"] = int(rid)
        rec["fid_count_road"] = int(road_max[rid])

    if n_orphans:
        C.log(f"    sotto-segmenti senza catena (orphan): {n_orphans}")

    # ---- 7g) Pulisco colonne ausiliarie e costruisco il GeoDataFrame -----
    for rec in seg_records:
        rec.pop("_a", None)
        rec.pop("_b", None)

    if seg_records:
        mup_gdf = gpd.GeoDataFrame(
            seg_records, geometry="geometry", crs=C.METRIC_CRS
        ).to_crs(C.OUTPUT_CRS)
    else:
        mup_gdf = gpd.GeoDataFrame(
            columns=["osm_id", "name", "highway", "length",
                     "fid_count", "road_id", "fid_count_road", "geometry"],
            geometry="geometry", crs=C.OUTPUT_CRS,
        )

    # Report compatto stile vecchia pipeline
    n_roads = len(set(edge_to_road.values()))
    avg_seg_per_road = (len(edge_to_road) / n_roads) if n_roads else 0
    C.log(f"    {len(seg_records)} sotto-segmenti raggruppati in "
          f"{n_roads} 'strade' (media {avg_seg_per_road:.2f} seg/strada)")
    if len(mup_gdf):
        C.log(f"    fid_count: min={int(mup_gdf['fid_count'].min())}, "
              f"median={int(mup_gdf['fid_count'].median())}, "
              f"max={int(mup_gdf['fid_count'].max())}")
        C.log(f"    fid_count_road: min={int(mup_gdf['fid_count_road'].min())}, "
              f"median={int(mup_gdf['fid_count_road'].median())}, "
              f"max={int(mup_gdf['fid_count_road'].max())}")


    # ------------------------------------------------------------------
    # 8) Footprints + join dei costi
    # ------------------------------------------------------------------
    C.log("  Join footprints + costi...")
    origins_df = pd.DataFrame([{
        "fid": r["fid"],
        "assigned_poi_id": r["assigned_poi_id"],
        "total_cost": r["total_cost"],
    } for r in rows_origins])

    fp_join = fp_gdf.merge(origins_df, how="left", on="fid")
    fp_join = fp_join.to_crs(C.OUTPUT_CRS)

    # ------------------------------------------------------------------
    # 9) Statistiche per ogni POI
    # ------------------------------------------------------------------
    C.log("  Calcolo statistiche per POI...")
    served = origins_df.dropna(subset=["assigned_poi_id"])
    if len(served):
        stats = served.groupby("assigned_poi_id").agg(
            n_buildings_served=("fid", "count"),
            avg_distance=("total_cost", "mean"),
            max_distance=("total_cost", "max"),
            min_distance=("total_cost", "min"),
        ).reset_index().rename(columns={"assigned_poi_id": "poi_id"})
        stats["poi_id"] = stats["poi_id"].astype(int)
    else:
        stats = pd.DataFrame(columns=["poi_id", "n_buildings_served",
                                       "avg_distance", "max_distance",
                                       "min_distance"])

    pois_out = pois_gdf.merge(stats, how="left", on="poi_id")
    pois_out["n_buildings_served"] = pois_out["n_buildings_served"].fillna(0).astype(int)
    pois_out = pois_out.to_crs(C.OUTPUT_CRS)

    # ------------------------------------------------------------------
    # 10) Origin points (baricentri + assigned_poi_id)
    # ------------------------------------------------------------------
    op_gdf = gpd.GeoDataFrame(rows_origins, geometry="geometry",
                              crs=C.METRIC_CRS).to_crs(C.OUTPUT_CRS)

    # ------------------------------------------------------------------
    # 11) Scrittura GPKG finale
    # ------------------------------------------------------------------
    if os.path.exists(out_gpkg):
        os.remove(out_gpkg)
    C.log(f"  Scrivo: {out_gpkg}")
    sp_gdf.to_file(out_gpkg, layer="shortest_path", driver="GPKG")
    mup_gdf.to_file(out_gpkg, layer="most_used_path", driver="GPKG")
    fp_join.to_file(out_gpkg, layer="footprints_with_shortest_path", driver="GPKG")
    # POI: dropno colonne che potrebbero non essere serializzabili
    pois_safe = pois_out.copy()
    for col in pois_safe.columns:
        if col in ("geometry",): continue
        if pois_safe[col].dtype == "object":
            pois_safe[col] = pois_safe[col].astype(str)
    pois_safe.to_file(out_gpkg, layer="pois", driver="GPKG")
    op_gdf.to_file(out_gpkg, layer="origin_points", driver="GPKG")

    # ------------------------------------------------------------------
    # REPORT
    # ------------------------------------------------------------------
    n_served_pois = int((pois_out["n_buildings_served"] > 0).sum())
    C.log(f"  REPORT categoria '{name}':")
    C.log(f"    POI totali           : {len(pois_out)}")
    C.log(f"    POI con almeno 1 ed. : {n_served_pois}")
    C.log(f"    POI fantasma (0 ed.) : {len(pois_out) - n_served_pois}")
    if len(served):
        C.log(f"    dist media (m)       : {served['total_cost'].mean():.1f}")
        C.log(f"    dist mediana (m)     : {served['total_cost'].median():.1f}")
        C.log(f"    dist max (m)         : {served['total_cost'].max():.1f}")
    C.log(f"  Categoria completata in {time.time()-t0:.1f} s", "##")


# -----------------------------------------------------------------------------
# MAIN
# -----------------------------------------------------------------------------
def main():
    t0 = time.time()
    C.log("=== 03 - Calcolo percorsi multi-source per categoria ===", "==")

    # ------------------------------------------------------------------
    # Caricamento cache (UNA volta sola)
    # ------------------------------------------------------------------
    C.log("Carico grafo preparato dalla cache...")
    if not os.path.exists(PREPARED_GRAPH_PKL):
        raise RuntimeError("Cache mancante. Esegui prima 02_prepare_graph.py")
    with open(PREPARED_GRAPH_PKL, "rb") as f:
        G_template = pickle.load(f)
    with open(SIMPLE_GRAPH_PKL, "rb") as f:
        H_template = pickle.load(f)
    b_meta = pd.read_pickle(BUILDINGS_META_PKL)
    C.log(f"  MultiDiGraph: {G_template.number_of_nodes()} nodi, "
          f"{G_template.number_of_edges()} archi")
    C.log(f"  DiGraph rid.: {H_template.number_of_nodes()} nodi, "
          f"{H_template.number_of_edges()} archi")
    C.log(f"  edifici meta: {len(b_meta)}")

    # ------------------------------------------------------------------
    # Carico footprints UNA volta (riusati in tutte le categorie)
    # ------------------------------------------------------------------
    C.log("Carico footprints (per join finali)...")
    fp_gdf = gpd.read_file(C.SOURCE_GPKG, layer=C.FOOTPRINTS_LAYER).to_crs(C.METRIC_CRS)

    # Filtro quartiere anche qui (coerente con 02_prepare_graph.py):
    # nel layer footprints_with_shortest_path scriviamo solo gli edifici
    # del quartiere, altrimenti uscirebbe un GPKG con 44k righe quasi tutte
    # NaN sui costi (spreco enorme e fuori scope per la dashboard di quartiere).
    quart_poly = C.load_quartiere_polygon()
    if quart_poly is not None:
        n_before = len(fp_gdf)
        mask = fp_gdf.geometry.centroid.within(quart_poly)
        fp_gdf = fp_gdf[mask].reset_index(drop=True)
        C.log(f"  filtro quartiere footprints: {n_before} -> {len(fp_gdf)}")

    # FIX BUG fid mismatch (23/06/2026):
    # Il GPKG sorgente non ha colonna 'fid' visibile a geopandas (rowid SQLite
    # implicito), quindi 02_prepare_graph.py cade nel branch "_autoid = range"
    # e b_meta['fid'] contiene 0..N-1 *dopo* il filtro quartiere. Qui dobbiamo
    # usare lo STESSO criterio: assegnare fid=0..N-1 DOPO il filtro quartiere,
    # cosi' il merge fp_gdf <-> origins_df su 'fid' funziona correttamente.
    # (Prima il fid veniva assegnato PRIMA del filtro -> mismatch totale,
    # matchavano solo 15 righe per puro caso e 2429 finivano con NULL.)
    fp_gdf = fp_gdf.drop(columns=["fid"], errors="ignore")
    fp_gdf["fid"] = range(len(fp_gdf))

    C.log(f"  edifici: {len(fp_gdf)}")

    # ------------------------------------------------------------------
    # Loop su tutte le categorie
    # ------------------------------------------------------------------
    for cat in C.POI_CATEGORIES:
        try:
            process_category(cat, G_template, H_template, b_meta, fp_gdf)
        except Exception as e:
            import traceback
            C.log(f"!!! ERRORE categoria '{cat['name']}': {e}", "!!")
            traceback.print_exc()

    C.log(f"\n=== TUTTE LE CATEGORIE COMPLETATE in {time.time()-t0:.1f} s ===", "==")
    C.log(f"Output in: {C.OUTPUT_DIR}")


if __name__ == "__main__":
    main()
