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

from __future__ import annotations

import argparse
import heapq
import json
import os
import sqlite3
import sys
import time
from collections import Counter
from collections.abc import Callable
from dataclasses import dataclass
from datetime import UTC, datetime
from pathlib import Path

from dataset_schema import (
    MAXIMUM_GEOMETRY_POINTS,
    OSM_ATTRIBUTION,
    OSM_LICENSE,
    OSM_LICENSE_URL,
    ROAD_FLAG_SPEED_CONTROL_AT_END,
    SCHEMA_VERSION,
    allowed_directions,
    create_indexes,
    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
SPEED_ENFORCEMENT_VALUES = frozenset({"maxspeed", "average_speed", "speed_camera"})
MAXIMUM_ENFORCEMENT_PATH_METERS = 10_000.0
MAXIMUM_ENFORCEMENT_PATH_SEGMENTS = 128


@dataclass(frozen=True)
class SpeedEnforcementRelation:
    relation_id: int
    from_node_ids: frozenset[int]
    to_node_ids: frozenset[int]
    device_node_ids: frozenset[int]


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")
    connection.execute("PRAGMA cache_size=-262144")


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,
        "speed_control_nodes": 0,
        "speed_control_relations": 0,
        "speed_control_flags": 0,
        "segments": 0,
    }
    topology_started = time.monotonic()
    junction_index = JunctionIndex(staging.with_suffix(".junctions.sqlite"))
    speed_control_node_ids: set[int] = set()
    related_device_node_ids: set[int] = set()
    enforcement_relations: dict[int, SpeedEnforcementRelation] = {}

    class TopologyHandler(osmium.SimpleHandler):
        def node(self, node: object) -> None:
            statistics["nodes"] += 1
            if node.tags.get("highway", "").strip().lower() == "speed_camera":
                speed_control_node_ids.add(int(node.id))
                statistics["speed_control_nodes"] += 1

        def relation(self, relation: object) -> None:
            statistics["relations"] += 1
            if relation.tags.get("type", "").strip().lower() != "enforcement":
                return
            enforcement_values = {
                value.strip().lower()
                for value in relation.tags.get("enforcement", "").split(";")
                if value.strip()
            }
            if not enforcement_values.intersection(SPEED_ENFORCEMENT_VALUES):
                return
            from_node_ids = frozenset(
                int(member.ref) for member in relation.members if member.type == "n" and member.role == "from"
            )
            to_node_ids = frozenset(
                int(member.ref) for member in relation.members if member.type == "n" and member.role == "to"
            )
            device_node_ids = frozenset(
                int(member.ref) for member in relation.members if member.type == "n" and member.role == "device"
            )
            if not from_node_ids or not (to_node_ids or device_node_ids):
                return
            enforcement_relations[int(relation.id)] = SpeedEnforcementRelation(
                relation_id=int(relation.id),
                from_node_ids=from_node_ids,
                to_node_ids=to_node_ids,
                device_node_ids=device_node_ids,
            )
            related_device_node_ids.update(device_node_ids)
            statistics["speed_control_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()
        junctions.update(speed_control_node_ids)
        for relation in enforcement_relations.values():
            junctions.update(relation.from_node_ids)
            junctions.update(relation.to_node_ids)
            junctions.update(relation.device_node_ids)
    finally:
        junction_index.close()
    topology_seconds = time.monotonic() - topology_started
    relations_by_node = index_enforcement_relations(enforcement_relations.values())
    standalone_speed_control_node_ids = speed_control_node_ids - related_device_node_ids
    resolved_enforcement_relation_ids: set[int] = set()

    connection = sqlite3.connect(staging)
    configure_database(connection)
    create_schema(connection, include_indexes=False)
    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)))
            directional_controls, resolved_relation_ids = speed_controls_for_way(
                [node[0] for node in nodes],
                relations_by_node,
            )
            resolved_enforcement_relation_ids.update(resolved_relation_ids)
            for part_index, part in enumerate(split_way_nodes(nodes, junctions)):
                self._insert_part(int(way.id), part_index, part, tags, directional_controls)

        def _insert_part(
            self,
            way_id: int,
            part_index: int,
            nodes: list[tuple[int, float, float]],
            tags: dict[str, str],
            directional_controls: set[tuple[int, int]],
        ) -> 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"
                has_speed_control_at_end = (
                    to_node_id in standalone_speed_control_node_ids
                    or (to_node_id, direction) in directional_controls
                )
                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, speed_control_at_end=has_speed_control_at_end),
                    "road_name": road_label(tags),
                    "points": points,
                }
                insert_segment(connection, values)
                if has_speed_control_at_end:
                    statistics["speed_control_flags"] += 1
                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()
        create_indexes(connection)
        resolved_enforcement_relation_ids.update(
            resolve_enforcement_relations(
                connection,
                enforcement_relations.values(),
                resolved_enforcement_relation_ids,
            )
        )
        statistics["speed_control_flags"] = int(
            connection.execute(
                "SELECT COUNT(*) FROM road_segments WHERE (road_flags & ?) != 0",
                (ROAD_FLAG_SPEED_CONTROL_AT_END,),
            ).fetchone()[0]
        )
        connection.execute("ANALYZE")
        connection.execute("PRAGMA optimize")
        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["resolved_speed_control_relations"] = len(resolved_enforcement_relation_ids)
    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 index_enforcement_relations(
    relations: object,
) -> dict[int, tuple[SpeedEnforcementRelation, ...]]:
    indexed: dict[int, list[SpeedEnforcementRelation]] = {}
    for relation in relations:
        for node_id in relation.from_node_ids | relation.to_node_ids | relation.device_node_ids:
            indexed.setdefault(node_id, []).append(relation)
    return {node_id: tuple(values) for node_id, values in indexed.items()}


def speed_controls_for_way(
    node_ids: list[int],
    relations_by_node: dict[int, tuple[SpeedEnforcementRelation, ...]],
) -> tuple[set[tuple[int, int]], set[int]]:
    positions = {node_id: index for index, node_id in enumerate(node_ids)}
    candidate_relations = {
        relation
        for node_id in positions
        for relation in relations_by_node.get(node_id, ())
    }
    controls: set[tuple[int, int]] = set()
    resolved_relation_ids: set[int] = set()
    for relation in candidate_relations:
        from_nodes = relation.from_node_ids.intersection(positions)
        device_nodes = relation.device_node_ids.intersection(positions)
        target_nodes = device_nodes or relation.to_node_ids.intersection(positions)
        for from_node_id in from_nodes:
            for target_node_id in target_nodes:
                from_position = positions[from_node_id]
                target_position = positions[target_node_id]
                if from_position == target_position:
                    continue
                direction = 1 if from_position < target_position else -1
                controls.add((target_node_id, direction))
                resolved_relation_ids.add(relation.relation_id)
    return controls, resolved_relation_ids


def resolve_enforcement_relations(
    connection: sqlite3.Connection,
    relations: object,
    already_resolved: set[int],
) -> set[int]:
    resolved: set[int] = set()
    incoming_cache: dict[int, tuple[tuple[int, int, float], ...]] = {}
    node_presence_cache: dict[int, bool] = {}

    def incoming(node_id: int) -> tuple[tuple[int, int, float], ...]:
        if node_id not in incoming_cache:
            incoming_cache[node_id] = tuple(
                (int(row_id), int(from_node_id), float(length_meters))
                for row_id, from_node_id, length_meters in connection.execute(
                    "SELECT row_id, from_node_id, length_meters FROM road_segments WHERE to_node_id = ?",
                    (node_id,),
                )
            )
        return incoming_cache[node_id]

    def node_is_on_graph(node_id: int) -> bool:
        if node_id not in node_presence_cache:
            node_presence_cache[node_id] = bool(
                connection.execute(
                    "SELECT 1 FROM road_segments WHERE from_node_id = ? LIMIT 1",
                    (node_id,),
                ).fetchone()
                or connection.execute(
                    "SELECT 1 FROM road_segments WHERE to_node_id = ? LIMIT 1",
                    (node_id,),
                ).fetchone()
            )
        return node_presence_cache[node_id]

    for relation in relations:
        if relation.relation_id in already_resolved:
            continue
        device_targets = {node_id for node_id in relation.device_node_ids if node_is_on_graph(node_id)}
        target_node_ids = device_targets or {
            node_id for node_id in relation.to_node_ids if node_is_on_graph(node_id)
        }
        approach_row_id = find_approach_segment(
            relation.from_node_ids,
            target_node_ids,
            incoming,
        )
        if approach_row_id is None:
            continue
        connection.execute(
            "UPDATE road_segments SET road_flags = road_flags | ? WHERE row_id = ?",
            (ROAD_FLAG_SPEED_CONTROL_AT_END, approach_row_id),
        )
        resolved.add(relation.relation_id)
    connection.commit()
    return resolved


def find_approach_segment(
    from_node_ids: frozenset[int],
    target_node_ids: set[int],
    incoming_segments: Callable[[int], tuple[tuple[int, int, float], ...]],
) -> int | None:
    queue: list[tuple[float, int, int, int | None]] = [
        (0.0, 0, target_node_id, None)
        for target_node_id in target_node_ids
    ]
    heapq.heapify(queue)
    best_distance: dict[int, float] = {target_node_id: 0.0 for target_node_id in target_node_ids}
    while queue:
        distance_meters, segment_count, node_id, approach_row_id = heapq.heappop(queue)
        if distance_meters > best_distance.get(node_id, float("inf")):
            continue
        if node_id in from_node_ids and approach_row_id is not None:
            return approach_row_id
        if (
            distance_meters >= MAXIMUM_ENFORCEMENT_PATH_METERS
            or segment_count >= MAXIMUM_ENFORCEMENT_PATH_SEGMENTS
        ):
            continue
        for row_id, previous_node_id, length_meters in incoming_segments(node_id):
            next_distance = distance_meters + length_meters
            if next_distance > MAXIMUM_ENFORCEMENT_PATH_METERS:
                continue
            if next_distance >= best_distance.get(previous_node_id, float("inf")):
                continue
            best_distance[previous_node_id] = next_distance
            heapq.heappush(
                queue,
                (
                    next_distance,
                    segment_count + 1,
                    previous_node_id,
                    approach_row_id if approach_row_id is not None else row_id,
                ),
            )
    return 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"],
            "speedControlNodes": statistics["speed_control_nodes"],
            "speedControlRelations": statistics["speed_control_relations"],
            "resolvedSpeedControlRelations": statistics["resolved_speed_control_relations"],
        },
        "output": {
            "segments": statistics["segments"],
            "speedControlApproachSegments": statistics["speed_control_flags"],
            "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
