from __future__ import annotations

import json
import math
from dataclasses import dataclass
from pathlib import Path
from typing import Any

import numpy as np

from .ctelemetry import parse_tc_file

def _mean(values: list[float]) -> float:
    return float(sum(values) / len(values)) if values else 0.0


def _stddev(values: list[float]) -> float:
    if len(values) < 2:
        return 0.0
    return float(np.std(np.array(values, dtype=float), ddof=0))


def _first_true_position(
    positions: np.ndarray,
    values: np.ndarray,
    *,
    threshold: float,
) -> float | None:
    indices = np.where(values > threshold)[0]
    if indices.size == 0:
        return None
    return float(positions[int(indices[0])])


def _safe_interp(positions: np.ndarray, values: np.ndarray, grid: np.ndarray) -> np.ndarray:
    if positions.size < 2:
        return np.full_like(grid, np.nan, dtype=float)
    monotonic = np.maximum.accumulate(np.clip(positions, 0.0, 1.0))
    unique_positions, unique_indices = np.unique(monotonic, return_index=True)
    unique_values = values[unique_indices]
    if unique_positions.size < 2:
        return np.full_like(grid, np.nan, dtype=float)
    return np.interp(grid, unique_positions, unique_values)


@dataclass(slots=True)
class Lap:
    lap_number: int
    lap_time_ms: int
    lap_time_s: float
    lap_length_m: float
    start_distance_m: float
    sector_times_ms: list[int]
    samples: list[dict[str, Any]]


def load_session(session_dir: Path) -> tuple[dict[str, Any], list[dict[str, Any]]]:
    meta = json.loads((session_dir / "meta.json").read_text(encoding="utf-8"))
    samples = []
    with (session_dir / "samples.jsonl").open("r", encoding="utf-8") as handle:
        for line in handle:
            line = line.strip()
            if not line:
                continue
            samples.append(json.loads(line))
    return meta, samples


def segment_laps(samples: list[dict[str, Any]], started_mid_lap: bool = False) -> list[Lap]:
    if not samples:
        return []
    laps: list[Lap] = []
    current: list[dict[str, Any]] = []
    sector_times: list[int] = []
    lap_index = 1
    started = not started_mid_lap

    for sample in samples:
        if sample.get("status") not in {"live", "pause"}:
            continue
        current.append(sample)
        if len(current) < 2:
            continue
        prev = current[-2]
        lap_counter_increased = int(sample["completed_laps"]) > int(prev["completed_laps"])
        sector_changed = int(sample["current_sector_index"]) != int(prev["current_sector_index"])
        if sector_changed and int(sample["last_sector_time_ms"]) > 0:
            sector_times.append(int(sample["last_sector_time_ms"]))

        if not lap_counter_increased:
            continue

        lap_samples = current[:-1]
        current = [sample]
        if not started:
            started = True
            sector_times = []
            continue

        reported_lap_ms = int(sample["last_lap_time_ms"])
        if reported_lap_ms <= 0 and lap_samples:
            reported_lap_ms = int(
                round(
                    (
                        float(lap_samples[-1]["captured_at_unix_s"])
                        - float(lap_samples[0]["captured_at_unix_s"])
                    )
                    * 1000.0
                )
            )
        if reported_lap_ms <= 0 or len(lap_samples) < 12:
            sector_times = []
            continue

        start_distance = float(lap_samples[0]["distance_traveled_m"])
        end_distance = float(lap_samples[-1]["distance_traveled_m"])
        lap_length_m = max(0.0, end_distance - start_distance)
        laps.append(
            Lap(
                lap_number=lap_index,
                lap_time_ms=reported_lap_ms,
                lap_time_s=reported_lap_ms / 1000.0,
                lap_length_m=lap_length_m,
                start_distance_m=start_distance,
                sector_times_ms=sector_times[:],
                samples=lap_samples,
            )
        )
        lap_index += 1
        sector_times = []

    return laps


def _lap_progress_positions(lap: Lap) -> np.ndarray:
    if lap.lap_length_m > 1.0 and all("distance_traveled_m" in sample for sample in lap.samples):
        positions = np.array(
            [
                max(0.0, (float(sample["distance_traveled_m"]) - lap.start_distance_m) / lap.lap_length_m)
                for sample in lap.samples
            ],
            dtype=float,
        )
        if np.nanmax(positions) > 0.2:
            return np.clip(positions, 0.0, 1.0)
    raw_positions = np.array(
        [float(sample["normalized_car_position"]) for sample in lap.samples],
        dtype=float,
    )
    unwrapped = raw_positions.copy()
    wraps = 0.0
    for idx in range(1, unwrapped.size):
        if unwrapped[idx] < unwrapped[idx - 1] - 0.5:
            wraps += 1.0
        unwrapped[idx] += wraps
    start = float(unwrapped[0])
    end = float(unwrapped[-1])
    span = max(end - start, 1e-6)
    return np.clip((unwrapped - start) / span, 0.0, 1.0)


def _lap_series(lap: Lap, field: str) -> tuple[np.ndarray, np.ndarray]:
    positions = _lap_progress_positions(lap)
    if field == "front_slip_abs":
        values = np.array(
            [
                (
                    abs(float(sample["slip_angle_rad"][0]))
                    + abs(float(sample["slip_angle_rad"][1]))
                )
                / 2.0
                for sample in lap.samples
            ],
            dtype=float,
        )
        return positions, values

    return positions, np.array([float(sample[field]) for sample in lap.samples], dtype=float)


def segment_metrics(lap: Lap, start_pos: float, end_pos: float) -> dict[str, float | None]:
    positions = _lap_progress_positions(lap)
    mask = (positions >= start_pos) & (positions <= end_pos)
    if not np.any(mask):
        return {}
    segment = [sample for sample, keep in zip(lap.samples, mask, strict=False) if keep]
    if not segment:
        return {}

    seg_positions = positions[mask]
    speed = np.array([float(sample["speed_kmh"]) for sample in segment], dtype=float)
    brake = np.array([float(sample["brake"]) for sample in segment], dtype=float)
    throttle = np.array([float(sample["throttle"]) for sample in segment], dtype=float)
    steer = np.array([abs(float(sample["steer_angle_deg"])) for sample in segment], dtype=float)
    front_slip = np.array(
        [
            (
                abs(float(sample["slip_angle_rad"][0]))
                + abs(float(sample["slip_angle_rad"][1]))
            )
            / 2.0
            for sample in segment
        ],
        dtype=float,
    )

    apex_index = int(np.argmin(speed))
    throttle_reapply_pos = _first_true_position(
        seg_positions[apex_index:],
        throttle[apex_index:],
        threshold=0.40,
    )
    brake_onset_pos = _first_true_position(seg_positions, brake, threshold=0.10)

    return {
        "entry_speed_kmh": float(speed[0]),
        "apex_speed_kmh": float(np.min(speed)),
        "exit_speed_kmh": float(speed[-1]),
        "brake_peak": float(np.max(brake)),
        "brake_mean": float(np.mean(brake)),
        "brake_duration_ratio": float(np.mean(brake > 0.10)),
        "brake_onset_pos": brake_onset_pos,
        "throttle_mean": float(np.mean(throttle)),
        "throttle_reapply_pos": throttle_reapply_pos,
        "steer_peak_deg": float(np.max(steer)),
        "front_slip_abs": float(np.mean(front_slip)),
    }


def _interp_array(positions: np.ndarray, values: np.ndarray, target_positions: np.ndarray) -> np.ndarray:
    monotonic = np.maximum.accumulate(np.clip(positions, 0.0, 1.0))
    unique_positions, unique_indices = np.unique(monotonic, return_index=True)
    if unique_positions.size < 2:
        return np.full((target_positions.size, values.shape[1]), np.nan, dtype=float)
    out = np.empty((target_positions.size, values.shape[1]), dtype=float)
    for axis in range(values.shape[1]):
        out[:, axis] = np.interp(target_positions, unique_positions, values[unique_indices, axis])
    return out


def line_deviation_stats(reference_lap: Lap, candidate_lap: Lap, start_pos: float, end_pos: float) -> dict[str, float]:
    if not all("car_coordinates" in sample for sample in reference_lap.samples) or not all(
        "car_coordinates" in sample for sample in candidate_lap.samples
    ):
        return {"mean_m": 0.0, "max_m": 0.0}
    ref_positions = _lap_progress_positions(reference_lap)
    cand_positions = _lap_progress_positions(candidate_lap)
    mask = (cand_positions >= start_pos) & (cand_positions <= end_pos)
    if not np.any(mask):
        return {"mean_m": 0.0, "max_m": 0.0}
    target_positions = cand_positions[mask]
    ref_coords = np.array([sample["car_coordinates"] for sample in reference_lap.samples], dtype=float)
    cand_coords = np.array([sample["car_coordinates"] for sample in candidate_lap.samples], dtype=float)[mask]
    ref_interp = _interp_array(ref_positions, ref_coords, target_positions)
    distances = np.linalg.norm(cand_coords - ref_interp, axis=1)
    return {
        "mean_m": float(np.nanmean(distances)),
        "max_m": float(np.nanmax(distances)),
    }


def telemetry_quality(samples: list[dict[str, Any]]) -> dict[str, Any]:
    def scalar_series(key: str) -> list[float]:
        values = []
        for sample in samples:
            value = sample.get(key)
            if value is None:
                continue
            values.append(float(value))
        return values

    def vector_abs_peak(key: str) -> list[float]:
        peaks = []
        for sample in samples:
            vector = sample.get(key)
            if not vector:
                continue
            peaks.append(max(abs(float(item)) for item in vector))
        return peaks

    checks = {
        "slip_angle_rad": vector_abs_peak("slip_angle_rad"),
        "tyre_force_x_n": vector_abs_peak("tyre_force_x_n"),
        "tyre_force_y_n": vector_abs_peak("tyre_force_y_n"),
        "tyre_self_align_torque_nm": vector_abs_peak("tyre_self_align_torque_nm"),
        "brake_pressure": vector_abs_peak("brake_pressure"),
        "current_max_rpm": scalar_series("current_max_rpm"),
        "water_temp_c": scalar_series("water_temp_c"),
    }
    report: dict[str, Any] = {}
    for key, values in checks.items():
        if not values:
            report[key] = {"status": "missing", "nonzero_ratio": 0.0}
            continue
        nonzero = [value for value in values if abs(value) > 1e-6]
        ratio = len(nonzero) / len(values)
        report[key] = {
            "status": "usable" if ratio >= 0.05 else "flat_or_invalid",
            "nonzero_ratio": round(ratio, 3),
            "peak": round(max(values), 5),
        }
    return report


def detect_incidents(samples: list[dict[str, Any]]) -> list[dict[str, Any]]:
    incidents: list[dict[str, Any]] = []
    last_event_time = -math.inf
    for sample in samples:
        if sample.get("status") != "live":
            continue
        t = float(sample["captured_at_unix_s"])
        if t - last_event_time < 2.5:
            continue
        speed = float(sample["speed_kmh"])
        yaw_rate = abs(float(sample["local_angular_velocity"][1]))
        local_v = sample.get("local_velocity_mps", [0.0, 0.0, 0.0])
        forward = abs(float(local_v[2]))
        lateral = abs(float(local_v[0]))
        tyres_out = int(sample.get("number_of_tyres_out", 0))
        if speed > 35.0 and (yaw_rate > 1.0 or (forward > 8.0 and lateral / max(forward, 1.0) > 0.65)):
            incidents.append(
                {
                    "type": "spin_or_big_slide",
                    "captured_at_unix_s": t,
                    "lap_time_ms": int(sample.get("current_lap_time_ms", 0)),
                    "progress": round(float(sample.get("normalized_car_position", 0.0)), 4),
                    "speed_kmh": round(speed, 1),
                    "yaw_rate": round(yaw_rate, 3),
                    "lateral_to_forward_ratio": round(lateral / max(forward, 1.0), 3),
                }
            )
            last_event_time = t
            continue
        if speed > 40.0 and tyres_out >= 2:
            incidents.append(
                {
                    "type": "off_track",
                    "captured_at_unix_s": t,
                    "lap_time_ms": int(sample.get("current_lap_time_ms", 0)),
                    "progress": round(float(sample.get("normalized_car_position", 0.0)), 4),
                    "speed_kmh": round(speed, 1),
                    "tyres_out": tyres_out,
                }
            )
            last_event_time = t
    return incidents


def classify_segment_issue(
    slower: dict[str, float | None],
    reference: dict[str, float | None],
    *,
    lap_length_m: float,
) -> tuple[str, str]:
    slower_brake_onset = slower.get("brake_onset_pos")
    reference_brake_onset = reference.get("brake_onset_pos")
    if (
        slower_brake_onset is not None
        and reference_brake_onset is not None
        and float(reference_brake_onset) - float(slower_brake_onset) > 0.01
        and float(slower["brake_duration_ratio"]) > float(reference["brake_duration_ratio"]) + 0.08
    ):
        meters = (float(reference_brake_onset) - float(slower_brake_onset)) * lap_length_m
        return (
            "Braking too early",
            f"You are getting on the brake about {meters:.1f} m earlier and staying there longer than the reference lap.",
        )

    if (
        float(slower["entry_speed_kmh"]) > float(reference["entry_speed_kmh"]) + 4.0
        and float(slower["apex_speed_kmh"]) < float(reference["apex_speed_kmh"]) - 4.0
        and float(slower["front_slip_abs"]) > float(reference["front_slip_abs"]) + 0.01
    ):
        return (
            "Overdriving entry",
            "You carry more speed into the segment but give it back at apex with extra front slip and a lower minimum speed.",
        )

    slower_reapply = slower.get("throttle_reapply_pos")
    reference_reapply = reference.get("throttle_reapply_pos")
    if (
        slower_reapply is not None
        and reference_reapply is not None
        and float(slower_reapply) - float(reference_reapply) > 0.01
        and float(slower["exit_speed_kmh"]) < float(reference["exit_speed_kmh"]) - 3.0
    ):
        meters = (float(slower_reapply) - float(reference_reapply)) * lap_length_m
        return (
            "Late throttle on exit",
            f"The reference lap is back to power roughly {meters:.1f} m earlier, which is costing exit speed.",
        )

    if float(slower["apex_speed_kmh"]) < float(reference["apex_speed_kmh"]) - 4.0:
        return (
            "Leaving apex speed on the table",
            "Minimum speed is lower here without a compensating gain on exit. This looks like a line or confidence loss, not a setup limit.",
        )

    return (
        "Low-confidence segment",
        "This segment loses time without one dominant signature. Focus on repeatability and compare steering and throttle timing against the best lap.",
    )


def compare_laps(reference_lap: Lap, slower_lap: Lap) -> dict[str, Any]:
    grid = np.linspace(0.0, 0.995, 240)

    ref_positions, ref_times = _lap_series(reference_lap, "captured_at_unix_s")
    slow_positions, slow_times = _lap_series(slower_lap, "captured_at_unix_s")
    ref_time_grid = _safe_interp(ref_positions, ref_times - ref_times[0], grid)
    slow_time_grid = _safe_interp(slow_positions, slow_times - slow_times[0], grid)
    delta = slow_time_grid - ref_time_grid

    window = 12
    growth = delta[window:] - delta[:-window]
    ranked = np.argsort(growth)[::-1]
    used: list[tuple[int, int]] = []
    findings: list[dict[str, Any]] = []
    lap_length_m = max(reference_lap.lap_length_m, slower_lap.lap_length_m, 1.0)

    for idx in ranked:
        gain_ms = float(growth[idx] * 1000.0)
        if gain_ms < 35.0:
            continue
        start_idx = int(idx)
        end_idx = start_idx + window
        if any(not (end_idx <= used_start or start_idx >= used_end) for used_start, used_end in used):
            continue
        start_pos = float(grid[start_idx])
        end_pos = float(grid[end_idx])
        used.append((start_idx, end_idx))

        slower_metrics = segment_metrics(slower_lap, start_pos, end_pos)
        reference_metrics = segment_metrics(reference_lap, start_pos, end_pos)
        if not slower_metrics or not reference_metrics:
            continue
        line_stats = line_deviation_stats(reference_lap, slower_lap, start_pos, end_pos)
        title, detail = classify_segment_issue(
            slower_metrics,
            reference_metrics,
            lap_length_m=lap_length_m,
        )
        findings.append(
            {
                "segment_start_pct": round(start_pos * 100.0, 1),
                "segment_end_pct": round(end_pos * 100.0, 1),
                "segment_start_m": round(start_pos * lap_length_m, 1),
                "segment_end_m": round(end_pos * lap_length_m, 1),
                "time_lost_ms": round(gain_ms, 1),
                "issue": title,
                "detail": detail if line_stats["mean_m"] < 2.5 else f"{detail} Average line deviation here is {line_stats['mean_m']:.1f} m.",
                "slower_metrics": slower_metrics,
                "reference_metrics": reference_metrics,
                "line_deviation_m": {
                    "mean": round(line_stats["mean_m"], 2),
                    "max": round(line_stats["max_m"], 2),
                },
            }
        )
        if len(findings) >= 3:
            break

    return {
        "reference_lap_number": reference_lap.lap_number,
        "reference_lap_time_ms": reference_lap.lap_time_ms,
        "slower_lap_number": slower_lap.lap_number,
        "slower_lap_time_ms": slower_lap.lap_time_ms,
        "delta_ms": slower_lap.lap_time_ms - reference_lap.lap_time_ms,
        "findings": findings,
    }


def analyze_session_dir(session_dir: Path) -> dict[str, Any]:
    meta, samples = load_session(session_dir)
    laps = segment_laps(samples, started_mid_lap=bool(meta.get("started_mid_lap")))
    if bool(meta.get("started_mid_lap")) and len(laps) < 2:
        override_laps = segment_laps(samples, started_mid_lap=False)
        if len(override_laps) > len(laps):
            laps = override_laps
    lap_times = [lap.lap_time_ms for lap in laps]

    ctelemetry_summary = None
    for artifact_path in (meta.get("artifacts") or {}).values():
        if not str(artifact_path).lower().endswith(".tc"):
            continue
        try:
            parsed_tc = parse_tc_file(Path(artifact_path))
            ctelemetry_summary = {
                "source_path": parsed_tc["source_path"],
                "lap_time_ms": parsed_tc["lap_time_ms"],
                "num_data_points": parsed_tc["num_data_points"],
                "track": parsed_tc["track"],
                "track_layout": parsed_tc["track_layout"],
                "car_model": parsed_tc["car_model"],
            }
            break
        except Exception:
            continue

    response: dict[str, Any] = {
        "session_id": meta["session_id"],
        "sample_count": meta.get("sample_count", len(samples)),
        "lap_count": len(laps),
        "telemetry_quality": telemetry_quality(samples),
        "incidents": detect_incidents(samples),
        "ctelemetry_reference": ctelemetry_summary,
        "laps": [
            {
                "lap_number": lap.lap_number,
                "lap_time_ms": lap.lap_time_ms,
                "lap_length_m": round(lap.lap_length_m, 1),
                "sector_times_ms": lap.sector_times_ms,
                "max_speed_kmh": round(max(float(s["speed_kmh"]) for s in lap.samples), 1),
                "avg_throttle": round(_mean([float(s["throttle"]) for s in lap.samples]), 3),
                "avg_brake": round(_mean([float(s["brake"]) for s in lap.samples]), 3),
                "max_wheel_slip": round(max(float(s.get("wheel_slip_peak", 0.0)) for s in lap.samples), 3),
            }
            for lap in laps
        ],
        "consistency_std_ms": round(_stddev(lap_times), 1),
        "best_lap_number": None,
        "best_lap_time_ms": None,
        "focus_comparison": None,
        "overall_notes": [],
    }

    if not laps:
        response["overall_notes"].append(
            "No complete laps were segmented from this capture. Start the recorder before the start/finish line or record more laps."
        )
        return response

    best_lap = min(laps, key=lambda lap: lap.lap_time_ms)
    response["best_lap_number"] = best_lap.lap_number
    response["best_lap_time_ms"] = best_lap.lap_time_ms

    if len(laps) >= 2:
        latest_non_best = laps[-1]
        if latest_non_best.lap_number == best_lap.lap_number:
            candidates = [lap for lap in laps if lap.lap_number != best_lap.lap_number]
            if candidates:
                latest_non_best = candidates[-1]
        if latest_non_best.lap_number != best_lap.lap_number:
            response["focus_comparison"] = compare_laps(best_lap, latest_non_best)

    if len(laps) >= 3 and _stddev(lap_times) > 700.0:
        response["overall_notes"].append(
            "Lap-time spread is large enough that consistency is probably a bigger lever than setup changes right now."
        )
    invalid_channels = [
        key for key, value in response["telemetry_quality"].items() if value["status"] != "usable"
    ]
    if invalid_channels:
        response["overall_notes"].append(
            "Some advanced shared-memory channels are flat or invalid on this build: " + ", ".join(invalid_channels) + ". Line, speed, load and motion channels are still usable."
        )
    if response["incidents"]:
        response["overall_notes"].append(
            f"Detected {len(response['incidents'])} instability/off-track event(s) in the session."
        )
    if response["focus_comparison"] is None:
        response["overall_notes"].append(
            "Record at least two complete laps in one capture to get corner-by-corner coaching."
        )
    return response
