"""Route-level solar conditions backed by reusable terrain horizon profiles."""

from __future__ import annotations

import json
import time
from dataclasses import asdict, dataclass
from datetime import datetime, timezone
from enum import Enum
from typing import Any, Mapping, Sequence

from mountain_twin.exposure import RoutePoint
from mountain_twin.route_analysis import prepare_route
from mountain_twin.route_solar_exposure import VisibilityStatus
from mountain_twin.solar.query_engine import SolarQueryEngine
from mountain_twin.solar.states import (
    AstronomicalState,
    TerrainSolarVisibility,
    classify_astronomical_state,
)
from mountain_twin.solar.sun_engine import solar_position
from mountain_twin.terrain.cache import CACHE_FORMAT, CacheIdentity, HorizonProfileCache
from mountain_twin.terrain.profile import HorizonProfile
from mountain_twin.weather.temporal import PlanningScenario, assign_timeline

ROUTE_SOLAR_METHOD = "cached_terrain_horizon_profile_route_solar_conditions_v0_3"
UNKNOWN_PROFILE_REASON = "TERRAIN_PROFILE_NOT_AVAILABLE"


@dataclass(frozen=True)
class RouteSolarPointResult:
    """One planned route point; terrain provenance is referenced by profile key."""

    route_id: str
    point_index: int
    route_distance_m: float
    latitude: float
    longitude: float
    route_elevation_m: float | None
    planned_arrival_time: str
    timezone: str
    solar_azimuth_deg: float
    solar_elevation_deg: float
    status: VisibilityStatus
    horizon_angle_deg: float | None
    solar_horizon_margin_deg: float | None
    controlling_distance_m: float | None
    final_range_m: float | None
    profile_key: str | None
    terrain_context_id: str
    observer_height_m: float | None
    observer_model: str | None
    approximation_method: str
    convergence_state: str
    reason_codes: tuple[str, ...]
    astronomical_state: AstronomicalState | None = None
    terrain_solar_visibility: TerrainSolarVisibility | None = None

    def to_dict(self) -> dict[str, Any]:
        return _json_value(asdict(self))


@dataclass(frozen=True)
class RouteSolarInterval:
    """A contiguous run of sampled point states, bounded by route points."""

    status: VisibilityStatus | AstronomicalState
    start_point_index: int
    end_point_index: int
    start_route_distance_m: float
    end_route_distance_m: float
    start_planned_time: str
    end_planned_time: str
    interval_distance_m: float
    interval_duration_s: float
    axis: str = "TERRAIN_VISIBILITY"

    def to_dict(self) -> dict[str, Any]:
        return _json_value(asdict(self))


@dataclass(frozen=True)
class RouteSolarTransition:
    """A state change bounded by two evaluated route points."""

    before_state: VisibilityStatus | AstronomicalState
    after_state: VisibilityStatus | AstronomicalState
    before_point_index: int
    after_point_index: int
    before_route_distance_m: float
    after_route_distance_m: float
    before_planned_time: str
    after_planned_time: str
    event_family: str = "TERRAIN_VISIBILITY_CHANGE"

    def to_dict(self) -> dict[str, Any]:
        return _json_value(asdict(self))


@dataclass(frozen=True)
class RouteSolarRuntime:
    profile_load_seconds: float
    timeline_seconds: float
    cached_query_seconds: float
    derivation_seconds: float
    total_analysis_seconds: float
    terrain_interpolation_queries: int = 0

    def to_dict(self) -> dict[str, Any]:
        return _json_value(asdict(self))


@dataclass(frozen=True)
class RouteSolarSeries:
    """Ordered route solar analysis with bounded intervals and transitions."""

    route_id: str
    planning_date: str
    timezone: str
    scenario_name: str
    scenario_start_time: str
    finish_time: str | None
    point_count: int
    evaluated_count: int
    evaluated_point_indices: tuple[int, ...]
    coverage_mode: str
    direct_count: int
    shadow_count: int
    unknown_count: int
    terrain_context_id: str
    terrain_provenance: dict[str, Any]
    points: tuple[RouteSolarPointResult, ...]
    intervals: tuple[RouteSolarInterval, ...]
    transitions: tuple[RouteSolarTransition, ...]
    summary: dict[str, Any]
    runtime: RouteSolarRuntime
    not_applicable_count: int = 0
    astronomical_state_counts: dict[str, int] | None = None
    astronomical_intervals: tuple[RouteSolarInterval, ...] = ()
    astronomical_transitions: tuple[RouteSolarTransition, ...] = ()

    def to_dict(self) -> dict[str, Any]:
        return _json_value(asdict(self))


def derive_intervals(rows: Sequence[RouteSolarPointResult]) -> tuple[RouteSolarInterval, ...]:
    """Group adjacent evaluated point states without smoothing or UNKNOWN merging."""
    if not rows:
        return ()
    intervals = []
    start = 0
    for index in range(1, len(rows) + 1):
        if index == len(rows) or rows[index].status is not rows[start].status:
            first, last = rows[start], rows[index - 1]
            first_time = datetime.fromisoformat(first.planned_arrival_time)
            last_time = datetime.fromisoformat(last.planned_arrival_time)
            intervals.append(
                RouteSolarInterval(
                    status=first.status,
                    start_point_index=first.point_index,
                    end_point_index=last.point_index,
                    start_route_distance_m=first.route_distance_m,
                    end_route_distance_m=last.route_distance_m,
                    start_planned_time=first.planned_arrival_time,
                    end_planned_time=last.planned_arrival_time,
                    interval_distance_m=last.route_distance_m - first.route_distance_m,
                    interval_duration_s=_elapsed_seconds(last_time, first_time),
                )
            )
            start = index
    return tuple(intervals)


def derive_transitions(rows: Sequence[RouteSolarPointResult]) -> tuple[RouteSolarTransition, ...]:
    """Return only sampled-point state changes; no exact boundary is inferred."""
    return tuple(
        RouteSolarTransition(
            before_state=before.status,
            after_state=after.status,
            before_point_index=before.point_index,
            after_point_index=after.point_index,
            before_route_distance_m=before.route_distance_m,
            after_route_distance_m=after.route_distance_m,
            before_planned_time=before.planned_arrival_time,
            after_planned_time=after.planned_arrival_time,
            event_family="TERRAIN_VISIBILITY_CHANGE",
        )
        for before, after in zip(rows, rows[1:])
        if before.status is not after.status
    )


def derive_astronomical_intervals(
    rows: Sequence[RouteSolarPointResult],
) -> tuple[RouteSolarInterval, ...]:
    """Group adjacent route points by their independent astronomical state."""
    if not rows or any(row.astronomical_state is None for row in rows):
        return ()
    intervals = []
    start = 0
    for index in range(1, len(rows) + 1):
        if (
            index == len(rows)
            or rows[index].astronomical_state is not rows[start].astronomical_state
        ):
            first, last = rows[start], rows[index - 1]
            first_time = datetime.fromisoformat(first.planned_arrival_time)
            last_time = datetime.fromisoformat(last.planned_arrival_time)
            intervals.append(
                RouteSolarInterval(
                    status=first.astronomical_state,
                    start_point_index=first.point_index,
                    end_point_index=last.point_index,
                    start_route_distance_m=first.route_distance_m,
                    end_route_distance_m=last.route_distance_m,
                    start_planned_time=first.planned_arrival_time,
                    end_planned_time=last.planned_arrival_time,
                    interval_distance_m=last.route_distance_m - first.route_distance_m,
                    interval_duration_s=_elapsed_seconds(last_time, first_time),
                    axis="ASTRONOMICAL_STATE",
                )
            )
            start = index
    return tuple(intervals)


def derive_astronomical_transitions(
    rows: Sequence[RouteSolarPointResult],
) -> tuple[RouteSolarTransition, ...]:
    """Return sampled astronomical-state changes, without exact boundaries."""
    return tuple(
        RouteSolarTransition(
            before_state=before.astronomical_state,
            after_state=after.astronomical_state,
            before_point_index=before.point_index,
            after_point_index=after.point_index,
            before_route_distance_m=before.route_distance_m,
            after_route_distance_m=after.route_distance_m,
            before_planned_time=before.planned_arrival_time,
            after_planned_time=after.planned_arrival_time,
            event_family="ASTRONOMICAL_STATE_CHANGE",
        )
        for before, after in zip(rows, rows[1:])
        if before.astronomical_state is not None
        and after.astronomical_state is not None
        and before.astronomical_state is not after.astronomical_state
    )


def summarize_route_solar(rows: Sequence[RouteSolarPointResult]) -> dict[str, Any]:
    """Summarize homogeneous route segments using conservative endpoint weighting.

    A segment whose endpoint states agree is assigned to that state. A segment
    crossing a sampled state transition is assigned to UNKNOWN rather than
    pretending to know the physical transition location.
    """
    distance = {status.value: 0.0 for status in VisibilityStatus}
    duration = {status.value: 0.0 for status in VisibilityStatus}
    astronomical_distance = {state.value: 0.0 for state in AstronomicalState}
    astronomical_duration = {state.value: 0.0 for state in AstronomicalState}
    transition_distance = transition_duration = 0.0
    terrain_transition_distance = terrain_transition_duration = 0.0
    astronomical_transition_distance = astronomical_transition_duration = 0.0
    for before, after in zip(rows, rows[1:]):
        segment_distance = after.route_distance_m - before.route_distance_m
        before_time = datetime.fromisoformat(before.planned_arrival_time)
        after_time = datetime.fromisoformat(after.planned_arrival_time)
        segment_duration = _elapsed_seconds(after_time, before_time)
        if before.astronomical_state is not None and after.astronomical_state is not None:
            if before.astronomical_state is after.astronomical_state:
                astronomical_distance[before.astronomical_state.value] += segment_distance
                astronomical_duration[before.astronomical_state.value] += segment_duration
            else:
                astronomical_transition_distance += segment_distance
                astronomical_transition_duration += segment_duration
        if before.status is after.status:
            distance[before.status.value] += segment_distance
            duration[before.status.value] += segment_duration
        else:
            terrain_transition_distance += segment_distance
            terrain_transition_duration += segment_duration
            if (
                before.status is not VisibilityStatus.NOT_APPLICABLE
                and after.status is not VisibilityStatus.NOT_APPLICABLE
            ):
                distance[VisibilityStatus.UNKNOWN.value] += segment_distance
                duration[VisibilityStatus.UNKNOWN.value] += segment_duration
                transition_distance += segment_distance
                transition_duration += segment_duration
    known_distance = (
        distance[VisibilityStatus.DIRECT.value] + distance[VisibilityStatus.SHADOW.value]
    )
    known_duration = (
        duration[VisibilityStatus.DIRECT.value] + duration[VisibilityStatus.SHADOW.value]
    )
    total_distance = sum(
        after.route_distance_m - before.route_distance_m for before, after in zip(rows, rows[1:])
    )
    total_duration = sum(
        _elapsed_seconds(
            datetime.fromisoformat(after.planned_arrival_time),
            datetime.fromisoformat(before.planned_arrival_time),
        )
        for before, after in zip(rows, rows[1:])
    )
    margins = sorted(
        abs(row.solar_horizon_margin_deg)
        for row in rows
        if row.solar_horizon_margin_deg is not None and row.solar_elevation_deg > 0
    )
    return {
        "route_distance_weighting": "same-state evaluated-point segments; transition segments are conservatively UNKNOWN",
        "planned_duration_weighting": "same-state evaluated-point time segments; transition segments are conservatively UNKNOWN",
        "direct_distance_m": distance[VisibilityStatus.DIRECT.value],
        "shadow_distance_m": distance[VisibilityStatus.SHADOW.value],
        "unknown_distance_m": distance[VisibilityStatus.UNKNOWN.value],
        "direct_duration_s": duration[VisibilityStatus.DIRECT.value],
        "shadow_duration_s": duration[VisibilityStatus.SHADOW.value],
        "unknown_duration_s": duration[VisibilityStatus.UNKNOWN.value],
        "not_applicable_distance_m": distance[VisibilityStatus.NOT_APPLICABLE.value],
        "not_applicable_duration_s": duration[VisibilityStatus.NOT_APPLICABLE.value],
        "known_distance_fraction": known_distance / total_distance if total_distance else None,
        "known_direct_distance_fraction": distance[VisibilityStatus.DIRECT.value] / known_distance
        if known_distance
        else None,
        "known_shadow_distance_fraction": distance[VisibilityStatus.SHADOW.value] / known_distance
        if known_distance
        else None,
        "known_duration_fraction": known_duration / total_duration if total_duration else None,
        "known_direct_duration_fraction": duration[VisibilityStatus.DIRECT.value] / known_duration
        if known_duration
        else None,
        "known_shadow_duration_fraction": duration[VisibilityStatus.SHADOW.value] / known_duration
        if known_duration
        else None,
        "transition_boundary_distance_m": transition_distance,
        "transition_boundary_duration_s": transition_duration,
        "terrain_transition_boundary_distance_m": terrain_transition_distance,
        "terrain_transition_boundary_duration_s": terrain_transition_duration,
        "astronomical_distance_m": astronomical_distance,
        "astronomical_duration_s": astronomical_duration,
        "astronomical_transition_boundary_distance_m": astronomical_transition_distance,
        "astronomical_transition_boundary_duration_s": astronomical_transition_duration,
        "astronomical_transition_count": len(derive_astronomical_transitions(rows)),
        "first_direct_interval": _first_interval(rows, VisibilityStatus.DIRECT),
        "first_shadow_interval": _first_interval(rows, VisibilityStatus.SHADOW),
        "last_direct_interval": _last_interval(rows, VisibilityStatus.DIRECT),
        "last_shadow_interval": _last_interval(rows, VisibilityStatus.SHADOW),
        "transition_count": len(derive_transitions(rows)),
        "smallest_known_absolute_margin_deg": margins[0] if margins else None,
    }


def analyze_route_solar(
    points: Sequence[RoutePoint],
    scenario: PlanningScenario,
    cache: HorizonProfileCache,
    identities_by_index: Mapping[int, CacheIdentity],
    *,
    evaluated_indices: Sequence[int] | None = None,
    loaded_profiles: Mapping[int, HorizonProfile] | None = None,
    pace_result: Any | None = None,
) -> RouteSolarSeries:
    """Evaluate planned solar state using cached profiles at route arrival times."""
    started = time.perf_counter()
    if not points:
        raise ValueError("route must contain at least one point")
    prepared = prepare_route(points).points
    timeline_started = time.perf_counter()
    timeline = assign_timeline(prepared, scenario, pace_result)
    timeline_seconds = time.perf_counter() - timeline_started
    by_index = {point.point_index: point for point in points}
    timeline_by_index = {point.point_index: point for point in timeline}
    selected = (
        tuple(point.point_index for point in points)
        if evaluated_indices is None
        else tuple(evaluated_indices)
    )
    if selected != tuple(sorted(set(selected))):
        raise ValueError("evaluated indices must be sorted and unique")
    if any(index not in by_index for index in selected):
        raise ValueError("evaluated index is not part of the route")
    loaded_profiles = loaded_profiles or {}
    profile_load_started = time.perf_counter()
    profiles = dict(loaded_profiles)
    for index in selected:
        if index not in profiles and index in identities_by_index:
            profiles[index] = cache.get_profile(identities_by_index[index])
    profile_load_seconds = time.perf_counter() - profile_load_started
    context_id, provenance = _terrain_context(identities_by_index, selected)
    engine = SolarQueryEngine(cache)
    query_started = time.perf_counter()
    rows = []
    terrain_interpolation_queries = 0
    for index in selected:
        point = by_index[index]
        arrival = timeline_by_index[index]
        identity = identities_by_index.get(index)
        profile = profiles.get(index)
        if identity is not None and profile is not None:
            result = engine.query_loaded(
                profile,
                identity,
                arrival.planned_arrival,
                latitude=point.latitude,
                longitude=point.longitude,
            )
            terrain_interpolation_queries += int(result.terrain_interpolation_performed)
            status = VisibilityStatus(result.status)
            horizon = result.horizon_angle_deg
            margin = result.solar_elevation_deg - horizon if horizon is not None else None
            row = RouteSolarPointResult(
                route_id=point.route_id,
                point_index=index,
                route_distance_m=arrival.route_distance_m,
                latitude=point.latitude,
                longitude=point.longitude,
                route_elevation_m=point.elevation_m,
                planned_arrival_time=arrival.planned_arrival.isoformat(),
                timezone=_timezone_name(arrival.planned_arrival),
                solar_azimuth_deg=result.solar_azimuth_deg,
                solar_elevation_deg=result.solar_elevation_deg,
                status=status,
                horizon_angle_deg=horizon,
                solar_horizon_margin_deg=margin,
                controlling_distance_m=result.controlling_distance_m,
                final_range_m=result.final_range_m,
                profile_key=result.profile_key,
                terrain_context_id=context_id,
                observer_height_m=result.observer_height_m,
                observer_model=result.observer_model,
                approximation_method=result.approximation_method,
                convergence_state=result.convergence_state,
                reason_codes=result.reason_codes,
                astronomical_state=result.astronomical_state,
                terrain_solar_visibility=result.terrain_solar_visibility,
            )
        else:
            elevation, azimuth = solar_position(
                arrival.planned_arrival, point.latitude, point.longitude
            )
            astronomical_state = classify_astronomical_state(elevation)
            terrain_visibility = (
                TerrainSolarVisibility.NOT_APPLICABLE
                if astronomical_state is not AstronomicalState.DAY
                else TerrainSolarVisibility.UNKNOWN
            )
            row = RouteSolarPointResult(
                route_id=point.route_id,
                point_index=index,
                route_distance_m=arrival.route_distance_m,
                latitude=point.latitude,
                longitude=point.longitude,
                route_elevation_m=point.elevation_m,
                planned_arrival_time=arrival.planned_arrival.isoformat(),
                timezone=_timezone_name(arrival.planned_arrival),
                solar_azimuth_deg=azimuth,
                solar_elevation_deg=elevation,
                status=VisibilityStatus(terrain_visibility.value),
                horizon_angle_deg=None,
                solar_horizon_margin_deg=None,
                controlling_distance_m=None,
                final_range_m=None,
                profile_key=identity.key() if identity else None,
                terrain_context_id=context_id,
                observer_height_m=identity.observer_height_m if identity else None,
                observer_model=identity.observer_model if identity else None,
                approximation_method=ROUTE_SOLAR_METHOD,
                convergence_state="not_evaluated",
                reason_codes=(
                    ("TERRAIN_VISIBILITY_NOT_APPLICABLE",)
                    if terrain_visibility is TerrainSolarVisibility.NOT_APPLICABLE
                    else (UNKNOWN_PROFILE_REASON,)
                ),
                astronomical_state=astronomical_state,
                terrain_solar_visibility=terrain_visibility,
            )
        rows.append(row)
    cached_query_seconds = time.perf_counter() - query_started
    derivation_started = time.perf_counter()
    intervals = derive_intervals(rows)
    transitions = derive_transitions(rows)
    astronomical_intervals = derive_astronomical_intervals(rows)
    astronomical_transitions = derive_astronomical_transitions(rows)
    summary = summarize_route_solar(rows)
    derivation_seconds = time.perf_counter() - derivation_started
    ordered_rows = tuple(rows)
    counts = {
        status.value: sum(row.status is status for row in rows) for status in VisibilityStatus
    }
    finish = ordered_rows[-1].planned_arrival_time if ordered_rows else None
    result = RouteSolarSeries(
        route_id=points[0].route_id,
        planning_date=scenario.start_datetime.date().isoformat(),
        timezone=_timezone_name(scenario.start_datetime),
        scenario_name=scenario.name,
        scenario_start_time=scenario.start_datetime.isoformat(),
        finish_time=finish,
        point_count=len(points),
        evaluated_count=len(rows),
        evaluated_point_indices=selected,
        coverage_mode="full_route" if len(rows) == len(points) else "sampled_route_subset",
        direct_count=counts[VisibilityStatus.DIRECT.value],
        shadow_count=counts[VisibilityStatus.SHADOW.value],
        unknown_count=counts[VisibilityStatus.UNKNOWN.value],
        terrain_context_id=context_id,
        terrain_provenance=provenance,
        points=ordered_rows,
        intervals=intervals,
        transitions=transitions,
        summary=summary,
        runtime=RouteSolarRuntime(
            profile_load_seconds=profile_load_seconds,
            timeline_seconds=timeline_seconds,
            cached_query_seconds=cached_query_seconds,
            derivation_seconds=derivation_seconds,
            total_analysis_seconds=time.perf_counter() - started,
            terrain_interpolation_queries=terrain_interpolation_queries,
        ),
        not_applicable_count=counts[VisibilityStatus.NOT_APPLICABLE.value],
        astronomical_state_counts={
            state.value: sum(row.astronomical_state is state for row in rows)
            for state in AstronomicalState
        },
        astronomical_intervals=astronomical_intervals,
        astronomical_transitions=astronomical_transitions,
    )
    return result


def serialize_route_solar(series: RouteSolarSeries) -> str:
    """Serialize deterministic route data without profiles or volatile timings."""
    document = series.to_dict()
    document.pop("runtime", None)
    return json.dumps(document, sort_keys=True, separators=(",", ":"), allow_nan=False)


def _terrain_context(
    identities: Mapping[int, CacheIdentity], selected: Sequence[int]
) -> tuple[str, dict[str, Any]]:
    available = [identities[index] for index in selected if index in identities]
    if not available:
        context_id = f"{CACHE_FORMAT}:no-profile-context"
        return context_id, {"cache_format": CACHE_FORMAT, "profile_count": 0}
    first = available[0]
    common_fields = (
        "provider",
        "product",
        "source_id",
        "source_sha256",
        "surface_semantics",
        "observer_model",
        "observer_height_m",
        "angular_resolution_deg",
        "interpolation_method",
        "horizon_range_m",
        "range_policy",
        "method_version",
    )
    if any(
        any(getattr(identity, field) != getattr(first, field) for field in common_fields)
        for identity in available[1:]
    ):
        raise ValueError("evaluated profiles do not share one semantic terrain context")
    context_id = f"{CACHE_FORMAT}:{first.source_sha256}:{first.method_version}"
    return context_id, {
        "cache_format": CACHE_FORMAT,
        "profile_count": len(available),
        **{field: getattr(first, field) for field in common_fields},
    }


def _first_interval(rows, status: VisibilityStatus):
    return next(
        (interval.to_dict() for interval in derive_intervals(rows) if interval.status is status),
        None,
    )


def _last_interval(rows, status: VisibilityStatus):
    matches = [
        interval.to_dict() for interval in derive_intervals(rows) if interval.status is status
    ]
    return matches[-1] if matches else None


def _timezone_name(instant: datetime) -> str:
    return getattr(instant.tzinfo, "key", str(instant.tzinfo))


def _elapsed_seconds(later: datetime, earlier: datetime) -> float:
    return (later.astimezone(timezone.utc) - earlier.astimezone(timezone.utc)).total_seconds()


def _json_value(value: Any) -> Any:
    if isinstance(value, Enum):
        return value.value
    if isinstance(value, dict):
        return {str(key): _json_value(item) for key, item in value.items()}
    if isinstance(value, (tuple, list)):
        return [_json_value(item) for item in value]
    return value
