"""Distance-domain route elevation preprocessing and terrain pace analysis."""

from __future__ import annotations

import bisect
import math
import time
from dataclasses import asdict
from typing import Any, Sequence

from mountain_twin.pace.contract import PaceEvaluationState, PaceInterval, PacePoint, PaceResult
from mountain_twin.pace.model import (
    MODEL_NAME,
    MODEL_VERSION,
    nominal_terrain_speed_kmh,
    terrain_speed_kmh,
)

DISTANCE_STEP_M = 50.0
PREPROCESSING_POLICY = "fixed_distance_linear_route_elevation_resampling"
ELEVATION_SOURCE = "route_following_gpx_elevation_m_raw"
PAUSE_STRATEGY_ID = "distance_triggered_pauses"
PAUSE_STRATEGY_VERSION = "v0_1"


def analyze_route_pace(
    prepared_points: Sequence[Any],
    *,
    scenario_name: str,
    scenario_factor: float,
    pauses: Sequence[Any] = (),
    source_identity: str | None = None,
    distance_step_m: float = DISTANCE_STEP_M,
) -> PaceResult:
    """Build a terrain-aware moving-time result for prepared route points."""
    route_id = getattr(prepared_points[0], "route_id", None) if prepared_points else None
    if not prepared_points:
        return _unresolved_result(None, scenario_name, scenario_factor, "PACE_ROUTE_EMPTY")
    if distance_step_m <= 0 or not math.isfinite(distance_step_m):
        return _unresolved_result(
            route_id, scenario_name, scenario_factor, "PACE_PREPROCESSING_DISTANCE_INVALID"
        )
    if any(
        getattr(point, "point_index", None) != index for index, point in enumerate(prepared_points)
    ):
        return _unresolved_result(
            route_id, scenario_name, scenario_factor, "PACE_ROUTE_ORDER_INVALID"
        )

    blocks = _segment_blocks(prepared_points)
    if blocks is None:
        return _unresolved_result(
            route_id, scenario_name, scenario_factor, "PACE_ROUTE_GEOMETRY_INVALID"
        )

    intervals: list[PaceInterval] = []
    pace_points: list[PacePoint] = []
    ascent = descent = moving_seconds = horizontal_distance = 0.0
    reasons: list[str] = []
    preprocessing_seconds = 0.0
    calculation_seconds = 0.0

    for segment_index, block in blocks:
        preprocessing_started = time.perf_counter()
        anchors, error = _anchors(block)
        if error:
            return _unresolved_result(route_id, scenario_name, scenario_factor, error)
        if anchors is None:
            return _unresolved_result(
                route_id, scenario_name, scenario_factor, "PACE_ELEVATION_UNRESOLVED"
            )
        start_distance = anchors[0][0]
        end_distance = anchors[-1][0]
        length = end_distance - start_distance
        if length < 0 or not math.isfinite(length):
            return _unresolved_result(
                route_id, scenario_name, scenario_factor, "PACE_ROUTE_GEOMETRY_INVALID"
            )
        grid = _grid(length, distance_step_m)
        anchor_distances = [distance for distance, _ in anchors]
        elevations = [_interpolate(local, anchors, anchor_distances) for local in grid]
        preprocessing_seconds += time.perf_counter() - preprocessing_started
        interval_elapsed = [0.0]
        segment_intervals: list[PaceInterval] = []
        calculation_started = time.perf_counter()
        for interval_number, (local_start, local_end, elevation_start, elevation_end) in enumerate(
            zip(grid, grid[1:], elevations, elevations[1:])
        ):
            distance = local_end - local_start
            if distance <= 0 or not math.isfinite(distance):
                return _unresolved_result(
                    route_id, scenario_name, scenario_factor, "PACE_INTERVAL_INVALID"
                )
            elevation_change = elevation_end - elevation_start
            grade = elevation_change / distance
            try:
                nominal_speed = nominal_terrain_speed_kmh(grade)
                scenario_speed = terrain_speed_kmh(grade, scenario_factor)
            except ValueError as exc:
                return _unresolved_result(route_id, scenario_name, scenario_factor, str(exc))
            duration = distance / 1000.0 / scenario_speed * 3600.0
            if not math.isfinite(duration) or duration < 0:
                return _unresolved_result(
                    route_id, scenario_name, scenario_factor, "PACE_DURATION_INVALID"
                )
            interval_id = f"segment:{segment_index}:interval:{interval_number}"
            # interval_elapsed[-1] is this segment's durations so far, added in
            # the same order a sum over them would (AV-048: no O(n^2) re-sum).
            cumulative = moving_seconds + interval_elapsed[-1] + duration
            segment_intervals.append(
                PaceInterval(
                    interval_id,
                    segment_index,
                    start_distance + local_start,
                    start_distance + local_end,
                    distance,
                    elevation_start,
                    elevation_end,
                    elevation_change,
                    grade,
                    nominal_speed,
                    scenario_speed,
                    duration,
                    cumulative,
                    PaceEvaluationState.COMPLETE,
                )
            )
            interval_elapsed.append(interval_elapsed[-1] + duration)
            ascent += max(elevation_change, 0.0)
            descent += max(-elevation_change, 0.0)
            horizontal_distance += distance
        calculation_seconds += time.perf_counter() - calculation_started
        intervals.extend(segment_intervals)
        for point in block:
            local_distance = point.cumulative_distance_m - start_distance
            local_elapsed = _interpolate_scalar(local_distance, grid, interval_elapsed)
            interval_id = _interval_for_distance(local_distance, grid, segment_intervals)
            pace_points.append(
                PacePoint(
                    point.point_index,
                    point.cumulative_distance_m,
                    moving_seconds + local_elapsed,
                    interval_id,
                )
            )
        moving_seconds += interval_elapsed[-1]

    pause_time = sum(float(pause.duration_minutes) * 60.0 for pause in pauses)
    total_time = moving_seconds + pause_time
    provenance = {
        "model": MODEL_NAME,
        "model_version": MODEL_VERSION,
        "elevation_source": ELEVATION_SOURCE,
        "elevation_source_identity": source_identity,
        "preprocessing_policy": PREPROCESSING_POLICY,
        "preprocessing_distance_m": distance_step_m,
        "grade_definition": "elevation_change_m / horizontal_distance_m",
        "grade_domain": "all finite grades producing finite positive model speed",
        "scenario_name": scenario_name,
        "scenario_factor": scenario_factor,
        "pause_strategy": {
            "id": PAUSE_STRATEGY_ID,
            "version": PAUSE_STRATEGY_VERSION,
            "pauses": [asdict(pause) for pause in pauses],
        },
    }
    return PaceResult(
        route_id=route_id,
        model_name=MODEL_NAME,
        model_version=MODEL_VERSION,
        elevation_source=ELEVATION_SOURCE,
        preprocessing_policy=PREPROCESSING_POLICY,
        preprocessing_distance_m=distance_step_m,
        scenario_name=scenario_name,
        scenario_factor=scenario_factor,
        pause_strategy_id=PAUSE_STRATEGY_ID,
        pause_strategy_version=PAUSE_STRATEGY_VERSION,
        horizontal_distance_m=horizontal_distance,
        ascent_m=ascent,
        descent_m=descent,
        moving_time_s=moving_seconds,
        pause_time_s=pause_time,
        total_planned_time_s=total_time,
        state=PaceEvaluationState.COMPLETE,
        reason_codes=tuple(reasons),
        intervals=tuple(intervals),
        points=tuple(sorted(pace_points, key=lambda item: item.point_index)),
        provenance=provenance,
        runtime={
            "preprocessing_seconds": preprocessing_seconds,
            "calculation_seconds": calculation_seconds,
            "total_seconds": preprocessing_seconds + calculation_seconds,
        },
    )


def _segment_blocks(points: Sequence[Any]):
    blocks = []
    current = []
    current_key = None
    for point in points:
        key = getattr(point, "segment_index", None)
        if current and key != current_key:
            blocks.append((current_key, current))
            current = []
        current_key = key
        current.append(point)
    if current:
        blocks.append((current_key, current))
    return blocks


def _anchors(block: Sequence[Any]):
    anchors: list[tuple[float, float]] = []
    previous_distance = None
    for point in block:
        distance = float(point.cumulative_distance_m)
        elevation = getattr(point, "gpx_elevation_m", None)
        if not math.isfinite(distance) or distance < 0:
            return None, "PACE_ROUTE_GEOMETRY_INVALID"
        if elevation is None or not math.isfinite(float(elevation)):
            return None, "PACE_ELEVATION_UNRESOLVED"
        if previous_distance is not None and distance < previous_distance:
            return None, "PACE_ROUTE_GEOMETRY_INVALID"
        if anchors and math.isclose(distance, anchors[-1][0], abs_tol=1e-9):
            if not math.isclose(float(elevation), anchors[-1][1], abs_tol=1e-9):
                return None, "PACE_EQUAL_DISTANCE_ELEVATION_CONFLICT"
        else:
            anchors.append((distance, float(elevation)))
        previous_distance = distance
    return anchors, None


def _grid(length: float, step: float) -> list[float]:
    values = [0.0]
    index = 1
    while index * step < length:
        values.append(index * step)
        index += 1
    if length > 0 and not math.isclose(values[-1], length, abs_tol=1e-9):
        values.append(length)
    return values


# AV-048: the three lookups below find their pair by binary search
# (bisect) instead of a scan from the start -- the same pair as the scan,
# the same arithmetic, so the same result to the bit; a 237 km route went
# from ~3 s to a fraction of a second per analysis.


def _interpolate(
    local_distance: float,
    anchors: Sequence[tuple[float, float]],
    anchor_distances: Sequence[float] | None = None,
) -> float:
    if len(anchors) == 1:
        return anchors[0][1]
    local_start = anchors[0][0]
    target = local_start + local_distance
    if target <= anchors[0][0]:
        return anchors[0][1]
    distances = anchor_distances or [distance for distance, _ in anchors]
    # The first pair whose right end is at or past the target.
    right = bisect.bisect_left(distances, target, 1)
    if right >= len(anchors):
        return anchors[-1][1]
    (left_distance, left_elevation), (right_distance, right_elevation) = (
        anchors[right - 1],
        anchors[right],
    )
    fraction = (target - left_distance) / (right_distance - left_distance)
    return left_elevation + fraction * (right_elevation - left_elevation)


def _interpolate_scalar(
    target: float, positions: Sequence[float], values: Sequence[float]
) -> float:
    if target <= positions[0]:
        return values[0]
    # The first position at or past the target.
    index = bisect.bisect_left(positions, target, 1)
    if index >= len(positions):
        return values[-1]
    left, right = positions[index - 1], positions[index]
    fraction = (target - left) / (right - left)
    return values[index - 1] + fraction * (values[index] - values[index - 1])


def _interval_for_distance(
    target: float, positions: Sequence[float], intervals: Sequence[PaceInterval]
):
    if not intervals:
        return None
    # The first interval whose right end is past the target.
    index = bisect.bisect_right(positions, target, 1) - 1
    if index < min(len(intervals), len(positions) - 1):
        return intervals[index].interval_id
    return intervals[-1].interval_id


def _unresolved_result(route_id, scenario_name, scenario_factor, reason):
    return PaceResult(
        route_id=route_id,
        model_name=MODEL_NAME,
        model_version=MODEL_VERSION,
        elevation_source=ELEVATION_SOURCE,
        preprocessing_policy=PREPROCESSING_POLICY,
        preprocessing_distance_m=DISTANCE_STEP_M,
        scenario_name=scenario_name,
        scenario_factor=scenario_factor,
        pause_strategy_id=PAUSE_STRATEGY_ID,
        pause_strategy_version=PAUSE_STRATEGY_VERSION,
        horizontal_distance_m=0.0,
        ascent_m=None,
        descent_m=None,
        moving_time_s=None,
        pause_time_s=None,
        total_planned_time_s=None,
        state=PaceEvaluationState.UNRESOLVED,
        reason_codes=(reason,),
        intervals=(),
        points=(),
        provenance={
            "model": MODEL_NAME,
            "model_version": MODEL_VERSION,
            "elevation_source": ELEVATION_SOURCE,
            "preprocessing_policy": PREPROCESSING_POLICY,
            "preprocessing_distance_m": DISTANCE_STEP_M,
            "scenario_name": scenario_name,
            "scenario_factor": scenario_factor,
        },
        runtime={},
    )
