import sys
import threading

import depthai as dai
import numpy as np
import pyvista as pv

from pyvistaqt import BackgroundPlotter
from PySide6.QtWidgets import QApplication

from vtkmodules.util.numpy_support import vtk_to_numpy


# ============================================================
# Settings
# ============================================================

FRAME_SIZE = (640, 400)
FPS = 30

MIN_DEPTH_M = 0.3
MAX_DEPTH_M = 5.0

# 640x400 = 256000 points
#
# 2 -> 128000 points
# 4 ->  64000 points
DISPLAY_STRIDE = 2

# GUI update rate
VIEWER_INTERVAL_MS = 33


# ============================================================
# Convert DepthAI point cloud
# ============================================================

def convert_pointcloud(pcl_data, unit_scale):
    """
    Convert DepthAI point cloud to fixed-size arrays.

    Return:
        xyz  : Nx3 float32 [m]
        rgba : Nx4 uint8
        valid_count : int

    Viewer coordinates:
        X : right
        Y : forward
        Z : up
    """

    raw = np.asarray(
        pcl_data.getPoints(),
        dtype=np.float32
    )

    # --------------------------------------------------------
    # Fixed decimation
    #
    # Do NOT remove invalid points because the VTK point count
    # must remain constant between frames.
    # --------------------------------------------------------

    raw = raw[::DISPLAY_STRIDE]

    points_m = raw * unit_scale

    z = points_m[:, 2]

    valid = np.isfinite(points_m).all(axis=1)

    valid &= z >= MIN_DEPTH_M
    valid &= z <= MAX_DEPTH_M

    # --------------------------------------------------------
    # DepthAI:
    #
    # X = right
    # Y = down
    # Z = forward
    #
    # Viewer:
    #
    # X = right
    # Y = forward
    # Z = up
    # --------------------------------------------------------

    xyz = np.empty_like(
        points_m,
        dtype=np.float32
    )

    xyz[:, 0] = points_m[:, 0]
    xyz[:, 1] = points_m[:, 2]
    xyz[:, 2] = -points_m[:, 1]

    # --------------------------------------------------------
    # RGBA color
    # --------------------------------------------------------

    rgba = np.zeros(
        (len(xyz), 4),
        dtype=np.uint8
    )

    if np.any(valid):

        # Depth normalized to 0 ... 1
        t = (
            z[valid] - MIN_DEPTH_M
        ) / (
            MAX_DEPTH_M - MIN_DEPTH_M
        )

        t = np.clip(
            t,
            0.0,
            1.0
        )

        # Simple near/far color
        rgba[valid, 0] = (
            255 * (1.0 - t)
        ).astype(np.uint8)

        rgba[valid, 1] = 180

        rgba[valid, 2] = (
            255 * t
        ).astype(np.uint8)

        rgba[valid, 3] = 255

        # ----------------------------------------------------
        # Invalid points:
        #
        # Place them at an existing valid point and alpha=0.
        #
        # This avoids NaN and avoids changing the point count.
        # ----------------------------------------------------

        first_valid = np.flatnonzero(valid)[0]

        xyz[~valid] = xyz[first_valid]

    else:

        xyz[:] = 0.0

    return (
        xyz,
        rgba,
        int(np.count_nonzero(valid))
    )


# ============================================================
# DepthAI
# ============================================================

with dai.Pipeline() as pipeline:

    # --------------------------------------------------------
    # Cameras
    # --------------------------------------------------------

    color = pipeline.create(
        dai.node.Camera
    ).build(
        dai.CameraBoardSocket.CAM_A,
        sensorFps=FPS
    )

    left = pipeline.create(
        dai.node.Camera
    ).build(
        dai.CameraBoardSocket.CAM_B,
        sensorFps=FPS
    )

    right = pipeline.create(
        dai.node.Camera
    ).build(
        dai.CameraBoardSocket.CAM_C,
        sensorFps=FPS
    )

    # --------------------------------------------------------
    # StereoDepth
    # --------------------------------------------------------

    stereo = pipeline.create(
        dai.node.StereoDepth
    )

    stereo.setDefaultProfilePreset(
        dai.node.StereoDepth.PresetMode.DEFAULT
    )

    stereo.setRectifyEdgeFillColor(0)

    stereo.enableDistortionCorrection(
        True
    )

    left.requestOutput(
        FRAME_SIZE
    ).link(
        stereo.left
    )

    right.requestOutput(
        FRAME_SIZE
    ).link(
        stereo.right
    )

    # --------------------------------------------------------
    # RGBD
    # --------------------------------------------------------

    rgbd = pipeline.create(
        dai.node.RGBD
    ).build(
        color,
        stereo,
        FRAME_SIZE,
        FPS
    )

    pcl_queue = (
        rgbd.pcl.createOutputQueue(
            maxSize=2,
            blocking=False
        )
    )

    pipeline.start()

    print("DepthAI started")
    print("Waiting for first point cloud...")

    # ========================================================
    # First frame
    # ========================================================

    first_data = pcl_queue.get()

    raw = np.asarray(
        first_data.getPoints(),
        dtype=np.float32
    )

    raw_z = raw[:, 2]

    positive = (
        np.isfinite(raw_z)
        & (raw_z > 0)
    )

    median_z = np.median(
        raw_z[positive]
    )

    # --------------------------------------------------------
    # Detect XYZ unit once
    # --------------------------------------------------------

    if median_z > 100.0:

        unit_scale = 0.001
        unit_name = "millimeter"

    else:

        unit_scale = 1.0
        unit_name = "meter"

    print(
        f"Raw Z range: "
        f"{raw_z[positive].min():.3f} - "
        f"{raw_z[positive].max():.3f}"
    )

    print(
        f"Median Z: {median_z:.3f}"
    )

    print(
        f"Detected unit: {unit_name}"
    )

    first_xyz, first_rgba, valid_count = (
        convert_pointcloud(
            first_data,
            unit_scale
        )
    )

    print(
        f"Viewer points: {len(first_xyz)}"
    )

    print(
        f"Valid points: {valid_count}"
    )

    # ========================================================
    # Shared latest-frame buffer
    #
    # Camera thread -> Qt GUI thread
    # ========================================================

    latest_lock = threading.Lock()

    latest_xyz = first_xyz.copy()
    latest_rgba = first_rgba.copy()

    latest_frame_id = 0

    stop_event = threading.Event()

    # ========================================================
    # Capture thread
    #
    # IMPORTANT:
    # DepthAI acquisition is completely separated from Qt/VTK.
    # ========================================================

    def capture_loop():

        nonlocal_vars = {
            "frame_id": 0
        }

        while not stop_event.is_set():

            try:

                pcl_data = pcl_queue.get()

            except Exception:

                break

            xyz, rgba, count = (
                convert_pointcloud(
                    pcl_data,
                    unit_scale
                )
            )

            with latest_lock:

                # In-place replacement of shared references
                globals_placeholder[0] = xyz
                globals_placeholder[1] = rgba

                nonlocal_vars["frame_id"] += 1

                globals_placeholder[2] = (
                    nonlocal_vars["frame_id"]
                )

    # --------------------------------------------------------
    # Python doesn't have writable nonlocal variables at module
    # scope, therefore store current frame data in this list.
    # --------------------------------------------------------

    globals_placeholder = [
        latest_xyz,
        latest_rgba,
        latest_frame_id
    ]

    capture_thread = threading.Thread(
        target=capture_loop,
        daemon=True
    )

    capture_thread.start()

    # ========================================================
    # Qt
    # ========================================================

    app = QApplication.instance()

    if app is None:

        app = QApplication(sys.argv)

    app.setQuitOnLastWindowClosed(True)

    # ========================================================
    # PyVistaQt BackgroundPlotter
    # ========================================================

    plotter = BackgroundPlotter(
        app=app,
        show=True,
        window_size=(1280, 720),
        title="OAK-D Point Cloud"
    )

    # --------------------------------------------------------
    # Create ONE PolyData
    # --------------------------------------------------------

    cloud = pv.PolyData(
        first_xyz.copy()
    )

    cloud.point_data["rgba"] = (
        first_rgba.copy()
    )

    actor = plotter.add_points(
        cloud,
        scalars="rgba",
        rgba=True,
        style="points",
        point_size=3,
        lighting=False,
        name="oak_pointcloud",
        reset_camera=True
    )

    plotter.show_axes()
    plotter.show_grid()

    plotter.add_text(
        "OAK-D Real-time Point Cloud",
        position="upper_left",
        font_size=10
    )

    plotter.reset_camera()

    # ========================================================
    # IMPORTANT:
    #
    # Get direct NumPy views into VTK memory.
    #
    # We will update THESE arrays, rather than assigning new
    # PyVista arrays every frame.
    # ========================================================

    vtk_points_data = (
        cloud.GetPoints().GetData()
    )

    vtk_rgba_data = (
        cloud.GetPointData().GetArray(
            "rgba"
        )
    )

    vtk_points_np = vtk_to_numpy(
        vtk_points_data
    )

    vtk_rgba_np = vtk_to_numpy(
        vtk_rgba_data
    )

    print()
    print("VTK buffers")
    print(
        "points:",
        vtk_points_np.shape,
        vtk_points_np.dtype
    )

    print(
        "rgba:",
        vtk_rgba_np.shape,
        vtk_rgba_np.dtype
    )

    # ========================================================
    # Qt GUI update callback
    # ========================================================

    displayed_frame_id = [-1]

    update_counter = [0]

    def update_view():

        # Window already closed
        if plotter._closed:
            return

        # --------------------------------------------
        # Get newest frame from capture thread
        # --------------------------------------------

        with latest_lock:

            frame_id = globals_placeholder[2]

            if frame_id == displayed_frame_id[0]:
                return

            xyz = globals_placeholder[0]
            rgba = globals_placeholder[1]

            # Copy because capture thread may replace
            # references while VTK is drawing.
            xyz = xyz.copy()
            rgba = rgba.copy()

        # --------------------------------------------
        # Sanity check
        # --------------------------------------------

        if xyz.shape != vtk_points_np.shape:

            print(
                "Point array shape changed:",
                xyz.shape,
                vtk_points_np.shape
            )

            return

        # ====================================================
        # KEY POINT:
        #
        # Write DIRECTLY into the existing VTK arrays.
        # ====================================================

        vtk_points_np[:] = xyz
        vtk_rgba_np[:] = rgba

        # ----------------------------------------------------
        # Explicitly tell VTK which buffers changed
        # ----------------------------------------------------

        vtk_points_data.Modified()
        vtk_rgba_data.Modified()

        cloud.GetPoints().Modified()
        cloud.GetPointData().Modified()
        cloud.Modified()

        # Mapper update
        actor.mapper.Update()

        # Render
        plotter.render()

        displayed_frame_id[0] = frame_id

        update_counter[0] += 1

        if update_counter[0] % 60 == 0:

            print(
                f"Viewer updates: "
                f"{update_counter[0]}"
            )

    # ========================================================
    # Qt QTimer
    #
    # BackgroundPlotter.add_callback() uses QTimer.
    # ========================================================

    plotter.add_callback(
        update_view,
        interval=VIEWER_INTERVAL_MS
    )

    # ========================================================
    # Closing
    # ========================================================

    def window_closed():

        print("Viewer closing...")

        stop_event.set()

        app.quit()

    plotter.app_window.signal_close.connect(
        window_closed
    )

    print()
    print("Viewer started")
    print("-----------------------------")
    print("Left drag   : rotate")
    print("Mouse wheel : zoom")
    print("Middle drag : pan")
    print("Q / X       : close")
    print()

    # ========================================================
    # Qt event loop
    #
    # This is the ONLY GUI event loop.
    # ========================================================

    try:

        app.exec()

    except KeyboardInterrupt:

        print("Interrupted")

    finally:

        # Stop capture thread
        stop_event.set()

        # Stop DepthAI first so blocking get() is released
        pipeline.stop()

        capture_thread.join(
            timeout=1.0
        )

        print("DepthAI stopped")
        print("Finished")