diff --git a/.gitignore b/.gitignore index a2796e5..04bafdc 100644 --- a/.gitignore +++ b/.gitignore @@ -4,3 +4,4 @@ __pycache__/ .env pbt.env instance/ +data/.osm-cache/ diff --git a/README.md b/README.md index d1cbae6..66db3e3 100644 --- a/README.md +++ b/README.md @@ -20,15 +20,18 @@ Beim ersten Start legt die App die Tabellen (`users`, `trips`, `stops`) automati ## Haltestellen-Lookup -Die Felder „Von" und „Nach" haben eine Autovervollständigung aus allen -österreichischen Haltestellen. Die Daten stammen aus OpenStreetMap -(© OpenStreetMap-Mitwirkende, ODbL) und liegen als Snapshot im Repo -(`data/stops_at.csv.gz`, ~37 000 Haltestellen). +Die Felder „Von" und „Nach" haben eine Autovervollständigung aus den +Haltestellen in Österreich, Deutschland und der Schweiz (DACH) – aktuell +**306 954 Haltestellen** (DE 243 278, AT 37 454, CH 26 222). Die Daten stammen +aus OpenStreetMap (© OpenStreetMap-Mitwirkende, ODbL) und liegen als Snapshot +im Repo (`data/stops_dach.csv.gz`, ~7,5 MB gepackt). Zusätzlich ist pro Haltestelle die Menge der Linien hinterlegt (aus den -OSM-`route`-Relationen, ~87 % der Halte). Das Feld „Linie" schlägt daraus die -Linien vor, die an den gewählten Halten verkehren, und füllt sich selbst aus, -wenn nur eine Linie in Frage kommt. Freitext bleibt immer möglich. +OSM-`route`-Relationen). Das Feld „Linie" schlägt daraus die Linien vor, die +an den gewählten Halten verkehren, und füllt sich selbst aus, wenn nur eine +Linie in Frage kommt. Freitext bleibt immer möglich. Da mehrere Länder +gleichnamige Haltestellen haben können (mehrere „Hauptbahnhof"), zeigt die +Vorschlagsliste zusätzlich das Land. Nach dem Anlegen der Tabellen den Snapshot in die DB laden: @@ -40,14 +43,28 @@ Der Import ist idempotent (Upsert) und braucht kein Internet. Im Produktivbetrie übernimmt das die Unit `pbt-import-stops.service` (wird bei jedem Deploy angestoßen, s. u.). -Snapshot neu von OpenStreetMap holen (Halte + Linien, ~5 min, braucht Netzugang -und ~1 GB RAM): +### Snapshot aktualisieren + +Bei dieser Größenordnung (AT+DE+CH) geht sich das nicht mehr über die +öffentliche Overpass-API aus – schon eine reine Zählabfrage für Deutschland +oder auch nur die Schweiz läuft dort in den Timeout. Stattdessen lädt +`fetch_stops.py` die offiziellen Geofabrik-Extrakte pro Land herunter +(`data/.osm-cache/`, nicht committed, ~6 GB) und wertet sie lokal mit +`pyosmium` aus – keine Last auf einer geteilten API. ```bash -python scripts/fetch_stops.py # überschreibt data/stops_at.csv.gz +pip install osmium # einmalig, nur für dieses Script +python scripts/fetch_stops.py # überschreibt data/stops_dach.csv.gz flask import-stops ``` +Braucht Netzugang zum Download, ein paar GB RAM und in der Praxis eher +1–1,5 Stunden als die 5 Minuten von früher (der Großteil davon ist +Deutschland – Österreich allein dauert ca. 11 Minuten). Bereits +heruntergeladene Extrakte werden für spätere Läufe wiederverwendet, solange +sie in `data/.osm-cache/` liegen (~6 GB, `rm -rf data/.osm-cache/` gibt den +Platz wieder frei). + ## Kartenvorschau (optional, standardmäßig aus) Trip-Formular und „Meine Fahrten" können eine kleine Leaflet/OSM-Karte mit den diff --git a/app.py b/app.py index 6026723..c914a16 100644 --- a/app.py +++ b/app.py @@ -142,6 +142,7 @@ def register_routes(app): "id": stop.id, "name": stop.name, "type": stop.stop_type, + "country": stop.country, "municipality": stop.municipality, "lat": stop.latitude, "lon": stop.longitude, diff --git a/data/stops_at.csv.gz b/data/stops_at.csv.gz deleted file mode 100644 index 03bfef0..0000000 Binary files a/data/stops_at.csv.gz and /dev/null differ diff --git a/data/stops_dach.csv.gz b/data/stops_dach.csv.gz new file mode 100644 index 0000000..c3f35e1 Binary files /dev/null and b/data/stops_dach.csv.gz differ diff --git a/models.py b/models.py index 48a4009..fad5ec7 100644 --- a/models.py +++ b/models.py @@ -37,7 +37,8 @@ VERKEHRSMITTEL_OPTIONEN = [ class Stop(db.Model): - """A boardable public-transport stop ("Haltestelle") in Austria. + """A boardable public-transport stop ("Haltestelle") in Austria, Germany + or Switzerland (DACH). Sourced from OpenStreetMap; see scripts/fetch_stops.py and stops_import.py. """ @@ -52,6 +53,9 @@ class Stop(db.Model): # lowercased, accent-folded copy of name for diacritic-insensitive search name_normalized = db.Column(db.String(200), nullable=False, index=True) stop_type = db.Column(db.String(20), nullable=False, default="other") + # ISO 3166-1 alpha-2 of the fetch query that found this stop (AT/DE/CH) - + # disambiguates same-named stops across borders (several "Hauptbahnhof"). + country = db.Column(db.String(2), nullable=False, default="AT") municipality = db.Column(db.String(120)) latitude = db.Column(db.Float, nullable=False) longitude = db.Column(db.Float, nullable=False) diff --git a/scripts/fetch_stops.py b/scripts/fetch_stops.py index 13060c0..4d8ce3a 100644 --- a/scripts/fetch_stops.py +++ b/scripts/fetch_stops.py @@ -1,59 +1,51 @@ #!/usr/bin/env python3 -"""Fetch all Austrian public-transport stops (and the lines serving them) from -OpenStreetMap via the Overpass API and write a deduplicated snapshot to -data/stops_at.csv.gz. +"""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. -Run this only when you want to refresh the snapshot; the import into the -database reads the checked-in CSV and needs no network access. The route pass -downloads a few hundred MB in ~25 tiles and needs roughly 1 GB of RAM. +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 json import math -import re import sys -import time import unicodedata -import urllib.parse import urllib.request from pathlib import Path -OUT_PATH = Path(__file__).resolve().parent.parent / "data" / "stops_at.csv.gz" +import osmium -OVERPASS_ENDPOINTS = [ - "https://overpass-api.de/api/interpreter", - "https://overpass.kumi.systems/api/interpreter", - "https://overpass.private.coffee/api/interpreter", -] +OUT_PATH = Path(__file__).resolve().parent.parent / "data" / "stops_dach.csv.gz" +CACHE_DIR = Path(__file__).resolve().parent.parent / "data" / ".osm-cache" -# Named nodes in Austria that represent a boardable stop ("Haltestelle"). -STOPS_QUERY = """ -[out:json][timeout:600]; -area["ISO3166-1"="AT"][admin_level=2]->.at; -( - node(area.at)["highway"="bus_stop"]["name"]; - node(area.at)["railway"="tram_stop"]["name"]; - node(area.at)["railway"="station"]["name"]; - node(area.at)["railway"="halt"]["name"]; - node(area.at)["public_transport"="station"]["name"]; - node(area.at)["amenity"="bus_station"]["name"]; -); -out body; -""" +# 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" - -# Austria bbox, tiled so each Overpass response stays small enough to parse. -BBOX = (46.35, 9.50, 49.05, 17.20) # south, west, north, east -LAT_STEP = 0.9 -LON_STEP = 1.0 +ROUTE_MODES = { + "bus", "trolleybus", "tram", "light_rail", "subway", "train", + "monorail", "share_taxi", +} STOP_ROLES = { "stop", "platform", @@ -62,7 +54,7 @@ STOP_ROLES = { } CSV_FIELDS = [ - "osm_type", "osm_id", "name", "stop_type", + "osm_type", "osm_id", "name", "stop_type", "country", "municipality", "latitude", "longitude", "lines", ] @@ -72,6 +64,8 @@ TYPE_RANK = {"subway": 5, "train": 4, "bus_station": 3, "tram": 2, "bus": 1, "ot # 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. @@ -112,46 +106,76 @@ def municipality_of(tags: dict) -> str: return "" -def overpass(query: str, *, tries: int = 4) -> dict: - body = urllib.parse.urlencode({"data": query}).encode() - last_error: Exception | None = None - for attempt in range(tries): - endpoint = OVERPASS_ENDPOINTS[attempt % len(OVERPASS_ENDPOINTS)] - try: - req = urllib.request.Request( - endpoint, data=body, headers={"User-Agent": "pbt-stop-import/1.0"} - ) - with urllib.request.urlopen(req, timeout=600) as resp: - return json.load(resp) - except Exception as exc: # noqa: BLE001 - retry on another mirror - last_error = exc - print(f" {endpoint} failed ({exc}); retrying", file=sys.stderr) - time.sleep(10 * (attempt + 1)) - raise SystemExit(f"Overpass failed after {tries} tries; last error: {last_error}") +# --------------------------------------------------------------------------- # +# 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 build_stop_rows(payload: dict) -> list[dict]: - by_name: dict[str, list[dict]] = {} - for element in payload.get("elements", []): - if element.get("type") != "node": +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 - tags = element.get("tags", {}) name = (tags.get("name") or "").strip() - lat, lon = element.get("lat"), element.get("lon") - if not name or lat is None or lon is None: + if not name: continue - node = { - "osm_id": int(element["id"]), + tags_dict = {t.k: t.v for t in tags} + nodes.append({ + "osm_id": int(obj.id), "name": name, - "stop_type": classify(tags), - "municipality": municipality_of(tags), - "latitude": round(float(lat), 6), - "longitude": round(float(lon), 6), - } - by_name.setdefault(normalize_name(name), []).append(node) + "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(): @@ -180,6 +204,7 @@ def build_stop_rows(payload: dict) -> list[dict]: "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), @@ -194,54 +219,52 @@ def build_stop_rows(payload: dict) -> list[dict]: # --------------------------------------------------------------------------- # # Lines (route relations -> stops) # --------------------------------------------------------------------------- # -def fetch_route_members() -> tuple[list[tuple[str, int]], dict[int, tuple[str, float, float]]]: - """Return (ref, node_id) pairs plus node_id -> (name, lat, lon).""" - seen_rel: set[int] = set() - ref_members: list[tuple[str, int]] = [] +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) - south, west, north, east = BBOX - lat = south - tiles: list[tuple[float, float, float, float]] = [] - while lat < north: - lon = west - while lon < east: - tiles.append((lat, lon, min(lat + LAT_STEP, north), min(lon + LON_STEP, east))) - lon += LON_STEP - lat += LAT_STEP - - for i, (s, w, n, e) in enumerate(tiles, 1): - print(f" route tile {i}/{len(tiles)} ({s:.1f},{w:.1f})", file=sys.stderr) - query = ( - f"[out:json][timeout:400];" - f'relation["type"="route"]["route"~"^({ROUTE_MODES})$"]["ref"]' - f"({s},{w},{n},{e})->.r;" - f".r out body;" - f"node(r.r);out body;" - ) - payload = overpass(query) - for el in payload.get("elements", []): - if el["type"] == "node": - tags = el.get("tags", {}) - node_pos[el["id"]] = (tags.get("name", ""), el["lat"], el["lon"]) - elif el["type"] == "relation": - if el["id"] in seen_rel: - continue - seen_rel.add(el["id"]) - # A few relations pack several numbers into one ref ("407, 413"). - refs = [r.strip() for r in re.split(r"[;,]", el["tags"]["ref"]) if r.strip()] - if not refs: - continue - members = [ - m["ref"] for m in el.get("members", []) - if m["type"] == "node" and m["role"] in STOP_ROLES - ] - for ref in refs: - for node_id in members: - ref_members.append((ref, node_id)) - del payload - - print(f" {len(seen_rel)} routes, {len(ref_members)} stop memberships", file=sys.stderr) + 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 @@ -302,16 +325,31 @@ def line_sort_key(ref: str): # --------------------------------------------------------------------------- # def main() -> None: - print("Fetching stops ...", file=sys.stderr) - rows = build_stop_rows(overpass(STOPS_QUERY)) - if len(rows) < 10_000: + 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 - refusing to overwrite the snapshot." + f"Only {len(rows)} stops parsed (< {MIN_STOPS}) - refusing to overwrite the snapshot." ) - print("Fetching routes ...", file=sys.stderr) - ref_members, node_pos = fetch_route_members() - matched, unmatched = assign_lines(rows, ref_members, node_pos) + 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: @@ -322,10 +360,13 @@ def main() -> None: 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 " diff --git a/static/stops.js b/static/stops.js index 90cdf43..833f964 100644 --- a/static/stops.js +++ b/static/stops.js @@ -152,7 +152,7 @@ } var kind = document.createElement("span"); kind.className = "kind"; - kind.textContent = TYPE_LABEL[stop.type] || ""; + kind.textContent = (TYPE_LABEL[stop.type] || "") + " · " + stop.country; el.appendChild(kind); }, onChoose: function (stop) { input.value = stop.name; markMatch(stop); }, diff --git a/stops_import.py b/stops_import.py index 36036ea..72970f2 100644 --- a/stops_import.py +++ b/stops_import.py @@ -1,8 +1,8 @@ -"""Load the checked-in Austrian stop snapshot into the ``stops`` table. +"""Load the checked-in DACH stop snapshot into the ``stops`` table. -The snapshot (data/stops_at.csv.gz) is produced by scripts/fetch_stops.py from -OpenStreetMap data (© OpenStreetMap contributors, ODbL). Importing needs no -network access. +The snapshot (data/stops_dach.csv.gz) is produced by scripts/fetch_stops.py +from OpenStreetMap data (© OpenStreetMap contributors, ODbL) for Austria, +Germany and Switzerland. Importing needs no network access. flask import-stops # import the checked-in snapshot flask import-stops --file X # import an alternative CSV(.gz) @@ -21,7 +21,7 @@ from flask.cli import with_appcontext from models import Stop, Trip, db -SNAPSHOT_PATH = Path(__file__).resolve().parent / "data" / "stops_at.csv.gz" +SNAPSHOT_PATH = Path(__file__).resolve().parent / "data" / "stops_dach.csv.gz" def normalize_name(value: str) -> str: @@ -69,6 +69,7 @@ def import_stops(path: Path | str = SNAPSHOT_PATH) -> dict[str, int]: name=name, name_normalized=normalize_name(name), stop_type=row["stop_type"] or "other", + country=(row.get("country") or "AT").strip().upper(), municipality=(row.get("municipality") or "").strip() or None, latitude=float(row["latitude"]), longitude=float(row["longitude"]), @@ -120,7 +121,7 @@ def import_stops(path: Path | str = SNAPSHOT_PATH) -> dict[str, int]: ) @with_appcontext def import_stops_command(file_path: str | None) -> None: - """Import Austrian public-transport stops into the database.""" + """Import DACH public-transport stops into the database.""" stats = import_stops(file_path or SNAPSHOT_PATH) click.echo( "Stops imported: "