#!/usr/bin/env python3
"""Build a compact regional NextLimit mobile dataset from an OSM .osm.pbf extract."""

from __future__ import annotations

import argparse
import json
import os
import sqlite3
import sys
import time
from collections import Counter
from datetime import UTC, datetime
from pathlib import Path

from dataset_schema import (
    MAXIMUM_GEOMETRY_POINTS,
    OSM_ATTRIBUTION,
    OSM_LICENSE,
    OSM_LICENSE_URL,
    SCHEMA_VERSION,
    allowed_directions,
    create_schema,
    encode_geometry,
    haversine_meters,
    insert_segment,
    is_drivable,
    resolve_speed_limit,
    road_class_code,
    road_flags,
    road_label,
)

EXTRACTED_TAGS = (
    "highway",
    "maxspeed",
    "maxspeed:forward",
    "maxspeed:backward",
    "oneway",
    "junction",
    "zone:maxspeed",
    "source:maxspeed",
    "maxspeed:type",
    "access",
    "motor_vehicle",
    "motorcar",
    "service",
    "ref",
    "name",
)
INSERT_BATCH_SIZE = 20_000
NODE_USAGE_BATCH_SIZE = 100_000


def parse_arguments() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("input", type=Path, help="regional .osm.pbf input")
    parser.add_argument("output", type=Path, help="output .sqlite mobile dataset")
    parser.add_argument("--dataset-id", required=True, help="stable identifier, for example FR-IDF-2026.08")
    parser.add_argument("--country", default="FR", help="ISO 3166-1 alpha-2 code (default: FR)")
    parser.add_argument("--region", required=True, help="region identifier from regions.json")
    parser.add_argument("--source-date", help="OSM extract timestamp in ISO-8601")
    parser.add_argument("--report", type=Path, help="JSON report path; defaults next to the dataset")
    parser.add_argument(
        "--location-index",
        default="sparse_file_array",
        help="libosmium node-location index; disk-backed by default for large regions",
    )
    parser.add_argument("--force", action="store_true", help="replace the exact output and report files")
    return parser.parse_args()


def configure_database(connection: sqlite3.Connection) -> None:
    connection.execute("PRAGMA journal_mode=OFF")
    connection.execute("PRAGMA synchronous=OFF")
    connection.execute("PRAGMA temp_store=MEMORY")
    connection.execute("PRAGMA locking_mode=EXCLUSIVE")


class JunctionIndex:
    """Disk-backed reference counter; only the much smaller junction set is kept in RAM."""

    def __init__(self, path: Path) -> None:
        path.unlink(missing_ok=True)
        self.path = path
        self.connection = sqlite3.connect(path)
        configure_database(self.connection)
        self.connection.execute(
            "CREATE TABLE node_usage (node_id INTEGER PRIMARY KEY, use_count INTEGER NOT NULL) WITHOUT ROWID"
        )
        self.pending: list[tuple[int]] = []

    def record_way(self, node_ids: list[int]) -> None:
        self.pending.extend((node_id,) for node_id in set(node_ids))
        if len(self.pending) >= NODE_USAGE_BATCH_SIZE:
            self.flush()

    def flush(self) -> None:
        if not self.pending:
            return
        self.connection.executemany(
            "INSERT INTO node_usage VALUES (?, 1) "
            "ON CONFLICT(node_id) DO UPDATE SET use_count = MIN(2, use_count + 1)",
            self.pending,
        )
        self.connection.commit()
        self.pending.clear()

    def junctions(self) -> set[int]:
        self.flush()
        return {int(row[0]) for row in self.connection.execute("SELECT node_id FROM node_usage WHERE use_count > 1")}

    def close(self) -> None:
        self.connection.close()
        self.path.unlink(missing_ok=True)


def run_import(arguments: argparse.Namespace, staging: Path) -> dict[str, object]:
    try:
        import osmium
    except ImportError as error:
        raise RuntimeError("pyosmium is required: python -m pip install -r requirements.txt") from error

    statistics: dict[str, int | float] = {
        "nodes": 0,
        "ways": 0,
        "relations": 0,
        "highway_ways": 0,
        "drivable_ways": 0,
        "filtered_highway_ways": 0,
        "highway_directed_edges": 0,
        "drivable_directed_edges": 0,
        "filtered_directed_edges": 0,
        "missing_location_ways": 0,
        "segments": 0,
    }
    topology_started = time.monotonic()
    junction_index = JunctionIndex(staging.with_suffix(".junctions.sqlite"))

    class TopologyHandler(osmium.SimpleHandler):
        def node(self, node: object) -> None:
            statistics["nodes"] += 1

        def relation(self, relation: object) -> None:
            statistics["relations"] += 1

        def way(self, way: object) -> None:
            statistics["ways"] += 1
            highway = way.tags.get("highway")
            if highway is None:
                return
            statistics["highway_ways"] += 1
            tags = extracted_tags(way)
            edge_count = max(0, len(way.nodes) - 1)
            directed_edges = edge_count * len(allowed_directions(tags))
            statistics["highway_directed_edges"] += directed_edges
            if not is_drivable(tags):
                statistics["filtered_highway_ways"] += 1
                statistics["filtered_directed_edges"] += directed_edges
                return
            statistics["drivable_ways"] += 1
            statistics["drivable_directed_edges"] += directed_edges
            junction_index.record_way([int(node.ref) for node in way.nodes])

    try:
        TopologyHandler().apply_file(str(arguments.input))
        junctions = junction_index.junctions()
    finally:
        junction_index.close()
    topology_seconds = time.monotonic() - topology_started

    connection = sqlite3.connect(staging)
    configure_database(connection)
    create_schema(connection)
    build_started = time.monotonic()

    class RoadHandler(osmium.SimpleHandler):
        def __init__(self) -> None:
            super().__init__()
            self.minimum_latitude_e7: int | None = None
            self.maximum_latitude_e7: int | None = None
            self.minimum_longitude_e7: int | None = None
            self.maximum_longitude_e7: int | None = None

        def way(self, way: object) -> None:
            tags = extracted_tags(way)
            if not is_drivable(tags) or len(way.nodes) < 2:
                return
            nodes = []
            for node in way.nodes:
                if not node.location.valid():
                    statistics["missing_location_ways"] += 1
                    return
                nodes.append((int(node.ref), float(node.location.lat), float(node.location.lon)))
            for part_index, part in enumerate(split_way_nodes(nodes, junctions)):
                self._insert_part(int(way.id), part_index, part, tags)

        def _insert_part(
            self,
            way_id: int,
            part_index: int,
            nodes: list[tuple[int, float, float]],
            tags: dict[str, str],
        ) -> None:
            forward_points = tuple((node[1], node[2]) for node in nodes)
            length_meters = sum(haversine_meters(first, second) for first, second in zip(forward_points, forward_points[1:]))
            if length_meters <= 0.01:
                return
            for direction in allowed_directions(tags):
                if direction == 1:
                    from_node_id, to_node_id = nodes[0][0], nodes[-1][0]
                    points = forward_points
                    suffix = "f"
                else:
                    from_node_id, to_node_id = nodes[-1][0], nodes[0][0]
                    points = tuple(reversed(forward_points))
                    suffix = "b"
                speed = resolve_speed_limit(tags, arguments.country, direction)
                values: dict[str, object] = {
                    "segment_id": f"{way_id}:{part_index}:{suffix}",
                    "osm_way_id": way_id,
                    "from_node_id": from_node_id,
                    "to_node_id": to_node_id,
                    "direction": direction,
                    "length_meters": length_meters,
                    "geometry": encode_geometry(points),
                    "road_class": road_class_code(tags["highway"]),
                    "speed_limit_kmh": speed.value_kmh,
                    "speed_source": speed.source,
                    "speed_confidence": speed.confidence_percent,
                    "road_flags": road_flags(tags),
                    "road_name": road_label(tags),
                    "points": points,
                }
                insert_segment(connection, values)
                self._extend_bounds(points)
                statistics["segments"] += 1
                if statistics["segments"] % INSERT_BATCH_SIZE == 0:
                    connection.commit()

        def _extend_bounds(self, points: tuple[tuple[float, float], ...]) -> None:
            for latitude, longitude in points:
                latitude_e7 = round(latitude * 10_000_000)
                longitude_e7 = round(longitude * 10_000_000)
                self.minimum_latitude_e7 = _minimum(self.minimum_latitude_e7, latitude_e7)
                self.maximum_latitude_e7 = _maximum(self.maximum_latitude_e7, latitude_e7)
                self.minimum_longitude_e7 = _minimum(self.minimum_longitude_e7, longitude_e7)
                self.maximum_longitude_e7 = _maximum(self.maximum_longitude_e7, longitude_e7)

    handler = RoadHandler()
    try:
        handler.apply_file(str(arguments.input), locations=True, idx=arguments.location_index)
        if statistics["segments"] == 0:
            raise RuntimeError("the extract produced no drivable road segment")
        generated_at = datetime.now(UTC).isoformat()
        source_timestamp = arguments.source_date or datetime.fromtimestamp(arguments.input.stat().st_mtime, UTC).isoformat()
        connection.execute(
            "INSERT INTO dataset_metadata VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)",
            (
                SCHEMA_VERSION,
                arguments.dataset_id,
                arguments.country.upper(),
                arguments.region,
                arguments.input.name,
                source_timestamp,
                generated_at,
                handler.minimum_latitude_e7,
                handler.maximum_latitude_e7,
                handler.minimum_longitude_e7,
                handler.maximum_longitude_e7,
                statistics["segments"],
                OSM_ATTRIBUTION,
                OSM_LICENSE,
                OSM_LICENSE_URL,
            ),
        )
        connection.commit()
        connection.execute("ANALYZE")
        integrity = connection.execute("PRAGMA integrity_check").fetchone()
        if integrity != ("ok",):
            raise RuntimeError(f"SQLite integrity check failed: {integrity}")
        rtree_integrity = connection.execute("SELECT rtreecheck('road_segment_rtree')").fetchone()
        if rtree_integrity != ("ok",):
            raise RuntimeError(f"RTree integrity check failed: {rtree_integrity}")
        connection.commit()
    finally:
        connection.close()
    statistics["junction_nodes"] = len(junctions)
    statistics["topology_seconds"] = round(topology_seconds, 3)
    statistics["build_seconds"] = round(time.monotonic() - build_started, 3)
    return statistics


def extracted_tags(way: object) -> dict[str, str]:
    return {key: way.tags.get(key) for key in EXTRACTED_TAGS if way.tags.get(key) is not None}


def split_way_nodes(
    nodes: list[tuple[int, float, float]],
    junctions: set[int],
) -> list[list[tuple[int, float, float]]]:
    if len(nodes) < 2:
        return []
    local_usage = Counter(node[0] for node in nodes)
    result: list[list[tuple[int, float, float]]] = []
    current = [nodes[0]]
    for index, node in enumerate(nodes[1:], start=1):
        current.append(node)
        is_last = index == len(nodes) - 1
        must_split = node[0] in junctions or local_usage[node[0]] > 1 or len(current) >= MAXIMUM_GEOMETRY_POINTS
        if is_last or must_split:
            if len(current) >= 2:
                result.append(current)
            current = [node]
    return result


def build_report(
    arguments: argparse.Namespace,
    statistics: dict[str, object],
    elapsed_seconds: float,
) -> dict[str, object]:
    source_size = arguments.input.stat().st_size
    dataset_size = arguments.output.stat().st_size
    total_objects = int(statistics["nodes"]) + int(statistics["ways"]) + int(statistics["relations"])
    collapsed = int(statistics["drivable_directed_edges"]) - int(statistics["segments"])
    return {
        "region": arguments.region,
        "country": arguments.country.upper(),
        "schemaVersion": SCHEMA_VERSION,
        "source": {
            "file": arguments.input.name,
            "sizeBytes": source_size,
            "totalObjects": total_objects,
            "nodes": statistics["nodes"],
            "ways": statistics["ways"],
            "relations": statistics["relations"],
            "highwayWays": statistics["highway_ways"],
        },
        "filtering": {
            "drivableWays": statistics["drivable_ways"],
            "filteredHighwayWays": statistics["filtered_highway_ways"],
            "sourceHighwayDirectedEdges": statistics["highway_directed_edges"],
            "filteredDirectedEdges": statistics["filtered_directed_edges"],
            "drivableDirectedEdges": statistics["drivable_directed_edges"],
            "junctionNodes": statistics["junction_nodes"],
            "missingLocationWays": statistics["missing_location_ways"],
        },
        "output": {
            "segments": statistics["segments"],
            "filteredSegments": statistics["filtered_directed_edges"],
            "collapsedGeometryEdges": max(0, collapsed),
            "sizeBytes": dataset_size,
            "datasetToPbfRatio": round(dataset_size / source_size, 4),
            "reductionPercentVsPbf": round((1 - dataset_size / source_size) * 100, 2),
        },
        "generation": {
            "seconds": round(elapsed_seconds, 3),
            "topologyPassSeconds": statistics["topology_seconds"],
            "datasetPassSeconds": statistics["build_seconds"],
        },
    }


def _minimum(current: int | None, value: int) -> int:
    return value if current is None else min(current, value)


def _maximum(current: int | None, value: int) -> int:
    return value if current is None else max(current, value)


def _write_report(path: Path, value: dict[str, object]) -> None:
    path.parent.mkdir(parents=True, exist_ok=True)
    staging = path.with_suffix(path.suffix + ".staging")
    staging.write_text(json.dumps(value, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
    os.replace(staging, path)


def main() -> int:
    arguments = parse_arguments()
    if not arguments.input.is_file():
        raise SystemExit(f"input not found: {arguments.input}")
    if not arguments.input.name.endswith(".osm.pbf"):
        raise SystemExit("input must be an .osm.pbf extract")
    arguments.output.parent.mkdir(parents=True, exist_ok=True)
    report_path = arguments.report or arguments.output.with_suffix(".report.json")
    if (arguments.output.exists() or report_path.exists()) and not arguments.force:
        raise SystemExit(f"output or report already exists; use --force to replace it")
    staging = arguments.output.with_suffix(arguments.output.suffix + ".staging")
    staging.unlink(missing_ok=True)
    started = time.monotonic()
    try:
        statistics = run_import(arguments, staging)
        os.replace(staging, arguments.output)
        report = build_report(arguments, statistics, time.monotonic() - started)
        _write_report(report_path, report)
    except Exception:
        staging.unlink(missing_ok=True)
        raise
    print(
        json.dumps(
            {
                "dataset": str(arguments.output),
                "report": str(report_path),
                "segments": statistics["segments"],
                "sizeBytes": arguments.output.stat().st_size,
            }
        )
    )
    return 0


if __name__ == "__main__":
    try:
        raise SystemExit(main())
    except RuntimeError as error:
        print(str(error), file=sys.stderr)
        raise SystemExit(1) from error
