from __future__ import annotations

import json
import os
import subprocess
import sys
import time
from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Any

import numpy as np

from .artifacts import build_session_paths, copy_session_artifacts, write_json
from .analysis import Lap, classify_segment_issue, load_session, segment_laps, segment_metrics
from .ctelemetry import find_matching_tc_file, tc_as_lap
from .paths import EnvironmentPaths, read_race_context
from .shared_memory import AssettoCorsaReader


try:  # pragma: no cover
    import pythoncom
    import win32com.client
except Exception:  # pragma: no cover
    pythoncom = None
    win32com = None


DETACHED_FLAGS = subprocess.CREATE_NEW_PROCESS_GROUP | subprocess.DETACHED_PROCESS


@dataclass(slots=True)
class LiveCoachConfig:
    car_model: str | None = None
    track: str | None = None
    track_layout: str | None = None
    mode: str = "consistency"
    reference_session_id: str | None = None
    sample_hz: float = 20.0
    tts_enabled: bool = True
    muted: bool = False
    max_callouts_per_lap: int = 3
    zone_count: int = 4
    setup_profile: str | None = None
    plan_label: str | None = None

    def to_dict(self) -> dict[str, Any]:
        return asdict(self)


def _runtime_files(paths: EnvironmentPaths) -> dict[str, Path]:
    return {
        "config": paths.runtime_root / "live_coach_config.json",
        "status": paths.runtime_root / "live_coach_status.json",
        "events": paths.runtime_root / "live_coach_events.jsonl",
        "pid": paths.runtime_root / "live_coach.pid",
        "stdout": paths.runtime_root / "live_coach_stdout.log",
    }


def _write_json(path: Path, payload: dict[str, Any]) -> None:
    path.parent.mkdir(parents=True, exist_ok=True)
    path.write_text(json.dumps(payload, indent=2), encoding="utf-8")


def _read_json(path: Path) -> dict[str, Any] | None:
    if not path.exists():
        return None
    try:
        return json.loads(path.read_text(encoding="utf-8"))
    except Exception:
        return None


def _process_running(pid: int) -> bool:
    if pid <= 0:
        return False
    result = subprocess.run(
        ["tasklist", "/FI", f"PID eq {pid}", "/FO", "CSV", "/NH"],
        capture_output=True,
        text=True,
        check=False,
    )
    return str(pid) in result.stdout


def load_live_coach_status(paths: EnvironmentPaths) -> dict[str, Any]:
    status = _read_json(_runtime_files(paths)["status"])
    if status is None:
        return {"state": "stopped"}
    pid = int(status.get("pid", 0) or 0)
    status["process_alive"] = _process_running(pid) if pid else False
    return status


def start_live_coach_background(paths: EnvironmentPaths, config: LiveCoachConfig) -> dict[str, Any]:
    files = _runtime_files(paths)
    paths.runtime_root.mkdir(parents=True, exist_ok=True)
    existing = load_live_coach_status(paths)
    if existing.get("process_alive"):
        raise RuntimeError("Live coach is already running")
    _write_json(files["config"], config.to_dict())
    files["events"].write_text("", encoding="utf-8")
    with files["stdout"].open("a", encoding="utf-8") as handle:
        process = subprocess.Popen(
            [sys.executable, "-m", "ac_coach_mcp.cli", "run-live-coach", "--config", str(files["config"])],
            cwd=paths.project_root,
            stdout=handle,
            stderr=handle,
            stdin=subprocess.DEVNULL,
            creationflags=DETACHED_FLAGS,
            close_fds=False,
        )
    files["pid"].write_text(str(process.pid), encoding="utf-8")
    status = {
        "state": "starting",
        "pid": process.pid,
        "config": config.to_dict(),
        "started_at_unix_s": time.time(),
        "status_path": str(files["status"]),
        "events_path": str(files["events"]),
    }
    _write_json(files["status"], status)
    return status


def stop_live_coach_background(paths: EnvironmentPaths) -> dict[str, Any]:
    files = _runtime_files(paths)
    status = load_live_coach_status(paths)
    pid = int(status.get("pid", 0) or 0)
    if not pid or not _process_running(pid):
        status.update({"state": "stopped", "process_alive": False})
        _write_json(files["status"], status)
        return status
    subprocess.run(["taskkill", "/PID", str(pid), "/T", "/F"], check=False, capture_output=True, text=True)
    status.update({"state": "stopped", "stopped_at_unix_s": time.time(), "process_alive": False})
    _write_json(files["status"], status)
    return status


class Speaker:
    def __init__(self, enabled: bool) -> None:
        self.enabled = enabled and pythoncom is not None and win32com is not None
        self.voice = None
        if self.enabled:  # pragma: no branch
            pythoncom.CoInitialize()
            self.voice = win32com.client.Dispatch("SAPI.SpVoice")
            self.voice.Rate = -1
            self.voice.Volume = 100

    def speak(self, text: str) -> None:
        if not text:
            return
        if self.enabled and self.voice is not None:  # pragma: no branch
            try:
                self.voice.Speak(text, 1)
                return
            except Exception:
                self.enabled = False
        try:
            escaped = text.replace("'", "''")
            subprocess.Popen(
                [
                    "powershell.exe",
                    "-NoProfile",
                    "-ExecutionPolicy",
                    "Bypass",
                    "-Command",
                    (
                        "Add-Type -AssemblyName System.Speech; "
                        "$speak = New-Object System.Speech.Synthesis.SpeechSynthesizer; "
                        f"$speak.Speak('{escaped}')"
                    ),
                ],
                creationflags=DETACHED_FLAGS,
                stdout=subprocess.DEVNULL,
                stderr=subprocess.DEVNULL,
                stdin=subprocess.DEVNULL,
            )
            return
        except Exception:
            pass
        print(text, flush=True)


def _latest_matching_reference(
    paths: EnvironmentPaths,
    *,
    car_model: str | None,
    track: str | None,
    track_layout: str | None,
    reference_session_id: str | None,
) -> tuple[str | None, Lap | None]:
    candidates = []
    if not paths.sessions_root.exists():
        return None, None
    for session_dir in paths.sessions_root.iterdir():
        meta_path = session_dir / "meta.json"
        if not meta_path.exists():
            continue
        meta = _read_json(meta_path) or {}
        initial = meta.get("initial_snapshot", {})
        race_context = meta.get("race_context", {})
        if reference_session_id and meta.get("session_id") != reference_session_id:
            continue
        if car_model and initial.get("car_model") != car_model:
            continue
        if track and initial.get("track") != track:
            continue
        if track_layout and race_context.get("track_config") != track_layout:
            continue
        samples_path = session_dir / "samples.jsonl"
        if not samples_path.exists():
            continue
        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:
            laps = segment_laps(samples, started_mid_lap=False)
        if not laps:
            continue
        best_lap = min(laps, key=lambda lap: lap.lap_time_ms)
        candidates.append((session_dir.stat().st_mtime, meta.get("session_id"), best_lap))
        if reference_session_id:
            break
    if not candidates:
        ctelemetry_root = paths.documents_root / "ctelemetry" / "player"
        tc_path = find_matching_tc_file(
            ctelemetry_root,
            track=track or "",
            track_layout=track_layout,
            car_model=car_model,
        ) if track else None
        if tc_path is not None:
            try:
                return "ac-ctelemetry", tc_as_lap(tc_path)
            except Exception:
                pass
        return None, None
    _, session_id, best_lap = max(candidates, key=lambda item: item[0])
    return session_id, best_lap


def _contiguous_regions(mask: np.ndarray) -> list[tuple[int, int]]:
    regions: list[tuple[int, int]] = []
    start = None
    for idx, value in enumerate(mask):
        if value and start is None:
            start = idx
        elif not value and start is not None:
            regions.append((start, idx - 1))
            start = None
    if start is not None:
        regions.append((start, len(mask) - 1))
    return regions


def derive_reference_zones(reference_lap: Lap, *, max_zones: int = 4) -> list[dict[str, Any]]:
    positions = np.array(
        [
            max(0.0, (float(sample["distance_traveled_m"]) - reference_lap.start_distance_m) / max(reference_lap.lap_length_m, 1.0))
            for sample in reference_lap.samples
        ],
        dtype=float,
    )
    brake = np.array([float(sample["brake"]) for sample in reference_lap.samples], dtype=float)
    speed = np.array([float(sample["speed_kmh"]) for sample in reference_lap.samples], dtype=float)
    throttle = np.array([float(sample["throttle"]) for sample in reference_lap.samples], dtype=float)
    brake_regions = _contiguous_regions(brake > 0.12)

    zones: list[dict[str, Any]] = []
    for start_idx, end_idx in brake_regions:
        if end_idx - start_idx < 20:
            continue
        search_end = min(len(speed) - 1, end_idx + 150)
        apex_local = int(np.argmin(speed[start_idx : search_end + 1]))
        apex_idx = start_idx + apex_local
        zone_start = max(0.0, float(positions[start_idx]) - 0.01)
        zone_end = min(0.99, float(positions[min(search_end, apex_idx + 80)]))
        metrics = segment_metrics(reference_lap, zone_start, zone_end)
        if not metrics:
            continue
        priority = (
            float(np.max(brake[start_idx : end_idx + 1])) * 100.0
            + max(float(speed[start_idx]) - float(speed[apex_idx]), 0.0)
            + float(np.mean(throttle[start_idx : min(search_end, apex_idx + 80) + 1])) * -10.0
        )
        zones.append(
            {
                "label": f"zone_{len(zones) + 1}",
                "start_pos": round(zone_start, 4),
                "end_pos": round(zone_end, 4),
                "priority": round(priority, 2),
                "reference_metrics": metrics,
            }
        )

    if not zones:
        thirds = np.linspace(0.05, 0.85, min(max_zones, 3))
        for idx, start in enumerate(thirds, start=1):
            zone_start = max(0.0, float(start) - 0.04)
            zone_end = min(0.99, float(start) + 0.06)
            metrics = segment_metrics(reference_lap, zone_start, zone_end)
            if metrics:
                zones.append(
                    {
                        "label": f"zone_{idx}",
                        "start_pos": round(zone_start, 4),
                        "end_pos": round(zone_end, 4),
                        "priority": 1.0,
                        "reference_metrics": metrics,
                    }
                )

    zones.sort(key=lambda item: item["priority"], reverse=True)
    return zones[:max_zones]


def _progress(snapshot: dict[str, Any], lap_start_distance_m: float, reference_lap_length_m: float) -> float:
    return max(
        0.0,
        min(
            0.999,
            (float(snapshot["distance_traveled_m"]) - lap_start_distance_m) / max(reference_lap_length_m, 1.0),
        ),
    )


def _pseudo_lap(
    lap_number: int,
    samples: list[dict[str, Any]],
    lap_start_distance_m: float,
    reference_lap_length_m: float,
) -> Lap:
    last_time_ms = int(samples[-1].get("current_lap_time_ms", 0)) if samples else 0
    return Lap(
        lap_number=lap_number,
        lap_time_ms=last_time_ms,
        lap_time_s=last_time_ms / 1000.0,
        lap_length_m=reference_lap_length_m,
        start_distance_m=lap_start_distance_m,
        sector_times_ms=[],
        samples=samples[:],
    )


def _zone_message(zone_index: int, issue: str) -> str:
    if issue == "Braking too early":
        return f"Zone {zone_index}: brake later."
    if issue == "Late throttle on exit":
        return f"Zone {zone_index}: throttle earlier on exit."
    if issue == "Overdriving entry":
        return f"Zone {zone_index}: less entry speed. Clean the turn in."
    if issue == "Leaving apex speed on the table":
        return f"Zone {zone_index}: carry more minimum speed."
    return f"Zone {zone_index}: repeat that cleaner."


def _evaluate_zone(mode: str, current_metrics: dict[str, float | None], reference_metrics: dict[str, float | None], lap_length_m: float) -> tuple[bool, str | None, dict[str, Any]]:
    brake_gap = None
    if current_metrics.get("brake_onset_pos") is not None and reference_metrics.get("brake_onset_pos") is not None:
        brake_gap = float(reference_metrics["brake_onset_pos"]) - float(current_metrics["brake_onset_pos"])
    throttle_gap = None
    if current_metrics.get("throttle_reapply_pos") is not None and reference_metrics.get("throttle_reapply_pos") is not None:
        throttle_gap = float(current_metrics["throttle_reapply_pos"]) - float(reference_metrics["throttle_reapply_pos"])
    apex_gap = float(reference_metrics["apex_speed_kmh"]) - float(current_metrics["apex_speed_kmh"])
    exit_gap = float(reference_metrics["exit_speed_kmh"]) - float(current_metrics["exit_speed_kmh"])

    pass_zone = True
    if mode == "braking":
        pass_zone = not ((brake_gap is not None and brake_gap > 0.01) or apex_gap > 5.0)
    elif mode == "exit":
        pass_zone = not ((throttle_gap is not None and throttle_gap > 0.01) or exit_gap > 4.0)
    elif mode == "smoothness":
        pass_zone = not (apex_gap > 4.0 or current_metrics["steer_peak_deg"] > reference_metrics["steer_peak_deg"] * 1.5)
    else:
        pass_zone = not (
            apex_gap > 4.0
            or exit_gap > 4.0
            or (brake_gap is not None and brake_gap > 0.01)
            or (throttle_gap is not None and throttle_gap > 0.01)
        )

    if pass_zone:
        return True, None, {
            "apex_gap_kmh": round(apex_gap, 2),
            "exit_gap_kmh": round(exit_gap, 2),
            "brake_gap_progress": round(brake_gap, 4) if brake_gap is not None else None,
            "throttle_gap_progress": round(throttle_gap, 4) if throttle_gap is not None else None,
        }
    issue, detail = classify_segment_issue(current_metrics, reference_metrics, lap_length_m=lap_length_m)
    return False, f"{issue}: {detail}", {
        "issue": issue,
        "detail": detail,
        "apex_gap_kmh": round(apex_gap, 2),
        "exit_gap_kmh": round(exit_gap, 2),
        "brake_gap_progress": round(brake_gap, 4) if brake_gap is not None else None,
        "throttle_gap_progress": round(throttle_gap, 4) if throttle_gap is not None else None,
    }


def _log_event(events_path: Path, event: dict[str, Any]) -> None:
    events_path.parent.mkdir(parents=True, exist_ok=True)
    with events_path.open("a", encoding="utf-8") as handle:
        handle.write(json.dumps(event, separators=(",", ":")) + "\n")


def run_live_coach(paths: EnvironmentPaths, config: LiveCoachConfig) -> None:
    files = _runtime_files(paths)
    paths.runtime_root.mkdir(parents=True, exist_ok=True)
    pid = os.getpid()
    files["pid"].write_text(str(pid), encoding="utf-8")
    speaker = Speaker(enabled=config.tts_enabled and not config.muted)
    reader = AssettoCorsaReader()
    session_id: str | None = None
    session_dir: Path | None = None
    session_samples_path: Path | None = None
    session_meta_path: Path | None = None
    session_events_path: Path | None = None
    session_handle = None
    session_sample_count = 0
    session_started_at: float | None = None
    session_initial_snapshot: dict[str, Any] | None = None

    def update_status(payload: dict[str, Any]) -> None:
        state = {
            "pid": pid,
            "process_alive": True,
            "updated_at_unix_s": time.time(),
            "config": config.to_dict(),
        }
        state.update(payload)
        _write_json(files["status"], state)

    speaker.speak("Live coach armed.")
    reference_session_id, reference_lap = _latest_matching_reference(
        paths,
        car_model=config.car_model,
        track=config.track,
        track_layout=config.track_layout,
        reference_session_id=config.reference_session_id,
    )
    zones: list[dict[str, Any]] = derive_reference_zones(reference_lap, max_zones=config.zone_count) if reference_lap else []
    if reference_lap:
        speaker.speak("Reference loaded. Coach is live.")
    else:
        speaker.speak("No reference session found. First two complete laps will seed the reference.")

    update_status(
        {
            "state": "waiting_for_live" if reference_lap is None else "running",
            "reference_session_id": reference_session_id,
            "reference_lap_time_ms": reference_lap.lap_time_ms if reference_lap else None,
            "zones": zones,
        }
    )

    active_lap_samples: list[dict[str, Any]] = []
    lap_start_distance_m: float | None = None
    current_lap_number = 0
    evaluated_zones: set[int] = set()
    lap_misses = 0
    callouts_this_lap = 0
    consecutive_off = 0
    seen_live = False
    pending_reference_laps: list[Lap] = []

    while True:
        snapshot = reader.read_snapshot()
        if snapshot is None:
            update_status({"state": "waiting_for_shared_memory"})
            time.sleep(0.25)
            continue
        if snapshot.get("status") != "live":
            consecutive_off += 1
            update_status({"state": "waiting_for_live", "last_snapshot": snapshot})
            if seen_live and consecutive_off >= max(8, int(config.sample_hz * 2.0)):
                break
            time.sleep(1.0 / config.sample_hz)
            continue
        consecutive_off = 0
        seen_live = True

        race_match = True
        if config.car_model and snapshot.get("car_model") != config.car_model:
            race_match = False
        if config.track and snapshot.get("track") != config.track:
            race_match = False
        if not race_match:
            update_status({"state": "waiting_for_target_combo", "last_snapshot": snapshot})
            time.sleep(1.0 / config.sample_hz)
            continue

        if session_dir is None:
            session_id, session_dir, session_samples_path, session_meta_path = build_session_paths(
                paths,
                snapshot,
                config.plan_label or "live-coach",
            )
            session_events_path = session_dir / "events.jsonl"
            session_handle = session_samples_path.open("w", encoding="utf-8")
            session_started_at = time.time()
            session_initial_snapshot = snapshot
        if session_handle is not None:
            session_handle.write(json.dumps(snapshot, separators=(",", ":")) + "\n")
            session_sample_count += 1

        if lap_start_distance_m is None:
            lap_start_distance_m = float(snapshot["distance_traveled_m"])
            current_lap_number = int(snapshot["completed_laps"]) + 1
            active_lap_samples = [snapshot]
            evaluated_zones = set()
            lap_misses = 0
            callouts_this_lap = 0
            update_status({"state": "running", "lap_number": current_lap_number, "last_snapshot": snapshot})
            time.sleep(1.0 / config.sample_hz)
            continue

        active_lap_samples.append(snapshot)
        if reference_lap and zones:
            progress = _progress(snapshot, lap_start_distance_m, reference_lap.lap_length_m)
            for idx, zone in enumerate(zones):
                if idx in evaluated_zones or progress < float(zone["end_pos"]):
                    continue
                evaluated_zones.add(idx)
                live_lap = _pseudo_lap(current_lap_number, active_lap_samples, lap_start_distance_m, reference_lap.lap_length_m)
                current_metrics = segment_metrics(live_lap, float(zone["start_pos"]), float(zone["end_pos"]))
                if not current_metrics:
                    continue
                passed, detail, telemetry = _evaluate_zone(
                    config.mode,
                    current_metrics,
                    zone["reference_metrics"],
                    reference_lap.lap_length_m,
                )
                event = {
                    "captured_at_unix_s": time.time(),
                    "type": "zone_feedback",
                    "lap_number": current_lap_number,
                    "zone_index": idx + 1,
                    "passed": passed,
                    "detail": detail,
                    "telemetry": telemetry,
                }
                _log_event(files["events"], event)
                if session_events_path is not None:
                    _log_event(session_events_path, event)
                if not passed:
                    lap_misses += 1
                    if callouts_this_lap < config.max_callouts_per_lap:
                        message = _zone_message(idx + 1, telemetry.get("issue", "zone"))
                        speaker.speak(message)
                        callouts_this_lap += 1

        if len(active_lap_samples) >= 2 and int(snapshot["completed_laps"]) > int(active_lap_samples[-2]["completed_laps"]):
            boundary = snapshot
            lap_time_ms = int(boundary["last_lap_time_ms"])
            lap_samples = active_lap_samples[:-1]
            lap_length = max(0.0, float(lap_samples[-1]["distance_traveled_m"]) - lap_start_distance_m)
            finished_lap = Lap(
                lap_number=current_lap_number,
                lap_time_ms=lap_time_ms,
                lap_time_s=lap_time_ms / 1000.0,
                lap_length_m=lap_length if lap_length > 1.0 else (reference_lap.lap_length_m if reference_lap else 1.0),
                start_distance_m=lap_start_distance_m,
                sector_times_ms=[],
                samples=lap_samples,
            )
            if reference_lap is None and lap_time_ms > 0:
                pending_reference_laps.append(finished_lap)
                if len(pending_reference_laps) >= 2:
                    reference_lap = min(pending_reference_laps, key=lambda lap: lap.lap_time_ms)
                    zones = derive_reference_zones(reference_lap, max_zones=config.zone_count)
                    reference_session_id = "live-self-reference"
                    speaker.speak("Reference lap locked. Coaching starts now.")
                else:
                    speaker.speak("Reference still seeding. Complete one more lap.")
            else:
                if lap_misses == 0:
                    speaker.speak(f"Lap {current_lap_number} passed.")
                else:
                    speaker.speak(f"Lap {current_lap_number} failed. {lap_misses} missed zones.")
            _log_event(
                files["events"],
                {
                    "captured_at_unix_s": time.time(),
                    "type": "lap_summary",
                    "lap_number": current_lap_number,
                    "lap_time_ms": lap_time_ms,
                    "missed_zones": lap_misses,
                },
            )
            if session_events_path is not None:
                _log_event(
                    session_events_path,
                    {
                        "captured_at_unix_s": time.time(),
                        "type": "lap_summary",
                        "lap_number": current_lap_number,
                        "lap_time_ms": lap_time_ms,
                        "missed_zones": lap_misses,
                    },
                )
            current_lap_number += 1
            lap_start_distance_m = float(boundary["distance_traveled_m"])
            active_lap_samples = [boundary]
            evaluated_zones = set()
            lap_misses = 0
            callouts_this_lap = 0

        update_status(
            {
                "state": "running",
                "reference_session_id": reference_session_id,
                "reference_lap_time_ms": reference_lap.lap_time_ms if reference_lap else None,
                "lap_number": current_lap_number,
                "missed_zones_this_lap": lap_misses,
                "zones": zones,
                "last_snapshot": snapshot,
            }
        )
        time.sleep(1.0 / config.sample_hz)

    if session_handle is not None:
        session_handle.close()
    if session_dir is not None and session_meta_path is not None and session_started_at is not None:
        meta = {
            "session_id": session_id,
            "created_at_unix_s": session_started_at,
            "capture_finished_at_unix_s": time.time(),
            "sample_hz": config.sample_hz,
            "capture_mode": "live_coach",
            "requested_lap_count": None,
            "requested_duration_s": None,
            "label": config.plan_label or "live-coach",
            "paths": paths.to_dict(),
            "race_context": read_race_context(paths),
            "initial_snapshot": session_initial_snapshot,
            "final_snapshot": snapshot if snapshot is not None else None,
            "sample_count": session_sample_count,
            "events_path": str(session_events_path) if session_events_path is not None else None,
            "started_mid_lap": bool(session_initial_snapshot and float(session_initial_snapshot.get("normalized_car_position", 0.0)) > 0.15),
            "reference_session_id": reference_session_id,
            "reference_lap_time_ms": reference_lap.lap_time_ms if reference_lap else None,
            "live_coach_config": config.to_dict(),
            "artifacts": copy_session_artifacts(paths, session_dir, since_unix_s=session_started_at),
        }
        write_json(session_meta_path, meta)
    update_status({"state": "stopped", "stopped_at_unix_s": time.time(), "process_alive": False, "session_id": session_id})
