"""Stretches of a route that are not walked or ridden (AV-064): a ``dojazd``
-- a train, a car, a ferry, a flight -- between two parts of a trip.

A route revision stores only its transfer spans (``segments_json``,
migration 014): ``{start_point_index, end_point_index, kind: "TRANSFER",
label}``, the stretch from one route point to a later one. Everything else
is movement (walking, running or riding, by the route's activity); a route
with no spans -- every route saved before AV-064 -- is one movement stretch,
so nothing about it changes.

**The movement view.** Every computation of distance, time, days, ascent and
descent reads the route through route_points_from_geometry_geojson(), which
puts each step of a transfer into a segment of its own
(``RoutePoint.segment_index``). mountain_twin.route_analysis.prepare_route
never measures distance across a segment boundary, and the pace engine
times each segment on its own -- so a transfer adds no metre and no second,
and the point indexes (camps, days, anchors) stay those of the whole route.

**Short elevation gaps** (``fill_short_gaps``): for timing only, a run of
points without elevation shorter than SHORT_GAP_M between two known heights
is filled by linear interpolation by distance -- in memory, never stored
(ADR-002: the database keeps "no data"). A longer gap still blocks the
timing, and says where it is. Ascent never sees the filled heights
(elevation_gain.py counts runs of known heights only).
"""

from __future__ import annotations

import math
from dataclasses import replace
from typing import Any, Sequence

TRANSFER = "TRANSFER"
SEGMENT_KINDS = (TRANSFER,)
SHORT_GAP_M = 1000.0
LABEL_MAX = 80


def validate_segments(raw: Any, point_count: int) -> tuple[dict[str, Any], ...]:
    """A route save's ``segments``: transfer spans in route order, each at
    least one step long, never overlapping (they may touch)."""
    if raw is None:
        return ()
    if not isinstance(raw, list) or len(raw) > 1000:
        raise ValueError("segments must be a list (at most 1000)")
    spans = []
    for item in raw:
        if not isinstance(item, dict):
            raise ValueError("a segment must be an object")
        kind = item.get("kind", TRANSFER)
        if kind not in SEGMENT_KINDS:
            raise ValueError(f"segment kind must be one of {SEGMENT_KINDS}")
        start, end = item.get("start_point_index"), item.get("end_point_index")
        if not all(isinstance(v, int) and not isinstance(v, bool) for v in (start, end)):
            raise ValueError("segment point indexes must be integers")
        if not 0 <= start < end < point_count:
            raise ValueError("a segment must run forward within the route")
        label = item.get("label")
        if label is not None:
            label = str(label).strip()[:LABEL_MAX] or None
        spans.append(
            {"start_point_index": start, "end_point_index": end, "kind": kind, "label": label}
        )
    spans.sort(key=lambda span: span["start_point_index"])
    for earlier, later in zip(spans, spans[1:]):
        if later["start_point_index"] < earlier["end_point_index"]:
            raise ValueError("segments must not overlap")
    return tuple(spans)


def transfer_steps(spans: Sequence[dict[str, Any]], point_count: int) -> list[bool]:
    """For each point index i > 0: whether the step i-1 -> i is a transfer
    (index 0 is always False)."""
    steps = [False] * point_count
    for span in spans or ():
        for index in range(span["start_point_index"] + 1, span["end_point_index"] + 1):
            if 0 < index < point_count:
                steps[index] = True
    return steps


def movement_segment_indexes(spans: Sequence[dict[str, Any]], point_count: int) -> list[int]:
    """``segment_index`` per point: a new segment begins at every point a
    transfer step leads to, so no distance is measured across a transfer."""
    steps = transfer_steps(spans, point_count)
    indexes, current = [], 0
    for index in range(point_count):
        if steps[index]:
            current += 1
        indexes.append(current)
    return indexes


def _metres(lat1, lon1, lat2, lon2) -> float:
    p1, p2 = math.radians(lat1), math.radians(lat2)
    h = (
        math.sin((p2 - p1) / 2) ** 2
        + math.cos(p1) * math.cos(p2) * math.sin(math.radians(lon2 - lon1) / 2) ** 2
    )
    return 2 * 6_371_000.0 * math.asin(min(1.0, math.sqrt(h)))


def transfer_distance_m(coordinates: Sequence[Sequence[float]], spans) -> float:
    """The straight-line length of the transfer steps (as drawn on the map)."""
    steps = transfer_steps(spans, len(coordinates))
    return sum(
        _metres(coordinates[i - 1][1], coordinates[i - 1][0], coordinates[i][1], coordinates[i][0])
        for i in range(1, len(coordinates))
        if steps[i]
    )


def transfers_document(coordinates: Sequence[Sequence[float]], spans) -> list[dict[str, Any]]:
    """The transfers as the page shows them: where, how long, their label."""
    out = []
    for span in spans or ():
        part = coordinates[span["start_point_index"] : span["end_point_index"] + 1]
        out.append(
            {
                **span,
                "distance_m": round(
                    sum(_metres(a[1], a[0], b[1], b[0]) for a, b in zip(part, part[1:]))
                ),
            }
        )
    return out


# --- short elevation gaps (timing only) ------------------------------------------
def elevation_gaps(prepared: Sequence[Any]) -> list[dict[str, Any]]:
    """Runs of movement points without elevation: where they are (route
    distance from the last known height before to the first after) and
    whether they are short enough to fill for timing."""
    gaps, run = [], []
    known_before = None
    for point in prepared:
        if point.gpx_elevation_m is None:
            run.append(point)
            continue
        if run:
            gaps.append(_gap(run, known_before, point))
            run = []
        known_before = point
    if run:
        gaps.append(_gap(run, known_before, None))
    return gaps


def _gap(run, before, after) -> dict[str, Any]:
    start = before.cumulative_distance_m if before is not None else run[0].cumulative_distance_m
    end = after.cumulative_distance_m if after is not None else run[-1].cumulative_distance_m
    bounded = before is not None and after is not None
    # Only within one movement segment: a transfer between is no slope.
    same_segment = bounded and before.segment_index == after.segment_index
    return {
        "point_indexes": [point.point_index for point in run],
        "from_m": start,
        "to_m": end,
        "length_m": end - start,
        "fillable": bool(same_segment and end - start <= SHORT_GAP_M),
    }


def fill_short_gaps(prepared: Sequence[Any]) -> tuple[tuple[Any, ...], frozenset[int]]:
    """The prepared points with every short gap filled by linear
    interpolation by distance between its two known neighbours, and the
    indexes filled. In memory only -- for the pace model, never stored or
    shown as a height. Longer gaps and gaps at the route's ends stay empty."""
    points = list(prepared)
    filled: set[int] = set()
    position = {point.point_index: k for k, point in enumerate(points)}
    for gap in elevation_gaps(points):
        if not gap["fillable"]:
            continue
        first = position[gap["point_indexes"][0]]
        last = position[gap["point_indexes"][-1]]
        before, after = points[first - 1], points[last + 1]
        span = after.cumulative_distance_m - before.cumulative_distance_m
        for k in range(first, last + 1):
            t = (
                0.0
                if span <= 0
                else (points[k].cumulative_distance_m - before.cumulative_distance_m) / span
            )
            height = before.gpx_elevation_m + t * (after.gpx_elevation_m - before.gpx_elevation_m)
            points[k] = replace(points[k], gpx_elevation_m=height)
            filled.add(points[k].point_index)
    # A one-point transfer step has no length: it is never timed, so its
    # missing height is the previous point's (nothing is climbed or shown).
    for k in range(1, len(points)):
        point, previous = points[k], points[k - 1]
        if (
            point.gpx_elevation_m is None
            and point.segment_index != previous.segment_index
            and (k + 1 == len(points) or points[k + 1].segment_index != point.segment_index)
            and previous.gpx_elevation_m is not None
        ):
            points[k] = replace(point, gpx_elevation_m=previous.gpx_elevation_m)
    return tuple(points), frozenset(filled)
