# Copyright 2025 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
#     http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

"""Highlight Reel strategy: on-demand recording via ring buffer."""

from __future__ import annotations

import contextlib
import logging
import time
from concurrent.futures import Future, ThreadPoolExecutor
from threading import Event as ThreadingEvent, Lock

from lerobot.datasets import LeRobotDataset, VideoEncodingManager
from lerobot.utils.constants import ACTION, OBS_STR
from lerobot.utils.cycle_timer import CycleTimer
from lerobot.utils.feature_utils import build_dataset_frame
from lerobot.utils.keyboard_input import create_key_listener
from lerobot.utils.utils import log_say

from ..configs import HighlightStrategyConfig, RolloutConfig
from ..context import RolloutContext
from ..ring_buffer import RolloutRingBuffer
from .core import RolloutStrategy, safe_push_to_hub, send_next_action

logger = logging.getLogger(__name__)


class HighlightStrategy(RolloutStrategy):
    """Autonomous rollout with on-demand recording via ring buffer.

    The robot runs autonomously while a memory-bounded ring buffer
    captures continuous telemetry.  When the user presses the save key:

    1. The ring buffer is flushed to the dataset (last *Z* seconds).
    2. Live recording continues until the save key is pressed again.
    3. The episode is saved and the ring buffer resumes capturing.

    Requires ``streaming_encoding=True`` (enforced in config validation)
    so that ``dataset.add_frame`` is a non-blocking queue put — flushing
    the entire ring buffer in one tick must not stall the control loop.
    """

    config: HighlightStrategyConfig

    def __init__(self, config: HighlightStrategyConfig) -> None:
        super().__init__(config)
        self._ring: RolloutRingBuffer | None = None
        self._listener = None
        self._save_requested = ThreadingEvent()
        self._recording_live = ThreadingEvent()
        self._push_requested = ThreadingEvent()
        self._push_executor: ThreadPoolExecutor | None = None
        self._pending_push: Future[None] | None = None
        self._episode_lock = Lock()

    def setup(self, ctx: RolloutContext) -> None:
        """Initialise the inference engine, ring buffer, and keyboard listener."""
        self._init_engine(ctx)

        self._ring = RolloutRingBuffer(
            max_seconds=self.config.ring_buffer_seconds,
            max_memory_mb=self.config.ring_buffer_max_memory_mb,
            fps=ctx.runtime.cfg.fps,
        )

        self._push_executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix="highlight-push")
        logger.info(
            "Ring buffer initialized (max_seconds=%.0f, max_memory=%.0fMB)",
            self.config.ring_buffer_seconds,
            self.config.ring_buffer_max_memory_mb,
        )
        self._setup_keyboard(ctx.runtime.shutdown_event)
        logger.info(
            "Highlight strategy ready (buffer=%.0fs, save='%s', push='%s')",
            self.config.ring_buffer_seconds,
            self.config.save_key,
            self.config.push_key,
        )

    def run(self, ctx: RolloutContext) -> None:
        """Run the autonomous loop, buffering frames and recording on demand."""
        engine = self._require_engine()
        interpolator = self._require_interpolator()
        ring = self._ring
        if ring is None:
            raise RuntimeError(f"{type(self).__name__}: ring buffer not attached; call setup() first")
        cfg = ctx.runtime.cfg
        robot = ctx.hardware.robot_wrapper
        dataset = self._require_dataset(ctx.data)
        features = ctx.data.dataset_features

        timer = CycleTimer(cfg.fps, interpolator.multiplier)

        engine.resume()
        play_sounds = cfg.play_sounds

        start_time = time.perf_counter()
        task_str = cfg.dataset.single_task if cfg.dataset else cfg.task
        logger.info("Highlight strategy recording started (press '%s' to save)", self.config.save_key)

        with VideoEncodingManager(dataset):
            try:
                while not ctx.runtime.shutdown_event.is_set():
                    timer.tick(new_cycle=interpolator.needs_new_action())

                    if cfg.duration > 0 and (time.perf_counter() - start_time) >= cfg.duration:
                        logger.info("Duration limit reached (%.0fs)", cfg.duration)
                        break

                    with timer.section("observe"):
                        obs = robot.get_observation()
                    with timer.section("process_obs"):
                        obs_processed = self._process_observation_and_notify(ctx.processors, obs)

                    if self._handle_warmup(cfg.use_torch_compile, timer):
                        continue

                    action_dict = send_next_action(obs_processed, obs, ctx, interpolator, timer)

                    if action_dict is not None:
                        with timer.section("telemetry"):
                            self._log_telemetry(obs_processed, action_dict, ctx.runtime)

                        if self._push_requested.is_set():
                            self._push_requested.clear()
                            logger.info("Push requested by user")
                            self._background_push(dataset, cfg)

                        # Record (buffer or live) once per interpolation cycle so
                        # the frame cadence matches the dataset's declared fps and
                        # the ring buffer's fps-based sizing; interpolated ticks
                        # only send commands to the robot.  Save toggles are also
                        # handled here so an episode boundary always lands on a
                        # recorded frame.
                        if interpolator.emitted_policy_action:
                            with timer.section("record"):
                                obs_frame = build_dataset_frame(features, obs_processed, prefix=OBS_STR)
                                action_frame = build_dataset_frame(features, action_dict, prefix=ACTION)
                                frame = {**obs_frame, **action_frame, "task": task_str}

                                toggled = False
                                frame_consumed = False
                                # NOTE: ``is_set()`` then ``clear()`` is not atomic
                                # against the keyboard thread setting the flag again
                                # in between — but that is benign: we lose at most one
                                # toggle, processed on the next iteration.
                                if self._save_requested.is_set():
                                    self._save_requested.clear()
                                    toggled = True
                                    if not self._recording_live.is_set():
                                        logger.info(
                                            "Flushing ring buffer (%d frames) + starting live recording",
                                            len(ring),
                                        )
                                        for buffered_frame in ring.drain():
                                            dataset.add_frame(buffered_frame)
                                        self._recording_live.set()
                                    else:
                                        dataset.add_frame(frame)
                                        with self._episode_lock:
                                            dataset.save_episode()
                                        logger.info("Episode saved (total: %d)", dataset.num_episodes)
                                        log_say(
                                            f"Episode {dataset.num_episodes} saved",
                                            play_sounds,
                                        )
                                        self._recording_live.clear()
                                        frame_consumed = True

                                if not frame_consumed:
                                    if self._recording_live.is_set():
                                        dataset.add_frame(frame)
                                    else:
                                        ring.append(frame)

                            # Draining the ring buffer and finalising an episode both
                            # block for a good fraction of a second inside the timed
                            # loop body.  They are operator events, not steady-state
                            # cost, so drop the partial group and the gap they opened.
                            if toggled:
                                if frame_consumed:
                                    timer.log_episode_summary(f"episode {dataset.num_episodes}")
                                timer.restart()

                    timer.wait()

            finally:
                logger.info("Highlight control loop ended")
                timer.log_run_summary()
                if self._recording_live.is_set():
                    logger.info("Saving in-progress live episode")
                    with contextlib.suppress(Exception), self._episode_lock:
                        dataset.save_episode()

    def teardown(self, ctx: RolloutContext) -> None:
        """Stop listeners, finalise the dataset, and disconnect hardware."""
        play_sounds = ctx.runtime.cfg.play_sounds
        logger.info("Stopping highlight recording")
        log_say("Stopping highlight recording", play_sounds)

        if self._listener is not None:
            logger.info("Stopping keyboard listener")
            self._listener.stop()

        if self._push_executor is not None:
            logger.info("Shutting down push executor (waiting for pending pushes)...")
            self._push_executor.shutdown(wait=True)
            self._push_executor = None

        if ctx.data.dataset is not None:
            logger.info("Finalizing dataset...")
            ctx.data.dataset.finalize()
            if ctx.runtime.cfg.dataset and ctx.runtime.cfg.dataset.push_to_hub:
                logger.info("Pushing final dataset to hub...")
                if safe_push_to_hub(
                    ctx.data.dataset,
                    tags=ctx.runtime.cfg.dataset.tags,
                    private=ctx.runtime.cfg.dataset.private,
                ):
                    logger.info("Dataset uploaded to hub")
                    log_say("Dataset uploaded to hub", play_sounds)

        self._teardown_hardware(
            ctx.hardware,
            return_to_initial_position=ctx.runtime.cfg.return_to_initial_position,
        )
        logger.info("Highlight strategy teardown complete")

    def _setup_keyboard(self, shutdown_event: ThreadingEvent) -> None:
        """Set up a keyboard listener for the save and push keys.

        Backend selection (pynput on X11 / trusted-macOS / Windows, a terminal reader on
        Wayland / headless TTY) is delegated to :func:`create_key_listener`.
        """
        save_key = self.config.save_key
        push_key = self.config.push_key

        def dispatch(name: str) -> None:
            """Apply a resolved key name to the highlight events."""
            if name == save_key:
                self._save_requested.set()
            elif name == push_key:
                self._push_requested.set()
            elif name == "esc":
                self._save_requested.clear()
                shutdown_event.set()

        self._listener = create_key_listener(
            dispatch, controls_help=f"save='{save_key}', push='{push_key}', ESC=stop"
        )

    def _background_push(self, dataset: LeRobotDataset, cfg: RolloutConfig) -> None:
        """Queue a Hub push on the single-worker executor."""
        if self._push_executor is None:
            return

        if self._pending_push is not None and not self._pending_push.done():
            logger.info("Previous push still in progress; queueing next")

        def _push() -> None:
            try:
                with self._episode_lock:
                    if safe_push_to_hub(
                        dataset,
                        tags=cfg.dataset.tags if cfg.dataset else None,
                        private=cfg.dataset.private if cfg.dataset else False,
                    ):
                        logger.info("Background push to hub complete")
            except Exception as e:
                logger.error("Background push failed: %s", e)

        self._pending_push = self._push_executor.submit(_push)
        logger.info("Background push task submitted")
