from __future__ import annotations

import sqlite3
import struct
import unittest

from dataset_schema import (
    ROAD_CLASS_CODES,
    SPEED_SOURCE_OSM_DIRECTIONAL,
    SPEED_SOURCE_OSM_EXPLICIT,
    allowed_directions,
    create_schema,
    encode_geometry,
    insert_segment,
    is_drivable,
    resolve_speed_limit,
)


class DatasetSchemaTest(unittest.TestCase):
    def test_schema_indexes_candidates_and_keeps_graph_nodes(self) -> None:
        connection = sqlite3.connect(":memory:")
        create_schema(connection)
        points = ((48.85, 2.34), (48.851, 2.341))
        row_id = insert_segment(connection, _segment(points))

        indexed = connection.execute(
            "SELECT row_id FROM road_segment_rtree WHERE min_lat_e7 <= ? AND max_lat_e7 >= ?",
            (488505000, 488505000),
        ).fetchall()
        graph_node = connection.execute("SELECT from_node_id, to_node_id FROM road_segments").fetchone()
        normalized = connection.execute(
            "SELECT road_class, speed_limit_kmh, speed_source FROM road_segments"
        ).fetchone()

        self.assertEqual([(row_id,)], indexed)
        self.assertEqual((10, 11), graph_node)
        self.assertEqual((ROAD_CLASS_CODES["primary"], 80, SPEED_SOURCE_OSM_EXPLICIT), normalized)
        self.assertEqual(("ok",), connection.execute("SELECT rtreecheck('road_segment_rtree')").fetchone())

    def test_geometry_uses_big_endian_e7_coordinates(self) -> None:
        value = encode_geometry(((48.85, 2.34), (48.851, 2.341)))

        decoded = struct.unpack(">5i", value)

        self.assertEqual((2, 488500000, 23400000, 488510000, 23410000), decoded)

    def test_motorway_and_residential_are_kept(self) -> None:
        self.assertTrue(is_drivable({"highway": "motorway", "maxspeed": "130"}))
        self.assertTrue(is_drivable({"highway": "residential"}))

    def test_non_motorized_highways_are_removed(self) -> None:
        for highway in ("footway", "path", "steps", "bridleway", "cycleway", "pedestrian"):
            with self.subTest(highway=highway):
                self.assertFalse(is_drivable({"highway": highway}))

    def test_motor_vehicle_restriction_and_motorcar_override_are_respected(self) -> None:
        self.assertFalse(is_drivable({"highway": "residential", "motor_vehicle": "no"}))
        self.assertTrue(
            is_drivable({"highway": "service", "access": "private", "motor_vehicle": "no", "motorcar": "yes"})
        )
        self.assertFalse(is_drivable({"highway": "service", "service": "parking_aisle"}))

    def test_directional_speed_is_normalized_per_directed_edge(self) -> None:
        tags = {
            "highway": "primary",
            "maxspeed": "80",
            "maxspeed:forward": "90",
            "maxspeed:backward": "70",
        }

        forward = resolve_speed_limit(tags, "FR", 1)
        backward = resolve_speed_limit(tags, "FR", -1)

        self.assertEqual((90, SPEED_SOURCE_OSM_DIRECTIONAL, 99), (forward.value_kmh, forward.source, forward.confidence_percent))
        self.assertEqual((70, SPEED_SOURCE_OSM_DIRECTIONAL, 99), (backward.value_kmh, backward.source, backward.confidence_percent))

    def test_oneway_and_roundabout_create_only_valid_directions(self) -> None:
        self.assertEqual((1,), allowed_directions({"highway": "primary", "oneway": "yes"}))
        self.assertEqual((-1,), allowed_directions({"highway": "primary", "oneway": "-1"}))
        self.assertEqual((1,), allowed_directions({"highway": "primary", "junction": "roundabout"}))
        self.assertEqual((1, -1), allowed_directions({"highway": "primary", "oneway": "no"}))


def _segment(points: tuple[tuple[float, float], ...]) -> dict[str, object]:
    return {
        "segment_id": "1:0:f",
        "osm_way_id": 1,
        "from_node_id": 10,
        "to_node_id": 11,
        "direction": 1,
        "length_meters": 100.0,
        "geometry": encode_geometry(points),
        "road_class": ROAD_CLASS_CODES["primary"],
        "speed_limit_kmh": 80,
        "speed_source": SPEED_SOURCE_OSM_EXPLICIT,
        "speed_confidence": 99,
        "road_flags": 0,
        "road_name": "D1",
        "points": points,
    }


if __name__ == "__main__":
    unittest.main()
