"""Narrow, deterministic mapped-character composition for A5 WP3."""

from __future__ import annotations

import math
from collections import defaultdict
from dataclasses import replace
from typing import Sequence

from mountain_twin.trail_character.contracts import (
    AssociationState,
    MappedCharacterExtension,
    MappedCharacteristicSegment,
    MappedCharacteristicSummary,
    MappedCharacteristicTotal,
    MappedFeature,
    MappedNormalizationPolicy,
    MappedValueState,
    RouteAssociationContinuityResult,
    RouteAssociationResult,
)

DEFAULT_MAPPED_NORMALIZATION_POLICY = MappedNormalizationPolicy(
    policy_id="mt_mapped_character_normalization",
    version="v0_1",
    route_class_vocabulary=(
        "PATH", "FOOTWAY", "TRACK", "STEPS", "BRIDLEWAY", "SERVICE", "UNCLASSIFIED", "RESIDENTIAL",
    ),
    surface_vocabulary=("GROUND", "ROCK", "UNPAVED", "GRAVEL", "OTHER_TAGGED"),
)

_ROUTE_CLASS = {value.lower(): value for value in DEFAULT_MAPPED_NORMALIZATION_POLICY.route_class_vocabulary}
_SURFACE = {value.lower(): value for value in DEFAULT_MAPPED_NORMALIZATION_POLICY.surface_vocabulary if value != "OTHER_TAGGED"}
_CHARACTERISTICS = ("route_class", "surface")


def compose_mapped_character(
    association: RouteAssociationResult,
    features: Sequence[MappedFeature],
    *,
    continuity: RouteAssociationContinuityResult | None = None,
    policy: MappedNormalizationPolicy = DEFAULT_MAPPED_NORMALIZATION_POLICY,
) -> tuple[MappedCharacterExtension, tuple[MappedCharacteristicSegment, ...]]:
    """Characterize raw association spans without changing their semantics.

    Short unmatched gaps represented by WP2C continuity are never turned into a
    mapped value here.  Every segment therefore references raw association
    span indices and keeps source feature IDs and raw source values.
    """
    _validate_policy(policy)
    feature_by_id = {feature.feature_id: feature for feature in features}
    if len(feature_by_id) != len(features):
        raise ValueError("mapped feature identities must be unique")
    segments = tuple(
        segment
        for characteristic_id in _CHARACTERISTICS
        for segment in _segments_for_characteristic(
            characteristic_id, association, feature_by_id, policy
        )
    )
    summaries = tuple(
        _summary(characteristic_id, segments, association) for characteristic_id in _CHARACTERISTICS
    )
    return (
        MappedCharacterExtension(
            coverage=association.coverage,
            populated=True,
            data_gaps=(),
            association=association,
            continuity=continuity,
            normalization_policy=policy,
            characteristics=summaries,
        ),
        segments,
    )


def apply_mapped_character(
    result,
    association: RouteAssociationResult,
    features: Sequence[MappedFeature],
    *,
    continuity: RouteAssociationContinuityResult | None = None,
    policy: MappedNormalizationPolicy = DEFAULT_MAPPED_NORMALIZATION_POLICY,
):
    """Populate the canonical result, rejecting a mismatched distance domain."""
    if association.spans:
        route_distance = result.route.distance_m
        if route_distance is None or not math.isclose(
            association.spans[-1].end_route_distance_m, route_distance, abs_tol=1e-6
        ):
            raise ValueError("mapped association and result route distance domains must match")
    extension, segments = compose_mapped_character(
        association, features, continuity=continuity, policy=policy
    )
    return replace(
        result,
        mapped_character=extension,
        mapped_character_segments=segments,
        data_gaps=tuple(item for item in result.data_gaps if item != "MAPPED_ROUTE_CHARACTER_NOT_IMPLEMENTED_WP1"),
        diagnostics=tuple(item for item in result.diagnostics if item != "NO_TERRAIN_SEGMENTATION_IN_WP1"),
    )


def _segments_for_characteristic(characteristic_id, association, features, policy):
    result = []
    for index, span in enumerate(association.spans):
        state, normalized, raw_values, reasons = _value_for_span(
            characteristic_id, span, features
        )
        result.append(
            MappedCharacteristicSegment(
                characteristic_id,
                span.start_route_distance_m,
                span.end_route_distance_m,
                state,
                normalized,
                raw_values,
                span.candidate_feature_ids,
                (index,),
                policy.policy_id,
                policy.version,
                reasons,
            )
        )
    return _coalesce_same_characteristic(result)


def _coalesce_same_characteristic(segments):
    """Coalesce only adjacent identical known facts, retaining every raw reference.

    This is deliberately characteristic-local: a surface transition on a new
    source way does not split an unchanged route class.  Unknown, unmatched,
    and ambiguous evidence is always a hard boundary.
    """
    coalesced = []
    for segment in segments:
        previous = coalesced[-1] if coalesced else None
        if (
            previous is not None
            and previous.value_state is MappedValueState.KNOWN
            and segment.value_state is MappedValueState.KNOWN
            and previous.normalized_value == segment.normalized_value
            and previous.raw_values == segment.raw_values
            and math.isclose(previous.end_route_distance_m, segment.start_route_distance_m, abs_tol=1e-9)
        ):
            coalesced[-1] = replace(
                previous,
                end_route_distance_m=segment.end_route_distance_m,
                source_feature_ids=tuple(sorted(set(previous.source_feature_ids + segment.source_feature_ids))),
                raw_association_span_indices=(
                    previous.raw_association_span_indices + segment.raw_association_span_indices
                ),
            )
        else:
            coalesced.append(segment)
    return tuple(coalesced)


def _value_for_span(characteristic_id, span, features):
    if span.state is AssociationState.UNMATCHED:
        return MappedValueState.UNMATCHED, None, (), ("RAW_ASSOCIATION_UNMATCHED",)
    if span.state is AssociationState.UNKNOWN:
        return MappedValueState.UNKNOWN, None, (), ("RAW_ASSOCIATION_UNKNOWN",)
    candidate_values = tuple(
        _raw_value(characteristic_id, features[item]) for item in span.candidate_feature_ids
    )
    values = tuple(sorted({value for value in candidate_values if value is not None}))
    normalized = tuple(sorted({_normalise(characteristic_id, value) for value in values}))
    if span.state is AssociationState.AMBIGUOUS and len(set(candidate_values)) > 1:
        return (
            MappedValueState.AMBIGUOUS,
            None,
            values,
            (
                "COMPETING_SOURCE_VALUES_INCLUDING_ABSENT_TAG"
                if None in candidate_values
                else "COMPETING_SOURCE_VALUES",
            ),
        )
    # Geometrically ambiguous features can still provide the identical raw map fact.
    value = candidate_values[0] if candidate_values else None
    normalized_value = normalized[0] if normalized else None
    if value is None:
        return MappedValueState.UNKNOWN, None, (), ("SOURCE_TAG_ABSENT",)
    if normalized_value is None:
        return MappedValueState.UNKNOWN, None, values, ("SOURCE_VALUE_OUTSIDE_V0_1_VOCABULARY",)
    return MappedValueState.KNOWN, normalized_value, values, ()


def _raw_value(characteristic_id, feature):
    key = "highway" if characteristic_id == "route_class" else "surface"
    value = feature.raw_properties.get(key)
    return str(value) if value is not None else None


def _normalise(characteristic_id, raw_value):
    if raw_value is None:
        return None
    lowered = raw_value.lower()
    if characteristic_id == "route_class":
        return _ROUTE_CLASS.get(lowered)
    return _SURFACE.get(lowered, "OTHER_TAGGED")


def _summary(characteristic_id, segments, association):
    selected = [item for item in segments if item.characteristic_id == characteristic_id]
    totals = defaultdict(float)
    for segment in selected:
        totals[(segment.value_state, segment.normalized_value)] += (
            segment.end_route_distance_m - segment.start_route_distance_m
        )
    return MappedCharacteristicSummary(
        characteristic_id,
        len(selected),
        association.spans[-1].end_route_distance_m if association.spans else 0.0,
        tuple(
            MappedCharacteristicTotal(state, value, distance)
            for (state, value), distance in sorted(
                totals.items(), key=lambda item: (item[0][0].value, item[0][1] or "")
            )
        ),
        tuple(sorted({raw for segment in selected for raw in segment.raw_values})),
    )


def _validate_policy(policy):
    if (
        not policy.policy_id
        or not policy.version
        or len(set(policy.route_class_vocabulary)) != len(policy.route_class_vocabulary)
        or len(set(policy.surface_vocabulary)) != len(policy.surface_vocabulary)
    ):
        raise ValueError("mapped normalization policy is invalid")
