from __future__ import annotations

import unittest

from build_dataset import (
    SpeedEnforcementRelation,
    find_approach_segment,
    index_enforcement_relations,
    speed_controls_for_way,
    split_way_nodes,
)


class BuildDatasetTest(unittest.TestCase):
    def test_junction_split_preserves_graph_connectivity(self) -> None:
        nodes = [
            (10, 48.0, 2.0),
            (11, 48.0, 2.1),
            (12, 48.0, 2.2),
            (13, 48.0, 2.3),
        ]

        parts = split_way_nodes(nodes, junctions={12})

        self.assertEqual([[10, 11, 12], [12, 13]], [[node[0] for node in part] for part in parts])
        self.assertEqual(parts[0][-1][0], parts[1][0][0])

    def test_enforcement_relation_marks_only_the_approaching_direction(self) -> None:
        relation = SpeedEnforcementRelation(
            relation_id=42,
            from_node_ids=frozenset({10}),
            to_node_ids=frozenset({13}),
            device_node_ids=frozenset({12}),
        )

        controls, resolved = speed_controls_for_way(
            [10, 11, 12, 13],
            index_enforcement_relations([relation]),
        )

        self.assertEqual({(12, 1)}, controls)
        self.assertEqual({42}, resolved)

    def test_reverse_enforcement_relation_marks_the_backward_edge(self) -> None:
        relation = SpeedEnforcementRelation(
            relation_id=43,
            from_node_ids=frozenset({13}),
            to_node_ids=frozenset({10}),
            device_node_ids=frozenset({11}),
        )

        controls, resolved = speed_controls_for_way(
            [10, 11, 12, 13],
            index_enforcement_relations([relation]),
        )

        self.assertEqual({(11, -1)}, controls)
        self.assertEqual({43}, resolved)

    def test_enforcement_relation_follows_a_directed_path_across_ways(self) -> None:
        incoming = {
            20: ((2, 10, 100.0),),
            30: ((3, 20, 100.0),),
        }

        approach_row_id = find_approach_segment(
            frozenset({10}),
            {30},
            lambda node_id: incoming.get(node_id, ()),
        )

        self.assertEqual(3, approach_row_id)

    def test_enforcement_relation_does_not_invent_a_disconnected_path(self) -> None:
        incoming = {
            30: ((3, 20, 100.0),),
        }

        approach_row_id = find_approach_segment(
            frozenset({10}),
            {30},
            lambda node_id: incoming.get(node_id, ()),
        )

        self.assertIsNone(approach_row_id)


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