"""Versioned Strava-to-Activity normalizer; only supplied source fields are mapped."""

from __future__ import annotations

from datetime import datetime, timezone
from typing import Any

from mountain_twin.activity.contracts import (
    Activity, ActivityProvenance, ActivitySourceKind, ActivityTelemetry, ActivityType,
    ExternalActivityReference, FieldAvailability, GeoPoint, RawProviderRecord,
    TelemetryProvenance, TelemetrySample, TrackReference, canonical_activity_type,
)

ACTIVITY_NORMALIZER_VERSION = "strava_activity_normalizer_v0_1"
TELEMETRY_NORMALIZER_VERSION = "strava_telemetry_normalizer_v0_1"
PROVIDER = "strava"


def strava_activity_type(sport_type: str | None) -> ActivityType:
    return {
        "Walk": ActivityType.WALKING, "Hike": ActivityType.HIKING,
        "Run": ActivityType.RUNNING, "TrailRun": ActivityType.TRAIL_RUNNING,
        "Ride": ActivityType.CYCLING, "GravelRide": ActivityType.GRAVEL_CYCLING,
        "MountainBikeRide": ActivityType.MOUNTAIN_BIKING, "EMountainBikeRide": ActivityType.MOUNTAIN_BIKING,
        "EBikeRide": ActivityType.E_BIKING, "Handcycle": ActivityType.ADAPTIVE_SPORT,
        "Velomobile": ActivityType.CYCLING, "VirtualRide": ActivityType.INDOOR_CYCLING,
        "WeightTraining": ActivityType.STRENGTH_TRAINING, "Crossfit": ActivityType.CROSSFIT,
        "Yoga": ActivityType.YOGA, "Pilates": ActivityType.PILATES,
        "RockClimbing": ActivityType.ROCK_CLIMBING,
        "Canoeing": ActivityType.CANOEING, "Kayaking": ActivityType.KAYAKING,
        "Rowing": ActivityType.ROWING, "StandUpPaddling": ActivityType.STAND_UP_PADDLING,
        "Swim": ActivityType.SWIMMING,
        "AlpineSki": ActivityType.ALPINE_SKIING, "BackcountrySki": ActivityType.SKI_TOURING,
        "NordicSki": ActivityType.CROSS_COUNTRY_SKIING, "Snowboard": ActivityType.SNOWBOARDING,
        "Snowshoe": ActivityType.SNOWSHOEING, "RollerSki": ActivityType.ROLLER_SKIING,
        "Sail": ActivityType.SAILING, "Surfing": ActivityType.SURFING,
        "Kitesurf": ActivityType.KITESURFING, "Windsurf": ActivityType.WINDSURFING,
        "IceSkate": ActivityType.ICE_SKATING, "InlineSkate": ActivityType.INLINE_SKATING,
        "Skateboard": ActivityType.SKATEBOARDING, "Workout": ActivityType.GENERAL_WORKOUT,
        "HighIntensityIntervalTraining": ActivityType.HIIT, "Elliptical": ActivityType.ELLIPTICAL,
        "StairStepper": ActivityType.STAIR_STEPPER, "PhysicalTherapy": ActivityType.REHABILITATION,
        "Wheelchair": ActivityType.ADAPTIVE_SPORT, "VirtualRun": ActivityType.RUNNING,
        "VirtualRow": ActivityType.ROWING,
        "Badminton": ActivityType.BADMINTON, "Basketball": ActivityType.BASKETBALL,
        "Cricket": ActivityType.CRICKET, "Golf": ActivityType.GOLF, "Padel": ActivityType.PADEL,
        "Pickleball": ActivityType.PICKLEBALL, "Racquetball": ActivityType.RACQUETBALL,
        "Soccer": ActivityType.SOCCER, "Squash": ActivityType.SQUASH,
        "TableTennis": ActivityType.TABLE_TENNIS, "Tennis": ActivityType.TENNIS,
        "Volleyball": ActivityType.VOLLEYBALL, "Dance": ActivityType.DANCE,
    }.get(sport_type or "", canonical_activity_type(sport_type))


def normalize_activity(raw: RawProviderRecord, payload: dict[str, Any]) -> Activity:
    activity_id = str(payload["id"])
    source = ExternalActivityReference(ActivitySourceKind.PROVIDER, PROVIDER, activity_id, raw.raw_record_id)
    track_id = payload.get("map", {}).get("id") if isinstance(payload.get("map"), dict) else None
    return Activity(
        activity_id=f"strava:{raw.user_id}:{activity_id}:{ACTIVITY_NORMALIZER_VERSION}",
        user_id=raw.user_id,
        activity_type=strava_activity_type(payload.get("sport_type") or payload.get("type")),
        provenance=ActivityProvenance((source,), (raw.raw_record_id,), ACTIVITY_NORMALIZER_VERSION, _now(), payload.get("sport_type") or payload.get("type")),
        started_at=_aware_iso(payload.get("start_date")), timezone=_timezone_name(payload.get("timezone")), title=payload.get("name"),
        elapsed_duration_s=_number(payload.get("elapsed_time")), moving_duration_s=_number(payload.get("moving_time")),
        distance_m=_number(payload.get("distance")), elevation_gain_m=_number(payload.get("total_elevation_gain")),
        min_elevation_m=_number(payload.get("elev_low")), max_elevation_m=_number(payload.get("elev_high")),
        average_speed_mps=_number(payload.get("average_speed")), max_speed_mps=_number(payload.get("max_speed")),
        average_heart_rate_bpm=_number(payload.get("average_heartrate")), max_heart_rate_bpm=_number(payload.get("max_heartrate")),
        track=TrackReference(f"strava-map:{track_id}") if track_id else None,
        start_location=_point(payload.get("start_latlng")), end_location=_point(payload.get("end_latlng")),
        field_availability=_availability(payload),
    )


def normalize_telemetry(raw: RawProviderRecord, activity: Activity, streams: dict[str, Any]) -> ActivityTelemetry:
    source = ExternalActivityReference(ActivitySourceKind.PROVIDER, PROVIDER, raw.source_activity_id, raw.raw_record_id)
    stream_values = {key: value.get("data", ()) for key, value in streams.items() if isinstance(value, dict)}
    time_values = stream_values.get("time", ())
    samples = tuple(
        sample
        for stream_name, values in stream_values.items()
        if stream_name != "time" and len(values) == len(time_values)
        for index, value in enumerate(values)
        if time_values[index] is not None
        for sample in (_stream_sample(stream_name, index, time_values[index], value),)
    )
    return ActivityTelemetry(
        telemetry_id=f"strava:{raw.raw_record_id}:{TELEMETRY_NORMALIZER_VERSION}", activity_id=activity.activity_id,
        provenance=TelemetryProvenance(source, raw.raw_record_id, TELEMETRY_NORMALIZER_VERSION, _now()), samples=samples,
    )


def _stream_sample(stream_name: str, index: int, elapsed_offset_s: float, value: Any) -> TelemetrySample:
    values = {"distance_m": None, "latitude": None, "longitude": None, "elevation_m": None,
              "heart_rate_bpm": None, "speed_mps": None, "cadence_rpm": None, "temperature_c": None}
    if stream_name == "latlng" and isinstance(value, list) and len(value) == 2:
        values["latitude"], values["longitude"] = value
    elif stream_name == "distance": values["distance_m"] = _number(value)
    elif stream_name == "altitude": values["elevation_m"] = _number(value)
    elif stream_name == "heartrate": values["heart_rate_bpm"] = _number(value)
    elif stream_name == "velocity_smooth": values["speed_mps"] = _number(value)
    elif stream_name == "cadence": values["cadence_rpm"] = _number(value)
    elif stream_name == "temp": values["temperature_c"] = _number(value)
    return TelemetrySample(f"strava:{stream_name}", index, elapsed_offset_s=_number(elapsed_offset_s), **values)


def _availability(payload: dict[str, Any]) -> tuple[FieldAvailability, ...]:
    fields = ("elapsed_duration_s", "moving_duration_s", "distance_m", "elevation_gain_m", "min_elevation_m", "max_elevation_m", "average_speed_mps", "max_speed_mps", "average_heart_rate_bpm", "max_heart_rate_bpm")
    source_names = {field: {"elapsed_duration_s": "elapsed_time", "moving_duration_s": "moving_time", "distance_m": "distance", "elevation_gain_m": "total_elevation_gain", "min_elevation_m": "elev_low", "max_elevation_m": "elev_high", "average_speed_mps": "average_speed", "max_speed_mps": "max_speed", "average_heart_rate_bpm": "average_heartrate", "max_heart_rate_bpm": "max_heartrate"}[field] for field in fields}
    from mountain_twin.activity.contracts import AvailabilityState
    return tuple(FieldAvailability(field, AvailabilityState.AVAILABLE if payload.get(source_names[field]) is not None else AvailabilityState.UNAVAILABLE) for field in fields)


def _number(value): return float(value) if value is not None else None
def _point(value): return GeoPoint(value[0], value[1]) if isinstance(value, list) and len(value) == 2 else None
def _timezone_name(value): return value.rsplit(" ", 1)[-1] if isinstance(value, str) and "/" in value else None
def _aware_iso(value): return value[:-1] + "+00:00" if isinstance(value, str) and value.endswith("Z") else value
def _now(): return datetime.now(timezone.utc).isoformat()
