"""Surface along a drawn route (AV-031, WP2 of docs/design_reference/
activity_types_spike_v0_1.md): what OSM says each stretch is made of, read
from the BRouter answer the route was snapped with (its ``messages`` table,
``WayTags``) -- data BRouter already sends and the app used to drop.

Nothing is guessed. A way with no ``surface`` tag in OSM is UNKNOWN, a class
of its own, never "probably asphalt" from the road type (BRouter's profiles
guess that for costing; the map must not). A stretch for which no BRouter
answer was kept (a route saved before AV-031, or an imported one) has no
class at all: NOT_RECORDED, said as such.

The classes only group OSM's own values for the legend; the raw ``surface``
value is kept with every span.
"""

from __future__ import annotations

from typing import Any, Sequence

from mountain_twin.route_analysis import prepare_route

SURFACE_CONTRACT = "route_surface_v0_1"

# Display class -> the OSM surface values it groups (wiki Key:surface).
SURFACE_CLASSES: dict[str, tuple[str, ...]] = {
    "PAVED": ("asphalt", "concrete", "concrete:plates", "concrete:lanes", "paved", "chipseal"),
    "SETT": ("paving_stones", "sett", "cobblestone", "unhewn_cobblestone", "grass_paver", "bricks"),
    "COMPACTED": ("compacted", "fine_gravel"),
    "GRAVEL": ("gravel", "pebblestone", "unpaved", "rock", "shells"),
    "NATURAL": ("ground", "dirt", "earth", "grass", "sand", "mud", "clay", "woodchips", "soil"),
}
CLASS_BY_SURFACE = {value: name for name, values in SURFACE_CLASSES.items() for value in values}
# A tagged surface outside the groups above (wood, metal, ...).
OTHER = "OTHER"
# No surface tag on the OSM way.
UNKNOWN = "UNKNOWN"
CLASSES = (*SURFACE_CLASSES, OTHER, UNKNOWN)


def surface_class(surface: str | None) -> str:
    if not surface:
        return UNKNOWN
    return CLASS_BY_SURFACE.get(surface, OTHER)


def point_surfaces_from_brouter(
    coordinates: Sequence[Sequence[float]], messages: Sequence[Sequence[str]]
) -> list[dict[str, Any] | None]:
    """One entry per geometry coordinate: the surface of the way that leads
    to it (the first point takes the first way's). BRouter's ``messages``
    rows are the points where the way changes, at the same coordinates as
    the geometry (micro-degrees); each row's ``WayTags`` belong to the way
    ending there. A coordinate no row covers stays None (not recorded)."""
    if not messages or len(messages) < 2:
        return [None] * len(coordinates)
    header = list(messages[0])
    try:
        lon_at, lat_at, tags_at = (
            header.index("Longitude"),
            header.index("Latitude"),
            header.index("WayTags"),
        )
    except ValueError:
        return [None] * len(coordinates)
    keys = [(round(c[0] * 1e6), round(c[1] * 1e6)) for c in coordinates]
    result: list[dict[str, Any] | None] = [None] * len(coordinates)
    position = 0  # the next coordinate a row may end at
    for row in messages[1:]:
        try:
            key = (int(row[lon_at]), int(row[lat_at]))
        except (ValueError, IndexError):
            continue
        end = next((j for j in range(position, len(keys)) if keys[j] == key), None)
        if end is None:
            continue
        tags = dict(item.split("=", 1) for item in str(row[tags_at]).split() if "=" in item)
        entry = {"class": surface_class(tags.get("surface")), "surface": tags.get("surface")}
        for j in range(position, end + 1):
            result[j] = entry
        position = end + 1
    return result


def validate_point_surfaces(values: Any, point_count: int) -> tuple[dict[str, Any] | None, ...]:
    """A route save's ``point_surfaces``: one entry per point, each None or
    {"class": one of CLASSES, "surface": the raw OSM value or None}."""
    if not isinstance(values, list) or len(values) != point_count:
        raise ValueError("point_surfaces must give one entry per route point")
    checked = []
    for value in values:
        if value is None:
            checked.append(None)
            continue
        if not isinstance(value, dict) or value.get("class") not in CLASSES:
            raise ValueError("point surface class is invalid")
        raw = value.get("surface")
        if raw is not None and (not isinstance(raw, str) or len(raw) > 64):
            raise ValueError("point surface value is invalid")
        if surface_class(raw) != value["class"]:
            raise ValueError("point surface class does not match its OSM value")
        checked.append({"class": value["class"], "surface": raw})
    return tuple(checked)


def surface_spans(point_surfaces: Sequence[dict[str, Any] | None]) -> list[dict[str, Any]]:
    """Run-length spans over point indexes: [start, end] share one entry
    (None runs are NOT_RECORDED stretches)."""
    spans: list[dict[str, Any]] = []
    for index, value in enumerate(point_surfaces):
        if spans and spans[-1]["value"] == value:
            spans[-1]["end"] = index
        else:
            spans.append({"start": index, "end": index, "value": value})
    return [
        {
            "start_point_index": span["start"],
            "end_point_index": span["end"],
            "class": None if span["value"] is None else span["value"]["class"],
            "surface": None if span["value"] is None else span["value"]["surface"],
        }
        for span in spans
    ]


def route_surface_document(
    route_points: Sequence[Any], spans: Sequence[dict[str, Any]] | None
) -> dict[str, Any]:
    """The surface the page shows: the spans and each class's share of the
    route's length (``parts``; a point's class is the way leading to it, so each step
    between two points counts for the class of its end point). No spans at
    all: NOT_EVALUATED -- the route was saved before surfaces were kept."""
    if not spans:
        return {
            "contract": SURFACE_CONTRACT,
            "state": "NOT_EVALUATED",
            "reason": "SURFACE_NOT_RECORDED",
        }
    prepared = prepare_route(route_points).points
    distance = {point.point_index: point.cumulative_distance_m for point in prepared}
    lengths: dict[str, float] = {}
    total = 0.0
    for span in spans:
        name = span["class"] or "NOT_RECORDED"
        for index in range(max(span["start_point_index"], 1), span["end_point_index"] + 1):
            if index in distance and index - 1 in distance:
                step = distance[index] - distance[index - 1]
                lengths[name] = lengths.get(name, 0.0) + step
                total += step
    return {
        "contract": SURFACE_CONTRACT,
        "state": "AVAILABLE",
        "source": "OpenStreetMap surface=* via BRouter WayTags",
        "spans": list(spans),
        "parts": [
            {
                "class": name,
                "length_m": round(length, 1),
                "fraction": round(length / total, 4) if total else 0.0,
            }
            for name, length in sorted(lengths.items(), key=lambda item: -item[1])
        ],
    }
