"""Versioned, recomputable aligned projections over source-faithful telemetry."""

from __future__ import annotations

from dataclasses import dataclass
from datetime import datetime, timedelta
from enum import Enum

from .contracts import Activity, ActivityTelemetry, GeoPoint, TelemetrySample, _aware


ALIGNED_TELEMETRY_DERIVATION_VERSION = "aligned_telemetry_exact_source_time_v0_1"


class TelemetryAlignmentPolicy(str, Enum):
    """Declared evidence rule used to create an aligned telemetry projection."""

    EXACT_SOURCE_TIME = "EXACT_SOURCE_TIME"


class TelemetryAlignmentAmbiguityError(ValueError):
    """A source dataset supplies competing values for one field at one source time."""


_FIELD_VALUES = (
    ("distance_m", "distance_m"),
    ("position", "position"),
    ("elevation_m", "elevation_m"),
    ("heart_rate_bpm", "heart_rate_bpm"),
    ("speed_mps", "speed_mps"),
    ("pace_s_per_km", "pace_s_per_km"),
    ("cadence_rpm", "cadence_rpm"),
    ("temperature_c", "temperature_c"),
)
_FIELD_NAMES = frozenset(name for name, _ in _FIELD_VALUES)


@dataclass(frozen=True)
class AlignedTelemetryFieldEvidence:
    """Traceability for one populated derived field back to its source sample."""

    field_name: str
    source_stream: str
    source_sample_index: int
    source_elapsed_offset_s: float

    def __post_init__(self) -> None:
        if self.field_name not in _FIELD_NAMES:
            raise ValueError("aligned telemetry evidence names an unknown field")
        if not self.source_stream or self.source_sample_index < 0:
            raise ValueError("aligned telemetry evidence requires a source stream and non-negative index")
        if self.source_elapsed_offset_s < 0:
            raise ValueError("aligned telemetry evidence elapsed offset cannot be negative")


@dataclass(frozen=True)
class AlignedTelemetryObservation:
    """Partial multi-field observation grouped only by an exact source time coordinate."""

    elapsed_offset_s: float
    field_evidence: tuple[AlignedTelemetryFieldEvidence, ...]
    derived_observed_at: str | None = None
    distance_m: float | None = None
    latitude: float | None = None
    longitude: float | None = None
    elevation_m: float | None = None
    heart_rate_bpm: float | None = None
    speed_mps: float | None = None
    pace_s_per_km: float | None = None
    cadence_rpm: float | None = None
    temperature_c: float | None = None

    def __post_init__(self) -> None:
        if self.elapsed_offset_s < 0:
            raise ValueError("aligned telemetry elapsed offset cannot be negative")
        if self.derived_observed_at is not None:
            _aware(self.derived_observed_at, "derived telemetry observation")
        if (self.latitude is None) != (self.longitude is None):
            raise ValueError("aligned telemetry position requires latitude and longitude together")
        if self.latitude is not None:
            GeoPoint(self.latitude, self.longitude)
        values = _observation_values(self)
        populated = {name for name, value in values.items() if value is not None}
        evidence = {item.field_name: item for item in self.field_evidence}
        if len(evidence) != len(self.field_evidence):
            raise ValueError("aligned telemetry field evidence must be unique")
        if set(evidence) != populated:
            raise ValueError("every populated aligned field requires exactly one source evidence record")
        if any(item.source_elapsed_offset_s != self.elapsed_offset_s for item in evidence.values()):
            raise ValueError("aligned telemetry evidence must match the observation source time")


@dataclass(frozen=True)
class AlignedTelemetryDerivationProvenance:
    """Versioned provenance for a disposable projection of one source dataset."""

    source_telemetry_id: str
    source_normalization_version: str
    alignment_policy: TelemetryAlignmentPolicy
    derivation_version: str

    def __post_init__(self) -> None:
        if not all((self.source_telemetry_id, self.source_normalization_version, self.derivation_version)):
            raise ValueError("aligned telemetry derivation requires source dataset and version identities")


@dataclass(frozen=True)
class AlignedTelemetryProjection:
    """One recomputable aligned projection; it never replaces source telemetry."""

    projection_id: str
    activity_id: str
    provenance: AlignedTelemetryDerivationProvenance
    observations: tuple[AlignedTelemetryObservation, ...]

    def __post_init__(self) -> None:
        if not self.projection_id or not self.activity_id:
            raise ValueError("aligned telemetry projection requires identity and canonical activity")
        offsets = tuple(item.elapsed_offset_s for item in self.observations)
        if offsets != tuple(sorted(offsets)) or len(set(offsets)) != len(offsets):
            raise ValueError("aligned telemetry observations must have unique ascending source times")


def derive_aligned_telemetry(
    activity: Activity, source_telemetry: ActivityTelemetry
) -> AlignedTelemetryProjection:
    """Align one dataset by identical supplied elapsed offsets, without resampling."""
    if source_telemetry.activity_id != activity.activity_id:
        raise ValueError("source telemetry belongs to a different canonical activity")

    grouped: dict[float, dict[str, tuple[object, AlignedTelemetryFieldEvidence]]] = {}
    for sample in source_telemetry.samples:
        if sample.elapsed_offset_s is None:
            continue
        offset = sample.elapsed_offset_s
        values = _sample_values(sample)
        bucket = grouped.setdefault(offset, {})
        for field_name, value in values.items():
            if value is None:
                continue
            if field_name in bucket:
                raise TelemetryAlignmentAmbiguityError(
                    f"multiple source values for {field_name} at elapsed offset {offset}"
                )
            bucket[field_name] = (
                value,
                AlignedTelemetryFieldEvidence(
                    field_name, sample.source_stream, sample.source_sample_index, offset
                ),
            )

    observations = tuple(
        _aligned_observation(activity, offset, grouped[offset])
        for offset in sorted(grouped)
    )
    provenance = AlignedTelemetryDerivationProvenance(
        source_telemetry_id=source_telemetry.telemetry_id,
        source_normalization_version=source_telemetry.provenance.normalization_version,
        alignment_policy=TelemetryAlignmentPolicy.EXACT_SOURCE_TIME,
        derivation_version=ALIGNED_TELEMETRY_DERIVATION_VERSION,
    )
    return AlignedTelemetryProjection(
        projection_id=f"{source_telemetry.telemetry_id}:{ALIGNED_TELEMETRY_DERIVATION_VERSION}",
        activity_id=activity.activity_id,
        provenance=provenance,
        observations=observations,
    )


def _aligned_observation(
    activity: Activity,
    elapsed_offset_s: float,
    fields: dict[str, tuple[object, AlignedTelemetryFieldEvidence]],
) -> AlignedTelemetryObservation:
    values = {name: None for name in _FIELD_NAMES}
    values.update({name: value for name, (value, _) in fields.items()})
    position = values.pop("position")
    if position is not None:
        values["latitude"], values["longitude"] = position
    return AlignedTelemetryObservation(
        elapsed_offset_s=elapsed_offset_s,
        derived_observed_at=_derived_observed_at(activity, elapsed_offset_s),
        field_evidence=tuple(fields[name][1] for name in sorted(fields)),
        **values,
    )


def _derived_observed_at(activity: Activity, elapsed_offset_s: float) -> str | None:
    if activity.started_at is None:
        return None
    return (datetime.fromisoformat(activity.started_at) + timedelta(seconds=elapsed_offset_s)).isoformat()


def _sample_values(sample: TelemetrySample) -> dict[str, object | None]:
    return {
        "distance_m": sample.distance_m,
        "position": (sample.latitude, sample.longitude) if sample.latitude is not None else None,
        "elevation_m": sample.elevation_m,
        "heart_rate_bpm": sample.heart_rate_bpm,
        "speed_mps": sample.speed_mps,
        "pace_s_per_km": sample.pace_s_per_km,
        "cadence_rpm": sample.cadence_rpm,
        "temperature_c": sample.temperature_c,
    }


def _observation_values(observation: AlignedTelemetryObservation) -> dict[str, object | None]:
    return {
        "distance_m": observation.distance_m,
        "position": (observation.latitude, observation.longitude) if observation.latitude is not None else None,
        "elevation_m": observation.elevation_m,
        "heart_rate_bpm": observation.heart_rate_bpm,
        "speed_mps": observation.speed_mps,
        "pace_s_per_km": observation.pace_s_per_km,
        "cadence_rpm": observation.cadence_rpm,
        "temperature_c": observation.temperature_c,
    }
