from __future__ import annotations

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

from .artifacts import build_session_paths, copy_session_artifacts, write_json
from .paths import EnvironmentPaths, read_race_context
from .shared_memory import AssettoCorsaReader


@dataclass(slots=True)
class CaptureResult:
    session_id: str
    session_dir: Path
    samples_path: Path
    meta_path: Path
    sample_count: int
    start_lap_counter: int
    end_lap_counter: int
    started_mid_lap: bool

    def to_dict(self) -> dict[str, Any]:
        return {
            "session_id": self.session_id,
            "session_dir": str(self.session_dir),
            "samples_path": str(self.samples_path),
            "meta_path": str(self.meta_path),
            "sample_count": self.sample_count,
            "start_lap_counter": self.start_lap_counter,
            "end_lap_counter": self.end_lap_counter,
            "started_mid_lap": self.started_mid_lap,
        }


class SessionRecorder:
    def __init__(self, paths: EnvironmentPaths, reader: AssettoCorsaReader) -> None:
        self.paths = paths
        self.reader = reader

    def _wait_for_snapshot(self, timeout_s: float) -> dict[str, Any]:
        deadline = time.monotonic() + timeout_s
        while time.monotonic() < deadline:
            snapshot = self.reader.read_snapshot()
            if snapshot is not None and snapshot.get("status") == "live":
                return snapshot
            time.sleep(0.25)
        raise RuntimeError("Assetto Corsa shared memory is not available")

    def capture(
        self,
        *,
        duration_s: float | None = None,
        lap_count: int | None = None,
        sample_hz: float = 50.0,
        max_wait_s: float = 30.0,
        max_duration_s: float | None = None,
        label: str | None = None,
    ) -> CaptureResult:
        if duration_s is None and lap_count is None:
            raise ValueError("capture requires duration_s or lap_count")
        if sample_hz <= 0:
            raise ValueError("sample_hz must be positive")

        self.paths.sessions_root.mkdir(parents=True, exist_ok=True)
        initial = self._wait_for_snapshot(timeout_s=max_wait_s)
        session_id, session_dir, samples_path, meta_path = build_session_paths(self.paths, initial, label)
        race_context = read_race_context(self.paths)

        started_mid_lap = float(initial["normalized_car_position"]) > 0.15
        armed = not started_mid_lap
        lap_baseline = int(initial["completed_laps"])
        capture_started_monotonic = time.monotonic()
        stop_deadline = (
            capture_started_monotonic + max_duration_s
            if max_duration_s is not None
            else math.inf
        )
        if duration_s is not None:
            stop_deadline = min(stop_deadline, capture_started_monotonic + duration_s)

        meta: dict[str, Any] = {
            "session_id": session_id,
            "created_at_unix_s": time.time(),
            "sample_hz": float(sample_hz),
            "capture_mode": "laps" if lap_count is not None else "duration",
            "requested_duration_s": duration_s,
            "requested_lap_count": lap_count,
            "label": label,
            "paths": self.paths.to_dict(),
            "race_context": race_context,
            "initial_snapshot": initial,
            "started_mid_lap": started_mid_lap,
        }

        sample_count = 0
        next_tick = time.monotonic()
        last_snapshot = initial
        consecutive_misses = 0
        consecutive_off = 0

        with samples_path.open("w", encoding="utf-8") as handle:
            while time.monotonic() <= stop_deadline:
                snapshot = self.reader.read_snapshot()
                if snapshot is None:
                    consecutive_misses += 1
                    if sample_count > 0 and consecutive_misses >= 25:
                        break
                else:
                    consecutive_misses = 0
                    if snapshot.get("status") == "off":
                        consecutive_off += 1
                        if sample_count > 0 and consecutive_off >= max(5, int(sample_hz * 2.0)):
                            break
                        next_tick += 1.0 / sample_hz
                        time.sleep(max(0.0, next_tick - time.monotonic()))
                        continue
                    consecutive_off = 0
                    if lap_count is not None and not armed:
                        if int(snapshot["completed_laps"]) > lap_baseline:
                            lap_baseline = int(snapshot["completed_laps"])
                            armed = True
                    handle.write(json.dumps(snapshot, separators=(",", ":")) + "\n")
                    sample_count += 1
                    last_snapshot = snapshot
                    if (
                        lap_count is not None
                        and armed
                        and int(snapshot["completed_laps"]) >= lap_baseline + lap_count
                    ):
                        break

                next_tick += 1.0 / sample_hz
                time.sleep(max(0.0, next_tick - time.monotonic()))

        meta.update(
            {
                "capture_finished_at_unix_s": time.time(),
                "sample_count": sample_count,
                "start_lap_counter": int(initial["completed_laps"]),
                "end_lap_counter": int(last_snapshot["completed_laps"]),
                "final_snapshot": last_snapshot,
                "samples_path": str(samples_path),
            }
        )
        meta["artifacts"] = copy_session_artifacts(
            self.paths,
            session_dir,
            since_unix_s=meta["created_at_unix_s"],
        )
        write_json(meta_path, meta)

        return CaptureResult(
            session_id=session_id,
            session_dir=session_dir,
            samples_path=samples_path,
            meta_path=meta_path,
            sample_count=sample_count,
            start_lap_counter=int(initial["completed_laps"]),
            end_lap_counter=int(last_snapshot["completed_laps"]),
            started_mid_lap=started_mid_lap,
        )
