"""Conservative offline association of a route to supplied mapped geometries."""

from __future__ import annotations

import math
from typing import Sequence

from mountain_twin.trail_character.contracts import (
    AssociationSpan,
    AssociationState,
    CandidateAssociationEvidence,
    CharacterCoverage,
    GeoCoordinate,
    MappedFeature,
    RouteAssociationPolicy,
    RouteAssociationResult,
)
from mountain_twin.trails.trail_engine import bearing_deg, haversine_m

DEFAULT_ASSOCIATION_POLICY = RouteAssociationPolicy(
    policy_id="mt_route_association_rnd",
    version="v0_1",
    threshold_origin="MT_RND_POLICY",
    maximum_lateral_distance_m=15.0,
    maximum_direction_difference_deg=30.0,
)


def associate_route(
    route_geometry: Sequence[GeoCoordinate],
    candidates: Sequence[MappedFeature],
    *,
    candidate_coverage: CharacterCoverage,
    policy: RouteAssociationPolicy = DEFAULT_ASSOCIATION_POLICY,
) -> RouteAssociationResult:
    """Associate finite supplied candidates using explicit lateral and direction predicates.

    This is not nearest-way matching.  A route leg is associated only where a
    candidate satisfies both predicates; equally qualifying candidates remain
    AMBIGUOUS.  Adjacent spans are coalesced only when their state and exact
    candidate set agree, so data/source gaps are never bridged.
    """
    if candidate_coverage in {CharacterCoverage.UNAVAILABLE, CharacterCoverage.UNKNOWN}:
        return _unknown_result(policy, "MAPPED_CANDIDATE_COVERAGE_UNAVAILABLE")
    if not _valid_policy(policy):
        return _unknown_result(policy, "ASSOCIATION_POLICY_INVALID")
    if not _valid_geometry(route_geometry):
        return _unknown_result(policy, "ROUTE_GEOMETRY_INSUFFICIENT_OR_INVALID")

    route_distances = _route_distances(route_geometry)
    valid, invalid = _partition_candidates(candidates)
    if not valid and invalid:
        reason = (
            "MAPPED_CANDIDATE_CRS_UNSUPPORTED"
            if all(candidate.crs != "EPSG:4326" for candidate in invalid)
            else "MAPPED_CANDIDATE_GEOMETRY_INVALID"
        )
        return _unknown_result(policy, reason, invalid)
    if not valid:
        return RouteAssociationResult(
            AssociationState.UNMATCHED,
            candidate_coverage,
            policy,
            (),
            _unmatched_spans(route_distances, "NO_MAPPED_CANDIDATES_IN_EVALUATED_COVERAGE"),
            ("NO_MAPPED_CANDIDATES_IN_EVALUATED_COVERAGE",),
        )

    support: dict[str, list[tuple[bool, float, float | None]]] = {
        candidate.feature_id: [] for candidate in valid
    }
    leg_candidates: list[tuple[str, ...]] = []
    for start, end in zip(route_geometry, route_geometry[1:]):
        length = haversine_m(start.latitude, start.longitude, end.latitude, end.longitude)
        qualifying = []
        for candidate in valid:
            lateral, direction_ok = _candidate_leg_evidence(start, end, candidate, policy)
            compatible = (
                lateral is not None
                and lateral <= policy.maximum_lateral_distance_m
                and direction_ok
            )
            support[candidate.feature_id].append((compatible, length, lateral))
            if compatible:
                qualifying.append(candidate.feature_id)
        leg_candidates.append(tuple(sorted(qualifying)))

    evidence = tuple(
        _candidate_evidence(candidate, support[candidate.feature_id], route_distances[-1])
        for candidate in sorted(valid, key=lambda item: item.feature_id)
    )
    spans = _spans(route_distances, leg_candidates)
    state = _overall_state(spans)
    coverage = _association_coverage(spans, route_distances[-1])
    diagnostics = tuple(sorted({"MAPPED_CANDIDATE_GEOMETRY_INVALID" for _ in invalid}))
    return RouteAssociationResult(state, coverage, policy, evidence, spans, diagnostics)


def mapped_extension(association: RouteAssociationResult):
    """Build the canonical mapped extension without adding tags or normalisation."""
    from mountain_twin.trail_character.contracts import MappedCharacterExtension

    return MappedCharacterExtension(
        coverage=association.coverage,
        populated=True,
        data_gaps=(),
        association=association,
    )


def _candidate_leg_evidence(
    start: GeoCoordinate,
    end: GeoCoordinate,
    candidate: MappedFeature,
    policy: RouteAssociationPolicy,
) -> tuple[float | None, bool]:
    midpoint = GeoCoordinate(
        (start.latitude + end.latitude) / 2, (start.longitude + end.longitude) / 2
    )
    route_bearing = bearing_deg(start.latitude, start.longitude, end.latitude, end.longitude)
    best_distance = None
    direction_ok = False
    for left, right in zip(candidate.geometry, candidate.geometry[1:]):
        distance = _point_segment_distance_m(midpoint, left, right)
        if best_distance is None or distance < best_distance:
            best_distance = distance
        candidate_bearing = bearing_deg(
            left.latitude, left.longitude, right.latitude, right.longitude
        )
        difference = _undirected_bearing_difference(route_bearing, candidate_bearing)
        if (
            distance <= policy.maximum_lateral_distance_m
            and difference <= policy.maximum_direction_difference_deg
        ):
            direction_ok = True
    return best_distance, direction_ok


def _candidate_evidence(
    candidate: MappedFeature,
    support: Sequence[tuple[bool, float, float | None]],
    route_length: float,
) -> CandidateAssociationEvidence:
    supported = sum(length for compatible, length, _ in support if compatible)
    distances = [
        distance for compatible, _, distance in support if compatible and distance is not None
    ]
    return CandidateAssociationEvidence(
        candidate.feature_id,
        candidate.source.source_id,
        supported,
        supported / route_length if route_length else None,
        max(distances) if distances else None,
        supported,
        True,
    )


def _spans(
    distances: Sequence[float], leg_candidates: Sequence[tuple[str, ...]]
) -> tuple[AssociationSpan, ...]:
    spans = []
    start = 0
    for index in range(1, len(leg_candidates) + 1):
        if index == len(leg_candidates) or leg_candidates[index] != leg_candidates[start]:
            candidates = leg_candidates[start]
            if len(candidates) == 1:
                state, reasons = AssociationState.MATCHED, ()
            elif len(candidates) > 1:
                state, reasons = (
                    AssociationState.AMBIGUOUS,
                    ("MULTIPLE_CANDIDATES_SATISFY_PREDICATES",),
                )
            else:
                state, reasons = AssociationState.UNMATCHED, ("NO_CANDIDATE_SATISFIES_PREDICATES",)
            spans.append(
                AssociationSpan(state, distances[start], distances[index], candidates, reasons)
            )
            start = index
    return tuple(spans)


def _overall_state(spans: Sequence[AssociationSpan]) -> AssociationState:
    states = {span.state for span in spans}
    if AssociationState.AMBIGUOUS in states:
        return AssociationState.AMBIGUOUS
    if AssociationState.MATCHED in states:
        return AssociationState.MATCHED
    return AssociationState.UNMATCHED


def _association_coverage(
    spans: Sequence[AssociationSpan], total_distance: float
) -> CharacterCoverage:
    matched = sum(
        span.end_route_distance_m - span.start_route_distance_m
        for span in spans
        if span.state is AssociationState.MATCHED
    )
    if matched:
        return (
            CharacterCoverage.FULL
            if math.isclose(matched, total_distance)
            else CharacterCoverage.PARTIAL
        )
    if any(span.state is AssociationState.UNMATCHED for span in spans) and any(
        span.state is AssociationState.AMBIGUOUS for span in spans
    ):
        return CharacterCoverage.PARTIAL
    return CharacterCoverage.FULL


def _unmatched_spans(distances: Sequence[float], reason: str) -> tuple[AssociationSpan, ...]:
    return (
        AssociationSpan(AssociationState.UNMATCHED, distances[0], distances[-1], (), (reason,)),
    )


def _unknown_result(
    policy: RouteAssociationPolicy, reason: str, invalid: Sequence[MappedFeature] = ()
) -> RouteAssociationResult:
    evidence = tuple(
        CandidateAssociationEvidence(
            item.feature_id, item.source.source_id, 0.0, None, None, 0.0, False, (reason,)
        )
        for item in sorted(invalid, key=lambda item: item.feature_id)
    )
    return RouteAssociationResult(
        AssociationState.UNKNOWN,
        CharacterCoverage.UNKNOWN,
        policy,
        evidence,
        (),
        (reason,),
    )


def _partition_candidates(candidates: Sequence[MappedFeature]):
    valid, invalid = [], []
    for candidate in candidates:
        (
            valid
            if candidate.crs == "EPSG:4326" and _valid_geometry(candidate.geometry)
            else invalid
        ).append(candidate)
    return valid, invalid


def _valid_geometry(geometry: Sequence[GeoCoordinate]) -> bool:
    return len(geometry) >= 2 and all(
        math.isfinite(point.latitude)
        and math.isfinite(point.longitude)
        and -90 <= point.latitude <= 90
        and -180 <= point.longitude <= 180
        for point in geometry
    )


def _valid_policy(policy: RouteAssociationPolicy) -> bool:
    return (
        policy.threshold_origin == "MT_RND_POLICY"
        and math.isfinite(policy.maximum_lateral_distance_m)
        and policy.maximum_lateral_distance_m >= 0
        and math.isfinite(policy.maximum_direction_difference_deg)
        and 0 <= policy.maximum_direction_difference_deg <= 90
    )


def _route_distances(route: Sequence[GeoCoordinate]) -> tuple[float, ...]:
    values = [0.0]
    for left, right in zip(route, route[1:]):
        values.append(
            values[-1] + haversine_m(left.latitude, left.longitude, right.latitude, right.longitude)
        )
    return tuple(values)


def _undirected_bearing_difference(first: float, second: float) -> float:
    difference = abs((first - second + 180) % 360 - 180)
    return min(difference, abs(180 - difference))


def _point_segment_distance_m(
    point: GeoCoordinate, left: GeoCoordinate, right: GeoCoordinate
) -> float:
    """Local equirectangular point-to-segment metric approximation in metres."""
    latitude = math.radians(point.latitude)
    scale_x = 6371008.8 * math.cos(latitude) * math.pi / 180
    scale_y = 6371008.8 * math.pi / 180
    ax, ay = (
        (left.longitude - point.longitude) * scale_x,
        (left.latitude - point.latitude) * scale_y,
    )
    bx, by = (
        (right.longitude - point.longitude) * scale_x,
        (right.latitude - point.latitude) * scale_y,
    )
    dx, dy = bx - ax, by - ay
    denominator = dx * dx + dy * dy
    if denominator == 0:
        return math.hypot(ax, ay)
    fraction = max(0.0, min(1.0, -(ax * dx + ay * dy) / denominator))
    return math.hypot(ax + fraction * dx, ay + fraction * dy)
