"""Composition of route, planned timeline, weather, and terrain-aware solar.

This module is deliberately a composition boundary.  It does not calculate
weather, terrain, solar geometry, or risk; it validates that existing domain
results refer to the same planned route points and keeps their provenance and
coverage independent.
"""

from __future__ import annotations

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

from mountain_twin.analysis_contract import AnalysisIdentity
from mountain_twin.exposure import RoutePoint
from mountain_twin.route_conditions import RouteConditionResult
from mountain_twin.route_solar_conditions import RouteSolarSeries
from mountain_twin.weather.temporal import PlanningScenario, TimelinePoint


class CoverageMode(str, Enum):
    FULL = "FULL"
    SAMPLED = "SAMPLED"
    PARTIAL = "PARTIAL"
    UNAVAILABLE = "UNAVAILABLE"


@dataclass(frozen=True)
class DomainCoverage:
    """Coverage for one independently evaluated analysis domain."""

    evaluated_points: int
    total_points: int
    mode: CoverageMode

    def __post_init__(self) -> None:
        if self.total_points < 0 or self.evaluated_points < 0:
            raise ValueError("coverage counts must be non-negative")
        if self.evaluated_points > self.total_points:
            raise ValueError("evaluated points cannot exceed total points")
        if self.total_points == 0 and self.mode is not CoverageMode.UNAVAILABLE:
            raise ValueError("empty coverage must be UNAVAILABLE")
        if (
            self.total_points
            and self.mode is CoverageMode.FULL
            and (self.evaluated_points != self.total_points)
        ):
            raise ValueError("FULL coverage requires every route point")

    @property
    def fraction(self) -> float | None:
        return self.evaluated_points / self.total_points if self.total_points else None

    def to_dict(self) -> dict[str, Any]:
        return {
            "evaluated_points": self.evaluated_points,
            "total_points": self.total_points,
            "fraction": self.fraction,
            "mode": self.mode.value,
        }


@dataclass(frozen=True)
class UnifiedPointResult:
    """One aligned route point with independent weather and solar domains."""

    route_id: str
    point_index: int
    route_distance_m: float
    latitude: float
    longitude: float
    route_elevation_m: float | None
    planned_arrival_time: str
    elapsed_route_seconds: float
    weather: dict[str, Any]
    weather_quality: dict[str, Any]
    weather_provenance_ref: str
    solar: dict[str, Any]
    solar_quality: dict[str, Any]
    solar_provenance_ref: str | None
    pace: dict[str, Any] | None = None

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


@dataclass(frozen=True)
class UnifiedRouteEvent:
    """An analytical event bounded by sampled route points."""

    event_type: str
    event_family: str
    domain: str
    before_state: str
    after_state: str
    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
    method: str
    provenance_ref: str
    limitations: tuple[str, ...]

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


@dataclass(frozen=True)
class UnifiedRouteRuntime:
    timeline_seconds: float
    weather_composition_seconds: float
    solar_composition_seconds: float
    summary_event_seconds: float
    serialization_seconds: float | None
    total_analysis_seconds: float

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


@dataclass(frozen=True)
class UnifiedRouteAnalysis:
    """Route-level composition with truthful domain coverage."""

    identity: AnalysisIdentity
    route_id: str
    scenario: dict[str, Any]
    route: dict[str, Any]
    coverage: dict[str, Any]
    points: tuple[UnifiedPointResult, ...]
    events: tuple[UnifiedRouteEvent, ...]
    summary: dict[str, Any]
    provenance: dict[str, Any]
    limitations: tuple[str, ...]
    runtime: UnifiedRouteRuntime

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


def compose_unified_route_analysis(
    *,
    route_points: Sequence[RoutePoint],
    timeline: Sequence[TimelinePoint],
    scenario: PlanningScenario,
    condition_results: Sequence[RouteConditionResult],
    solar_series: RouteSolarSeries,
    weather_coverage: DomainCoverage,
    provenance: Mapping[str, Any],
    weather_summary: Mapping[str, Any],
    weather_resolution: Mapping[str, Any] | None = None,
    timing: Mapping[str, float] | None = None,
    pace_result: Any | None = None,
) -> UnifiedRouteAnalysis:
    """Validate and compose already evaluated domain results.

    ``condition_results`` is normally the intersection of weather and solar
    coverage.  Unprofiled route points are therefore not fabricated as solar
    UNKNOWN values; the top-level coverage says exactly how many were queried.
    """
    if not route_points:
        raise ValueError("route must contain at least one point")
    if len(timeline) != len(route_points):
        raise ValueError("route points and timeline must have equal length")
    if timeline and any(
        current.point_index != index
        or current.route_distance_m < previous.route_distance_m
        or current.planned_arrival.astimezone(timezone.utc)
        < previous.planned_arrival.astimezone(timezone.utc)
        for index, (previous, current) in enumerate(zip(timeline, timeline[1:]), start=1)
    ):
        raise ValueError("timeline point order, distance, or planned time is invalid")
    if tuple(row.point_index for row in condition_results) != tuple(
        row.point_index for row in solar_series.points
    ):
        raise ValueError("weather and solar point identities are not aligned")
    if len(condition_results) != solar_series.evaluated_count:
        raise ValueError("weather and solar evaluated counts are not aligned")

    route_by_index = {point.point_index: point for point in route_points}
    timeline_by_index = {point.point_index: point for point in timeline}
    unified_points = []
    weather_provenance = dict(provenance.get("weather", {}))
    if weather_resolution is not None:
        weather_provenance["resolution"] = _json_value(dict(weather_resolution))
    weather_point_provenance = weather_provenance.setdefault("point_provenance", {})
    for condition, solar in zip(condition_results, solar_series.points):
        point = route_by_index.get(condition.point_index)
        planned = timeline_by_index.get(condition.point_index)
        if point is None or planned is None:
            raise ValueError("domain result refers to a route point outside the route")
        if condition.route_id != point.route_id or solar.route_id != point.route_id:
            raise ValueError("route identity mismatch between composed domains")
        if condition.route_distance_m != planned.route_distance_m:
            raise ValueError("weather route distance does not match planned timeline")
        if condition.planned_arrival_time != planned.planned_arrival.isoformat():
            raise ValueError("weather planned arrival does not match timeline")
        if solar.planned_arrival_time != planned.planned_arrival.isoformat():
            raise ValueError("solar planned arrival does not match timeline")
        if condition.latitude != point.latitude or condition.longitude != point.longitude:
            raise ValueError("weather coordinates do not match route point")
        if solar.latitude != point.latitude or solar.longitude != point.longitude:
            raise ValueError("solar coordinates do not match route point")
        weather_ref = f"weather:{scenario.name}:{condition.point_index}"
        weather_point_provenance[str(condition.point_index)] = _json_value(
            condition.weather_provenance
        )
        solar_ref = solar.profile_key
        pace = {
            "cumulative_moving_seconds": planned.cumulative_moving_seconds,
            "cumulative_pause_seconds": planned.cumulative_pause_seconds,
            "interval_id": planned.pace_interval_id,
            "model_reference": (
                f"pace:{pace_result.model_version}:{scenario.name}"
                if pace_result is not None
                else None
            ),
        }
        unified_points.append(
            UnifiedPointResult(
                route_id=point.route_id,
                point_index=point.point_index,
                route_distance_m=planned.route_distance_m,
                latitude=point.latitude,
                longitude=point.longitude,
                route_elevation_m=point.elevation_m,
                planned_arrival_time=planned.planned_arrival.isoformat(),
                elapsed_route_seconds=planned.elapsed_route_seconds,
                weather=_compact_weather(condition.weather),
                weather_quality={
                    "state": condition.quality["weather_state"],
                    "reason_codes": condition.quality["weather_reason_codes"],
                },
                weather_provenance_ref=weather_ref,
                solar=condition.solar,
                solar_quality={
                    "state": condition.quality["solar_state"],
                    "reason_codes": condition.quality["solar_reason_codes"],
                },
                solar_provenance_ref=solar_ref,
                pace=pace,
            )
        )

    point_indices = tuple(row.point_index for row in unified_points)
    if point_indices != tuple(sorted(point_indices)) or len(point_indices) != len(
        set(point_indices)
    ):
        raise ValueError("unified point results must be ordered and unique")
    solar_coverage = DomainCoverage(
        evaluated_points=solar_series.evaluated_count,
        total_points=len(route_points),
        mode=(
            CoverageMode.FULL
            if solar_series.evaluated_count == len(route_points)
            else CoverageMode.SAMPLED
        ),
    )
    route_start = timeline[0].planned_arrival
    route_finish = timeline[-1].planned_arrival
    all_transitions = (*solar_series.transitions, *solar_series.astronomical_transitions)
    events = tuple(
        sorted(
            (
                UnifiedRouteEvent(
                    event_type=(
                        f"SOLAR_{_state_value(transition.before_state)}_TO_"
                        f"{_state_value(transition.after_state)}"
                        if transition.event_family == "TERRAIN_VISIBILITY_CHANGE"
                        else f"ASTRONOMICAL_{_state_value(transition.before_state)}_TO_"
                        f"{_state_value(transition.after_state)}"
                    ),
                    event_family=transition.event_family,
                    domain=(
                        "solar"
                        if transition.event_family == "TERRAIN_VISIBILITY_CHANGE"
                        else "solar_astronomical"
                    ),
                    before_state=_state_value(transition.before_state),
                    after_state=_state_value(transition.after_state),
                    before_point_index=transition.before_point_index,
                    after_point_index=transition.after_point_index,
                    before_route_distance_m=transition.before_route_distance_m,
                    after_route_distance_m=transition.after_route_distance_m,
                    before_planned_time=transition.before_planned_time,
                    after_planned_time=transition.after_planned_time,
                    method="sampled_route_point_transition",
                    provenance_ref=solar_series.terrain_context_id,
                    limitations=("TRANSITION_BOUNDED_BY_EVALUATED_ROUTE_POINTS",),
                )
                for transition in all_transitions
            ),
            key=lambda event: (event.after_point_index, event.event_family),
        )
    )
    common_coverage = DomainCoverage(
        evaluated_points=len(unified_points),
        total_points=len(route_points),
        mode=(
            CoverageMode.FULL if len(unified_points) == len(route_points) else CoverageMode.SAMPLED
        ),
    )
    summary = {
        "route_distance_m": timeline[-1].route_distance_m,
        "planned_duration_s": (
            route_finish.astimezone(timezone.utc) - route_start.astimezone(timezone.utc)
        ).total_seconds(),
        "elevation_min_m": min(
            (point.elevation_m for point in route_points if point.elevation_m is not None),
            default=None,
        ),
        "elevation_max_m": max(
            (point.elevation_m for point in route_points if point.elevation_m is not None),
            default=None,
        ),
        "weather": dict(weather_summary),
        "weather_resolution": _json_value(dict(weather_resolution or {})),
        "pace": _pace_summary(pace_result),
        "solar": solar_series.summary,
        "solar_direct_count": solar_series.direct_count,
        "solar_shadow_count": solar_series.shadow_count,
        "solar_unknown_count": solar_series.unknown_count,
        "solar_not_applicable_count": solar_series.not_applicable_count,
        "solar_interval_count": len(solar_series.intervals),
        "solar_transition_count": len(events),
        "weather_coverage": weather_coverage.to_dict(),
        "solar_coverage": solar_coverage.to_dict(),
    }
    route_id = route_points[0].route_id
    timing = timing or {}
    runtime = UnifiedRouteRuntime(
        timeline_seconds=float(timing.get("timeline_seconds", 0.0)),
        weather_composition_seconds=float(timing.get("weather_composition_seconds", 0.0)),
        solar_composition_seconds=float(timing.get("solar_composition_seconds", 0.0)),
        summary_event_seconds=float(timing.get("summary_event_seconds", 0.0)),
        serialization_seconds=None,
        total_analysis_seconds=float(timing.get("total_analysis_seconds", 0.0)),
    )
    return UnifiedRouteAnalysis(
        identity=AnalysisIdentity(
            analysis_type="unified_route_analysis",
            semantic_type="route_environment_context",
            version="0.4",
            subject_id=route_id,
            scenario_datetime=scenario.start_datetime.isoformat(),
            timezone=getattr(
                scenario.start_datetime.tzinfo, "key", str(scenario.start_datetime.tzinfo)
            ),
        ),
        route_id=route_id,
        scenario={
            "name": scenario.name,
            "start_datetime": scenario.start_datetime.isoformat(),
            "moving_speed_mps": scenario.moving_speed_mps,
            "pace_factor": scenario.pace_factor,
            "pace_model": pace_result.model_name if pace_result is not None else None,
            "pace_model_version": pace_result.model_version if pace_result is not None else None,
            "pauses": [asdict(pause) for pause in scenario.pauses],
            "finish_datetime": route_finish.isoformat(),
            "time_semantics": "planned synthetic scenario; GPX timestamps are not chronology",
        },
        route={
            "route_id": route_id,
            "point_count": len(route_points),
            "timezone": getattr(route_start.tzinfo, "key", str(route_start.tzinfo)),
            "start_datetime": route_start.isoformat(),
            "finish_datetime": route_finish.isoformat(),
        },
        coverage={
            "route_points_total": len(route_points),
            "weather": weather_coverage.to_dict(),
            "solar": solar_coverage.to_dict(),
            "common_evaluated_points": common_coverage.evaluated_points,
            "common_mode": common_coverage.mode.value,
        },
        points=tuple(unified_points),
        events=events,
        summary=summary,
        provenance=_json_value(
            {
                **dict(provenance),
                "weather": weather_provenance,
                "pace": pace_result.provenance if pace_result is not None else None,
            }
        ),
        limitations=(
            "GPX_EMBEDDED_TIMESTAMPS_NOT_USED_AS_HIKING_CHRONOLOGY",
            "SOLAR_COVERAGE_IS_SAMPLED_WHERE_NOT_FULL",
            "SOLAR_TRANSITIONS_ARE_BOUNDED_BY_EVALUATED_ROUTE_POINTS",
            "WEATHER_MODEL_GRID_IS_COARSER_THAN_ROUTE_GEOMETRY",
            "GLO30_DSM_IS_NOT_BARE_EARTH_PHYSICAL_TERRAIN_TRUTH",
            "RESULT_IS_ANALYTICAL_CONTEXT_NOT_SAFETY_ADVICE",
        ),
        runtime=runtime,
    )


def serialize_unified_route_analysis(result: UnifiedRouteAnalysis) -> str:
    """Serialize stable route/domain content without volatile runtime timings."""
    document = result.to_dict()
    document.pop("runtime", None)
    return json.dumps(document, sort_keys=True, separators=(",", ":"), allow_nan=False)


def _json_value(value: Any) -> Any:
    if isinstance(value, Enum):
        return value.value
    if isinstance(value, Mapping):
        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]
    if hasattr(value, "__dataclass_fields__"):
        return _json_value(asdict(value))
    return value


def _state_value(value: Any) -> str:
    return str(value.value if isinstance(value, Enum) else value)


def _compact_weather(weather: Mapping[str, Any]) -> dict[str, Any]:
    """Keep route values and representation states; share provider metadata above."""
    values = weather.get("values", {})
    compact_values = {}
    for name, item in values.items():
        compact_values[name] = {
            key: item[key]
            for key in (
                "value",
                "unit",
                "representation_state",
                "temporal_representation_state",
                "spatial_method",
                "spatial_contributing_sample_ids",
                "spatial_support_distances_m",
                "temporal_contributing_valid_times",
                "temporal_interpolation_method",
                "field_semantics",
                "projection_policy_id",
                "snow_spatial_representation",
                "snow_temporal_projection",
                "coverage",
                "limitation_codes",
            )
            if key in item
        }
    return {
        "values": compact_values,
        "spatial_representation": weather.get("spatial_representation", {}),
        "temporal_representation": weather.get("temporal_representation", {}),
    }


def _pace_summary(pace_result: Any | None) -> dict[str, Any]:
    if pace_result is None:
        return {
            "state": "LEGACY_CONSTANT_SPEED",
            "moving_time_s": None,
            "pause_time_s": None,
            "total_planned_time_s": None,
        }
    return {
        "state": pace_result.state.value,
        "model": pace_result.model_name,
        "model_version": pace_result.model_version,
        "elevation_source": pace_result.elevation_source,
        "preprocessing_policy": pace_result.preprocessing_policy,
        "preprocessing_distance_m": pace_result.preprocessing_distance_m,
        "scenario_factor": pace_result.scenario_factor,
        "horizontal_distance_m": pace_result.horizontal_distance_m,
        "ascent_m": pace_result.ascent_m,
        "descent_m": pace_result.descent_m,
        "moving_time_s": pace_result.moving_time_s,
        "pause_time_s": pace_result.pause_time_s,
        "total_planned_time_s": pace_result.total_planned_time_s,
        "pause_strategy_id": pace_result.pause_strategy_id,
        "pause_strategy_version": pace_result.pause_strategy_version,
        "reason_codes": list(pace_result.reason_codes),
        "runtime": dict(pace_result.runtime),
    }
