from __future__ import annotations

import argparse
import time

import cv2
import depthai as dai
import numpy as np

from pose_backends import (
    create_pose_backend
)

from pose3d_utils import (
    JOINT_NAMES,
    SKELETON,
    JointTracker3D,
    pixel_to_xyz,
    robust_joint_depth,
    xyz_to_pixel,
)


# ============================================================
# Configuration
# ============================================================

FRAME_SIZE = (
    640,
    400
)

CAMERA_FPS = 30

DISPLAY_SCALE = 1.5


PERSON_CONF = 0.30
KEYPOINT_CONF = 0.35


MIN_DEPTH_M = 0.30
MAX_DEPTH_M = 5.00


DEPTH_ROI_RADIUS = 4

MIN_DEPTH_PIXELS = 4

PREVIOUS_DEPTH_GATE_M = 0.40

CENTER_DEPTH_GATE_M = 0.25


FILTER_ALPHA = 0.55
FILTER_BETA = 0.08

MAX_JOINT_SPEED_MPS = 6.0

BASE_POSITION_GATE_M = 0.10

MAX_MISSED_FRAMES = 6


SHOW_XYZ = True


# ============================================================
# Colors
# ============================================================

COLOR_MEASURED = (
    0,
    255,
    0
)

COLOR_PREDICTED = (
    0,
    200,
    255
)

COLOR_SKELETON = (
    255,
    180,
    0
)

COLOR_TEXT = (
    255,
    255,
    255
)

COLOR_BOX = (
    255,
    0,
    255
)


# ============================================================
# Command line
# ============================================================

def parse_args():

    parser = argparse.ArgumentParser(
        description=(
            "OAK-D RGB-D "
            "3D pose estimation"
        )
    )

    parser.add_argument(
        "--backend",

        choices=[
            "ultralytics",
            "mediapipe"
        ],

        default="ultralytics",

        help=(
            "Pose estimation backend"
        ),
    )

    parser.add_argument(
        "--mediapipe-model",

        choices=[
            "lite",
            "full",
            "heavy"
        ],

        default="lite",

        help=(
            "MediaPipe model"
        ),
    )

    parser.add_argument(
        "--yolo-model",

        default=(
            "yolo26n-pose.pt"
        ),

        help=(
            "Ultralytics pose model"
        ),
    )

    parser.add_argument(
        "--yolo-imgsz",

        type=int,

        default=640,

        help=(
            "Ultralytics input size"
        ),
    )

    return parser.parse_args()


# ============================================================
# Trackers
# ============================================================

def create_joint_trackers():

    return [

        JointTracker3D(

            nominal_fps=
                CAMERA_FPS,

            alpha=
                FILTER_ALPHA,

            beta=
                FILTER_BETA,

            max_speed_mps=
                MAX_JOINT_SPEED_MPS,

            base_gate_m=
                BASE_POSITION_GATE_M,

            max_missed_frames=
                MAX_MISSED_FRAMES,
        )

        for _ in range(17)
    ]


# ============================================================
# Draw skeleton
# ============================================================

def draw_pose(
    frame,
    filtered_xyz,
    measurement_accepted,
    K,
):

    h, w = frame.shape[:2]

    display_uv = [

        xyz_to_pixel(
            xyz,
            K
        )

        for xyz
        in filtered_xyz
    ]

    # --------------------------------------------------------
    # Skeleton
    # --------------------------------------------------------

    for a, b in SKELETON:

        pa = display_uv[a]
        pb = display_uv[b]

        if (
            pa is None
            or pb is None
        ):
            continue

        if not (

            0 <= pa[0] < w
            and
            0 <= pa[1] < h

            and

            0 <= pb[0] < w
            and
            0 <= pb[1] < h

        ):
            continue

        cv2.line(
            frame,
            pa,
            pb,
            COLOR_SKELETON,
            2,
            cv2.LINE_AA,
        )

    # --------------------------------------------------------
    # Joints
    # --------------------------------------------------------

    for joint_id, xyz in enumerate(
        filtered_xyz
    ):

        uv = display_uv[
            joint_id
        ]

        if (
            xyz is None
            or uv is None
        ):
            continue

        u, v = uv

        if not (
            0 <= u < w
            and
            0 <= v < h
        ):
            continue

        if measurement_accepted[
            joint_id
        ]:

            joint_color = (
                COLOR_MEASURED
            )

        else:

            joint_color = (
                COLOR_PREDICTED
            )

        cv2.circle(
            frame,
            (u, v),
            4,
            joint_color,
            -1,
            cv2.LINE_AA,
        )

        # ----------------------------------------------------
        # XYZ text
        # ----------------------------------------------------

        if SHOW_XYZ:

            x, y, z = xyz

            text = (
                f"{joint_id}:"
                f"{JOINT_NAMES[joint_id]} "
                f"({x:+.2f},"
                f"{y:+.2f},"
                f"{z:.2f})"
            )

            dy = (
                -7
                if joint_id % 2 == 0
                else 12
            )

            tx = min(
                max(
                    u + 5,
                    0
                ),

                max(
                    0,
                    w - 230
                )
            )

            ty = min(
                max(
                    v + dy,
                    10
                ),

                h - 5
            )

            cv2.putText(
                frame,
                text,
                (tx, ty),

                cv2.FONT_HERSHEY_SIMPLEX,

                0.30,

                COLOR_TEXT,

                1,

                cv2.LINE_AA,
            )


# ============================================================
# Main
# ============================================================

def main():

    args = parse_args()

    # --------------------------------------------------------
    # Pose backend
    # --------------------------------------------------------

    pose_backend = (
        create_pose_backend(

            args.backend,

            ultralytics_model=
                args.yolo_model,

            ultralytics_imgsz=
                args.yolo_imgsz,

            person_conf=
                PERSON_CONF,

            mediapipe_model=
                args.mediapipe_model,
        )
    )

    joint_trackers = (
        create_joint_trackers()
    )

    try:

        with dai.Pipeline() as pipeline:

            # =================================================
            # RGB camera
            # =================================================

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

                sensorFps=
                    CAMERA_FPS,
            )

            # =================================================
            # Stereo cameras
            # =================================================

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

                sensorFps=
                    CAMERA_FPS,
            )

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

                sensorFps=
                    CAMERA_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,
                CAMERA_FPS,
            )

            rgbd_queue = (
                rgbd.rgbd
                .createOutputQueue(

                    maxSize=2,

                    blocking=False,
                )
            )

            pipeline.start()

            print()
            print(
                "OAK-D 3D Pose started"
            )

            print(
                f"Pose backend: "
                f"{pose_backend.name}"
            )

            print(
                "Q / ESC : quit"
            )

            print()

            intrinsic_printed = False

            fps_value = 0.0

            last_loop_time = (
                time.perf_counter()
            )

            frame_counter = 0

            # =================================================
            # Main loop
            # =================================================

            while pipeline.isRunning():

                # --------------------------------------------
                # Synchronized RGB-D
                # --------------------------------------------

                rgbd_data = (
                    rgbd_queue.get()
                )

                rgb_msg = (
                    rgbd_data
                    .getRGBFrame()
                )

                depth_msg = (
                    rgbd_data
                    .getDepthFrame()
                )

                if (
                    rgb_msg is None
                    or depth_msg is None
                ):
                    continue

                frame = (
                    rgb_msg.getCvFrame()
                )

                depth_mm = (
                    depth_msg.getCvFrame()
                )

                if (
                    depth_mm.shape[:2]
                    != frame.shape[:2]
                ):

                    raise RuntimeError(
                        "RGB/Depth size mismatch: "
                        f"RGB={frame.shape[:2]}, "
                        f"Depth={depth_mm.shape[:2]}"
                    )

                h, w = frame.shape[:2]

                # --------------------------------------------
                # RGB intrinsic matrix
                # --------------------------------------------

                K = np.asarray(

                    rgb_msg
                    .getTransformation()
                    .getIntrinsicMatrix(),

                    dtype=np.float64,
                )

                if not intrinsic_printed:

                    print(
                        f"RGB size: "
                        f"{w} x {h}"
                    )

                    print(
                        "Intrinsic matrix:"
                    )

                    print(K)

                    intrinsic_printed = True

                # =============================================
                # Pose inference
                # =============================================

                pose = (
                    pose_backend.infer(
                        frame
                    )
                )

                measurements = [
                    None
                ] * 17

                now = (
                    time.perf_counter()
                )

                # =============================================
                # 2D Pose -> Depth -> XYZ
                # =============================================

                if pose is not None:

                    # ----------------------------------------
                    # Bounding box
                    # ----------------------------------------

                    if pose.bbox is not None:

                        box = (
                            pose.bbox
                            .astype(int)
                        )

                        cv2.rectangle(
                            frame,

                            (
                                box[0],
                                box[1]
                            ),

                            (
                                box[2],
                                box[3]
                            ),

                            COLOR_BOX,

                            1,
                        )

                    # ----------------------------------------
                    # 17 joints
                    # ----------------------------------------

                    for joint_id in range(
                        17
                    ):

                        conf = float(
                            pose.conf[
                                joint_id
                            ]
                        )

                        if (
                            conf
                            < KEYPOINT_CONF
                        ):
                            continue

                        u = float(
                            pose.xy[
                                joint_id,
                                0
                            ]
                        )

                        v = float(
                            pose.xy[
                                joint_id,
                                1
                            ]
                        )

                        if not (
                            0 <= u < w
                            and
                            0 <= v < h
                        ):
                            continue

                        previous_z = (
                            joint_trackers[
                                joint_id
                            ]
                            .get_predicted_z()
                        )

                        # ------------------------------------
                        # Robust ROI depth
                        # ------------------------------------

                        z_m = robust_joint_depth(

                            depth_mm,

                            u,
                            v,

                            previous_z_m=
                                previous_z,

                            min_depth_m=
                                MIN_DEPTH_M,

                            max_depth_m=
                                MAX_DEPTH_M,

                            roi_radius=
                                DEPTH_ROI_RADIUS,

                            min_pixels=
                                MIN_DEPTH_PIXELS,

                            previous_depth_gate_m=
                                PREVIOUS_DEPTH_GATE_M,

                            center_depth_gate_m=
                                CENTER_DEPTH_GATE_M,
                        )

                        if z_m is None:
                            continue

                        # ------------------------------------
                        # Pixel -> XYZ
                        # ------------------------------------

                        measurements[
                            joint_id
                        ] = pixel_to_xyz(
                            u,
                            v,
                            z_m,
                            K,
                        )

                # =============================================
                # Temporal 3D filtering
                # =============================================

                filtered_xyz = [
                    None
                ] * 17

                measurement_accepted = [
                    False
                ] * 17

                for joint_id in range(
                    17
                ):

                    xyz, accepted = (
                        joint_trackers[
                            joint_id
                        ].update(
                            measurements[
                                joint_id
                            ],
                            now,
                        )
                    )

                    filtered_xyz[
                        joint_id
                    ] = xyz

                    measurement_accepted[
                        joint_id
                    ] = accepted

                # =============================================
                # Draw pose
                # =============================================

                draw_pose(
                    frame,
                    filtered_xyz,
                    measurement_accepted,
                    K,
                )

                # =============================================
                # FPS
                # =============================================

                loop_now = (
                    time.perf_counter()
                )

                dt_loop = (
                    loop_now
                    - last_loop_time
                )

                last_loop_time = (
                    loop_now
                )

                if dt_loop > 0:

                    instant_fps = (
                        1.0 / dt_loop
                    )

                    if fps_value == 0.0:

                        fps_value = (
                            instant_fps
                        )

                    else:

                        fps_value = (
                            0.9 * fps_value
                            +
                            0.1 * instant_fps
                        )

                # =============================================
                # Status
                # =============================================

                if (
                    args.backend
                    == "mediapipe"
                ):

                    backend_detail = (
                        args.mediapipe_model
                    )

                else:

                    backend_detail = (
                        args.yolo_model
                    )

                status = (
                    f"{pose_backend.name}/"
                    f"{backend_detail}  "
                    f"FPS {fps_value:.1f}  "
                    f"Pose "
                    f"{pose_backend.last_inference_ms:.1f} ms  "
                    f"XYZ[m]: "
                    f"X=right Y=down Z=forward"
                )

                cv2.rectangle(
                    frame,

                    (0, 0),

                    (w, 24),

                    (0, 0, 0),

                    -1,
                )

                cv2.putText(
                    frame,

                    status,

                    (8, 17),

                    cv2.FONT_HERSHEY_SIMPLEX,

                    0.40,

                    (
                        255,
                        255,
                        255
                    ),

                    1,

                    cv2.LINE_AA,
                )

                # =============================================
                # Display
                # =============================================

                if (
                    DISPLAY_SCALE
                    != 1.0
                ):

                    display_frame = (
                        cv2.resize(

                            frame,

                            None,

                            fx=
                                DISPLAY_SCALE,

                            fy=
                                DISPLAY_SCALE,

                            interpolation=
                                cv2.INTER_LINEAR,
                        )
                    )

                else:

                    display_frame = (
                        frame
                    )

                cv2.imshow(
                    "OAK-D 3D Pose",
                    display_frame,
                )

                # =============================================
                # Benchmark output
                # =============================================

                frame_counter += 1

                if (
                    frame_counter
                    % 60
                    == 0
                ):

                    print(
                        f"{pose_backend.name}: "
                        f"pose="
                        f"{pose_backend.last_inference_ms:.1f} ms, "
                        f"total="
                        f"{fps_value:.1f} fps"
                    )

                # =============================================
                # Keyboard
                # =============================================

                key = (
                    cv2.waitKey(1)
                    & 0xFF
                )

                if (
                    key == ord("q")
                    or key == 27
                ):
                    break

            pipeline.stop()

    finally:

        pose_backend.close()

        cv2.destroyAllWindows()


if __name__ == "__main__":

    main()