#!/usr/bin/env python3 """Fetch all public-transport stops (and the lines serving them) in Austria, Germany and Switzerland (DACH) from OpenStreetMap and write a deduplicated snapshot to data/stops_dach.csv.gz. Data © OpenStreetMap contributors, licensed under the ODbL. At this scale the public Overpass API can't be used directly - even a plain count query for one of these countries times out on it. Instead this downloads the official Geofabrik .osm.pbf extract per country (cached in data/.osm-cache/, never committed) and processes it locally with pyosmium. No load is put on any shared API. Run this only when you want to refresh the snapshot; the import into the database (`flask import-stops`) reads the checked-in CSV and needs no network access. Refreshing needs: `pip install osmium`, ~6 GB of disk for the extracts, a few GB of RAM, and maybe 15-30 minutes depending on disk speed (most of it is Germany). pip install osmium python scripts/fetch_stops.py """ from __future__ import annotations import csv import gzip import math import sys import unicodedata import urllib.request from pathlib import Path import osmium OUT_PATH = Path(__file__).resolve().parent.parent / "data" / "stops_dach.csv.gz" CACHE_DIR = Path(__file__).resolve().parent.parent / "data" / ".osm-cache" # ISO 3166-1 alpha-2 -> Geofabrik region file basename. GEOFABRIK_REGIONS = { "AT": "austria", "CH": "switzerland", "DE": "germany", } ROUTE_MODES = { "bus", "trolleybus", "tram", "light_rail", "subway", "train", "monorail", "share_taxi", } STOP_ROLES = { "stop", "platform", "stop_entry_only", "stop_exit_only", "platform_entry_only", "platform_exit_only", } CSV_FIELDS = [ "osm_type", "osm_id", "name", "stop_type", "country", "municipality", "latitude", "longitude", "lines", ] TYPE_RANK = {"subway": 5, "train": 4, "bus_station": 3, "tram": 2, "bus": 1, "other": 0} # OSM keeps a separate node per direction for most bus stops (~20-60 m apart); # nodes with the same name closer than this are treated as one stop. MERGE_METERS = 250.0 MIN_STOPS = 100_000 # DACH sanity floor; refuse to overwrite a good snapshot with junk def normalize_name(value: str) -> str: """Lowercase, fold accents and ß so search is diacritic-insensitive. The same function lives in stops_import.py; keep them in sync. """ value = value.strip().lower().replace("ß", "ss") decomposed = unicodedata.normalize("NFKD", value) return "".join(ch for ch in decomposed if not unicodedata.combining(ch)) def haversine_m(lat1: float, lon1: float, lat2: float, lon2: float) -> float: rlat1, rlon1, rlat2, rlon2 = map(math.radians, (lat1, lon1, lat2, lon2)) dlat, dlon = rlat2 - rlat1, rlon2 - rlon1 h = math.sin(dlat / 2) ** 2 + math.cos(rlat1) * math.cos(rlat2) * math.sin(dlon / 2) ** 2 return 2 * 6_371_000 * math.asin(min(1.0, math.sqrt(h))) def classify(tags: dict) -> str: if tags.get("station") == "subway" or tags.get("subway") == "yes": return "subway" if tags.get("railway") in {"station", "halt"}: return "train" if tags.get("railway") == "tram_stop" or tags.get("tram") == "yes": return "tram" if tags.get("amenity") == "bus_station": return "bus_station" if tags.get("highway") == "bus_stop" or tags.get("bus") == "yes": return "bus" return "other" def municipality_of(tags: dict) -> str: for key in ("addr:city", "is_in:municipality", "is_in:city", "is_in"): value = tags.get(key) if value: return value.split(",")[0].strip() return "" # --------------------------------------------------------------------------- # # Geofabrik download # --------------------------------------------------------------------------- # def download_extract(iso: str) -> Path: region = GEOFABRIK_REGIONS[iso] CACHE_DIR.mkdir(parents=True, exist_ok=True) dest = CACHE_DIR / f"{region}-latest.osm.pbf" if dest.exists() and dest.stat().st_size > 0: print(f" using cached {dest.name} ({dest.stat().st_size / 1e6:.0f} MB)", file=sys.stderr) return dest url = f"https://download.geofabrik.de/europe/{region}-latest.osm.pbf" tmp = dest.with_suffix(".pbf.part") print(f" downloading {url} ...", file=sys.stderr) with urllib.request.urlopen(url, timeout=1800) as resp, open(tmp, "wb") as fh: while True: chunk = resp.read(1 << 20) if not chunk: break fh.write(chunk) tmp.rename(dest) print(f" downloaded {dest.name} ({dest.stat().st_size / 1e6:.0f} MB)", file=sys.stderr) return dest # --------------------------------------------------------------------------- # # Stops # # All three passes below use osmium.FileProcessor with a native (C++-side) # filter rather than SimpleHandler.apply_file(): filtering there means most # objects never cross into Python at all. A plain SimpleHandler callback - # even one that returns immediately - has to be invoked for every single # node in the file, and that per-call overhead alone made a whole-country # extract (tens of millions of nodes) impractically slow in testing. # --------------------------------------------------------------------------- # def collect_stop_nodes(pbf_path: Path, iso: str) -> list[dict]: nodes: list[dict] = [] fp = osmium.FileProcessor(str(pbf_path)).with_filter(osmium.filter.EmptyTagFilter()) for obj in fp: if not obj.is_node() or not obj.location.valid(): continue tags = obj.tags if not ( tags.get("highway") == "bus_stop" or tags.get("railway") in ("tram_stop", "station", "halt") or tags.get("public_transport") == "station" or tags.get("amenity") == "bus_station" ): continue name = (tags.get("name") or "").strip() if not name: continue tags_dict = {t.k: t.v for t in tags} nodes.append({ "osm_id": int(obj.id), "name": name, "stop_type": classify(tags_dict), "country": iso, "municipality": municipality_of(tags_dict), "latitude": round(obj.location.lat, 6), "longitude": round(obj.location.lon, 6), }) print(f" {iso}: {len(nodes)} raw stop nodes", file=sys.stderr) return nodes def build_stop_rows(nodes: list[dict]) -> list[dict]: by_name: dict[str, list[dict]] = {} for node in nodes: by_name.setdefault(normalize_name(node["name"]), []).append(node) rows: list[dict] = [] for group in by_name.values(): clusters: list[list[dict]] = [] for node in sorted(group, key=lambda n: n["osm_id"]): for cluster in clusters: if any( haversine_m( node["latitude"], node["longitude"], other["latitude"], other["longitude"], ) <= MERGE_METERS for other in cluster ): cluster.append(node) break else: clusters.append([node]) for cluster in clusters: lead = max( cluster, key=lambda n: (TYPE_RANK[n["stop_type"]], n["municipality"] != "", -n["osm_id"]), ) rows.append({ "osm_type": "node", "osm_id": min(n["osm_id"] for n in cluster), "name": lead["name"], "stop_type": lead["stop_type"], "country": lead["country"], "municipality": lead["municipality"] or next((n["municipality"] for n in cluster if n["municipality"]), ""), "latitude": round(sum(n["latitude"] for n in cluster) / len(cluster), 6), "longitude": round(sum(n["longitude"] for n in cluster) / len(cluster), 6), "lines": set(), }) rows.sort(key=lambda r: (normalize_name(r["name"]), r["osm_id"])) return rows # --------------------------------------------------------------------------- # # Lines (route relations -> stops) # --------------------------------------------------------------------------- # def collect_route_members(pbf_path: Path) -> tuple[list[tuple[str, int]], dict[int, tuple[str, float, float]]]: """Pass 1: which route relations do we care about, and which node ids do their stop/platform members point at? Pass 2: resolve name/coordinates for exactly those node ids (they're often bare stop_position nodes with no tags at all, so this can't be narrowed with a tag filter like the stops pass - it has to look at every node, hence EntityFilter rather than EmptyTagFilter to at least skip all ways in the same pass).""" relations: list[tuple[list[str], list[int]]] = [] wanted_node_ids: set[int] = set() fp1 = osmium.FileProcessor(str(pbf_path)).with_filter(osmium.filter.EntityFilter(osmium.osm.RELATION)) for obj in fp1: tags = obj.tags if tags.get("route") not in ROUTE_MODES: continue raw_ref = tags.get("ref") if not raw_ref: continue # A few relations pack several numbers into one ref ("407, 413"). refs = [part.strip() for part in raw_ref.replace(",", ";").split(";") if part.strip()] if not refs: continue member_ids = [m.ref for m in obj.members if m.type == "n" and m.role in STOP_ROLES] if not member_ids: continue relations.append((refs, member_ids)) wanted_node_ids.update(member_ids) print( f" {len(relations)} route relations, " f"{len(wanted_node_ids)} distinct stop-member nodes to resolve", file=sys.stderr, ) node_pos: dict[int, tuple[str, float, float]] = {} fp2 = osmium.FileProcessor(str(pbf_path)).with_filter(osmium.filter.EntityFilter(osmium.osm.NODE)) for obj in fp2: if obj.id not in wanted_node_ids or not obj.location.valid(): continue node_pos[obj.id] = (obj.tags.get("name", ""), obj.location.lat, obj.location.lon) ref_members: list[tuple[str, int]] = [] for refs, member_ids in relations: for ref in refs: for node_id in member_ids: ref_members.append((ref, node_id)) return ref_members, node_pos def assign_lines(rows: list[dict], ref_members, node_pos) -> tuple[int, int]: by_name: dict[str, list[int]] = {} grid: dict[tuple[int, int], list[int]] = {} for idx, row in enumerate(rows): by_name.setdefault(normalize_name(row["name"]), []).append(idx) cell = (round(row["latitude"] / 0.01), round(row["longitude"] / 0.01)) grid.setdefault(cell, []).append(idx) def nearest(lat: float, lon: float, max_m: float) -> int | None: best, best_d = None, max_m base = (round(lat / 0.01), round(lon / 0.01)) for dr in (-1, 0, 1): for dc in (-1, 0, 1): for idx in grid.get((base[0] + dr, base[1] + dc), ()): d = haversine_m(lat, lon, rows[idx]["latitude"], rows[idx]["longitude"]) if d < best_d: best, best_d = idx, d return best matched = unmatched = 0 for ref, node_id in ref_members: pos = node_pos.get(node_id) if pos is None: unmatched += 1 continue name, lat, lon = pos idx = None key = normalize_name(name) if name else "" if key and key in by_name: idx = min( by_name[key], key=lambda i: haversine_m(lat, lon, rows[i]["latitude"], rows[i]["longitude"]), ) if haversine_m(lat, lon, rows[idx]["latitude"], rows[idx]["longitude"]) > 600: idx = None if idx is None: idx = nearest(lat, lon, 120) if idx is None: unmatched += 1 else: matched += 1 rows[idx]["lines"].add(ref) return matched, unmatched def line_sort_key(ref: str): head, i = "", 0 while i < len(ref) and not ref[i].isdigit(): head, i = head + ref[i], i + 1 digits = "" while i < len(ref) and ref[i].isdigit(): digits, i = digits + ref[i], i + 1 return (head.lower(), int(digits) if digits else -1, ref[i:].lower()) # --------------------------------------------------------------------------- # def main() -> None: all_nodes: list[dict] = [] all_ref_members: list[tuple[str, int]] = [] all_node_pos: dict[int, tuple[str, float, float]] = {} for iso in GEOFABRIK_REGIONS: print(f"=== {iso} ===", file=sys.stderr) pbf_path = download_extract(iso) print(" collecting stops ...", file=sys.stderr) all_nodes.extend(collect_stop_nodes(pbf_path, iso)) print(" collecting routes ...", file=sys.stderr) ref_members, node_pos = collect_route_members(pbf_path) all_ref_members.extend(ref_members) all_node_pos.update(node_pos) print("Clustering stops ...", file=sys.stderr) rows = build_stop_rows(all_nodes) if len(rows) < MIN_STOPS: raise SystemExit( f"Only {len(rows)} stops parsed (< {MIN_STOPS}) - refusing to overwrite the snapshot." ) print("Matching lines to stops ...", file=sys.stderr) matched, unmatched = assign_lines(rows, all_ref_members, all_node_pos) OUT_PATH.parent.mkdir(parents=True, exist_ok=True) with gzip.open(OUT_PATH, "wt", newline="", encoding="utf-8") as fh: writer = csv.DictWriter(fh, fieldnames=CSV_FIELDS) writer.writeheader() for row in rows: row["lines"] = ";".join(sorted(row["lines"], key=line_sort_key)) writer.writerow(row) by_type: dict[str, int] = {} by_country: dict[str, int] = {} for row in rows: by_type[row["stop_type"]] = by_type.get(row["stop_type"], 0) + 1 by_country[row["country"]] = by_country.get(row["country"], 0) + 1 with_lines = sum(1 for r in rows if r["lines"]) print(f"Wrote {len(rows)} stops to {OUT_PATH}", file=sys.stderr) print(f" by country: {by_country}", file=sys.stderr) print(f" by type: {by_type}", file=sys.stderr) print( f" lines: {with_lines} stops have >=1 line " f"(memberships matched {matched}, unmatched {unmatched})", file=sys.stderr, ) if __name__ == "__main__": main()