"""„Dopasuj do szlaków” (AV-053): an imported route re-routed through
BRouter along the file's own line, so the user can compare and pick.

The file's line is sampled every SAMPLE_SPACING_M (always keeping its start,
its end and every stage end); BRouter routes through those samples in
chunks of CHUNK_POINTS (one request each, through the same client as route
drawing -- counted, budgeted, offline in tests). The answer says how much
longer or shorter the matched line is and where it leaves the file's line
by more than DEVIATION_M -- the places to look at before choosing.

Heights of the matched line come from our DEM, as for the file's own line
(route_import.py), so the two versions differ only in where they go.
"""

from __future__ import annotations

from typing import Any, Callable, Sequence

import numpy as np

from .route_import import _anchors, _local_xy, cumulative_distances, dem_elevations
from .route_persistence import RoutePointInput
from .route_provider import fetch_route_segment_detail

SAMPLE_SPACING_M = 400.0
CHUNK_POINTS = 40
DEVIATION_M = 50.0
MAX_REQUESTS = 25  # ~400 km; a longer route is matched by the user stage by stage


class _Point:
    __slots__ = ("latitude", "longitude")

    def __init__(self, latitude: float, longitude: float):
        self.latitude, self.longitude = latitude, longitude


def sample_indexes(cumulative: Sequence[float], forced: set[int], spacing_m: float) -> list[int]:
    last = len(cumulative) - 1
    chosen, next_at = [0], spacing_m
    for i in range(1, last):
        if i in forced or cumulative[i] >= next_at:
            chosen.append(i)
            next_at = cumulative[i] + spacing_m
    chosen.append(last)
    return chosen


def distances_to_line(points: Sequence, line: Sequence) -> np.ndarray:
    """Each point's distance (m) to the polyline ``line``."""
    if len(line) < 2:
        return np.full(len(points), np.inf)
    origin = line[len(line) // 2]
    xy = _local_xy(list(points), origin)
    lxy = _local_xy(list(line), origin)
    a, b = lxy[:-1], lxy[1:]
    d = b - a
    length = np.einsum("ij,ij->i", d, d)
    safe = np.where(length > 0, length, 1.0)
    out = np.empty(len(points))
    for k, p in enumerate(xy):
        t = np.clip(((p - a) * d).sum(1) / safe, 0.0, 1.0)
        closest = a + t[:, None] * d
        out[k] = float(np.min(np.hypot(*(closest - p).T)))
    return out


def deviation_spans(cumulative: Sequence[float], distances: np.ndarray) -> list[dict[str, Any]]:
    spans, start, worst = [], None, 0.0
    for i, d in enumerate(distances):
        if d > DEVIATION_M:
            if start is None:
                start, worst = i, 0.0
            worst = max(worst, float(d))
        elif start is not None:
            spans.append(
                {
                    "from_m": round(cumulative[start]),
                    "to_m": round(cumulative[i - 1]),
                    "max_m": round(worst),
                }
            )
            start = None
    if start is not None:
        spans.append(
            {
                "from_m": round(cumulative[start]),
                "to_m": round(cumulative[-1]),
                "max_m": round(worst),
            }
        )
    return spans


def match_to_trails(
    points: Sequence[dict[str, Any]],
    stage_end_indexes: Sequence[int],
    *,
    profile: str,
    opener,
    dem=None,
    elevations: Callable | None = None,
) -> dict[str, Any]:
    """The matched version of an imported line (``points``: the preview's
    {latitude, longitude} list). Raises RuntimeError (BRouter failed) or
    ValueError (the route is too long to match in one go)."""
    line = [_Point(p["latitude"], p["longitude"]) for p in points]
    cumulative = cumulative_distances(line)
    samples = sample_indexes(cumulative, set(stage_end_indexes), SAMPLE_SPACING_M)
    chunks = [samples[i : i + CHUNK_POINTS] for i in range(0, len(samples) - 1, CHUNK_POINTS - 1)]
    chunks = [chunk for chunk in chunks if len(chunk) >= 2]
    if len(chunks) > MAX_REQUESTS:
        raise ValueError(
            f"Trasa jest za długa, by dopasować ją do szlaków naraz ({len(chunks)} zapytań; "
            f"najwyżej {MAX_REQUESTS})."
        )
    matched: list[list[float]] = []
    surfaces: list[dict[str, Any] | None] = []
    sample_at: dict[int, int] = {}  # sample's index in ``points`` -> index in ``matched``
    for chunk in chunks:
        geometry, chunk_surfaces = fetch_route_segment_detail(
            [RoutePointInput(line[i].latitude, line[i].longitude) for i in chunk],
            profile=profile,
            opener=opener,
        )
        coordinates = geometry["coordinates"]
        skip = 1 if matched else 0
        sample_at.setdefault(chunk[0], len(matched) - 1 if matched else 0)
        matched.extend(coordinates[skip:])
        surfaces.extend(chunk_surfaces[skip:])
        sample_at[chunk[-1]] = len(matched) - 1
    matched_line = [_Point(c[1], c[0]) for c in matched]
    matched_cumulative = cumulative_distances(matched_line)
    # A stage end lands on the matched point nearest to it, in route order.
    ends, floor = [], 0
    for index in sorted(stage_end_indexes):
        if index in sample_at:
            at = sample_at[index]
        else:
            xy = _local_xy(matched_line[floor:], line[index])
            at = floor + int(np.argmin(np.hypot(xy[:, 0], xy[:, 1])))
        if ends:
            at = max(at, floor + 1)
        ends.append(max(1, min(at, len(matched_line) - 2)))
        floor = ends[-1]
    heights = (
        elevations(matched_line) if elevations is not None else dem_elevations(matched_line, dem)
    )
    deviations = deviation_spans(cumulative, distances_to_line(line, matched_line))
    return {
        "points": [
            {"latitude": p.latitude, "longitude": p.longitude, "elevation_m": h}
            for p, h in zip(matched_line, heights)
        ],
        "point_surfaces": surfaces,
        "anchor_point_indexes": _anchors(matched_line, matched_cumulative, set(ends)),
        "stage_end_indexes": ends,
        "distance_m": round(matched_cumulative[-1]),
        "original_distance_m": round(cumulative[-1]),
        "deviations": deviations,
        "deviation_threshold_m": DEVIATION_M,
        "requests": len(chunks),
    }
