"""Pure provider-neutral Personal Intelligence computation over frozen canonical inputs."""

from __future__ import annotations

import hashlib
import json
import math
from datetime import datetime, timedelta, timezone
from typing import Iterable

from mountain_twin.activity import ACTIVITY_TAXONOMY_VERSION, AvailabilityState
from mountain_twin.activity.persistence import ActivityRevision

from .contracts import (
    BASELINE_POLICY_VERSION,
    PERSONAL_INTELLIGENCE_ALGORITHM_VERSION,
    QUANTILE_POLICY_VERSION,
    WINDOW_DAYS,
    WINDOW_POLICY_VERSION,
    ActivityComposition,
    BaselineDistribution,
    ComparisonRelation,
    CompositionEntry,
    ExposureDistribution,
    FrozenActivity,
    HistoricalRequirement,
    IntelligenceAvailability,
    PersonalIntelligenceInput,
    PersonalIntelligenceSnapshot,
    RequirementEvidence,
    SignalComparison,
    SignalCoverage,
    SignalMeasurement,
    SignalName,
    WindowSignalSet,
)

_METRIC_FIELDS = {
    SignalName.MOVING_DURATION: ("moving_duration_s", "s"),
    SignalName.ELAPSED_DURATION: ("elapsed_duration_s", "s"),
    SignalName.DISTANCE: ("distance_m", "m"),
    SignalName.ELEVATION_GAIN: ("elevation_gain_m", "m"),
}


def freeze_current_revisions(
    revisions: Iterable[ActivityRevision],
    *,
    user_id: str,
    as_of: datetime,
    taxonomy_version: str = ACTIVITY_TAXONOMY_VERSION,
) -> PersonalIntelligenceInput:
    """Freeze selected current revisions before any Personal Intelligence computation."""
    as_of_utc = _as_utc(as_of)
    frozen = []
    for revision in revisions:
        activity = revision.canonical_activity
        if activity.user_id != user_id:
            raise ValueError("Personal Intelligence revision owner differs from requested user")
        if revision.taxonomy_version != taxonomy_version:
            raise ValueError("Personal Intelligence requires one explicit taxonomy version")
        frozen.append(
            FrozenActivity(
                revision.activity_entity_id,
                revision.activity_revision_id,
                activity.activity_type,
                activity.activity_family,
                _as_utc(datetime.fromisoformat(activity.started_at))
                if activity.started_at is not None
                else None,
                activity.moving_duration_s,
                activity.elapsed_duration_s,
                activity.distance_m,
                activity.elevation_gain_m,
                tuple(
                    sorted((item.field_name, item.state) for item in activity.field_availability)
                ),
            )
        )
    return PersonalIntelligenceInput(
        user_id,
        as_of_utc,
        taxonomy_version,
        tuple(
            sorted(frozen, key=lambda item: (item.activity_entity_id, item.activity_revision_id))
        ),
    )


def compute_personal_intelligence(
    frozen_input: PersonalIntelligenceInput,
    *,
    requirement: HistoricalRequirement | None = None,
) -> PersonalIntelligenceSnapshot:
    """Compute reproducible descriptive signals without rereading persistence state."""
    as_of = _as_utc(frozen_input.as_of)
    dated = tuple(
        item
        for item in frozen_input.activities
        if item.started_at is not None and item.started_at < as_of
    )
    undated_count = sum(item.started_at is None for item in frozen_input.activities)
    future_count = sum(
        item.started_at is not None and item.started_at >= as_of for item in frozen_input.activities
    )
    earliest = min((item.started_at for item in dated), default=None)

    all_history = _measurements(dated, history_complete=True)
    windows = tuple(
        _window_signal_set(dated, earliest, as_of, window_days) for window_days in WINDOW_DAYS
    )
    composition = _composition(dated)
    moving_distribution = _exposure_distribution(dated, SignalName.MOVING_DURATION)
    elevation_distribution = _exposure_distribution(dated, SignalName.ELEVATION_GAIN)
    requirement_evidence = _requirement_evidence(dated, requirement) if requirement else None
    semantic = {
        "algorithm_version": PERSONAL_INTELLIGENCE_ALGORITHM_VERSION,
        "as_of": _instant_text(as_of),
        "baseline_policy_version": BASELINE_POLICY_VERSION,
        "input_revisions": [
            [item.activity_entity_id, item.activity_revision_id] for item in frozen_input.activities
        ],
        "parameters": _requirement_parameters(requirement),
        "taxonomy_version": frozen_input.taxonomy_version,
        "window_policy_version": WINDOW_POLICY_VERSION,
    }
    input_digest = _digest(semantic)
    return PersonalIntelligenceSnapshot(
        f"personal-intelligence-snapshot-v0_1:{input_digest}",
        input_digest,
        frozen_input.user_id,
        as_of,
        PERSONAL_INTELLIGENCE_ALGORITHM_VERSION,
        frozen_input.taxonomy_version,
        WINDOW_POLICY_VERSION,
        BASELINE_POLICY_VERSION,
        frozen_input,
        undated_count,
        future_count,
        all_history,
        windows,
        composition,
        moving_distribution,
        elevation_distribution,
        requirement_evidence,
    )


def _window_signal_set(
    dated: tuple[FrozenActivity, ...],
    earliest: datetime | None,
    as_of: datetime,
    window_days: int,
) -> WindowSignalSet:
    start = as_of - timedelta(days=window_days)
    history_complete = earliest is not None and start >= earliest
    current = _in_window(dated, start, as_of)
    measurements = _measurements(current, history_complete=history_complete)
    comparisons = tuple(
        _comparison(dated, earliest, as_of, window_days, measurement)
        for measurement in measurements
    )
    return WindowSignalSet(window_days, measurements, comparisons, history_complete)


def _measurements(
    activities: tuple[FrozenActivity, ...], *, history_complete: bool
) -> tuple[SignalMeasurement, ...]:
    return (
        _activity_count(activities, history_complete),
        *(_numeric_measurement(activities, signal, history_complete) for signal in _METRIC_FIELDS),
    )


def _activity_count(
    activities: tuple[FrozenActivity, ...], history_complete: bool
) -> SignalMeasurement:
    count = len(activities)
    coverage = SignalCoverage(count, count, 0, history_complete)
    return SignalMeasurement(
        SignalName.ACTIVITY_COUNT,
        float(count),
        "activities",
        IntelligenceAvailability.AVAILABLE
        if history_complete
        else IntelligenceAvailability.PARTIAL,
        coverage,
    )


def _numeric_measurement(
    activities: tuple[FrozenActivity, ...],
    signal: SignalName,
    history_complete: bool,
) -> SignalMeasurement:
    field, unit = _METRIC_FIELDS[signal]
    values = [
        float(getattr(item, field)) for item in activities if getattr(item, field) is not None
    ]
    candidate_count = len(activities)
    evaluable_count = len(values)
    missing_count = candidate_count - evaluable_count
    coverage = SignalCoverage(candidate_count, evaluable_count, missing_count, history_complete)
    quality_complete = all(
        _field_quality_complete(item, field)
        for item in activities
        if getattr(item, field) is not None
    )
    if not activities:
        if history_complete:
            return SignalMeasurement(
                signal, 0.0, unit, IntelligenceAvailability.AVAILABLE, coverage
            )
        return SignalMeasurement(signal, None, unit, IntelligenceAvailability.PARTIAL, coverage)
    if not values:
        return SignalMeasurement(signal, None, unit, IntelligenceAvailability.UNAVAILABLE, coverage)
    availability = (
        IntelligenceAvailability.AVAILABLE
        if history_complete and missing_count == 0 and quality_complete
        else IntelligenceAvailability.PARTIAL
    )
    return SignalMeasurement(signal, sum(values), unit, availability, coverage)


def _comparison(
    dated: tuple[FrozenActivity, ...],
    earliest: datetime | None,
    as_of: datetime,
    window_days: int,
    current: SignalMeasurement,
) -> SignalComparison:
    required_count = 3 if window_days == 365 else 6
    values = []
    for index in range(1, required_count + 1):
        end = as_of - timedelta(days=index * window_days)
        start = as_of - timedelta(days=(index + 1) * window_days)
        if earliest is None or start < earliest:
            continue
        measurement = _measurement_for_signal(_in_window(dated, start, end), True, current.signal)
        if (
            measurement.availability is IntelligenceAvailability.AVAILABLE
            and measurement.value is not None
        ):
            values.append(measurement.value)
    annual = window_days == 365
    baseline = _baseline_distribution(values, required_count, annual)
    if current.availability is not IntelligenceAvailability.AVAILABLE:
        return SignalComparison(current.signal, current.availability, baseline)
    if baseline.availability is not IntelligenceAvailability.AVAILABLE:
        return SignalComparison(
            current.signal, IntelligenceAvailability.INSUFFICIENT_HISTORY, baseline
        )
    assert current.value is not None
    lower = baseline.minimum if annual else baseline.lower_quartile
    upper = baseline.maximum if annual else baseline.upper_quartile
    assert lower is not None and upper is not None
    relation = (
        ComparisonRelation.BELOW_REFERENCE_RANGE
        if current.value < lower
        else ComparisonRelation.ABOVE_REFERENCE_RANGE
        if current.value > upper
        else ComparisonRelation.WITHIN_REFERENCE_RANGE
    )
    return SignalComparison(current.signal, IntelligenceAvailability.AVAILABLE, baseline, relation)


def _measurement_for_signal(
    activities: tuple[FrozenActivity, ...], history_complete: bool, signal: SignalName
) -> SignalMeasurement:
    if signal is SignalName.ACTIVITY_COUNT:
        return _activity_count(activities, history_complete)
    return _numeric_measurement(activities, signal, history_complete)


def _baseline_distribution(
    values: list[float], required_count: int, annual: bool
) -> BaselineDistribution:
    if len(values) < required_count:
        return BaselineDistribution(IntelligenceAvailability.INSUFFICIENT_HISTORY, len(values))
    sorted_values = sorted(values)
    return BaselineDistribution(
        IntelligenceAvailability.AVAILABLE,
        len(sorted_values),
        median=_quantile(sorted_values, 0.5),
        minimum=sorted_values[0],
        maximum=sorted_values[-1],
        lower_quartile=None if annual else _quantile(sorted_values, 0.25),
        upper_quartile=None if annual else _quantile(sorted_values, 0.75),
    )


def _composition(activities: tuple[FrozenActivity, ...]) -> ActivityComposition:
    return ActivityComposition(
        _composition_entries(activities, lambda item: item.activity_family.value),
        _composition_entries(activities, lambda item: item.activity_type.value),
    )


def _composition_entries(activities, selector) -> tuple[CompositionEntry, ...]:
    grouped = {}
    for item in activities:
        grouped.setdefault(selector(item), []).append(item)
    entries = []
    for classification, members in grouped.items():
        members_tuple = tuple(members)
        moving = _numeric_measurement(
            members_tuple, SignalName.MOVING_DURATION, history_complete=True
        )
        entries.append(
            CompositionEntry(
                classification,
                len(members_tuple),
                moving.value,
                moving.coverage,
            )
        )
    return tuple(sorted(entries, key=lambda item: item.classification))


def _exposure_distribution(
    activities: tuple[FrozenActivity, ...], signal: SignalName
) -> ExposureDistribution:
    measurement = _numeric_measurement(activities, signal, history_complete=True)
    field, _unit = _METRIC_FIELDS[signal]
    values = sorted(
        float(getattr(item, field)) for item in activities if getattr(item, field) is not None
    )
    if not values:
        return ExposureDistribution(signal, measurement.availability, measurement.coverage)
    return ExposureDistribution(
        signal,
        measurement.availability,
        measurement.coverage,
        values[0],
        _quantile(values, 0.25),
        _quantile(values, 0.5),
        _quantile(values, 0.75),
        values[-1],
    )


def _requirement_evidence(
    activities: tuple[FrozenActivity, ...], requirement: HistoricalRequirement
) -> RequirementEvidence:
    required_signals = []
    if requirement.min_moving_duration_s is not None:
        required_signals.append(SignalName.MOVING_DURATION)
    if requirement.min_elevation_gain_m is not None:
        required_signals.append(SignalName.ELEVATION_GAIN)
    coverage = tuple(
        (signal, _numeric_measurement(activities, signal, history_complete=True).coverage)
        for signal in required_signals
    )
    qualifying = []
    evaluable_count = 0
    quality_complete = True
    for item in activities:
        values = {
            SignalName.MOVING_DURATION: item.moving_duration_s,
            SignalName.ELEVATION_GAIN: item.elevation_gain_m,
        }
        if any(values[signal] is None for signal in required_signals):
            continue
        evaluable_count += 1
        quality_complete = quality_complete and all(
            _field_quality_complete(item, _METRIC_FIELDS[signal][0]) for signal in required_signals
        )
        matches = (
            requirement.min_moving_duration_s is None
            or values[SignalName.MOVING_DURATION] >= requirement.min_moving_duration_s
        ) and (
            requirement.min_elevation_gain_m is None
            or values[SignalName.ELEVATION_GAIN] >= requirement.min_elevation_gain_m
        )
        if matches:
            qualifying.append((item.activity_entity_id, item.activity_revision_id))
    total_count = len(activities)
    non_evaluable_count = total_count - evaluable_count
    if total_count and not evaluable_count:
        availability = IntelligenceAvailability.UNAVAILABLE
    elif non_evaluable_count or not quality_complete:
        availability = IntelligenceAvailability.PARTIAL
    else:
        availability = IntelligenceAvailability.AVAILABLE
    return RequirementEvidence(
        requirement,
        availability,
        total_count,
        evaluable_count,
        non_evaluable_count,
        len(qualifying),
        tuple(qualifying),
        coverage,
    )


def _field_quality_complete(activity: FrozenActivity, field_name: str) -> bool:
    availability = dict(activity.field_availability).get(field_name)
    return availability not in {AvailabilityState.PARTIAL, AvailabilityState.UNKNOWN}


def _in_window(
    activities: tuple[FrozenActivity, ...], start: datetime, end: datetime
) -> tuple[FrozenActivity, ...]:
    return tuple(
        item
        for item in activities
        if item.started_at is not None and start <= item.started_at < end
    )


def _quantile(values: list[float], probability: float) -> float:
    """R-7 linear interpolation, explicitly fixed by QUANTILE_POLICY_VERSION."""
    if not values:
        raise ValueError("quantile requires values")
    position = (len(values) - 1) * probability
    lower = math.floor(position)
    upper = math.ceil(position)
    if lower == upper:
        return values[lower]
    return values[lower] + (values[upper] - values[lower]) * (position - lower)


def _as_utc(value: datetime) -> datetime:
    if value.tzinfo is None or value.utcoffset() is None:
        raise ValueError("Personal Intelligence datetime must be timezone-aware")
    return value.astimezone(timezone.utc)


def _instant_text(value: datetime) -> str:
    return _as_utc(value).isoformat().replace("+00:00", "Z")


def _digest(value: object) -> str:
    return hashlib.sha256(
        json.dumps(value, sort_keys=True, separators=(",", ":"), ensure_ascii=True).encode()
    ).hexdigest()


def _requirement_parameters(
    requirement: HistoricalRequirement | None,
) -> dict[str, float | None] | None:
    if requirement is None:
        return None
    return {
        "min_elevation_gain_m": requirement.min_elevation_gain_m,
        "min_moving_duration_s": requirement.min_moving_duration_s,
        "quantile_policy_version": QUANTILE_POLICY_VERSION,
    }
