"""A route file (route_files.py) mapped onto the Journey model (AV-053):
what Planning Workspace opens for review before anything is saved.

* **One line.** Every segment of every track, in file order, joined into
  one route (a ``rte`` only when the file has no ``trk``). Where one segment
  ends and the next begins is a proposed stage boundary -- a camp at the
  end of each day, as Planning Workspace already models days; the user
  merges (drops a boundary) or splits (adds a camp) in the preview.
* **Simplified, not resampled.** Douglas-Peucker with SIMPLIFY_TOLERANCE_M
  keeps every turn wider than a few metres; segment ends are always kept.
  The original file stays with the route (route_imports.py) for a 1:1
  export.
* **Elevation from our DEM**, always: Copernicus GLO-90, bilinear between
  cell centres -- the same terrain the 3D view and "Pogoda" read. The
  file's own elevation is kept only for comparison. Where the DEM cannot be
  read, the point has no elevation (UNAVAILABLE), never a guess.
* **Editable** (AV-037): the route's placed points ("anchors") are its shape
  points -- Douglas-Peucker with ANCHOR_SHAPE_TOLERANCE_M -- plus one at
  least every ANCHOR_SPACING_M, never two within ANCHOR_MIN_GAP_M, and
  always every stage boundary. Between two anchors the file's geometry
  stays as it is until the user moves one of them; only then is that
  stretch routed again.
* **Waypoints** sit on the route's nearest point; one farther than
  OFF_ROUTE_M is "poza trasą" (kept as a point of interest, never a camp).
  A waypoint whose name or type reads like a place to sleep is proposed as
  a camp; the user decides.
* **Times** are metadata ("the original walk took ..."), never the plan.
"""

from __future__ import annotations

import math
import re
from dataclasses import dataclass
from typing import Any, Callable, Sequence

import numpy as np

from .elevation_gain import plausible_elevation
from .route_files import FilePoint, RouteFile

SIMPLIFY_TOLERANCE_M = 3.0
ANCHOR_SHAPE_TOLERANCE_M = 60.0
ANCHOR_SPACING_M = 1500.0
ANCHOR_MIN_GAP_M = 250.0
OFF_ROUTE_M = 300.0
DEM_CHUNK_DEGREES = 0.25
ELEVATION_SOURCE = "DEM_COPERNICUS_GLO90"
EARTH_RADIUS_M = 6_371_000.0

CAMP_WORDS = re.compile(
    r"camp|biwak|bivouac|bivacco|schronisk|nocleg|hotel|hostel|lodge|refug|rifugio|"
    r"hütte|hutte|hut\b|chata|chatka|cabin|campground|lodging|guest ?house|pensjonat|"
    r"teahouse|albergue|gîte|gite|hytte|koj|kemp",
    re.IGNORECASE,
)


def metres(a, b) -> float:
    """Haversine distance between two objects with latitude/longitude."""
    lat1, lat2 = math.radians(a.latitude), math.radians(b.latitude)
    dlat, dlon = lat2 - lat1, math.radians(b.longitude - a.longitude)
    h = math.sin(dlat / 2) ** 2 + math.cos(lat1) * math.cos(lat2) * math.sin(dlon / 2) ** 2
    return 2 * EARTH_RADIUS_M * math.asin(min(1.0, math.sqrt(h)))


def _local_xy(points: Sequence, origin) -> np.ndarray:
    """Metres east/north of ``origin`` (an equirectangular plane: exact enough
    for tolerances of metres over one track), as an (n, 2) array."""
    k = math.cos(math.radians(origin.latitude))
    lat = np.radians(np.array([p.latitude for p in points], dtype=float) - origin.latitude)
    lon = np.radians(np.array([p.longitude for p in points], dtype=float) - origin.longitude)
    return np.column_stack((lon * EARTH_RADIUS_M * k, lat * EARTH_RADIUS_M))


def _distances_to_segment(xy: np.ndarray, a: np.ndarray, b: np.ndarray) -> np.ndarray:
    """Each row of ``xy``'s distance to the segment a-b."""
    d = b - a
    length = float(d @ d)
    if length == 0.0:
        return np.hypot(*(xy - a).T)
    t = np.clip(((xy - a) @ d) / length, 0.0, 1.0)
    return np.hypot(*(xy - (a + t[:, None] * d)).T)


def douglas_peucker(points: Sequence, tolerance_m: float, keep=frozenset()) -> list[int]:
    """Indexes of the points kept: every point farther than ``tolerance_m``
    from the simplified line survives, and so does every index in ``keep``.
    Iterative (a long track would overflow recursion), numpy per stretch."""
    n = len(points)
    if n <= 2:
        return list(range(n))
    xy = _local_xy(points, points[n // 2])
    kept = np.zeros(n, dtype=bool)
    kept[0] = kept[-1] = True
    for index in keep:
        if 0 <= index < n:
            kept[index] = True
    forced = np.flatnonzero(kept).tolist()
    stack = list(zip(forced, forced[1:]))
    while stack:
        start, end = stack.pop()
        if end - start < 2:
            continue
        d = _distances_to_segment(xy[start + 1 : end], xy[start], xy[end])
        worst = int(np.argmax(d))
        if d[worst] > tolerance_m:
            index = start + 1 + worst
            kept[index] = True
            stack.append((start, index))
            stack.append((index, end))
    return np.flatnonzero(kept).tolist()


def cumulative_distances(points: Sequence) -> list[float]:
    out = [0.0]
    for a, b in zip(points, points[1:]):
        out.append(out[-1] + metres(a, b))
    return out


# --- elevation from the DEM -------------------------------------------------------
def dem_elevations(points: Sequence, dem) -> list[float | None]:
    """Bilinear DEM heights at the points, chunk by chunk (each chunk's box at
    most DEM_CHUNK_DEGREES a side, so a long route never asks for one huge
    grid). None where the DEM could not be read."""
    out: list[float | None] = [None] * len(points)
    if dem is None or not points:
        return out
    start = 0
    while start < len(points):
        south = north = points[start].latitude
        west = east = points[start].longitude
        end = start + 1
        while end < len(points):
            p = points[end]
            s, n_ = min(south, p.latitude), max(north, p.latitude)
            w, e = min(west, p.longitude), max(east, p.longitude)
            if n_ - s > DEM_CHUNK_DEGREES or e - w > DEM_CHUNK_DEGREES:
                break
            south, north, west, east = s, n_, w, e
            end += 1
        margin = 0.002
        try:
            grid = dem.grid(west - margin, south - margin, east + margin, north + margin)
        except Exception:  # noqa: BLE001 -- any DEM failure: these points stay unknown
            grid = None
        if grid is not None:
            for i in range(start, end):
                out[i] = _bilinear(grid, points[i].latitude, points[i].longitude)
        start = end
    return out


def _bilinear(grid, latitude: float, longitude: float) -> float | None:
    values = grid.elevation_m
    rows, cols = values.shape
    cell = grid.cell_degrees
    fy = (grid.north - latitude) / cell - 0.5
    fx = (longitude - grid.west) / cell - 0.5
    y0 = min(rows - 1, max(0, math.floor(fy)))
    x0 = min(cols - 1, max(0, math.floor(fx)))
    y1, x1 = min(rows - 1, y0 + 1), min(cols - 1, x0 + 1)
    ty, tx = min(1.0, max(0.0, fy - y0)), min(1.0, max(0.0, fx - x0))
    corners = [
        plausible_elevation(values[y, x]) for y, x in ((y0, x0), (y0, x1), (y1, x0), (y1, x1))
    ]
    if any(corner is None for corner in corners):
        return None  # a no-data cell nearby: this height is unknown, never blended in
    a, b, c, d = corners
    value = (a * (1 - tx) + b * tx) * (1 - ty) + (c * (1 - tx) + d * tx) * ty
    return round(value, 1)


# --- the mapping ---------------------------------------------------------------------
@dataclass(frozen=True)
class ImportedRoute:
    points: list[dict[str, Any]]  # {latitude, longitude, elevation_m}
    anchor_point_indexes: list[int]
    stage_boundaries: list[dict[str, Any]]
    waypoints: list[dict[str, Any]]
    summary: dict[str, Any]
    # AV-064: transfers (dojazd) the file marks, as route_segments spans.
    transfer_spans: list[dict[str, Any]] = ()

    def to_dict(self) -> dict[str, Any]:
        return {
            "points": self.points,
            "anchor_point_indexes": self.anchor_point_indexes,
            "stage_boundaries": self.stage_boundaries,
            "waypoints": self.waypoints,
            "summary": self.summary,
            "transfer_spans": list(self.transfer_spans),
        }


def _pieces(route_file: RouteFile) -> list[tuple[str | None, tuple[FilePoint, ...], bool]]:
    """(track name, points, is_transfer) per track segment, in file order."""
    tracks = [t for t in route_file.tracks if t.kind == "TRACK"] or list(route_file.tracks)
    pieces = []
    for track in tracks:
        for k, segment in enumerate(track.segments):
            if segment:
                transfer = k < len(track.transfer_segments) and track.transfer_segments[k]
                pieces.append((track.name, segment, transfer))
    return pieces


def _duration_text(seconds: float) -> str:
    hours, rest = divmod(int(round(seconds / 60.0)), 60)
    days, hours = divmod(hours, 24)
    if days:
        return f"{days} d {hours} h {rest} min"
    return f"{hours} h {rest:02d} min"


def _time_summary(points: Sequence[FilePoint]) -> dict[str, Any] | None:
    times = [p.time for p in points if p.time is not None]
    if len(times) < 2:
        return None
    start, end = min(times), max(times)
    seconds = (end - start).total_seconds()
    return {
        "start": start.isoformat(),
        "end": end.isoformat(),
        "duration_s": round(seconds),
        "duration_text": _duration_text(seconds),
        "points_with_time": len(times),
    }


def nearest_on_line(points: Sequence, target) -> tuple[int, float]:
    """The line's vertex nearest to ``target`` along its nearest segment, and
    the distance from ``target`` to that segment (metres)."""
    if len(points) == 1:
        return 0, metres(points[0], target)
    xy = _local_xy(points, target)  # target at (0, 0)
    a, b = xy[:-1], xy[1:]
    d = b - a
    length = np.einsum("ij,ij->i", d, d)
    t = np.where(
        length > 0, np.clip(-np.einsum("ij,ij->i", a, d) / np.where(length > 0, length, 1), 0, 1), 0
    )
    closest = a + t[:, None] * d
    distances = np.hypot(closest[:, 0], closest[:, 1])
    segment = int(np.argmin(distances))
    index = segment if t[segment] <= 0.5 else segment + 1
    return index, float(distances[segment])


def _anchors(points: Sequence, cumulative: Sequence[float], forced: set[int]) -> list[int]:
    shape = set(douglas_peucker(points, ANCHOR_SHAPE_TOLERANCE_M, forced))
    last = len(points) - 1
    chosen = sorted(shape | forced | {0, last})
    # At least one every ANCHOR_SPACING_M.
    filled = [chosen[0]]
    for nxt in chosen[1:]:
        prev = filled[-1]
        gap = cumulative[nxt] - cumulative[prev]
        if gap > ANCHOR_SPACING_M:
            steps = int(gap // ANCHOR_SPACING_M)
            for k in range(1, steps + 1):
                goal = cumulative[prev] + gap * k / (steps + 1)
                index = min(
                    range(prev + 1, nxt), key=lambda i: abs(cumulative[i] - goal), default=None
                )
                if index is not None and index > filled[-1]:
                    filled.append(index)
        filled.append(nxt)
    # Never two within ANCHOR_MIN_GAP_M -- unless forced (a stage end).
    thinned = [filled[0]]
    for index in filled[1:]:
        if index in forced or index == last:
            if (
                thinned[-1] not in forced
                and thinned[-1] != 0
                and cumulative[index] - cumulative[thinned[-1]] < ANCHOR_MIN_GAP_M
            ):
                thinned.pop()
            thinned.append(index)
        elif cumulative[index] - cumulative[thinned[-1]] >= ANCHOR_MIN_GAP_M:
            thinned.append(index)
    return sorted(set(thinned))


def import_route(
    route_file: RouteFile, *, dem=None, elevations: Callable | None = None
) -> ImportedRoute:
    """The file as Planning Workspace opens it (see the module docstring).
    ``elevations(points)`` overrides the DEM lookup (tests)."""
    pieces = _pieces(route_file)
    original: list[FilePoint] = []
    piece_ends: list[int] = []  # index in ``original`` of each piece's last point
    piece_names: list[str | None] = []
    piece_starts: list[int] = []
    piece_transfer: list[bool] = []
    joins_m: list[float] = []
    for name, segment, transfer in pieces:
        if original:
            joins_m.append(metres(original[-1], segment[0]))
        piece_starts.append(len(original))
        original.extend(segment)
        piece_ends.append(len(original) - 1)
        piece_names.append(name)
        piece_transfer.append(transfer)
    # Drop exact repeats (a GPS standing still writes the same point).
    cleaned: list[FilePoint] = []
    clean_index: list[int] = []
    for p in original:
        if cleaned and p.latitude == cleaned[-1].latitude and p.longitude == cleaned[-1].longitude:
            clean_index.append(len(cleaned) - 1)
            continue
        cleaned.append(p)
        clean_index.append(len(cleaned) - 1)
    # AV-064: an end next to a transfer piece is no camp (the user travels
    # on); the transfer's own ends are kept exactly.
    boundary_pieces = [
        k for k in range(len(pieces) - 1) if not piece_transfer[k] and not piece_transfer[k + 1]
    ]
    ends_clean = sorted({clean_index[piece_ends[k]] for k in boundary_pieces})
    transfer_clean = [
        (clean_index[piece_starts[k]], clean_index[piece_ends[k]])
        for k in range(len(pieces))
        if piece_transfer[k] and clean_index[piece_ends[k]] > clean_index[piece_starts[k]]
    ]
    forced = set(ends_clean) | {i for span in transfer_clean for i in span}
    kept = douglas_peucker(cleaned, SIMPLIFY_TOLERANCE_M, forced)
    position = {original_index: k for k, original_index in enumerate(kept)}
    line = [cleaned[i] for i in kept]
    cumulative = cumulative_distances(line)
    heights = elevations(line) if elevations is not None else dem_elevations(line, dem)
    points = [
        {"latitude": p.latitude, "longitude": p.longitude, "elevation_m": h}
        for p, h in zip(line, heights)
    ]
    stage_ends = [position[i] for i in ends_clean]
    stage_next_names = [piece_names[k + 1] for k in boundary_pieces]
    transfer_spans = [
        {
            "start_point_index": position[a],
            "end_point_index": position[b],
            "kind": "TRANSFER",
            "label": None,
        }
        for a, b in transfer_clean
    ]
    last = len(line) - 1
    waypoints = []
    for w in route_file.waypoints:
        index, distance = nearest_on_line(line, w)
        off_route = distance > OFF_ROUTE_M
        reads_like_camp = bool(CAMP_WORDS.search(" ".join(filter(None, (w.name, w.kind)))))
        waypoints.append(
            {
                "name": w.name,
                "latitude": w.latitude,
                "longitude": w.longitude,
                "file_type": w.kind,
                "route_point_index": index,
                "route_distance_m": round(cumulative[index]),
                "distance_from_route_m": round(distance),
                "off_route": off_route,
                # A camp ends a day: never at the start or the finish.
                "suggested_kind": "CAMP"
                if reads_like_camp and not off_route and 0 < index < last
                else "POI",
            }
        )
    # A stage ends where a file segment ends; a camp-like waypoint near that
    # end names the camp (and is that camp, not a second one).
    stage_boundaries = []
    for k, index in enumerate(stage_ends):
        near = [
            w
            for w in waypoints
            if w["suggested_kind"] == "CAMP"
            and abs(w["route_distance_m"] - cumulative[index]) <= OFF_ROUTE_M
        ]
        for w in near:
            w["suggested_kind"] = "STAGE_END"
        stage_boundaries.append(
            {
                "route_point_index": index,
                "route_distance_m": round(cumulative[index]),
                "label": next((w["name"] for w in near if w["name"]), None) or f"Nocleg {k + 1}",
                "next_stage_name": stage_next_names[k],
                "source": "FILE_SEGMENT_END",
            }
        )
    anchors = _anchors(
        line,
        cumulative,
        set(stage_ends)
        | {w["route_point_index"] for w in waypoints if w["suggested_kind"] == "CAMP"}
        | {
            i
            for span in transfer_spans
            for i in (span["start_point_index"], span["end_point_index"])
        },
    )
    file_elevations = [p.elevation_m for p in original if p.elevation_m is not None]
    dem_known = [h for h in heights if h is not None]
    paired = [
        abs(h - p.elevation_m)
        for p, h in zip(line, heights)
        if h is not None and p.elevation_m is not None
    ]
    piece_lengths = []
    for k, end in enumerate(stage_ends + [len(line) - 1]):
        begin = 0 if k == 0 else stage_ends[k - 1]
        piece_lengths.append(round(cumulative[end] - cumulative[begin]))
    summary = {
        "file_format": route_file.file_format,
        "name": route_file.name or next((n for n in piece_names if n), None),
        "creator": route_file.creator,
        "track_count": len({id(t) for t in route_file.tracks}),
        "piece_count": len(pieces),
        "piece_lengths_m": piece_lengths,
        "piece_time": [_time_summary(segment) for _, segment, _transfer in pieces],
        "joins_m": [round(j) for j in joins_m],
        "original_point_count": len(original),
        "point_count": len(line),
        "simplify_tolerance_m": SIMPLIFY_TOLERANCE_M,
        "distance_m": round(cumulative[-1]),
        "elevation_source": ELEVATION_SOURCE,
        "dem_unavailable_points": len(heights) - len(dem_known),
        "file_elevation": None
        if not file_elevations
        else {
            "points": len(file_elevations),
            "min_m": round(min(file_elevations), 1),
            "max_m": round(max(file_elevations), 1),
            "mean_abs_difference_to_dem_m": round(sum(paired) / len(paired), 1) if paired else None,
        },
        "time": _time_summary(original),
        "waypoint_count": len(waypoints),
    }
    return ImportedRoute(points, anchors, stage_boundaries, waypoints, summary, transfer_spans)
