#!/usr/bin/env python3
"""Create a compact transactional SQLite patch between two NextLimit datasets."""

from __future__ import annotations

import os
import sqlite3
from pathlib import Path

from dataset_schema import ROAD_SEGMENT_COLUMNS, SCHEMA_VERSION

LEGACY_SCHEMA_VERSION = 1
DIFF_FORMATS = {
    LEGACY_SCHEMA_VERSION: "nextlimit-sqlite-diff-v1",
    SCHEMA_VERSION: "nextlimit-sqlite-diff-v2",
}
LEGACY_ROAD_COLUMNS = (
    "segment_id",
    "osm_way_id",
    "from_node_id",
    "to_node_id",
    "direction",
    "length_meters",
    "geometry",
    "highway",
    "maxspeed",
    "maxspeed_forward",
    "maxspeed_backward",
    "oneway",
    "junction",
    "access",
    "motor_vehicle",
    "motorcar",
    "surface",
    "lanes",
    "ref",
    "name",
    "zone_maxspeed",
    "source_maxspeed",
    "maxspeed_type",
)
ROAD_COLUMNS_BY_SCHEMA = {
    LEGACY_SCHEMA_VERSION: LEGACY_ROAD_COLUMNS,
    SCHEMA_VERSION: ROAD_SEGMENT_COLUMNS,
}
BOUNDS_COLUMNS = ("min_lat_e7", "max_lat_e7", "min_lon_e7", "max_lon_e7")
INTEGER_COLUMNS = {
    "osm_way_id",
    "from_node_id",
    "to_node_id",
    "direction",
    "road_class",
    "speed_limit_kmh",
    "speed_source",
    "speed_confidence",
    "road_flags",
    *BOUNDS_COLUMNS,
}
REAL_COLUMNS = {"length_meters"}
BLOB_COLUMNS = {"geometry"}


def create_diff(
    previous: Path,
    current: Path,
    output: Path,
    from_version: str,
    to_version: str,
) -> tuple[int, int]:
    """Create `output` atomically and return deleted/upserted segment counts."""
    previous_metadata = read_metadata(previous)
    current_metadata = read_metadata(current)
    if previous_metadata["region_code"] != current_metadata["region_code"]:
        raise ValueError("datasets belong to different regions")
    if previous_metadata["schema_version"] != current_metadata["schema_version"]:
        raise ValueError("dataset schema changed; publish a full package instead")
    schema_version = int(current_metadata["schema_version"])
    road_columns = ROAD_COLUMNS_BY_SCHEMA.get(schema_version)
    diff_format = DIFF_FORMATS.get(schema_version)
    if road_columns is None or diff_format is None:
        raise ValueError(f"unsupported dataset schema: {schema_version}")

    output.parent.mkdir(parents=True, exist_ok=True)
    staging = output.with_suffix(output.suffix + ".staging")
    staging.unlink(missing_ok=True)
    connection = sqlite3.connect(staging)
    try:
        connection.execute("ATTACH DATABASE ? AS previous", (str(previous.resolve()),))
        connection.execute("ATTACH DATABASE ? AS current", (str(current.resolve()),))
        _create_patch_schema(connection, road_columns)
        connection.execute(
            "INSERT INTO diff_metadata VALUES (?, ?, ?, ?)",
            (diff_format, current_metadata["region_code"], from_version, to_version),
        )
        connection.execute(
            """
            INSERT INTO deleted_segments(segment_id)
            SELECT old.segment_id
            FROM previous.road_segments old
            LEFT JOIN current.road_segments new ON new.segment_id = old.segment_id
            WHERE new.segment_id IS NULL
            """
        )
        comparisons = " OR ".join(f"NOT (new.{column} IS old.{column})" for column in road_columns[1:])
        comparisons += " OR " + " OR ".join(
            f"NOT (new_bounds.{column} IS old_bounds.{column})" for column in BOUNDS_COLUMNS
        )
        selected = ", ".join(f"new.{column}" for column in road_columns)
        selected += ", " + ", ".join(f"new_bounds.{column}" for column in BOUNDS_COLUMNS)
        connection.execute(
            f"""
            INSERT INTO upsert_segments
            SELECT {selected}
            FROM current.road_segments new
            JOIN current.road_segment_rtree new_bounds ON new_bounds.row_id = new.row_id
            LEFT JOIN previous.road_segments old ON old.segment_id = new.segment_id
            LEFT JOIN previous.road_segment_rtree old_bounds ON old_bounds.row_id = old.row_id
            WHERE old.segment_id IS NULL OR {comparisons}
            """
        )
        metadata_columns = _table_columns(connection, "current", "dataset_metadata")
        columns = ", ".join(metadata_columns)
        connection.execute(f"INSERT INTO target_metadata SELECT {columns} FROM current.dataset_metadata LIMIT 1")
        connection.commit()
        deleted = int(connection.execute("SELECT COUNT(*) FROM deleted_segments").fetchone()[0])
        upserted = int(connection.execute("SELECT COUNT(*) FROM upsert_segments").fetchone()[0])
        integrity = connection.execute("PRAGMA quick_check(1)").fetchone()
        if integrity != ("ok",):
            raise RuntimeError(f"diff SQLite integrity check failed: {integrity}")
        connection.execute("DETACH DATABASE previous")
        connection.execute("DETACH DATABASE current")
        connection.execute("VACUUM")
        connection.close()
        os.replace(staging, output)
        return deleted, upserted
    except Exception:
        connection.close()
        staging.unlink(missing_ok=True)
        raise


def read_metadata(dataset: Path) -> dict[str, object]:
    connection = sqlite3.connect(f"file:{dataset.resolve()}?mode=ro", uri=True)
    try:
        columns = _table_columns(connection, "main", "dataset_metadata")
        row = connection.execute(f"SELECT {', '.join(columns)} FROM dataset_metadata LIMIT 1").fetchone()
        if row is None:
            raise ValueError(f"missing dataset metadata: {dataset}")
        return dict(zip(columns, row, strict=True))
    finally:
        connection.close()


def _create_patch_schema(connection: sqlite3.Connection, road_columns: tuple[str, ...]) -> None:
    upsert_columns = ",\n            ".join(_column_definition(column) for column in road_columns + BOUNDS_COLUMNS)
    connection.executescript(
        f"""
        CREATE TABLE diff_metadata (
            format TEXT NOT NULL,
            region_id TEXT NOT NULL,
            from_version TEXT NOT NULL,
            to_version TEXT NOT NULL
        );
        CREATE TABLE deleted_segments (segment_id TEXT PRIMARY KEY);
        CREATE TABLE upsert_segments (
            {upsert_columns}
        );
        CREATE TABLE target_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
        );
        """
    )


def _column_definition(column: str) -> str:
    if column == "segment_id":
        return "segment_id TEXT PRIMARY KEY"
    if column in INTEGER_COLUMNS:
        return f"{column} INTEGER"
    if column in REAL_COLUMNS:
        return f"{column} REAL"
    if column in BLOB_COLUMNS:
        return f"{column} BLOB"
    return f"{column} TEXT"


def _table_columns(connection: sqlite3.Connection, database: str, table: str) -> list[str]:
    return [str(row[1]) for row in connection.execute(f"PRAGMA {database}.table_info({table})")]
