#!/usr/bin/env python

# Copyright 2024 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.
import logging
import multiprocessing
import multiprocessing.queues
import queue
import threading
from pathlib import Path

import numpy as np
import PIL.Image
import torch

logger = logging.getLogger(__name__)

# One pending write: (image, destination path, PNG compress level). ``None`` is the stop sentinel.
ImageWriterItem = tuple[np.ndarray | PIL.Image.Image, Path, int]
ImageWriterQueue = (
    queue.Queue[ImageWriterItem | None] | multiprocessing.queues.JoinableQueue[ImageWriterItem | None]
)


def safe_stop_image_writer(func):
    def wrapper(*args, **kwargs):
        try:
            return func(*args, **kwargs)
        except BaseException:
            dataset = kwargs.get("dataset")
            writer = getattr(dataset, "writer", None) if dataset else None
            if writer is not None and writer.image_writer is not None:
                logger.warning("Waiting for image writer to terminate...")
                writer.image_writer.stop()
            raise

    return wrapper


def squeeze_single_channel(array: np.ndarray) -> np.ndarray:
    """Drop a leading or trailing singleton channel dim: ``(1, H, W)`` / ``(H, W, 1)`` -> ``(H, W)``.

    Unlike ``array.squeeze()``, this only removes the channel axis, never an ``H`` or ``W`` of size 1.
    """
    if array.ndim == 3:
        if array.shape[0] == 1:
            return array[0]
        if array.shape[-1] == 1:
            return array[..., 0]
    return array


def image_array_to_pil_image(image_array: np.ndarray, range_check: bool = True) -> PIL.Image.Image:
    """Convert a NumPy array to a PIL Image, preserving precision for grayscale.

    Behaviour by shape:

    - ``(H, W)`` or ``(1, H, W)`` / ``(H, W, 1)``: single-channel grayscale.
      The native dtype is preserved using the matching PIL mode
      (``I;16`` / ``F``). This is the path used for raw depth maps (no rescaling, clamping, or downcasting)
    - ``(3, H, W)`` / ``(H, W, 3)``: RGB. Channels-first inputs are transposed
      to channels-last. Float inputs in ``[0, 1]`` are scaled to ``uint8``
      (existing behaviour, gated by ``range_check``).

    Other shapes / channel counts raise ``NotImplementedError`` or
    ``ValueError``.
    """
    # TODO(CarolinePascal): 4 dimensions RGB-D images
    if image_array.ndim not in (2, 3):
        raise ValueError(f"The array has {image_array.ndim} dimensions, but 2 or 3 is expected for an image.")

    # Squeeze 3D single-channel inputs to 2D so depth maps work whether the
    # caller emits (H, W), (1, H, W), or (H, W, 1).
    image_array = squeeze_single_channel(image_array)

    if image_array.ndim == 2:
        if image_array.dtype not in [np.uint16, np.float32]:
            raise ValueError(
                f"Unsupported single-channel image dtype: {image_array.dtype}. "
                f"Supported dtypes: {sorted(str(d) for d in [np.uint16, np.float32])}."
            )
        return PIL.Image.fromarray(np.ascontiguousarray(image_array))

    # 3D path: must be RGB (3 channels), channels-first or channels-last.
    if image_array.shape[0] == 3:
        # Transpose from pytorch convention (C, H, W) to (H, W, C)
        image_array = image_array.transpose(1, 2, 0)

    elif image_array.shape[-1] != 3:
        raise NotImplementedError(
            f"The image has {image_array.shape[-1]} channels, but 3 is required for now."
        )

    if image_array.dtype != np.uint8:
        if range_check:
            max_ = image_array.max().item()
            min_ = image_array.min().item()
            if max_ > 1.0 or min_ < 0.0:
                raise ValueError(
                    "The image data type is float, which requires values in the range [0.0, 1.0]. "
                    f"However, the provided range is [{min_}, {max_}]. Please adjust the range or "
                    "provide a uint8 image with values in the range [0, 255]."
                )

        image_array = (image_array * 255).astype(np.uint8)

    return PIL.Image.fromarray(image_array)


def save_kwargs_for_path(fpath: Path, compress_level: int) -> dict:
    """Pick the right format-specific kwargs for :meth:`PIL.Image.Image.save`.

    PNG uses ``compress_level`` (0-9, zlib). TIFF uses ``compression`` (raw) for lossless raw depth maps.
    """
    suffix = Path(fpath).suffix.lower()
    if suffix == ".png":
        return {"compress_level": compress_level}
    if suffix in (".tif", ".tiff"):
        return {"compression": "raw"}
    else:
        raise ValueError(f"Unsupported image file extension: {suffix}")


def write_image(image: np.ndarray | PIL.Image.Image, fpath: Path, compress_level: int = 1):
    """
    Saves a NumPy array or PIL Image to a file.

    This function handles both NumPy arrays and PIL Image objects, converting
    the former to a PIL Image before saving. It includes error handling for
    the save operation. The output format is inferred from the *fpath*
    extension: ``.png`` → PNG with ``compress_level``, ``.tiff`` / ``.tif``
    → lossless raw depth maps (TIFF).

    Args:
        image (np.ndarray | PIL.Image.Image): The image data to save.
        fpath (Path): The destination file path for the image.
        compress_level (int, optional): The compression level for the saved
            image, as used by PIL.Image.save(). Defaults to 1.
            Refer to: https://github.com/huggingface/lerobot/pull/2135
            for more details on the default value rationale.

    Raises:
        TypeError: If the input 'image' is not a NumPy array or a
            PIL.Image.Image object.

    Side Effects:
        Logs an error message if the image writing process fails for any reason.
    """
    try:
        if isinstance(image, np.ndarray):
            img = image_array_to_pil_image(image)
        elif isinstance(image, PIL.Image.Image):
            img = image
        else:
            raise TypeError(f"Unsupported image type: {type(image)}")
        img.save(fpath, **save_kwargs_for_path(fpath, compress_level))
    except Exception as e:
        logger.error("Error writing image %s: %s", fpath, e)


def worker_thread_loop(queue: ImageWriterQueue) -> None:
    while True:
        item = queue.get()
        if item is None:
            queue.task_done()
            break
        image_array, fpath, compress_level = item
        write_image(image_array, fpath, compress_level)
        queue.task_done()


def worker_process(queue: ImageWriterQueue, num_threads: int) -> None:
    threads = []
    for _ in range(num_threads):
        t = threading.Thread(target=worker_thread_loop, args=(queue,))
        t.daemon = True
        t.start()
        threads.append(t)
    for t in threads:
        t.join()


class AsyncImageWriter:
    """
    This class abstract away the initialisation of processes or/and threads to
    save images on disk asynchronously, which is critical to control a robot and record data
    at a high frame rate.

    When `num_processes=0`, it creates a threads pool of size `num_threads`.
    When `num_processes>0`, it creates processes pool of size `num_processes`, where each subprocess starts
    their own threads pool of size `num_threads`.

    The optimal number of processes and threads depends on your computer capabilities.
    We advise to use 4 threads per camera with 0 processes. If the fps is not stable, try to increase or lower
    the number of threads. If it is still not stable, try to use 1 subprocess, or more.
    """

    def __init__(self, num_processes: int = 0, num_threads: int = 1) -> None:
        self.num_processes = num_processes
        self.num_threads = num_threads
        self.queue: ImageWriterQueue
        self.threads: list[threading.Thread] = []
        self.processes: list[multiprocessing.Process] = []
        self._stopped = False

        if num_threads <= 0 and num_processes <= 0:
            raise ValueError("Number of threads and processes must be greater than zero.")

        if self.num_processes == 0:
            # Use threading
            self.queue = queue.Queue()
            for _ in range(self.num_threads):
                t = threading.Thread(target=worker_thread_loop, args=(self.queue,))
                t.daemon = True
                t.start()
                self.threads.append(t)
        else:
            # Use multiprocessing
            self.queue = multiprocessing.JoinableQueue()
            for _ in range(self.num_processes):
                p = multiprocessing.Process(target=worker_process, args=(self.queue, self.num_threads))
                p.daemon = True
                p.start()
                self.processes.append(p)

    def save_image(
        self, image: torch.Tensor | np.ndarray | PIL.Image.Image, fpath: Path, compress_level: int = 1
    ) -> None:
        if isinstance(image, torch.Tensor):
            # Convert tensor to numpy array to minimize main process time
            image = image.cpu().numpy()
        self.queue.put((image, fpath, compress_level))

    def wait_until_done(self) -> None:
        self.queue.join()

    def stop(self) -> None:
        if self._stopped:
            return

        if isinstance(self.queue, queue.Queue):
            for _ in self.threads:
                self.queue.put(None)
            for t in self.threads:
                t.join()
        else:
            num_nones = self.num_processes * self.num_threads
            for _ in range(num_nones):
                self.queue.put(None)
            for p in self.processes:
                p.join()
                if p.is_alive():
                    p.terminate()
            self.queue.close()
            self.queue.join_thread()

        self._stopped = True
