"""The forecast for one place -- a peak, a pass, a town or a point on the map
-- outside any Journey (AV-062, "Pogoda"; docs/reports/AV-062_raport.md).

One document for the page and, later, the mobile app:

* ``levels`` -- the heights asked for (mountain_twin.places.elevation): a
  peak's summit, middle and base, else the place itself. ONE provider
  request carries all of them (a location per height; Open-Meteo corrects
  its values to the ``elevation`` asked for). Each level has its hourly
  provider series on the document's hour axis (``values_by_hour``), as
  they came -- nothing interpolated.
* ``hours`` -- from the current hour, as long as the horizon (24 h, 72 h, 7
  or 14 days) and the data reach; ``hour_states`` the sun's state each hour
  at the place (deterministic calculation, mountain_twin.solar), for the
  day/night bands.
* ``days`` -- per level and local day, over all fetched days (the daily
  cards show 7-14 days whatever the chart horizon): the provider hours'
  minimum/maximum/sum (``derived``: a plain aggregate of provider values,
  never a new value), sunrise/sunset (calculation) and the weather window.
* the weather window is a FILTER on provider values, not a verdict
  (ADR-001): the longest run of daylight hours where every criterion in
  WINDOW_CRITERIA holds -- no precipitation, low precipitation probability,
  gusts under a threshold, no thunderstorm weather code. A day where no hour
  meets them has no window (that is all it says); a day without the data
  is UNAVAILABLE.

Provider use: the basic request asks for BASIC_VARIABLES (ten: one call per
location); the further charts' EXTRA_VARIABLES are fetched only when asked
for; the model comparison only on request, for one level. Every answer goes
through the shared grid-cell cache (mountain_twin.weather.forecast_cache):
the same peak looked up by many people is one fetch per model run. The days
asked for are always FETCH_DAYS (Open-Meteo weighs up to 14 days the same),
so changing the horizon never fetches again.
"""

from __future__ import annotations

from datetime import datetime, timedelta
from typing import Any
from zoneinfo import ZoneInfo

from mountain_twin.journey.route_timezone import resolve_route_timezone
from mountain_twin.solar.states import classify_astronomical_state
from mountain_twin.solar.sun_engine import solar_position
from mountain_twin.weather.forecast_cache import call_source, call_weight
from mountain_twin.weather.live import LIVE_WEATHER_TTL_SECONDS
from mountain_twin.weather.sources import budget_document, sources_document

PLACE_FORECAST_CONTRACT = "place_forecast_v0_1"
BASIC_VARIABLES: tuple[str, ...] = (
    "temperature_2m",
    "apparent_temperature",
    "wind_speed_10m",
    "wind_gusts_10m",
    "wind_direction_10m",
    "precipitation",
    "precipitation_probability",
    "snowfall",
    "freezing_level_height",
    "weather_code",
)
EXTRA_VARIABLES: tuple[str, ...] = (
    "cloud_cover",
    "cloud_cover_low",
    "cloud_cover_mid",
    "cloud_cover_high",
    "relative_humidity_2m",
    "cape",
    "visibility",
    "snow_depth",
)
COMPARISON_MODELS: tuple[str, ...] = ("icon_seamless", "ecmwf_ifs025", "gfs_seamless")
FETCH_DAYS = 14
HORIZON_HOURS = {"24": 24, "72": 72, "7": 7 * 24, "14": 14 * 24}
DEFAULT_HORIZON = "72"
# The hour table (like Mountain-Forecast's): every TABLE_STEP_HOURS local
# hours for TABLE_DAYS, whatever the charts' horizon -- the provider's own
# values at those hours, picked, not averaged.
TABLE_DAYS = 7
TABLE_STEP_HOURS = 3
# The weather window (a filter, see the module docstring).
WINDOW_CRITERIA = {
    "precipitation_max_mm": 0.0,
    "precipitation_probability_max_pct": 30,
    "wind_gusts_max_kmh": 40,
    "no_thunderstorm_codes": [95, 96, 99],
    "daylight_only": True,
    "min_hours": 2,
}
# Sunrise/sunset: the sun's centre at -0.833 deg (standard refraction plus
# the disc's radius), the usual convention for published times.
SUN_HORIZON_DEG = -0.833


def place_forecast_document(
    *,
    place: dict[str, Any],
    levels: list[dict[str, Any]],
    single_level_reason: str | None,
    live_resolver,
    now: datetime | None = None,
    horizon: str | None = None,
    extras: bool = False,
    compare_level: str | None = None,
) -> dict[str, Any]:
    """``place``: name, kind (PEAK/PLACE), latitude, longitude and the rest
    of the search hit; ``levels``: Level.to_dict() of each height;
    ``compare_level``: the level whose named-model comparison is wanted
    (None: no comparison, nothing fetched for it)."""
    horizon = horizon or DEFAULT_HORIZON
    if horizon not in HORIZON_HOURS:
        raise ValueError("horizon must be 24, 72, 7 or 14")
    latitude, longitude = float(place["latitude"]), float(place["longitude"])
    zone_name = resolve_route_timezone(latitude, longitude)
    zone = ZoneInfo(zone_name)
    now_local = (now or datetime.now(zone)).astimezone(zone)
    first_hour = now_local.replace(minute=0, second=0, microsecond=0)
    start_date = now_local.date()
    end_date = start_date + timedelta(days=FETCH_DAYS - 1)
    points = [(latitude, longitude, level["elevation_m"]) for level in levels]

    freshness: dict[str, Any] = {}
    with call_source("PLACE"):
        basic = _resolve(
            live_resolver, points, start_date, end_date, zone_name, BASIC_VARIABLES, freshness
        )
        extra = (
            _resolve(
                live_resolver, points, start_date, end_date, zone_name, EXTRA_VARIABLES, freshness
            )
            if extras
            else None
        )
    fetched_hours = _hours(basic)
    hours = [
        hour
        for hour in fetched_hours
        if first_hour <= hour < first_hour + timedelta(hours=HORIZON_HOURS[horizon])
    ]
    failures = {failure for _, failure in basic if failure}
    axis_reason = None
    if not hours:
        axis_reason = next(
            (
                code
                for code in (
                    "PROVIDER_DAILY_LIMIT",
                    "PROVIDER_HOURLY_LIMIT",
                    "PROVIDER_RATE_LIMITED",
                    "PROVIDER_NETWORK_FAILURE",
                )
                if code in failures
            ),
            "PROVIDER_NO_DATA",
        )
    day_hours = [hour for hour in fetched_hours if hour >= first_hour.replace(hour=0)]
    table_hours = [
        hour
        for hour in fetched_hours
        if first_hour <= hour < first_hour + timedelta(days=TABLE_DAYS)
        and hour.hour % TABLE_STEP_HOURS == 0
    ]
    variables = BASIC_VARIABLES + (EXTRA_VARIABLES if extras else ())
    level_documents = []
    for position, level in enumerate(levels):
        resolved = [basic[position]] + ([extra[position]] if extra else [])
        level_documents.append(
            {
                **level,
                **_series(resolved, hours, variables),
                "table": {
                    variable: _aligned(basic[position][0], variable, table_hours)
                    for variable in BASIC_VARIABLES
                },
                "days": _days(resolved, day_hours, latitude, longitude, zone),
            }
        )
    comparison = None
    if compare_level is not None:
        comparison = _comparison(
            live_resolver,
            levels,
            compare_level,
            latitude,
            longitude,
            start_date,
            end_date,
            zone_name,
            hours,
            freshness,
        )
    spent = freshness.get("fetched_cells", 0)
    return {
        "contract": PLACE_FORECAST_CONTRACT,
        "place": place,
        "timezone": zone_name,
        "horizon": horizon,
        "axis": "TIME",
        "axis_state": "AVAILABLE" if hours else "UNAVAILABLE",
        "axis_unavailable_reason": axis_reason,
        "hours": [hour.isoformat() for hour in hours],
        "hour_states": [
            classify_astronomical_state(solar_position(hour, latitude, longitude)[0]).value
            for hour in hours
        ],
        "forecast_until": fetched_hours[-1].isoformat() if fetched_hours else None,
        "table_hours": [hour.isoformat() for hour in table_hours],
        "table_hour_states": [
            classify_astronomical_state(solar_position(hour, latitude, longitude)[0]).value
            for hour in table_hours
        ],
        "levels": level_documents,
        "single_level_reason": single_level_reason,
        "variables": list(variables),
        "extra_variables": list(EXTRA_VARIABLES),
        "extras_included": extras,
        "units": _units([basic] + ([extra] if extra else []), variables),
        "window_criteria": WINDOW_CRITERIA,
        "source": {
            "provider": "Open-Meteo",
            "model_selection": "best_match",
            "cache_ttl_seconds": LIVE_WEATHER_TTL_SECONDS,
        },
        # What this document weighs in the provider's quota (Open-Meteo:
        # per location and model, variables in tens) and how much of it
        # was actually fetched now -- the rest came from the shared cache.
        "provider_cost": {
            "basic_weight": call_weight(len(points), 1, len(BASIC_VARIABLES), FETCH_DAYS),
            "extras_weight": call_weight(len(points), 1, len(EXTRA_VARIABLES), FETCH_DAYS),
            "comparison_weight": call_weight(
                1, len(COMPARISON_MODELS), len(BASIC_VARIABLES), FETCH_DAYS
            ),
            "fetched_locations": spent,
            "cached_locations": freshness.get("cells", 0) - spent,
        },
        "forecast_freshness": _freshness(freshness, zone),
        # AV-051: a limit in force -- which one and when it comes back.
        "provider_limit": live_resolver.rate_limit_document()
        if callable(getattr(live_resolver, "rate_limit_document", None))
        else None,
        "model_comparison": comparison,
        **sources_document(basic),
        **budget_document(),
        "source_primary": getattr(live_resolver, "_primary", "open_meteo"),
    }


def _resolve(resolver, points, start_date, end_date, zone_name, variables, freshness):
    import inspect

    try:
        takes = "freshness" in inspect.signature(resolver.resolve_many).parameters
    except (TypeError, ValueError):
        takes = False
    return resolver.resolve_many(
        points,
        start_date=start_date,
        end_date=end_date,
        timezone_name=zone_name,
        variables=variables,
        **({"freshness": freshness} if takes else {}),
    )


def _hours(resolved) -> list[datetime]:
    return list(next((sample.times for sample, _ in resolved if sample and sample.times), ()))


def _aligned(sample, variable, hours) -> list[Any]:
    if sample is None:
        return [None] * len(hours)
    position = {time: i for i, time in enumerate(sample.times)}
    series = sample.variables_by_time.get(variable)
    if series is None:
        return [None] * len(hours)
    return [
        series[position[hour]] if hour in position and position[hour] < len(series) else None
        for hour in hours
    ]


def _series(resolved, hours, variables) -> dict[str, Any]:
    """One level's provider series on the hour axis (basic answer first, the
    extra one after), its gaps and which model gave each field."""
    values: dict[str, list[Any]] = {}
    missing: dict[str, str] = {}
    models: dict[str, str] = {}
    failure = None
    for sample, fail in resolved:
        failure = failure or fail
        for variable in variables:
            if variable in values and any(value is not None for value in values[variable]):
                continue
            if sample is None or variable not in sample.variables_by_time:
                continue
            values[variable] = _aligned(sample, variable, hours)
            models.update({k: v for k, v in dict(sample.variable_models).items() if k == variable})
    for variable in variables:
        values.setdefault(variable, [None] * len(hours))
        if not any(value is not None for value in values[variable]):
            reasons = [
                sample.base.missing_variable_reasons.get(variable)
                for sample, _ in resolved
                if sample is not None
            ]
            missing[variable] = next(
                (reason for reason in reasons if reason),
                failure or "WEATHER_PROVIDER_VALUE_MISSING",
            )
    present = sum(1 for series in values.values() for value in series if value is not None)
    total = sum(len(series) for series in values.values())
    state = "AVAILABLE" if total and present == total else "PARTIAL" if present else "UNAVAILABLE"
    return {
        "state": state,
        "unavailable_reason": None if present else failure or "PROVIDER_NO_DATA",
        "values_by_hour": values,
        "missing_reasons": missing,
        "variable_models": models,
    }


def _values_on(sample, variable, hours) -> list[Any]:
    return _aligned(sample, variable, hours)


def _days(resolved, day_hours, latitude, longitude, zone) -> list[dict[str, Any]]:
    """Per local day: aggregates of the provider's hours (labelled as such),
    the sun's times and the weather window."""
    sample = resolved[0][0]
    if sample is None or not day_hours:
        return []
    columns = {variable: _values_on(sample, variable, day_hours) for variable in BASIC_VARIABLES}
    by_day: dict[str, list[int]] = {}
    for index, hour in enumerate(day_hours):
        by_day.setdefault(hour.date().isoformat(), []).append(index)
    days = []
    for day, indexes in by_day.items():

        def pick(variable):
            return [columns[variable][i] for i in indexes if columns[variable][i] is not None]

        temperature = pick("temperature_2m")
        sunrise, sunset = sun_times(datetime.fromisoformat(day).date(), latitude, longitude, zone)
        days.append(
            {
                "date": day,
                "hours": len(indexes),
                "complete": len(indexes) == 24,
                "derived": "AGGREGATE_OF_PROVIDER_HOURS",
                "temperature_min": min(temperature) if temperature else None,
                "temperature_max": max(temperature) if temperature else None,
                "apparent_min": min(pick("apparent_temperature"), default=None),
                "precipitation_sum": _sum(pick("precipitation")),
                "snowfall_sum": _sum(pick("snowfall")),
                "precipitation_probability_max": max(
                    pick("precipitation_probability"), default=None
                ),
                "wind_speed_max": max(pick("wind_speed_10m"), default=None),
                "wind_gusts_max": max(pick("wind_gusts_10m"), default=None),
                "freezing_level_min": min(pick("freezing_level_height"), default=None),
                "weather_code_max": max(pick("weather_code"), default=None),
                "sunrise": sunrise.isoformat() if sunrise else None,
                "sunset": sunset.isoformat() if sunset else None,
                "window": _window(
                    [day_hours[i] for i in indexes],
                    {variable: [columns[variable][i] for i in indexes] for variable in columns},
                    latitude,
                    longitude,
                ),
            }
        )
    return days


def _sum(values) -> float | None:
    return round(sum(values), 2) if values else None


def _window(hours, columns, latitude, longitude) -> dict[str, Any]:
    """The longest run of hours meeting WINDOW_CRITERIA (see the module
    docstring); hours missing any criterion's value never count."""
    criteria = WINDOW_CRITERIA
    meets: list[bool | None] = []
    for i, hour in enumerate(hours):
        values = {variable: columns[variable][i] for variable in columns}
        needed = (
            values["precipitation"],
            values["precipitation_probability"],
            values["wind_gusts_10m"],
            values["weather_code"],
        )
        if any(value is None for value in needed):
            meets.append(None)
            continue
        daylight = solar_position(hour, latitude, longitude)[0] >= 0
        meets.append(
            (daylight or not criteria["daylight_only"])
            and values["precipitation"] <= criteria["precipitation_max_mm"]
            and values["precipitation_probability"] <= criteria["precipitation_probability_max_pct"]
            and values["wind_gusts_10m"] < criteria["wind_gusts_max_kmh"]
            and int(values["weather_code"]) not in criteria["no_thunderstorm_codes"]
        )
    if not any(value is not None for value in meets):
        return {"state": "UNAVAILABLE", "reason": "PROVIDER_NO_DATA"}
    best, run = None, None
    for i, ok in enumerate(meets):
        if ok:
            run = (run[0], i) if run else (i, i)
            if best is None or run[1] - run[0] > best[1] - best[0]:
                best = run
        else:
            run = None
    if best is None or best[1] - best[0] + 1 < criteria["min_hours"]:
        return {"state": "NONE", "reason": "NO_HOURS_MEET_CRITERIA"}
    return {
        "state": "FOUND",
        "start": hours[best[0]].isoformat(),
        "end": (hours[best[1]] + timedelta(hours=1)).isoformat(),
        "hours": best[1] - best[0] + 1,
    }


def sun_times(day, latitude, longitude, zone) -> tuple[datetime | None, datetime | None]:
    """Sunrise and sunset (the sun's centre crossing SUN_HORIZON_DEG) of a
    local day, to the minute; None for a polar day or night."""

    def elevation(moment):
        return solar_position(moment, latitude, longitude)[0] - SUN_HORIZON_DEG

    start = datetime(day.year, day.month, day.day, tzinfo=zone)
    steps = [start + timedelta(minutes=15 * i) for i in range(97)]
    heights = [elevation(moment) for moment in steps]
    sunrise = sunset = None
    for (a, ha), (b, hb) in zip(zip(steps, heights), zip(steps[1:], heights[1:])):
        if ha < 0 <= hb and sunrise is None:
            sunrise = _bisect(elevation, a, b, rising=True)
        if ha >= 0 > hb and sunset is None:
            sunset = _bisect(elevation, a, b, rising=False)
    return sunrise, sunset


def _bisect(height, low, high, rising):
    for _ in range(10):
        middle = low + (high - low) / 2
        if (height(middle) >= 0) == rising:
            high = middle
        else:
            low = middle
    return (low + (high - low) / 2).replace(second=0, microsecond=0)


def _comparison(
    resolver,
    levels,
    level_id,
    latitude,
    longitude,
    start_date,
    end_date,
    zone_name,
    hours,
    freshness,
) -> dict[str, Any]:
    """Named models side by side for one level -- never merged (ADR-003)."""
    level = next((item for item in levels if item["level_id"] == level_id), levels[0])
    with call_source("COMPARE"):
        answer = resolver.resolve_many_models(
            [(latitude, longitude, level["elevation_m"])],
            models=COMPARISON_MODELS,
            start_date=start_date,
            end_date=end_date,
            timezone_name=zone_name,
            variables=BASIC_VARIABLES,
            freshness=freshness,
        )[0]
    by_model, failure = answer
    models = {}
    for model in COMPARISON_MODELS:
        sample = (by_model or {}).get(model)
        values = {variable: _aligned(sample, variable, hours) for variable in BASIC_VARIABLES}
        available = any(value is not None for series in values.values() for value in series)
        models[model] = {
            "state": "AVAILABLE" if available else "UNAVAILABLE",
            "unavailable_reason": None if available else failure or "PROVIDER_NO_DATA",
            "values_by_hour": values,
        }
    return {
        "level_id": level["level_id"],
        "elevation_m": level["elevation_m"],
        "models": list(COMPARISON_MODELS),
        "state": "AVAILABLE"
        if any(entry["state"] == "AVAILABLE" for entry in models.values())
        else "UNAVAILABLE",
        "by_model": models,
    }


def _units(answers, variables) -> dict[str, str | None]:
    units: dict[str, str | None] = {variable: None for variable in variables}
    for answer in answers:
        for sample, _ in answer:
            if sample is None:
                continue
            for variable in variables:
                unit = sample.base.units.get(variable)
                if units[variable] is None and unit and unit != "undefined":
                    units[variable] = unit
    return units


def _freshness(freshness: dict[str, Any], zone) -> dict[str, Any] | None:
    oldest = freshness.get("oldest_fetched_at")
    if oldest is None:
        return None
    return {
        "fetched_at": datetime.fromtimestamp(oldest, zone).isoformat(),
        "refreshing": bool(freshness.get("refreshing")),
        "stale_cells": freshness.get("stale_cells", 0),
    }
