from __future__ import annotations

import sqlite3
import tempfile
import unittest
from pathlib import Path

from dataset_diff import create_diff
from dataset_schema import (
    OSM_ATTRIBUTION,
    OSM_LICENSE,
    OSM_LICENSE_URL,
    ROAD_CLASS_CODES,
    SCHEMA_VERSION,
    SPEED_SOURCE_OSM_EXPLICIT,
    create_schema,
    encode_geometry,
    insert_segment,
)


class DatasetDiffTest(unittest.TestCase):
    def test_diff_contains_only_deleted_and_changed_segments(self) -> None:
        with tempfile.TemporaryDirectory() as directory:
            root = Path(directory)
            previous = root / "previous.sqlite"
            current = root / "current.sqlite"
            patch = root / "patch.sqlite"
            _dataset(previous, "old", {"A": 90, "B": 90})
            _dataset(current, "new", {"B": 70, "C": 50})

            deleted, upserted = create_diff(previous, current, patch, "old", "new")

            connection = sqlite3.connect(patch)
            self.assertEqual(1, deleted)
            self.assertEqual(2, upserted)
            self.assertEqual([("A",)], connection.execute("SELECT segment_id FROM deleted_segments").fetchall())
            self.assertEqual(
                [("B", 70), ("C", 50)],
                connection.execute(
                    "SELECT segment_id, speed_limit_kmh FROM upsert_segments ORDER BY segment_id"
                ).fetchall(),
            )
            self.assertEqual(("new",), connection.execute("SELECT dataset_id FROM target_metadata").fetchone())
            connection.close()


def _dataset(path: Path, dataset_id: str, limits: dict[str, int]) -> None:
    connection = sqlite3.connect(path)
    create_schema(connection)
    for index, (segment_id, limit) in enumerate(limits.items()):
        points = ((48.0, 2.0 + index * 0.001), (48.0, 2.001 + index * 0.001))
        insert_segment(connection, _segment(segment_id, limit, index, points))
    connection.execute(
        "INSERT INTO dataset_metadata VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)",
        (
            SCHEMA_VERSION,
            dataset_id,
            "FR",
            "FR-IDF",
            "fixture.osm.pbf",
            "2026-08-09T00:00:00Z",
            "2026-08-09T00:00:00Z",
            480000000,
            480000000,
            20000000,
            20100000,
            len(limits),
            OSM_ATTRIBUTION,
            OSM_LICENSE,
            OSM_LICENSE_URL,
        ),
    )
    connection.commit()
    connection.close()


def _segment(
    segment_id: str,
    limit: int,
    index: int,
    points: tuple[tuple[float, float], ...],
) -> dict[str, object]:
    return {
        "segment_id": segment_id,
        "osm_way_id": index + 1,
        "from_node_id": index + 1,
        "to_node_id": index + 2,
        "direction": 1,
        "length_meters": 100.0,
        "geometry": encode_geometry(points),
        "road_class": ROAD_CLASS_CODES["primary"],
        "speed_limit_kmh": limit,
        "speed_source": SPEED_SOURCE_OSM_EXPLICIT,
        "speed_confidence": 99,
        "road_flags": 1,
        "road_name": segment_id,
        "points": points,
    }


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