#!/usr/bin/env python3
"""Validate and publish a compressed NextLimit regional mobile dataset."""

from __future__ import annotations

import argparse
import gzip
import hashlib
import json
import os
import shutil
import sqlite3
import subprocess
import time
from datetime import UTC, datetime
from pathlib import Path
from urllib.parse import urljoin, urlparse

from dataset_diff import create_diff, read_metadata
from dataset_schema import SCHEMA_VERSION

LEGACY_SCHEMA_VERSION = 1
DEFAULT_COMPRESSION = "gzip"
GZIP_LEVEL = 9
ZSTD_COMPARISON_LEVEL = 19
FORBIDDEN_MOBILE_SOURCE_HOSTS = frozenset({"download.geofabrik.de"})


def parse_arguments() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("dataset", type=Path)
    parser.add_argument("output", type=Path, help="publication root")
    parser.add_argument("--version", required=True)
    parser.add_argument("--region-name", required=True)
    parser.add_argument("--base-url", default="", help="optional HTTPS publication base URL")
    parser.add_argument("--previous", type=Path)
    parser.add_argument("--previous-version")
    parser.add_argument(
        "--compression",
        choices=("gzip", "none"),
        default=DEFAULT_COMPRESSION,
        help="download compression; the app extracts gzip once during installation",
    )
    parser.add_argument(
        "--compare-compression",
        action="store_true",
        help="measure raw, gzip and zstd without publishing a zstd package",
    )
    return parser.parse_args()


def validate_dataset(dataset: Path) -> dict[str, object]:
    metadata = read_metadata(dataset)
    schema_version = int(metadata["schema_version"])
    if schema_version not in {LEGACY_SCHEMA_VERSION, SCHEMA_VERSION}:
        raise RuntimeError(f"unsupported dataset schema: {schema_version}")
    connection = sqlite3.connect(f"file:{dataset.resolve()}?mode=ro", uri=True)
    try:
        checks = {
            "integrity": connection.execute("PRAGMA quick_check(1)").fetchone() == ("ok",),
            "rtree": connection.execute("SELECT rtreecheck('road_segment_rtree')").fetchone() == ("ok",),
            "segment_count": int(connection.execute("SELECT COUNT(*) FROM road_segments").fetchone()[0]),
            "rtree_count": int(connection.execute("SELECT COUNT(*) FROM road_segment_rtree").fetchone()[0]),
            "graph_indexes": _required_indexes(connection),
            "dangling_geometry": int(
                connection.execute(
                    "SELECT COUNT(*) FROM road_segments s LEFT JOIN road_segment_rtree r ON r.row_id=s.row_id "
                    "WHERE r.row_id IS NULL"
                ).fetchone()[0]
            ),
            "invalid_speed_limit": _invalid_speed_count(connection, schema_version),
        }
    finally:
        connection.close()
    if not checks["integrity"] or not checks["rtree"]:
        raise RuntimeError("dataset integrity validation failed")
    if checks["segment_count"] != checks["rtree_count"] or checks["dangling_geometry"] != 0:
        raise RuntimeError("dataset spatial index is inconsistent")
    if checks["segment_count"] != int(metadata["segment_count"]):
        raise RuntimeError("dataset metadata segment count is inconsistent")
    if not checks["graph_indexes"]:
        raise RuntimeError("dataset graph indexes are missing")
    if checks["invalid_speed_limit"] != 0:
        raise RuntimeError("dataset contains invalid normalized speed limits")
    return {"metadata": metadata, "checks": checks}


def publish(arguments: argparse.Namespace) -> dict[str, object]:
    if not arguments.dataset.is_file():
        raise ValueError(f"dataset not found: {arguments.dataset}")
    if arguments.previous is not None and not arguments.previous_version:
        raise ValueError("--previous-version is required with --previous")
    _validate_publication_base_url(arguments.base_url)
    compression = getattr(arguments, "compression", DEFAULT_COMPRESSION)
    compare_compression = getattr(arguments, "compare_compression", False)
    validation = validate_dataset(arguments.dataset)
    metadata = validation["metadata"]
    region = str(metadata["region_code"])
    country = str(metadata["country_code"])
    generated_at = datetime.now(UTC).isoformat().replace("+00:00", "Z")
    root = arguments.output / "v1" / "road-data"
    package_directory = root / "packages"
    manifest_directory = root / "regions" / region
    report_directory = root / "reports"
    for directory in (package_directory, manifest_directory, report_directory):
        directory.mkdir(parents=True, exist_ok=True)

    full_stem = f"{region}-{arguments.version}.sqlite"
    full_name = _compressed_name(full_stem, compression)
    full_package = package_directory / full_name
    full_compression_seconds = _write_package(arguments.dataset, full_package, compression)
    full_checksum = sha256(full_package)
    diffs: list[dict[str, object]] = []
    diff_counts: dict[str, int] | None = None
    if arguments.previous is not None:
        validate_dataset(arguments.previous)
        diff_stem = f"{region}-{arguments.previous_version}-to-{arguments.version}.patch.sqlite"
        raw_diff = package_directory / f".{diff_stem}.raw"
        raw_diff.unlink(missing_ok=True)
        try:
            deleted, upserted = create_diff(
                previous=arguments.previous,
                current=arguments.dataset,
                output=raw_diff,
                from_version=arguments.previous_version,
                to_version=arguments.version,
            )
            diff_name = _compressed_name(diff_stem, compression)
            diff_package = package_directory / diff_name
            _write_package(raw_diff, diff_package, compression)
            diffs.append(
                {
                    "fromVersion": arguments.previous_version,
                    "toVersion": arguments.version,
                    "sizeBytes": diff_package.stat().st_size,
                    "installedSizeBytes": raw_diff.stat().st_size,
                    "checksumSha256": sha256(diff_package),
                    "packageUrl": _publication_url(arguments.base_url, f"v1/road-data/packages/{diff_name}"),
                }
            )
            diff_counts = {"deletedSegments": deleted, "upsertedSegments": upserted}
        finally:
            raw_diff.unlink(missing_ok=True)

    manifest = {
        "region": region,
        "version": arguments.version,
        "sizeBytes": full_package.stat().st_size,
        "installedSizeBytes": arguments.dataset.stat().st_size,
        "checksumSha256": full_checksum,
        "packageUrl": _publication_url(arguments.base_url, f"v1/road-data/packages/{full_name}"),
        "generatedAt": generated_at,
        "schemaVersion": int(metadata["schema_version"]),
        "diffs": diffs,
    }
    manifest_path = manifest_directory / "manifest.json"
    _write_json_atomically(manifest_path, manifest)
    _update_catalog(
        root=root,
        region=region,
        region_name=arguments.region_name,
        country=country,
        manifest=manifest,
        manifest_url=_publication_url(arguments.base_url, f"v1/road-data/regions/{region}/manifest.json"),
        generated_at=generated_at,
    )
    compression_report = {
        "selected": compression,
        "raw": _compression_result(arguments.dataset.stat().st_size, arguments.dataset.stat().st_size, 0.0),
        compression: _compression_result(
            arguments.dataset.stat().st_size,
            full_package.stat().st_size,
            full_compression_seconds,
        ),
    }
    if compare_compression:
        compression_report.update(_compare_compression(arguments.dataset, report_directory, compression_report))
    build_report = _load_build_report(arguments.dataset)
    source_pbf_size = _source_pbf_size(build_report)
    report = {
        "region": region,
        "version": arguments.version,
        "generatedAt": generated_at,
        "fullPackage": {
            "path": str(full_package),
            "sizeBytes": full_package.stat().st_size,
            "installedSizeBytes": arguments.dataset.stat().st_size,
            "checksumSha256": full_checksum,
            "reductionPercentVsSourcePbf":
                None
                if source_pbf_size is None
                else round((1 - full_package.stat().st_size / source_pbf_size) * 100, 2),
        },
        "compression": compression_report,
        "diff": diff_counts,
        "build": build_report,
        **validation,
    }
    _write_json_atomically(report_directory / f"{region}-{arguments.version}.json", report)
    return report


def sha256(path: Path) -> str:
    digest = hashlib.sha256()
    with path.open("rb") as stream:
        for block in iter(lambda: stream.read(1024 * 1024), b""):
            digest.update(block)
    return digest.hexdigest()


def _required_indexes(connection: sqlite3.Connection) -> bool:
    names = {
        str(row[0])
        for row in connection.execute(
            "SELECT name FROM sqlite_master WHERE type='index' AND name IN "
            "('road_segments_from_node_idx','road_segments_to_node_idx','road_segments_way_idx')"
        )
    }
    return len(names) == 3


def _invalid_speed_count(connection: sqlite3.Connection, schema_version: int) -> int:
    if schema_version == SCHEMA_VERSION:
        return int(
            connection.execute(
                "SELECT COUNT(*) FROM road_segments WHERE "
                "(speed_limit_kmh IS NULL AND (speed_source != 0 OR speed_confidence != 0)) OR "
                "(speed_limit_kmh IS NOT NULL AND "
                "(speed_limit_kmh NOT BETWEEN 5 AND 250 OR speed_source = 0 OR speed_confidence NOT BETWEEN 1 AND 100))"
            ).fetchone()[0]
        )
    count = 0
    for (raw_value,) in connection.execute("SELECT maxspeed FROM road_segments WHERE maxspeed IS NOT NULL"):
        value = str(raw_value).strip().lower().removesuffix("km/h").removesuffix("kph").strip()
        if value.replace(".", "", 1).isdigit() and float(value) not in range(5, 251):
            count += 1
    return count


def _compressed_name(stem: str, compression: str) -> str:
    return f"{stem}.gz" if compression == "gzip" else stem


def _write_package(source: Path, destination: Path, compression: str) -> float:
    started = time.monotonic()
    if compression == "gzip":
        _gzip_atomically(source, destination)
    elif compression == "none":
        _copy_atomically(source, destination)
    else:
        raise ValueError(f"unsupported compression: {compression}")
    return time.monotonic() - started


def _gzip_atomically(source: Path, destination: Path) -> None:
    staging = destination.with_suffix(destination.suffix + ".staging")
    staging.unlink(missing_ok=True)
    try:
        with source.open("rb") as input_stream, staging.open("wb") as output_stream:
            with gzip.GzipFile(fileobj=output_stream, mode="wb", compresslevel=GZIP_LEVEL, mtime=0) as compressed:
                shutil.copyfileobj(input_stream, compressed, length=1024 * 1024)
        os.replace(staging, destination)
    except Exception:
        staging.unlink(missing_ok=True)
        raise


def _copy_atomically(source: Path, destination: Path) -> None:
    staging = destination.with_suffix(destination.suffix + ".staging")
    staging.unlink(missing_ok=True)
    shutil.copyfile(source, staging)
    os.replace(staging, destination)


def _compare_compression(
    dataset: Path,
    temporary_directory: Path,
    existing: dict[str, object],
) -> dict[str, object]:
    results: dict[str, object] = {}
    if "gzip" not in existing:
        gzip_path = temporary_directory / f".{dataset.name}.comparison.gz"
        try:
            elapsed = _write_package(dataset, gzip_path, "gzip")
            results["gzip"] = _compression_result(dataset.stat().st_size, gzip_path.stat().st_size, elapsed)
        finally:
            gzip_path.unlink(missing_ok=True)
    zstd = shutil.which("zstd")
    if zstd is None:
        results["zstd"] = {"available": False}
        return results
    zstd_path = temporary_directory / f".{dataset.name}.comparison.zst"
    zstd_path.unlink(missing_ok=True)
    started = time.monotonic()
    try:
        subprocess.run(
            [zstd, f"-{ZSTD_COMPARISON_LEVEL}", "--quiet", "--force", str(dataset), "-o", str(zstd_path)],
            check=True,
        )
        results["zstd"] = {
            "available": True,
            "level": ZSTD_COMPARISON_LEVEL,
            **_compression_result(dataset.stat().st_size, zstd_path.stat().st_size, time.monotonic() - started),
        }
    finally:
        zstd_path.unlink(missing_ok=True)
    return results


def _compression_result(raw_size: int, compressed_size: int, elapsed: float) -> dict[str, object]:
    return {
        "sizeBytes": compressed_size,
        "ratioVsRaw": round(compressed_size / raw_size, 4),
        "reductionPercentVsRaw": round((1 - compressed_size / raw_size) * 100, 2),
        "seconds": round(elapsed, 3),
    }


def _load_build_report(dataset: Path) -> dict[str, object] | None:
    path = dataset.with_suffix(".report.json")
    return json.loads(path.read_text(encoding="utf-8")) if path.is_file() else None


def _source_pbf_size(build_report: dict[str, object] | None) -> int | None:
    if build_report is None:
        return None
    source = build_report.get("source")
    if not isinstance(source, dict):
        return None
    value = source.get("sizeBytes")
    return value if isinstance(value, int) and value > 0 else None


def _write_json_atomically(path: Path, value: object) -> None:
    staging = path.with_suffix(path.suffix + ".staging")
    staging.write_text(json.dumps(value, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
    os.replace(staging, path)


def _update_catalog(
    root: Path,
    region: str,
    region_name: str,
    country: str,
    manifest: dict[str, object],
    manifest_url: str,
    generated_at: str,
) -> None:
    path = root / "catalog.json"
    catalog = json.loads(path.read_text(encoding="utf-8")) if path.is_file() else {"version": 1, "regions": []}
    entry = {
        "id": region,
        "name": region_name,
        "countryCode": country,
        "manifestUrl": manifest_url,
        "available": manifest,
    }
    regions = [item for item in catalog.get("regions", []) if item.get("id") != region] + [entry]
    catalog.update({"generatedAt": generated_at, "regions": sorted(regions, key=lambda item: item["id"])})
    _assert_no_geofabrik_urls(catalog)
    _write_json_atomically(path, catalog)


def _validate_publication_base_url(base_url: str) -> None:
    if not base_url:
        return
    parsed = urlparse(base_url)
    if parsed.scheme != "https" or not parsed.hostname:
        raise ValueError("publication base URL must use HTTPS")
    if parsed.hostname.lower() in FORBIDDEN_MOBILE_SOURCE_HOSTS:
        raise ValueError("Geofabrik is an import source, not a mobile publication endpoint")


def _assert_no_geofabrik_urls(value: object) -> None:
    if isinstance(value, dict):
        for child in value.values():
            _assert_no_geofabrik_urls(child)
    elif isinstance(value, list):
        for child in value:
            _assert_no_geofabrik_urls(child)
    elif isinstance(value, str):
        host = urlparse(value).hostname
        if host is not None and host.lower() in FORBIDDEN_MOBILE_SOURCE_HOSTS:
            raise ValueError("mobile catalog must not reference Geofabrik")


def _publication_url(base_url: str, relative_path: str) -> str:
    value = urljoin(base_url.rstrip("/") + "/", relative_path) if base_url else relative_path
    _assert_no_geofabrik_urls(value)
    return value


def main() -> int:
    arguments = parse_arguments()
    report = publish(arguments)
    print(json.dumps({"region": report["region"], "version": report["version"]}))
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
