"""Compose bounded terrain-light encounters from authoritative Unified facts."""

from __future__ import annotations

from dataclasses import asdict
from typing import Any, Mapping, Sequence

from mountain_twin.photographer.contracts import (
    BoundarySemantics,
    RouteBoundary,
    RouteSampleReference,
    TerrainLightEvent,
    TerrainLightEventType,
    TerrainLightResult,
    derive_terrain_light_event_id,
)
from mountain_twin.photographer.policy import PhotographerPolicy
from mountain_twin.rci.contracts import CoverageState
from mountain_twin.solar.states import TerrainSolarVisibility

_TERRAIN_SOURCE_REFERENCE = "unified_route_analysis.points[].solar"


def project_terrain_light(
    unified_analysis: Any,
    *,
    upstream_analysis_reference: str,
    policy: PhotographerPolicy,
) -> TerrainLightResult:
    """Project existing route-level terrain states without any new solar query.

    The Unified `terrain_solar_visibility` fact is preferred, with its
    compatibility `status` field accepted for existing Unified payloads.  Only
    adjacent SHADOW/DIRECT pairs produce transition events: UNKNOWN and
    NOT_APPLICABLE deliberately break possible transitions.
    """
    records = _records(_value(unified_analysis, "points", ()))
    coverage, reasons = _coverage(unified_analysis, records)
    counts = {state.value: 0 for state in TerrainSolarVisibility}
    for record in records:
        counts[record.state.value] += 1
    if coverage is CoverageState.UNAVAILABLE:
        return TerrainLightResult(
            coverage=coverage,
            source_reference=_TERRAIN_SOURCE_REFERENCE,
            state_counts=counts,
            reason_codes=reasons,
        )

    first_direct_index = next(
        (index for index, record in enumerate(records) if record.state is TerrainSolarVisibility.DIRECT),
        None,
    )
    last_direct_index = next(
        (index for index in range(len(records) - 1, -1, -1) if records[index].state is TerrainSolarVisibility.DIRECT),
        None,
    )
    events: list[TerrainLightEvent] = []
    if first_direct_index is not None:
        events.append(
            _direct_encounter_event(
                records,
                first_direct_index,
                TerrainLightEventType.FIRST_DIRECT_TERRAIN_LIGHT,
                upstream_analysis_reference,
                policy,
            )
        )
    for before, after in zip(records, records[1:]):
        if (before.state, after.state) == (
            TerrainSolarVisibility.SHADOW,
            TerrainSolarVisibility.DIRECT,
        ):
            events.append(
                _transition_event(
                    before,
                    after,
                    TerrainLightEventType.DIRECT_LIGHT_ENTRY,
                    upstream_analysis_reference,
                    policy,
                )
            )
        elif (before.state, after.state) == (
            TerrainSolarVisibility.DIRECT,
            TerrainSolarVisibility.SHADOW,
        ):
            events.append(
                _transition_event(
                    before,
                    after,
                    TerrainLightEventType.DIRECT_LIGHT_EXIT,
                    upstream_analysis_reference,
                    policy,
                )
            )
    if last_direct_index is not None:
        events.append(
            _direct_encounter_event(
                records,
                last_direct_index,
                TerrainLightEventType.LAST_DIRECT_TERRAIN_LIGHT,
                upstream_analysis_reference,
                policy,
            )
        )

    ordered_events = tuple(sorted(events, key=_event_order_key))
    first_event = next(
        (event for event in ordered_events if event.event_type is TerrainLightEventType.FIRST_DIRECT_TERRAIN_LIGHT),
        None,
    )
    last_event = next(
        (event for event in ordered_events if event.event_type is TerrainLightEventType.LAST_DIRECT_TERRAIN_LIGHT),
        None,
    )
    return TerrainLightResult(
        coverage=coverage,
        source_reference=_TERRAIN_SOURCE_REFERENCE,
        state_counts=counts,
        events=ordered_events,
        first_direct_event_id=first_event.event_id if first_event else None,
        last_direct_event_id=last_event.event_id if last_event else None,
        reason_codes=reasons,
    )


class _TerrainRecord:
    def __init__(self, sample: RouteSampleReference, state: TerrainSolarVisibility, source_reference: str | None, explicit_state: bool):
        self.sample = sample
        self.state = state
        self.source_reference = source_reference
        self.explicit_state = explicit_state


def _records(points: Sequence[Any]) -> tuple[_TerrainRecord, ...]:
    records = []
    previous_index = None
    for point in points:
        sample = _sample_reference(point)
        if previous_index is not None and sample.route_point_index <= previous_index:
            raise ValueError("Unified route points must be strictly ordered")
        previous_index = sample.route_point_index
        solar = _mapping(_value(point, "solar", {}))
        raw_state = solar.get("terrain_solar_visibility", solar.get("status"))
        state = _terrain_state(raw_state)
        records.append(
            _TerrainRecord(
                sample,
                state,
                _string_or_none(_value(point, "solar_provenance_ref")),
                _recognized_terrain_state(raw_state),
            )
        )
    return tuple(records)


def _terrain_state(value: Any) -> TerrainSolarVisibility:
    raw = getattr(value, "value", value)
    try:
        return TerrainSolarVisibility(raw)
    except (TypeError, ValueError):
        return TerrainSolarVisibility.UNKNOWN


def _recognized_terrain_state(value: Any) -> bool:
    raw = getattr(value, "value", value)
    try:
        TerrainSolarVisibility(raw)
    except (TypeError, ValueError):
        return False
    return True


def _coverage(unified_analysis: Any, records: Sequence[_TerrainRecord]) -> tuple[CoverageState, tuple[str, ...]]:
    declared = _value(_value(_value(unified_analysis, "coverage", {}), "solar", {}), "mode")
    declared = getattr(declared, "value", declared)
    declared_coverage = {
        "FULL": CoverageState.FULL,
        "SAMPLED": CoverageState.PARTIAL,
        "PARTIAL": CoverageState.PARTIAL,
        "UNAVAILABLE": CoverageState.UNAVAILABLE,
    }.get(declared, CoverageState.UNKNOWN)
    route_point_count = _value(_value(unified_analysis, "route", {}), "point_count")
    if declared_coverage is CoverageState.UNAVAILABLE:
        return CoverageState.UNAVAILABLE, ("TERRAIN_LIGHT_COVERAGE_UNAVAILABLE",)
    if not records:
        return CoverageState.UNAVAILABLE, ("TERRAIN_LIGHT_ROUTE_POINTS_UNAVAILABLE",)
    if (
        declared_coverage is CoverageState.FULL
        and isinstance(route_point_count, int)
        and route_point_count == len(records)
        and all(record.explicit_state for record in records)
    ):
        return CoverageState.FULL, ()
    reasons = []
    if declared_coverage is CoverageState.PARTIAL or route_point_count != len(records):
        reasons.append("TERRAIN_LIGHT_COVERAGE_PARTIAL")
    if not all(record.explicit_state for record in records):
        reasons.append("TERRAIN_LIGHT_STATE_UNAVAILABLE")
    if declared_coverage is CoverageState.UNKNOWN:
        reasons.append("TERRAIN_LIGHT_COVERAGE_UNKNOWN")
    return (CoverageState.PARTIAL if reasons else CoverageState.UNKNOWN), tuple(sorted(set(reasons)))


def _direct_encounter_event(records, index, event_type, upstream_reference, policy):
    record = records[index]
    edge = (
        "AT_ANALYSIS_WINDOW_START"
        if index == 0
        else "AT_ANALYSIS_WINDOW_END"
        if index == len(records) - 1
        else None
    )
    return TerrainLightEvent(
        event_id=derive_terrain_light_event_id(
            upstream_analysis_reference=upstream_reference,
            policy_id=policy.policy_id,
            policy_version=policy.version,
            event_type=event_type,
            route_point_indices=(record.sample.route_point_index,),
        ),
        event_type=event_type,
        boundary=RouteBoundary(BoundarySemantics.SAMPLED_POINT, (record.sample,)),
        before_state=None,
        after_state=TerrainSolarVisibility.DIRECT,
        source_references=_source_references(record),
        analysis_window_edge=edge,
    )


def _transition_event(before, after, event_type, upstream_reference, policy):
    return TerrainLightEvent(
        event_id=derive_terrain_light_event_id(
            upstream_analysis_reference=upstream_reference,
            policy_id=policy.policy_id,
            policy_version=policy.version,
            event_type=event_type,
            route_point_indices=(before.sample.route_point_index, after.sample.route_point_index),
        ),
        event_type=event_type,
        boundary=RouteBoundary(BoundarySemantics.TRANSITION_BRACKET, (before.sample, after.sample)),
        before_state=before.state,
        after_state=after.state,
        source_references=_source_references(before, after),
    )


def _event_order_key(event: TerrainLightEvent) -> tuple[int, int]:
    priority = {
        TerrainLightEventType.FIRST_DIRECT_TERRAIN_LIGHT: 0,
        TerrainLightEventType.DIRECT_LIGHT_ENTRY: 1,
        TerrainLightEventType.DIRECT_LIGHT_EXIT: 2,
        TerrainLightEventType.LAST_DIRECT_TERRAIN_LIGHT: 3,
    }
    return event.boundary.samples[-1].route_point_index, priority[event.event_type]


def _source_references(*records: _TerrainRecord) -> tuple[str, ...]:
    return tuple(sorted({record.source_reference for record in records if record.source_reference}))


def _sample_reference(point: Any) -> RouteSampleReference:
    return RouteSampleReference(
        route_point_index=_value(point, "point_index"),
        route_distance_m=_value(point, "route_distance_m"),
        planned_time=_value(point, "planned_arrival_time"),
        elapsed_planned_seconds=_value(point, "elapsed_route_seconds"),
    )


def _mapping(value: Any) -> Mapping[str, Any]:
    if isinstance(value, Mapping):
        return value
    if hasattr(value, "__dataclass_fields__"):
        return asdict(value)
    return {}


def _value(item: Any, key: str, default: Any = None) -> Any:
    return item.get(key, default) if isinstance(item, Mapping) else getattr(item, key, default)


def _string_or_none(value: Any) -> str | None:
    return str(value) if value is not None else None
