from __future__ import annotations

from ac_coach_mcp.analysis import Lap, compare_laps, detect_incidents, telemetry_quality


def _sample(
    pos: float,
    t: float,
    speed: float,
    brake: float,
    throttle: float,
    steer: float,
    slip: float,
) -> dict:
    return {
        "normalized_car_position": pos,
        "captured_at_unix_s": t,
        "speed_kmh": speed,
        "brake": brake,
        "throttle": throttle,
        "steer_angle_deg": steer,
        "slip_angle_rad": [slip, slip, slip, slip],
        "tyre_force_x_n": [0.0, 0.0, 0.0, 0.0],
        "tyre_force_y_n": [0.0, 0.0, 0.0, 0.0],
        "tyre_self_align_torque_nm": [0.0, 0.0, 0.0, 0.0],
        "brake_pressure": [0.0, 0.0, 0.0, 0.0],
        "current_max_rpm": None,
        "water_temp_c": None,
        "local_angular_velocity": [0.0, 0.0, 0.0],
        "local_velocity_mps": [0.0, 0.0, speed / 3.6],
        "number_of_tyres_out": 0,
        "status": "live",
        "final_ff": 0.5,
    }


def _build_reference_lap() -> Lap:
    positions = [0.00, 0.08, 0.12, 0.16, 0.20, 0.24, 0.28, 0.32, 0.36, 0.40, 0.44, 0.48, 0.52, 0.56]
    times = [i * 0.45 for i in range(len(positions))]
    samples = []
    for pos, t in zip(positions, times, strict=True):
        if pos < 0.16:
            samples.append(_sample(pos, t, 160 - pos * 120, 0.0, 1.0, 4.0, 0.03))
        elif pos < 0.30:
            samples.append(_sample(pos, t, 120 - (pos - 0.16) * 250, 0.55, 0.08, 10.0, 0.05))
        else:
            samples.append(_sample(pos, t, 88 + (pos - 0.30) * 170, 0.0, 0.85, 7.0, 0.04))
    return Lap(
        lap_number=1,
        lap_time_ms=60000,
        lap_time_s=60.0,
        lap_length_m=3000.0,
        start_distance_m=0.0,
        sector_times_ms=[20000, 20000, 20000],
        samples=samples,
    )


def _build_slower_lap() -> Lap:
    positions = [0.00, 0.08, 0.12, 0.16, 0.20, 0.24, 0.28, 0.32, 0.36, 0.40, 0.44, 0.48, 0.52, 0.56]
    times = [0.0, 0.47, 0.95, 1.45, 1.97, 2.55, 3.10, 3.68, 4.20, 4.72, 5.18, 5.64, 6.09, 6.55]
    samples = []
    for pos, t in zip(positions, times, strict=True):
        if pos < 0.12:
            samples.append(_sample(pos, t, 160 - pos * 110, 0.0, 1.0, 4.0, 0.03))
        elif pos < 0.32:
            samples.append(_sample(pos, t, 118 - (pos - 0.12) * 260, 0.72, 0.03, 11.5, 0.07))
        else:
            throttle = 0.18 if pos < 0.40 else 0.65
            samples.append(_sample(pos, t, 82 + (pos - 0.32) * 145, 0.0, throttle, 8.0, 0.05))
    return Lap(
        lap_number=2,
        lap_time_ms=62300,
        lap_time_s=62.3,
        lap_length_m=3000.0,
        start_distance_m=0.0,
        sector_times_ms=[20800, 20700, 20800],
        samples=samples,
    )


def test_compare_laps_finds_actionable_issue() -> None:
    comparison = compare_laps(_build_reference_lap(), _build_slower_lap())
    assert comparison["delta_ms"] == 2300
    assert comparison["findings"]
    assert any(
        finding["issue"] in {
            "Braking too early",
            "Late throttle on exit",
            "Overdriving entry",
            "Leaving apex speed on the table",
        }
        for finding in comparison["findings"]
    )


def test_quality_and_incident_detection_work() -> None:
    samples = [
        {
            **_sample(0.1, 1.0, 110.0, 0.0, 1.0, 2.0, 0.0),
            "local_angular_velocity": [0.0, 1.3, 0.0],
            "local_velocity_mps": [8.0, 0.0, 10.0],
            "number_of_tyres_out": 2,
        },
        {
            **_sample(0.2, 4.0, 90.0, 0.0, 0.8, 3.0, 0.0),
            "slip_angle_rad": [0.0, 0.1, 0.05, 0.02],
            "tyre_force_y_n": [10.0, 12.0, 8.0, 7.0],
            "current_max_rpm": 7200.0,
        },
    ]
    quality = telemetry_quality(samples)
    assert quality["slip_angle_rad"]["status"] == "usable"
    assert quality["tyre_force_y_n"]["status"] == "usable"
    incidents = detect_incidents(samples)
    assert incidents
