"""Registry-backed GPX parsing, preserving document order and source timestamps."""

from __future__ import annotations

import hashlib
import json
import math
import xml.etree.ElementTree as ET
from dataclasses import dataclass
from datetime import date, datetime, timedelta, timezone
from pathlib import Path


@dataclass(frozen=True)
class RouteMetadata:
    """Explicit route identity and provenance; dates never come from GPX time."""

    route_id: str
    source_filename: str
    name: str
    route_group: str
    tmb_day: int | None
    assigned_hike_date: str | None
    region: str | None
    notes: str
    timestamps_authoritative: bool


@dataclass(frozen=True)
class RawPoint:
    """One GPX track point; indices are zero-based and segment indices span tracks."""

    point_index: int
    track_index: int
    segment_index: int
    lat: float
    lon: float
    elevation_m_raw: float | None
    source_gpx_time: str | None
    actual_hike_time: str | None


@dataclass(frozen=True)
class ParsedRoute:
    """Immutable parsed input with hash of the exact source bytes."""

    metadata: RouteMetadata
    points: tuple[RawPoint, ...]
    source_sha256: str


def load_registry(path: Path) -> list[RouteMetadata]:
    """Validate metadata and return routes sorted by group, day, then ID."""
    try:
        document = json.loads(path.read_text(encoding="utf-8"))
        if document["schema_version"] != 1 or not document["routes"]:
            raise ValueError("expected schema_version 1 and a nonempty routes list")
        routes = [RouteMetadata(**row) for row in document["routes"]]
        ids, sources, days = set(), set(), set()
        for r in routes:
            for label, value in [("route_id", r.route_id), ("source_filename", r.source_filename)]:
                if (
                    not isinstance(value, str)
                    or not value.strip()
                    or value in {".", ".."}
                    or any(c in value for c in "/\\\x00\n\r")
                ):
                    raise ValueError(f"invalid {label}: {value!r}")
            if not r.source_filename.lower().endswith(".gpx"):
                raise ValueError(f"{r.route_id}: source_filename must end in .gpx")
            if (
                r.route_group not in {"tmb", "reference"}
                or type(r.timestamps_authoritative) is not bool
            ):
                raise ValueError(f"{r.route_id}: invalid group or timestamp authority")
            if (
                not isinstance(r.name, str)
                or not r.name.strip()
                or not isinstance(r.notes, str)
                or (r.region is not None and not isinstance(r.region, str))
            ):
                raise ValueError(f"{r.route_id}: invalid name, notes, or region")
            if r.assigned_hike_date is not None:
                if date.fromisoformat(r.assigned_hike_date).isoformat() != r.assigned_hike_date:
                    raise ValueError(f"{r.route_id}: date must be YYYY-MM-DD")
            if r.route_group == "tmb":
                if type(r.tmb_day) is not int or not 1 <= r.tmb_day <= 9 or r.tmb_day in days:
                    raise ValueError(f"{r.route_id}: invalid or duplicate TMB day")
                if r.assigned_hike_date != str(date(2026, 7, 2) + timedelta(days=r.tmb_day - 1)):
                    raise ValueError(f"{r.route_id}: incorrect TMB date mapping")
                days.add(r.tmb_day)
            elif r.tmb_day is not None:
                raise ValueError(f"{r.route_id}: reference route cannot have a TMB day")
            source = (r.route_group, r.source_filename)
            if r.route_id.casefold() in ids or (source[0], source[1].casefold()) in sources:
                raise ValueError(f"{r.route_id}: duplicate route ID or source")
            ids.add(r.route_id.casefold())
            sources.add((source[0], source[1].casefold()))
        return sorted(routes, key=lambda r: (r.route_group != "tmb", r.tmb_day or 0, r.route_id))
    except (KeyError, TypeError, ValueError) as exc:
        raise ValueError(f"{path}: invalid registry: {exc}") from exc


def parse_gpx(path: Path, metadata: RouteMetadata) -> ParsedRoute:
    """Read GPX 1.0/1.1 tracks only; reject invalid points/XML without dropping data.

    Elevation and time are optional. Raw time text is retained even when it is
    non-authoritative. Present timestamps must be ISO-8601 with a UTC offset.
    Segments and tracks are never connected; waypoints/routes are not ingested.
    """
    data = path.read_bytes()
    try:
        # ElementTree does not fetch external entities; reject all DTDs/entities
        # as well, including UTF-16/32 encodings, to prevent entity expansion.
        declaration_scan = data.replace(b"\x00", b"").upper()
        if b"<!DOCTYPE" in declaration_scan or b"<!ENTITY" in declaration_scan:
            raise ValueError("DTD/entity declarations are not supported")
        root = ET.fromstring(data)
        namespace = root.tag.partition("}")[0] + "}" if root.tag.startswith("{") else ""
        if (
            namespace
            not in {
                "",
                "{http://www.topografix.com/GPX/1/0}",
                "{http://www.topografix.com/GPX/1/1}",
            }
            or root.tag != namespace + "gpx"
        ):
            raise ValueError("expected a GPX 1.0/1.1 gpx root")
        points = []
        segment_index = 0
        for track_index, track in enumerate(root.findall(namespace + "trk")):
            for segment in track.findall(namespace + "trkseg"):
                previous_time = None
                for element in segment.findall(namespace + "trkpt"):
                    context = f"track {track_index}, segment {segment_index}, point {len(points)}"
                    try:
                        lat, lon = float(element.attrib["lat"]), float(element.attrib["lon"])
                        if (
                            not math.isfinite(lat)
                            or not math.isfinite(lon)
                            or not -90 <= lat <= 90
                            or not -180 <= lon <= 180
                        ):
                            raise ValueError(
                                "coordinates must be finite and within latitude/longitude bounds"
                            )
                        elevations, times = (
                            element.findall(namespace + "ele"),
                            element.findall(namespace + "time"),
                        )
                        if len(elevations) > 1 or len(times) > 1:
                            raise ValueError("duplicate elevation or timestamp")
                        elevation = float(elevations[0].text) if elevations else None
                        if elevation is not None and not math.isfinite(elevation):
                            raise ValueError("elevation must be finite")
                        raw_time = times[0].text if times else None
                        actual_time = None
                        if times:
                            if not raw_time or not raw_time.strip():
                                raise ValueError("empty timestamp")
                            timestamp = datetime.fromisoformat(
                                raw_time.strip().replace("Z", "+00:00")
                            )
                            if timestamp.utcoffset() is None:
                                raise ValueError("timestamp must include timezone")
                            if metadata.timestamps_authoritative:
                                if previous_time is not None and timestamp < previous_time:
                                    raise ValueError(
                                        "authoritative timestamps go backwards within segment"
                                    )
                                previous_time = timestamp
                                actual_time = timestamp.astimezone(timezone.utc).isoformat()
                        points.append(
                            RawPoint(
                                len(points),
                                track_index,
                                segment_index,
                                lat,
                                lon,
                                elevation,
                                raw_time,
                                actual_time,
                            )
                        )
                    except (KeyError, TypeError, ValueError) as exc:
                        raise ValueError(f"{context}: {exc}") from exc
                segment_index += 1
        if not points:
            raise ValueError("no track points found (route/waypoint-only GPX is unsupported)")
        return ParsedRoute(metadata, tuple(points), hashlib.sha256(data).hexdigest())
    except (ET.ParseError, ValueError) as exc:
        raise ValueError(f"{path}: malformed GPX: {exc}") from exc


def ingest_routes(registry: Path, raw_dir: Path) -> list[ParsedRoute]:
    """Require every registered source before parsing; never skip missing routes."""
    metadata = load_registry(registry)
    missing = [
        raw_dir / r.route_group / r.source_filename
        for r in metadata
        if not (raw_dir / r.route_group / r.source_filename).is_file()
    ]
    if missing:
        raise ValueError("Missing registered GPX files:\n" + "\n".join(str(p) for p in missing))
    return [parse_gpx(raw_dir / r.route_group / r.source_filename, r) for r in metadata]
