"""Composition of route timeline, weather, and optional solar knowledge."""

from __future__ import annotations

from dataclasses import asdict, dataclass
from typing import Any, Iterable, Sequence

from mountain_twin.analysis_contract import (
    IntendedUseContext,
    TemporalProvenance,
    WeatherSourceType,
)
from mountain_twin.weather.provenance import make_temporal_provenance


@dataclass(frozen=True)
class RouteConditionResult:
    """One planned route location with independently labelled domain payloads."""

    route_id: str
    point_index: int
    route_distance_m: float
    latitude: float
    longitude: float
    route_elevation_m: float | None
    scenario_name: str
    planned_arrival_time: str
    elapsed_route_seconds: float
    timezone: str
    weather: dict[str, Any]
    weather_provenance: TemporalProvenance
    solar: dict[str, Any]
    terrain: dict[str, Any]
    quality: dict[str, Any]
    limitations: tuple[str, ...]

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


def compose_route_conditions(
    *,
    route_id: str,
    scenario_name: str,
    timeline: Sequence[Any],
    weather_rows: Sequence[dict[str, Any]],
    source_type: WeatherSourceType | None,
    fetched_at: str | None,
    provider: str | None,
    weather_units: dict[str, str | None] | None = None,
    solar_by_index: dict[int, dict[str, Any]] | None = None,
    terrain: dict[str, Any] | None = None,
    weather_resolution: dict[str, Any] | None = None,
) -> tuple[RouteConditionResult, ...]:
    if len(timeline) != len(weather_rows):
        raise ValueError("timeline and weather rows must have equal length")
    solar_by_index = solar_by_index or {}
    weather_units = weather_units or {}
    terrain = terrain or {
        "provider": None,
        "product": None,
        "surface_semantics": None,
        "status": "NOT_EVALUATED",
    }
    results = []
    for point, row in zip(timeline, weather_rows):
        if point.point_index != row["point_index"]:
            raise ValueError("route point index mismatch between timeline and weather")
        values = {
            variable: {**item, "unit": weather_units.get(variable)}
            for variable, item in row["values"].items()
        }
        valid_times = tuple(
            sorted(
                {
                    valid_time
                    for item in values.values()
                    for valid_time in item.get("temporal_contributing_valid_times", ())
                }
            )
        )
        if source_type is None:
            weather_provenance = TemporalProvenance(
                fetched_at=None,
                valid_at=None,
                intended_use_context=IntendedUseContext.PLANNING_CONTEXT,
                consistency_reason_codes=tuple(
                    (weather_resolution or {}).get("reason_codes", ("WEATHER_SOURCE_UNAVAILABLE",))
                ),
            )
        else:
            weather_provenance = make_temporal_provenance(
                source_type=source_type,
                fetched_at=fetched_at,
                valid_at=valid_times[0] if len(valid_times) == 1 else None,
                contributing_valid_times=valid_times,
                model_run_at=(weather_resolution or {}).get("model_run_at"),
                evaluation_time=point.planned_arrival,
                intended_use_context=IntendedUseContext.PLANNING_CONTEXT,
            )
        solar = solar_by_index.get(
            point.point_index,
            {
                "status": "UNKNOWN",
                "reason_codes": ("SOLAR_ARBITRARY_TIME_NOT_EVALUATED",),
                "planned_arrival_time": point.planned_arrival.isoformat(),
            },
        )
        weather_known = any(item.get("value") is not None for item in values.values())
        solar_known = solar.get("status") in {"DIRECT", "SHADOW", "NOT_APPLICABLE"}
        results.append(
            RouteConditionResult(
                route_id=route_id,
                point_index=point.point_index,
                route_distance_m=point.route_distance_m,
                latitude=row["latitude"],
                longitude=row["longitude"],
                route_elevation_m=point.elevation_m,
                scenario_name=scenario_name,
                planned_arrival_time=point.planned_arrival.isoformat(),
                elapsed_route_seconds=point.elapsed_route_seconds,
                timezone=row["scenario_datetime_timezone"],
                weather={
                    "provider": provider,
                    "source_type": source_type.value if source_type is not None else None,
                    "resolution": weather_resolution,
                    "values": values,
                    "spatial_representation": _spatial_states(values),
                    "temporal_representation": _temporal_states(values),
                },
                weather_provenance=weather_provenance,
                solar=solar,
                terrain=terrain,
                quality={
                    "weather_state": "SUCCESS" if weather_known else "UNRESOLVED",
                    "solar_state": "SUCCESS" if solar_known else "UNRESOLVED",
                    "weather_reason_codes": () if weather_known else ("WEATHER_VALUE_UNAVAILABLE",),
                    "solar_reason_codes": tuple(solar.get("reason_codes", ())),
                },
                limitations=(
                    "PLANNED_ARRIVAL_IS_SYNTHETIC_SCENARIO_NOT_GPX_CHRONOLOGY",
                    "WEATHER_MODEL_GRID_IS_COARSER_THAN_ROUTE_GEOMETRY",
                ),
            )
        )
    return tuple(results)


def conditions_summary(results: Iterable[RouteConditionResult]) -> dict[str, Any]:
    rows = tuple(results)
    values = [
        float(row.weather["values"]["temperature_2m"]["value"])
        for row in rows
        if row.weather["values"].get("temperature_2m", {}).get("value") is not None
    ]
    gusts = [
        float(row.weather["values"]["wind_gusts_10m"]["value"])
        for row in rows
        if row.weather["values"].get("wind_gusts_10m", {}).get("value") is not None
    ]
    solar_counts = {
        status: sum(row.solar.get("status") == status for row in rows)
        for status in ("DIRECT", "SHADOW", "NOT_APPLICABLE", "UNKNOWN")
    }
    return {
        "route_id": rows[0].route_id if rows else None,
        "scenario_name": rows[0].scenario_name if rows else None,
        "point_count": len(rows),
        "temperature_range_c": [min(values), max(values)] if values else [None, None],
        "maximum_wind_gust": max(gusts) if gusts else None,
        "maximum_wind_gust_unit": next(
            (
                row.weather["values"]["wind_gusts_10m"].get("unit")
                for row in rows
                if row.weather["values"].get("wind_gusts_10m", {}).get("unit")
            ),
            None,
        ),
        "solar_counts": solar_counts,
        "weather_known_points": sum(row.quality["weather_state"] == "SUCCESS" for row in rows),
        "solar_evaluated_points": sum(row.quality["solar_state"] == "SUCCESS" for row in rows),
    }


def _spatial_states(values: dict[str, Any]) -> dict[str, str]:
    return {name: str(item.get("representation_state", "UNKNOWN")) for name, item in values.items()}


def _temporal_states(values: dict[str, Any]) -> dict[str, str]:
    return {
        name: str(item.get("temporal_representation_state", "UNKNOWN"))
        for name, item in values.items()
    }


def _json_value(value: Any) -> Any:
    if hasattr(value, "value"):
        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
