#!/usr/bin/env python3
"""Schema, filtering and normalization helpers for NextLimit mobile datasets."""

from __future__ import annotations

import math
import re
import sqlite3
import struct
from collections.abc import Mapping, Sequence
from dataclasses import dataclass

SCHEMA_VERSION = 2
OSM_ATTRIBUTION = "© OpenStreetMap contributors"
OSM_LICENSE = "ODbL-1.0"
OSM_LICENSE_URL = "https://www.openstreetmap.org/copyright"
COORDINATE_SCALE = 10_000_000
MAXIMUM_GEOMETRY_POINTS = 4_096

# Stable numeric values shared with AndroidRoadSegmentStore. Append only.
ROAD_CLASS_CODES = {
    "motorway": 1,
    "motorway_link": 2,
    "trunk": 3,
    "trunk_link": 4,
    "primary": 5,
    "primary_link": 6,
    "secondary": 7,
    "secondary_link": 8,
    "tertiary": 9,
    "tertiary_link": 10,
    "unclassified": 11,
    "residential": 12,
    "living_street": 13,
    "service": 14,
}
ROAD_CLASS_NAMES = {value: key for key, value in ROAD_CLASS_CODES.items()}
DRIVABLE_HIGHWAYS = frozenset(ROAD_CLASS_CODES)

RESTRICTED_ACCESS = frozenset({"no", "private", "agricultural", "forestry", "emergency", "military"})
ALLOWED_ACCESS = frozenset({"yes", "designated", "permissive", "destination", "customers", "delivery"})
EXCLUDED_SERVICE_TYPES = frozenset({"parking_aisle", "driveway", "drive-through", "emergency_access"})
TRUE_ONEWAY = frozenset({"yes", "1", "true"})
FALSE_ONEWAY = frozenset({"no", "0", "false"})

ROAD_FLAG_ONEWAY = 1
ROAD_FLAG_ROUNDABOUT = 1 << 1

SPEED_SOURCE_UNKNOWN = 0
SPEED_SOURCE_OSM_EXPLICIT = 1
SPEED_SOURCE_OSM_DIRECTIONAL = 2
SPEED_SOURCE_OSM_RULE = 3
SPEED_SOURCE_LEGAL_DEFAULT = 4

UNKNOWN_MAXSPEED_VALUES = frozenset({"none", "signals", "variable", "walk", "implicit", "unposted", "unknown"})
SPEED_PATTERN = re.compile(r"^([0-9]{1,3}(?:\.[0-9]+)?)\s*(mph|km/?h|kph)?$")
MINIMUM_SUPPORTED_SPEED_KPH = 5
MAXIMUM_SUPPORTED_SPEED_KPH = 250

ROAD_SEGMENT_COLUMNS = (
    "segment_id",
    "osm_way_id",
    "from_node_id",
    "to_node_id",
    "direction",
    "length_meters",
    "geometry",
    "road_class",
    "speed_limit_kmh",
    "speed_source",
    "speed_confidence",
    "road_flags",
    "road_name",
)

SCHEMA_SQL = """
CREATE TABLE dataset_metadata (
    schema_version INTEGER NOT NULL,
    dataset_id TEXT NOT NULL,
    country_code TEXT NOT NULL,
    region_code TEXT NOT NULL,
    source_file TEXT NOT NULL,
    source_timestamp TEXT NOT NULL,
    generated_at TEXT NOT NULL,
    min_lat_e7 INTEGER NOT NULL,
    max_lat_e7 INTEGER NOT NULL,
    min_lon_e7 INTEGER NOT NULL,
    max_lon_e7 INTEGER NOT NULL,
    segment_count INTEGER NOT NULL,
    osm_attribution TEXT NOT NULL,
    osm_license TEXT NOT NULL,
    osm_license_url TEXT NOT NULL
);

CREATE TABLE road_segments (
    row_id INTEGER PRIMARY KEY,
    segment_id TEXT NOT NULL UNIQUE,
    osm_way_id INTEGER NOT NULL,
    from_node_id INTEGER NOT NULL,
    to_node_id INTEGER NOT NULL,
    direction INTEGER NOT NULL CHECK (direction IN (-1, 1)),
    length_meters REAL NOT NULL CHECK (length_meters > 0),
    geometry BLOB NOT NULL,
    road_class INTEGER NOT NULL CHECK (road_class BETWEEN 1 AND 14),
    speed_limit_kmh INTEGER CHECK (speed_limit_kmh BETWEEN 5 AND 250),
    speed_source INTEGER NOT NULL CHECK (speed_source BETWEEN 0 AND 4),
    speed_confidence INTEGER NOT NULL CHECK (speed_confidence BETWEEN 0 AND 100),
    road_flags INTEGER NOT NULL DEFAULT 0,
    road_name TEXT,
    CHECK (
        (speed_limit_kmh IS NULL AND speed_source = 0 AND speed_confidence = 0)
        OR
        (speed_limit_kmh IS NOT NULL AND speed_source != 0 AND speed_confidence > 0)
    )
);

CREATE INDEX road_segments_from_node_idx ON road_segments(from_node_id);
CREATE INDEX road_segments_to_node_idx ON road_segments(to_node_id);
CREATE INDEX road_segments_way_idx ON road_segments(osm_way_id, direction);

CREATE VIRTUAL TABLE road_segment_rtree USING rtree_i32(
    row_id,
    min_lat_e7,
    max_lat_e7,
    min_lon_e7,
    max_lon_e7
);
"""

INSERT_SEGMENT_SQL = f"""
INSERT INTO road_segments ({", ".join(ROAD_SEGMENT_COLUMNS)})
VALUES ({", ".join(f":{column}" for column in ROAD_SEGMENT_COLUMNS)})
"""


@dataclass(frozen=True)
class NormalizedSpeedLimit:
    value_kmh: int | None
    source: int = SPEED_SOURCE_UNKNOWN
    confidence_percent: int = 0


UNKNOWN_SPEED_LIMIT = NormalizedSpeedLimit(None)


def create_schema(connection: sqlite3.Connection) -> None:
    connection.executescript(SCHEMA_SQL)


def encode_geometry(points: Sequence[tuple[float, float]]) -> bytes:
    if len(points) < 2:
        raise ValueError("a road geometry requires at least two points")
    if len(points) > MAXIMUM_GEOMETRY_POINTS:
        raise ValueError(f"a road geometry cannot exceed {MAXIMUM_GEOMETRY_POINTS} points")
    values: list[int] = [len(points)]
    for latitude, longitude in points:
        values.extend((round(latitude * COORDINATE_SCALE), round(longitude * COORDINATE_SCALE)))
    return struct.pack(f">{len(values)}i", *values)


def haversine_meters(first: tuple[float, float], second: tuple[float, float]) -> float:
    first_latitude = math.radians(first[0])
    second_latitude = math.radians(second[0])
    latitude_delta = math.radians(second[0] - first[0])
    longitude_delta = math.radians(second[1] - first[1])
    value = (
        math.sin(latitude_delta / 2) ** 2
        + math.cos(first_latitude) * math.cos(second_latitude) * math.sin(longitude_delta / 2) ** 2
    )
    return 6_371_000.0 * 2 * math.atan2(math.sqrt(value), math.sqrt(1 - value))


def is_drivable(tags: Mapping[str, str]) -> bool:
    highway = tags.get("highway", "").strip().lower()
    if highway not in DRIVABLE_HIGHWAYS:
        return False
    if highway == "service" and tags.get("service", "").strip().lower() in EXCLUDED_SERVICE_TYPES:
        return False
    for key in ("motorcar", "motor_vehicle"):
        value = tags.get(key, "").strip().lower()
        if value in ALLOWED_ACCESS:
            return True
        if value in RESTRICTED_ACCESS:
            return False
    return tags.get("access", "").strip().lower() not in RESTRICTED_ACCESS


def allowed_directions(tags: Mapping[str, str]) -> tuple[int, ...]:
    oneway = tags.get("oneway", "").strip().lower()
    if oneway == "-1":
        return (-1,)
    if oneway in TRUE_ONEWAY:
        return (1,)
    if oneway in FALSE_ONEWAY:
        return (1, -1)
    if tags.get("junction", "").lower() == "roundabout" or tags.get("highway", "").lower() == "motorway":
        return (1,)
    return (1, -1)


def road_flags(tags: Mapping[str, str]) -> int:
    value = 0
    if len(allowed_directions(tags)) == 1:
        value |= ROAD_FLAG_ONEWAY
    if tags.get("junction", "").strip().lower() == "roundabout":
        value |= ROAD_FLAG_ROUNDABOUT
    return value


def road_class_code(highway: str) -> int:
    try:
        return ROAD_CLASS_CODES[highway.strip().lower()]
    except KeyError as error:
        raise ValueError(f"unsupported road class: {highway}") from error


def road_label(tags: Mapping[str, str]) -> str | None:
    for key in ("name", "ref"):
        value = tags.get(key, "").strip()
        if value:
            return value[:160]
    return None


def resolve_speed_limit(
    tags: Mapping[str, str],
    country_code: str,
    direction: int,
) -> NormalizedSpeedLimit:
    directional_key = "maxspeed:forward" if direction == 1 else "maxspeed:backward"
    directional_value = tags.get(directional_key, "").strip()
    if directional_value:
        return _resolve_tagged_value(
            directional_value,
            country_code,
            SPEED_SOURCE_OSM_DIRECTIONAL,
        )

    explicit_value = tags.get("maxspeed", "").strip()
    if explicit_value:
        return _resolve_tagged_value(
            explicit_value,
            country_code,
            SPEED_SOURCE_OSM_EXPLICIT,
        )

    for key in ("zone:maxspeed", "maxspeed:type", "source:maxspeed"):
        rule = _country_rule(country_code, tags.get(key, ""))
        if rule is not None:
            return rule

    return _legal_default(country_code, tags.get("highway", "")) or UNKNOWN_SPEED_LIMIT


def parse_speed_kmh(raw_value: str) -> int | None:
    normalized = raw_value.strip().lower()
    if ";" in normalized or normalized in UNKNOWN_MAXSPEED_VALUES:
        return None
    match = SPEED_PATTERN.fullmatch(normalized)
    if match is None:
        return None
    numeric_value = float(match.group(1))
    rounded_value = _round_positive(numeric_value)
    speed_kmh = _round_positive(rounded_value * 1.609344) if match.group(2) == "mph" else rounded_value
    if speed_kmh not in range(MINIMUM_SUPPORTED_SPEED_KPH, MAXIMUM_SUPPORTED_SPEED_KPH + 1):
        return None
    return speed_kmh


def insert_segment(connection: sqlite3.Connection, values: dict[str, object]) -> int:
    cursor = connection.execute(INSERT_SEGMENT_SQL, values)
    row_id = int(cursor.lastrowid)
    points = values["points"]
    if not isinstance(points, Sequence):
        raise TypeError("points must be a coordinate sequence")
    latitudes = [round(float(point[0]) * COORDINATE_SCALE) for point in points]
    longitudes = [round(float(point[1]) * COORDINATE_SCALE) for point in points]
    connection.execute(
        "INSERT INTO road_segment_rtree VALUES (?, ?, ?, ?, ?)",
        (row_id, min(latitudes), max(latitudes), min(longitudes), max(longitudes)),
    )
    return row_id


def _resolve_tagged_value(
    raw_value: str,
    country_code: str,
    numeric_source: int,
) -> NormalizedSpeedLimit:
    numeric = parse_speed_kmh(raw_value)
    if numeric is not None:
        return NormalizedSpeedLimit(numeric, numeric_source, 99)
    return _country_rule(country_code, raw_value) or UNKNOWN_SPEED_LIMIT


def _country_rule(country_code: str, raw_value: str) -> NormalizedSpeedLimit | None:
    raw = raw_value.strip().lower()
    if not raw:
        return None
    country = country_code.upper()
    normalized = raw.removeprefix(f"{country.lower()}:")
    if country == "FR":
        normalized = normalized.removeprefix("zone")
        rules = {
            "30": (30, 90),
            "20": (20, 90),
            "living_street": (20, 82),
            "urban": (50, 82),
            "motorway": (130, 82),
            "rural": (80, 68),
        }
    elif country == "BE":
        rules = {"zone30": (30, 90), "30": (30, 90), "motorway": (120, 82)}
    elif country == "DE":
        rules = {"zone30": (30, 90), "30": (30, 90), "urban": (50, 82), "rural": (100, 82)}
    else:
        return None
    resolved = rules.get(normalized)
    return None if resolved is None else NormalizedSpeedLimit(resolved[0], SPEED_SOURCE_OSM_RULE, resolved[1])


def _legal_default(country_code: str, highway: str) -> NormalizedSpeedLimit | None:
    key = (country_code.upper(), highway.strip().lower())
    rules = {
        ("FR", "motorway"): 130,
        ("FR", "living_street"): 20,
        ("BE", "motorway"): 120,
    }
    value = rules.get(key)
    return None if value is None else NormalizedSpeedLimit(value, SPEED_SOURCE_LEGAL_DEFAULT, 55)


def _round_positive(value: float) -> int:
    return math.floor(value + 0.5)
